# process_manager.py — subprocess lifecycle for TDX BLACK HOST
import os
import sys
import signal
import logging
import subprocess
import psutil
from pathlib import Path
from typing import Optional, Dict, Tuple
from datetime import datetime

import database as db
from config import LOGS_DIR, PROJECTS_DIR

logger = logging.getLogger('TDX.Process')

# in-memory map: project_id -> subprocess.Popen
_processes: Dict[int, subprocess.Popen] = {}


def _log_path(project_id: int) -> Path:
    return LOGS_DIR / f"project_{project_id}.log"


def start_process(project_id: int) -> Tuple[bool, str]:
    project = db.get_project(project_id)
    if not project:
        return False, "Project not found."

    if project_id in _processes and _processes[project_id].poll() is None:
        return False, "Process already running."

    file_path = Path(project['file_path'])
    if not file_path.exists():
        return False, "Project file missing from disk."

    log_path = _log_path(project_id)
    suffix = file_path.suffix.lower()

    try:
        log_handle = open(log_path, 'a', buffering=1)

        if suffix == '.py':
            cmd = [sys.executable, str(file_path)]
        elif suffix == '.js':
            cmd = ['node', str(file_path)]
        elif suffix == '.zip':
            # look for main.py or index.js inside extracted dir
            extracted = PROJECTS_DIR / f"user_{project['owner_id']}" / f"zip_{project_id}"
            main_py = extracted / 'main.py'
            index_js = extracted / 'index.js'
            if main_py.exists():
                cmd = [sys.executable, str(main_py)]
            elif index_js.exists():
                cmd = ['node', str(index_js)]
            else:
                return False, "No entry point (main.py / index.js) found in ZIP."
        else:
            return False, f"Unsupported file type: {suffix}"

        proc = subprocess.Popen(
            cmd,
            stdout=log_handle,
            stderr=subprocess.STDOUT,
            cwd=str(file_path.parent),
            env={**os.environ},
            preexec_fn=os.setsid if os.name != 'nt' else None,
            creationflags=subprocess.CREATE_NEW_PROCESS_GROUP if os.name == 'nt' else 0,
        )

        _processes[project_id] = proc
        db.save_process(project_id, proc.pid, str(log_path))
        db.update_project_status(project_id, 'running')

        logger.info(f"Started project {project_id} PID={proc.pid}")
        return True, f"Process started (PID {proc.pid})."

    except FileNotFoundError as e:
        return False, f"Runtime not found: {e}"
    except Exception as e:
        logger.exception(f"Failed to start project {project_id}")
        return False, f"Launch error: {e}"


def stop_process(project_id: int) -> Tuple[bool, str]:
    proc = _processes.get(project_id)
    if proc is None or proc.poll() is not None:
        _processes.pop(project_id, None)
        db.update_project_status(project_id, 'stopped')
        db.delete_process(project_id)
        return True, "Process was not running."

    try:
        _kill_tree(proc.pid)
        proc.wait(timeout=5)
    except Exception:
        pass

    _processes.pop(project_id, None)
    db.update_project_status(project_id, 'stopped')
    db.delete_process(project_id)
    logger.info(f"Stopped project {project_id}")
    return True, "Process terminated."


def restart_process(project_id: int) -> Tuple[bool, str]:
    stop_process(project_id)
    return start_process(project_id)


def _kill_tree(pid: int) -> None:
    try:
        parent = psutil.Process(pid)
        children = parent.children(recursive=True)
        for child in children:
            try:
                child.kill()
            except psutil.NoSuchProcess:
                pass
        parent.kill()
    except psutil.NoSuchProcess:
        pass


def get_process_stats(project_id: int) -> Dict:
    proc = _processes.get(project_id)
    if proc is None or proc.poll() is not None:
        return {'alive': False, 'cpu': 0.0, 'ram_mb': 0.0, 'pid': None}
    try:
        p = psutil.Process(proc.pid)
        cpu = p.cpu_percent(interval=0.1)
        ram = p.memory_info().rss / 1024 / 1024
        return {'alive': True, 'cpu': round(cpu, 1), 'ram_mb': round(ram, 1), 'pid': proc.pid}
    except psutil.NoSuchProcess:
        return {'alive': False, 'cpu': 0.0, 'ram_mb': 0.0, 'pid': None}


def is_running(project_id: int) -> bool:
    proc = _processes.get(project_id)
    return proc is not None and proc.poll() is None


def get_logs(project_id: int, lines: int = 50) -> str:
    log_path = _log_path(project_id)
    if not log_path.exists():
        return "No logs available."
    try:
        all_lines = log_path.read_text(encoding='utf-8', errors='replace').splitlines()
        return '\n'.join(all_lines[-lines:]) or "Log file is empty."
    except Exception as e:
        return f"Error reading logs: {e}"


def stop_all() -> None:
    for pid in list(_processes.keys()):
        stop_process(pid)


def sync_status_on_startup() -> None:
    """Mark all projects as stopped on cold start."""
    for p in db.get_all_projects():
        if p['status'] == 'running':
            db.update_project_status(p['id'], 'stopped')
            db.delete_process(p['id'])