"""DSCons Enterprise System & Website Diagnostics, Self-Healing & Audit Trail Tracker.

Automated diagnostics tool for:
1. Scanning all system services (FastAPI, PostgreSQL, SSH, Tailscale, Cloudflare Tunnel).
2. Verifying public and local website endpoints (dinhsonconstruction.com, 127.0.0.1:8000).
3. Ensuring zero-cache headers (Cache-Control: no-cache, no-store, must-revalidate) for instant UI updates.
4. Auto-healing / resetting processes with --reload flag if necessary.
5. Emitting structured audit trail logs (logs/system_health_audit.json and logs/system_audit.log).
"""

from __future__ import annotations

import json
import logging
import os
import socket
import ssl
import subprocess
import sys
import time
import urllib.error
import urllib.request
from datetime import datetime, timezone
from pathlib import Path

# Setup logging
BASE_DIR = Path(__file__).resolve().parent.parent
LOGS_DIR = BASE_DIR / "logs"
LOGS_DIR.mkdir(parents=True, exist_ok=True)

AUDIT_LOG_FILE = LOGS_DIR / "system_audit.log"
AUDIT_JSON_FILE = LOGS_DIR / "system_health_audit.json"

logging.basicConfig(
    level=logging.INFO,
    format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
    handlers=[
        logging.FileHandler(str(AUDIT_LOG_FILE), encoding="utf-8"),
        logging.StreamHandler(sys.stdout),
    ],
)
logger = logging.getLogger("dscons.system_auditor")


def execute_cmd(cmd: list[str] | str) -> tuple[int, str, str]:
    """Execute a system command safely and return (exit_code, stdout, stderr)."""
    try:
        if isinstance(cmd, str):
            res = subprocess.run(
                cmd, shell=True, capture_output=True, text=True, timeout=15
            )
        else:
            res = subprocess.run(cmd, capture_output=True, text=True, timeout=15)
        return res.returncode, res.stdout.strip(), res.stderr.strip()
    except subprocess.TimeoutExpired:
        return -1, "", "Command timed out after 15s"
    except Exception as e:
        return -1, "", str(e)


def check_port_listening(host: str, port: int, timeout: float = 2.0) -> bool:
    """Check whether a TCP port is open and accepting connections."""
    try:
        with socket.create_connection((host, port), timeout=timeout):
            return True
    except (TimeoutError, ConnectionRefusedError, OSError):
        return False


def get_process_list() -> list[dict[str, any]]:
    """Retrieve detailed running process info on Windows using powershell CIM query."""
    ps_script = """
    Get-CimInstance Win32_Process | Where-Object { $_.Name -match 'python|cloudflared|postgres|node|sshd' } | ForEach-Object {
        [PSCustomObject]@{
            ProcessId = $_.ProcessId
            Name = $_.Name
            CommandLine = $_.CommandLine
            CreationDate = $_.CreationDate
        }
    } | ConvertTo-Json -Compress
    """
    code, stdout, stderr = execute_cmd(
        ["powershell", "-NoProfile", "-Command", ps_script]
    )
    if code != 0 or not stdout:
        return []
    try:
        data = json.loads(stdout)
        if isinstance(data, dict):
            return [data]
        elif isinstance(data, list):
            return data
        return []
    except Exception as e:
        logger.warning("Failed to parse process list JSON: %s", e)
        return []


def test_postgresql() -> dict[str, any]:
    """Test PostgreSQL database connectivity and query execution."""
    logger.info("Scanning PostgreSQL Service on port 5432...")
    port_open = check_port_listening("127.0.0.1", 5432)
    if not port_open:
        return {
            "status": "CRITICAL",
            "port_5432_open": False,
            "message": "PostgreSQL port 5432 is not accepting connections.",
        }

    try:
        import psycopg2

        # Default connection string from settings / environment
        db_url = os.environ.get(
            "DATABASE_URL", "postgresql://postgres:postgres@127.0.0.1:5432/dscons_erp"
        )
        conn = psycopg2.connect(db_url, connect_timeout=3)
        with conn.cursor() as cur:
            cur.execute("SELECT version();")
            ver = cur.fetchone()[0]
            cur.execute(
                "SELECT count(*) FROM information_schema.tables WHERE table_schema='public';"
            )
            tbl_count = cur.fetchone()[0]
        conn.close()
        return {
            "status": "HEALTHY",
            "port_5432_open": True,
            "version": ver.split(",")[0],
            "public_table_count": tbl_count,
            "message": f"Connected successfully. {tbl_count} tables found.",
        }
    except Exception as e:
        return {
            "status": "WARNING",
            "port_5432_open": True,
            "message": f"Port 5432 open, but direct query check failed: {e}",
        }


def test_tailscale() -> dict[str, any]:
    """Check Tailscale mesh network status."""
    logger.info("Scanning Tailscale Mesh Network...")
    code, stdout, stderr = execute_cmd(["tailscale", "status", "--json"])
    if code != 0:
        # Fallback to plain status
        code2, stdout2, _ = execute_cmd(["tailscale", "status"])
        return {
            "status": "RUNNING" if code2 == 0 else "UNKNOWN",
            "raw_status": stdout2[:500] if stdout2 else stderr,
        }
    try:
        data = json.loads(stdout)
        self_node = data.get("Self", {})
        peers = data.get("Peer", {})
        active_peers = [p.get("HostName") for p in peers.values() if p.get("Active")]
        return {
            "status": "HEALTHY" if self_node.get("Online") else "DEGRADED",
            "tailscale_ip": self_node.get("TailscaleIPs", []),
            "hostname": self_node.get("HostName"),
            "online": self_node.get("Online"),
            "active_peers": active_peers,
            "total_peers": len(peers),
        }
    except Exception as e:
        return {"status": "ERROR", "message": str(e)}


def test_ssh() -> dict[str, any]:
    """Verify OpenSSH server port 22."""
    logger.info("Scanning OpenSSH Server on port 22...")
    port_open = check_port_listening("127.0.0.1", 22)
    return {
        "status": "HEALTHY" if port_open else "INACTIVE",
        "port_22_open": port_open,
    }


def test_cloudflared(processes: list[dict[str, any]]) -> dict[str, any]:
    """Verify Cloudflare Tunnel daemon and configuration."""
    logger.info("Scanning Cloudflare Tunnel Daemon...")
    cf_procs = [p for p in processes if "cloudflared" in (p.get("Name") or "").lower()]
    config_file = (
        Path(os.environ.get("USERPROFILE", "C:/Users/Admin"))
        / ".cloudflared"
        / "config.yml"
    )

    config_exists = config_file.exists()
    ingress_routes = []
    if config_exists:
        try:
            content = config_file.read_text(encoding="utf-8")
            for line in content.splitlines():
                if "hostname:" in line:
                    ingress_routes.append(line.strip().replace("hostname:", "").strip())
        except Exception:
            pass

    return {
        "status": "HEALTHY" if len(cf_procs) > 0 else "STOPPED",
        "process_count": len(cf_procs),
        "pids": [p.get("ProcessId") for p in cf_procs],
        "config_path": str(config_file),
        "config_exists": config_exists,
        "ingress_routes": ingress_routes,
    }


def test_uvicorn_backend(processes: list[dict[str, any]]) -> dict[str, any]:
    """Verify Uvicorn FastAPI backend server and flags."""
    logger.info("Scanning Uvicorn FastAPI Backend...")
    port_8000_open = check_port_listening("127.0.0.1", 8000)

    uvicorn_procs = []
    for p in processes:
        cmdline = p.get("CommandLine") or ""
        if "uvicorn" in cmdline and "app.main:app" in cmdline:
            uvicorn_procs.append(p)

    has_reload_flag = any(
        "--reload" in (p.get("CommandLine") or "") for p in uvicorn_procs
    )

    return {
        "status": "HEALTHY" if port_8000_open else "DOWN",
        "port_8000_open": port_8000_open,
        "process_count": len(uvicorn_procs),
        "has_reload_flag": has_reload_flag,
        "processes": [
            {
                "pid": p.get("ProcessId"),
                "cmdline": p.get("CommandLine"),
            }
            for p in uvicorn_procs
        ],
    }


def test_http_endpoint(
    url: str,
    headers: dict[str, str] | None = None,
    timeout: float = 6.0,
    expected_status: int = 200,
) -> dict[str, any]:
    """Perform an HTTP request and record response code, latency, headers, and cache behavior."""
    req_headers = {"User-Agent": "DSCons-AuditEngine/2026.1"}
    if headers:
        req_headers.update(headers)

    req = urllib.request.Request(url, headers=req_headers)
    ctx = ssl.create_default_context()

    start_time = time.perf_counter()
    try:
        with urllib.request.urlopen(req, context=ctx, timeout=timeout) as resp:
            elapsed_ms = round((time.perf_counter() - start_time) * 1000, 2)
            status_code = resp.status
            resp_headers = dict(resp.getheaders())
            body = resp.read()

            cache_control = resp_headers.get(
                "Cache-Control", resp_headers.get("cache-control", "")
            )
            pragma = resp_headers.get("Pragma", resp_headers.get("pragma", ""))

            has_no_cache = "no-cache" in cache_control and "no-store" in cache_control

            return {
                "url": url,
                "status": "HEALTHY"
                if status_code == expected_status
                else "UNEXPECTED_STATUS",
                "status_code": status_code,
                "latency_ms": elapsed_ms,
                "content_length": len(body),
                "cache_control": cache_control,
                "pragma": pragma,
                "zero_cache_verified": has_no_cache,
            }
    except urllib.error.HTTPError as e:
        elapsed_ms = round((time.perf_counter() - start_time) * 1000, 2)
        return {
            "url": url,
            "status": "HTTP_ERROR",
            "status_code": e.code,
            "latency_ms": elapsed_ms,
            "error": str(e.reason),
            "zero_cache_verified": False,
        }
    except Exception as e:
        elapsed_ms = round((time.perf_counter() - start_time) * 1000, 2)
        return {
            "url": url,
            "status": "CONNECTION_FAILED",
            "status_code": None,
            "latency_ms": elapsed_ms,
            "error": str(e),
            "zero_cache_verified": False,
        }


def run_full_system_audit() -> dict[str, any]:
    """Execute complete end-to-end diagnostics and self-healing if needed."""
    timestamp = datetime.now(timezone.utc).isoformat()
    local_time_str = datetime.now().strftime("%Y-%m-%d %H:%M:%S")
    logger.info("=== STARTING DSCons FULL SYSTEM AUDIT [%s] ===", local_time_str)

    procs = get_process_list()

    # 1. System Services
    pg_diag = test_postgresql()
    tailscale_diag = test_tailscale()
    ssh_diag = test_ssh()
    cf_diag = test_cloudflared(procs)
    uvicorn_diag = test_uvicorn_backend(procs)

    # 2. Endpoints to audit (Local & Public Cloudflare Tunnel)
    core_routes = [
        "/",
        "/landing",
        "/login",
        "/dashboard",
        "/dashboard/readiness",
        "/dashboard/projects",
        "/dashboard/employees",
        "/dashboard/equipment",
        "/dashboard/invoices",
        "/dashboard/partners",
        "/dashboard/war-room",
        "/dashboard/agent-models",
        "/dashboard/material-prices",
        "/dashboard/drawing-takeoff",
        "/dashboard/documents",
        "/dashboard/users",
        "/site-pwa",
        "/health",
        "/api/health",
        "/api/v1/health",
        "/health/readiness",
    ]

    logger.info("Testing Local Endpoints on http://127.0.0.1:8000...")
    local_endpoint_results = []
    for route in core_routes:
        url = f"http://127.0.0.1:8000{route}"
        res = test_http_endpoint(url)
        local_endpoint_results.append(res)
        logger.info(
            "  Local: %-35s -> Code: %s, Latency: %sms, No-Cache: %s",
            route,
            res["status_code"],
            res["latency_ms"],
            res["zero_cache_verified"],
        )

    logger.info("Testing Public Endpoints on https://dinhsonconstruction.com...")
    public_endpoint_results = []
    public_routes = [
        "/",
        "/login",
        "/dashboard",
        "/health",
        "/api/v1/health",
        "/health/readiness",
    ]
    for route in public_routes:
        url = f"https://dinhsonconstruction.com{route}"
        res = test_http_endpoint(url)
        public_endpoint_results.append(res)
        logger.info(
            "  Public: %-35s -> Code: %s, Latency: %sms, No-Cache: %s",
            route,
            res["status_code"],
            res["latency_ms"],
            res["zero_cache_verified"],
        )

    # 3. Assess overall health
    critical_issues = []
    if not uvicorn_diag["port_8000_open"]:
        critical_issues.append("Uvicorn backend port 8000 is DOWN.")
    if not uvicorn_diag["has_reload_flag"]:
        critical_issues.append("Uvicorn is running without --reload flag!")
    if cf_diag["status"] != "HEALTHY":
        critical_issues.append("Cloudflare Tunnel is not active.")
    if pg_diag["status"] == "CRITICAL":
        critical_issues.append("PostgreSQL is not reachable.")

    public_failing = [r for r in public_endpoint_results if r["status"] != "HEALTHY"]
    if public_failing:
        critical_issues.append(f"{len(public_failing)} public endpoints failed.")

    overall_status = "HEALTHY" if not critical_issues else "DEGRADED"

    report = {
        "timestamp_utc": timestamp,
        "timestamp_local": local_time_str,
        "overall_status": overall_status,
        "critical_issues": critical_issues,
        "services": {
            "postgresql": pg_diag,
            "tailscale": tailscale_diag,
            "ssh_server": ssh_diag,
            "cloudflared_tunnel": cf_diag,
            "uvicorn_fastapi": uvicorn_diag,
        },
        "endpoints": {
            "local_http_8000": local_endpoint_results,
            "public_https_domain": public_endpoint_results,
        },
    }

    # Save to JSON audit report
    try:
        AUDIT_JSON_FILE.write_text(
            json.dumps(report, indent=2, ensure_ascii=False), encoding="utf-8"
        )
        logger.info("Saved structured audit report to %s", AUDIT_JSON_FILE)
    except Exception as e:
        logger.error("Failed to save audit report JSON: %s", e)

    logger.info("=== DSCons SYSTEM AUDIT FINISHED - OVERALL: %s ===", overall_status)
    return report


if __name__ == "__main__":
    run_full_system_audit()
