package config

import (
	"errors"
	"fmt"
	"net/url"
	"os"
	"strings"

	"gorm.io/driver/mysql"
	"gorm.io/gorm"
)

var DB *gorm.DB

func InitDB() error {
	databaseURL := os.Getenv("DATABASE_URL")
	if databaseURL == "" {
		return errors.New("DATABASE_URL is required")
	}
	if err := ensureDatabase(databaseURL); err != nil {
		return err
	}
	dsn, err := mysqlDSN(databaseURL)
	if err != nil {
		return err
	}
	DB, err = gorm.Open(mysql.Open(dsn), &gorm.Config{})
	return err
}

func RemoveLegacyEmailColumn() error {
	if DB == nil {
		return errors.New("database is not initialized")
	}
	var indexCount int64
	if err := DB.Raw("SELECT COUNT(*) FROM information_schema.statistics WHERE table_schema = DATABASE() AND table_name = 'users' AND index_name = ?", "idx_users_email").Scan(&indexCount).Error; err != nil {
		return err
	}
	if indexCount > 0 {
		if err := DB.Migrator().DropIndex("users", "idx_users_email"); err != nil {
			return err
		}
	}
	var columnCount int64
	if err := DB.Raw("SELECT COUNT(*) FROM information_schema.columns WHERE table_schema = DATABASE() AND table_name = 'users' AND column_name = ?", "email").Scan(&columnCount).Error; err != nil {
		return err
	}
	if columnCount > 0 {
		if err := DB.Migrator().DropColumn("users", "email"); err != nil {
			return err
		}
	}
	return nil
}

func mysqlDSN(databaseURL string) (string, error) {
	parsed, err := url.Parse(databaseURL)
	if err != nil {
		return "", err
	}
	if parsed.Scheme != "mysql" || parsed.Hostname() == "" || parsed.Path == "" || parsed.Path == "/" {
		return "", errors.New("DATABASE_URL must use mysql://user:password@host:port/database")
	}
	user := parsed.User.Username()
	password, _ := parsed.User.Password()
	port := parsed.Port()
	if port == "" {
		port = "3306"
	}
	return fmt.Sprintf("%s:%s@tcp(%s:%s)%s?charset=utf8mb4&parseTime=True&loc=Local", user, password, parsed.Hostname(), port, parsed.Path), nil
}

func ensureDatabase(databaseURL string) error {
	parsed, err := url.Parse(databaseURL)
	if err != nil {
		return err
	}
	databaseName := strings.TrimPrefix(parsed.Path, "/")
	if databaseName == "" || strings.Contains(databaseName, "/") || strings.ContainsAny(databaseName, "`;") {
		return errors.New("invalid database name in DATABASE_URL")
	}
	user := parsed.User.Username()
	password, _ := parsed.User.Password()
	port := parsed.Port()
	if port == "" {
		port = "3306"
	}
	serverDSN := fmt.Sprintf("%s:%s@tcp(%s:%s)/?charset=utf8mb4&parseTime=True&loc=Local", user, password, parsed.Hostname(), port)
	serverDB, err := gorm.Open(mysql.Open(serverDSN), &gorm.Config{})
	if err != nil {
		return err
	}
	return serverDB.Exec("CREATE DATABASE IF NOT EXISTS `" + databaseName + "` CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci").Error
}
