from __future__ import annotations

import os
import re
import shutil
import subprocess
from config import load_config
from logger import get_logger

log = get_logger("discovery.databases")


def run_mysql_query(query: str) -> str | None:
    """Run a MySQL query using the CLI and return output or None on failure."""
    if not shutil.which("mysql"):
        log.debug("mysql CLI not found on system")
        return None

    try:
        config = load_config()
    except Exception as exc:
        log.debug("failed to load config in database discovery: %s", exc)
        return None

    env = os.environ.copy()
    if config.mysql_password:
        env["MYSQL_PWD"] = config.mysql_password

    cmd = ["mysql", f"-u{config.mysql_user}"]
    if config.mysql_socket:
        cmd.append(f"--socket={config.mysql_socket}")
    
    # -N: skip column names, -B: batch mode (tab-separated)
    cmd.extend(["-N", "-B", "-e", query])

    try:
        res = subprocess.run(
            cmd,
            env=env,
            capture_output=True,
            text=True,
            timeout=10,
            check=False,
        )
        if res.returncode != 0:
            log.debug("mysql command returned non-zero code %d: %s", res.returncode, res.stderr)
            return None
        return res.stdout
    except (subprocess.SubprocessError, OSError) as exc:
        log.debug("mysql command execution failed: %s", exc)
        return None


def discover_database_schemas() -> list[dict]:
    """Discover logical database schemas, their metadata, and sizes."""
    db_output = run_mysql_query("SHOW DATABASES;")
    if not db_output:
        return []

    system_dbs = {"information_schema", "performance_schema", "sys", "mysql"}
    dbs = []
    for line in db_output.splitlines():
        db_name = line.strip()
        if not db_name or db_name in system_dbs:
            continue
        dbs.append(db_name)

    schemas = []
    for db in dbs:
        # Get charset & collation
        meta_query = f"SELECT DEFAULT_CHARACTER_SET_NAME, DEFAULT_COLLATION_NAME FROM information_schema.SCHEMATA WHERE SCHEMA_NAME='{db}';"
        meta_out = run_mysql_query(meta_query)
        charset, collation = None, None
        if meta_out:
            parts = meta_out.strip().split("\t")
            if len(parts) >= 2:
                charset, collation = parts[0], parts[1]
            elif len(parts) == 1:
                charset = parts[0]

        # Get size in MB
        size_query = f"SELECT SUM(data_length+index_length)/1024/1024 FROM information_schema.TABLES WHERE table_schema='{db}';"
        size_out = run_mysql_query(size_query)
        size_mb = 0.0
        if size_out:
            try:
                val = size_out.strip()
                if val and val != "NULL":
                    size_mb = float(val)
            except ValueError:
                pass

        schemas.append({
            "name": db,
            "charset": charset,
            "db_collation": collation,
            "size_mb": size_mb,
            "engine_type": "mysql",
            "category": "database_schemas",
        })

    return schemas


def discover_database_users() -> list[dict]:
    """Discover database users, hosts, database mappings, and grant option status."""
    users_out = run_mysql_query("SELECT User, Host FROM mysql.user;")
    if not users_out:
        return []

    system_users = {"root", "mysql.session", "mysql.sys", "mysql.infoschema", "debian-sys-maint"}
    users = []
    for line in users_out.splitlines():
        parts = line.strip().split("\t")
        if len(parts) < 2:
            continue
        username, host = parts[0], parts[1]
        if username in system_users or username.startswith("mysql."):
            continue
        users.append((username, host))

    db_users = []
    for username, host in users:
        grants_query = f"SHOW GRANTS FOR '{username}'@'{host}';"
        grants_out = run_mysql_query(grants_query)
        
        assigned_dbs = set()
        has_grant = False
        if grants_out:
            for grant_line in grants_out.splitlines():
                if "WITH GRANT OPTION" in grant_line.upper():
                    has_grant = True

                # Extract database name from grants (e.g. GRANT ALL PRIVILEGES ON `db`.* TO ...)
                match = re.search(r"ON\s+`?([^`\s\.]+)(?:`?\.`?[*`]|`?\.\*)", grant_line, re.IGNORECASE)
                if match:
                    db_name = match.group(1)
                    if db_name != "*":
                        assigned_dbs.add(db_name)

        db_users.append({
            "username": username,
            "host": host,
            "databases": sorted(list(assigned_dbs)),
            "has_grant_option": has_grant,
            "category": "database_users",
        })

    return db_users
