#!/usr/bin/env python3
"""Empirical Gate 2 Stress Test Suite for HAD Digital MVP.

Adversarially probes:
1. Standalone packaging integrity & working directory independence (dist/HAD Digital.exe)
2. Security boundaries (path traversal, XSS injection, HTTP headers)
3. Brute-force defense & 5-failure account lockout (30-min lockout & restart persistence)
4. Concurrent request resilience (25 concurrent read/write threads, SQLite WAL validation)
5. Port conflict & startup binding resilience

Zero external test runner dependencies; runnable via Python standard library + requests.
"""

import concurrent.futures
import json
import os
import shutil
import socket
import sqlite3
import subprocess
import sys
import tempfile
import time
from pathlib import Path
import requests

PROJECT_ROOT = Path(__file__).resolve().parent.parent
EXE_PATH = PROJECT_ROOT / "dist" / "HAD Digital.exe"


def allocate_free_port() -> int:
    s = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
    s.bind(("127.0.0.1", 0))
    port = s.getsockname()[1]
    s.close()
    return port


class ExeServerController:
    """Manages spawning and lifecycle of dist/HAD Digital.exe instances."""

    def __init__(self, port: int = None, db_path: str = None, cwd: str = None, env_override: dict = None):
        self.port = port or allocate_free_port()
        self.temp_dir = tempfile.mkdtemp(prefix="had_exe_test_")
        self.db_path = db_path or os.path.join(self.temp_dir, "test_had.db")
        self.cwd = cwd or self.temp_dir
        self.base_url = f"http://127.0.0.1:{self.port}"
        self.proc = None
        self.env_override = env_override

    def start(self, timeout: float = 25.0) -> bool:
        if not EXE_PATH.exists():
            raise FileNotFoundError(f"Executable not found at: {EXE_PATH}")

        if self.env_override is not None:
            env = self.env_override.copy()
        else:
            env = os.environ.copy()
            env["HAD_DB_PATH"] = str(self.db_path)
            env["HAD_PORT"] = str(self.port)
            env["HAD_HOST"] = "127.0.0.1"
            env["HAD_DEBUG"] = "false"

        cmd = [str(EXE_PATH), "--port", str(self.port)]
        self.proc = subprocess.Popen(
            cmd,
            cwd=str(self.cwd),
            env=env,
            stdout=subprocess.PIPE,
            stderr=subprocess.PIPE,
            text=True,
        )

        start = time.time()
        while time.time() - start < timeout:
            if self.proc.poll() is not None:
                time.sleep(0.5)
                stdout, stderr = self.proc.communicate()
                raise RuntimeError(
                    f"Executable terminated prematurely (exit code {self.proc.returncode}):\nSTDOUT:\n{stdout}\nSTDERR:\n{stderr}"
                )
            try:
                r = requests.get(f"{self.base_url}/api/whoami", timeout=1.0)
                if r.status_code in (200, 401):
                    return True
            except Exception:
                pass
            time.sleep(0.3)
        raise TimeoutError(f"Server did not respond on port {self.port} within {timeout}s")

    def stop(self):
        if self.proc:
            try:
                subprocess.run(
                    ["taskkill", "/F", "/T", "/PID", str(self.proc.pid)],
                    capture_output=True,
                    check=False,
                )
            except Exception:
                pass
            self.proc = None
        shutil.rmtree(self.temp_dir, ignore_errors=True)

    def query_db(self, query: str, params: tuple = ()) -> list[dict]:
        conn = sqlite3.connect(self.db_path)
        conn.row_factory = sqlite3.Row
        cur = conn.cursor()
        cur.execute(query, params)
        rows = [dict(r) for r in cur.fetchall()]
        conn.close()
        return rows


# =====================================================================
# PROBE 1: Standalone Packaging & Working Directory Independence
# =====================================================================
def test_probe_1_standalone_packaging() -> dict:
    print("\n" + "=" * 70)
    print("PROBE 1: Standalone Packaging Integrity & Working Directory Independence")
    print("=" * 70)
    results = {}

    isolated_cwd = tempfile.mkdtemp(prefix="had_isolated_cwd_")
    port = allocate_free_port()

    # Scrub environment of all HAD_* and PYTHON* variables
    clean_env = {
        k: v
        for k, v in os.environ.items()
        if not k.startswith("HAD_")
        and not k.startswith("PYTHON")
        and k
        in (
            "SYSTEMROOT",
            "COMSPEC",
            "TEMP",
            "TMP",
            "PATH",
            "PATHEXT",
            "USERPROFILE",
            "LOCALAPPDATA",
            "APPDATA",
        )
    }

    server = ExeServerController(port=port, cwd=isolated_cwd, env_override=clean_env)
    try:
        print(f"[*] Launching dist/HAD Digital.exe from isolated cwd: {isolated_cwd}")
        print(f"[*] Port: {port}, Scrubbed env vars: {len(clean_env)} keys")
        server.start(timeout=30.0)
        print(f"[+] Executable successfully bound to port {port} and responded to health check")
        results["launch_isolated_cwd"] = "PASS"

        # Check static HTML
        r_index = requests.get(f"{server.base_url}/index.html", timeout=5)
        assert r_index.status_code == 200, f"Expected 200, got {r_index.status_code}"
        assert "HAD Digital" in r_index.text, "index.html missing brand title"
        assert "text/html" in r_index.headers.get("Content-Type", ""), "Bad Content-Type for index.html"
        results["static_html_served"] = "PASS"

        # Check static JS and CSS
        r_js = requests.get(f"{server.base_url}/static/app.js", timeout=5)
        assert r_js.status_code == 200, f"Expected 200, got {r_js.status_code}"
        assert "HAD Digital MVP" in r_js.text, "app.js content missing"
        results["static_js_served"] = "PASS"

        # Check bundled CTCAE rules loaded via API
        r_symptoms = requests.get(f"{server.base_url}/api/symptoms", timeout=5)
        assert r_symptoms.status_code == 200, f"Expected 200, got {r_symptoms.status_code}"
        symptoms_data = r_symptoms.json()
        assert "symptoms" in symptoms_data, "No symptoms key in response"
        assert len(symptoms_data["symptoms"]) >= 10, f"Expected >=10 symptoms, got {len(symptoms_data['symptoms'])}"
        results["bundled_ctcae_rules"] = "PASS"

        # Check bundled French guidance loaded via API
        r_guidance = requests.get(f"{server.base_url}/api/guidance?symptom_id=nausea&grade=1", timeout=5)
        assert r_guidance.status_code == 200, f"Expected 200, got {r_guidance.status_code}"
        guidance_data = r_guidance.json()
        assert guidance_data.get("grade") == 1, "Guidance grade mismatch"
        assert "guidance" in guidance_data, "Guidance text missing"
        results["bundled_french_guidance"] = "PASS"

        print("[+] Probe 1: All standalone packaging checks PASSED!")
    finally:
        server.stop()
        shutil.rmtree(isolated_cwd, ignore_errors=True)

    return results


# =====================================================================
# PROBE 2: Security Boundaries & Attack Vectors
# =====================================================================
def test_probe_2_security_boundaries() -> dict:
    print("\n" + "=" * 70)
    print("PROBE 2: Security Boundaries & Attack Vectors")
    print("=" * 70)
    results = {}

    server = ExeServerController()
    server.start()

    try:
        # --- 2A: Path Traversal Vectors ---
        traversal_payloads = [
            "/static/../../MVP/app.py",
            "/static/..%2f..%2fMVP/app.py",
            "/static/....//....//MVP/app.py",
            r"/static/..\..\MVP\app.py",
            "/static/../data/had.db",
            "/static/..%5c..%5cMVP/app.py",
            "/static/%2e%2e/%2e%2e/MVP/app.py",
            "/static/../../../../../../Windows/win.ini",
            "/static/..\\..\\config.json",
            "/static/..%2f..%2f..%2f..%2fWindows%2fwin.ini",
        ]

        blocked_count = 0
        for payload in traversal_payloads:
            url = f"{server.base_url}{payload}"
            try:
                # Use raw session to prevent automatic URL normalization
                req = requests.Request("GET", url).prepare()
                s = requests.Session()
                r = s.send(req, timeout=5)
                status = r.status_code
                body = r.text
            except Exception as e:
                # Connection abort or socket rejection is also a secure defense
                status = 400
                body = str(e)

            # Assert traversal blocked: status must be 400, 403, or 404
            # AND must never leak sensitive file tokens
            leaks = [
                "HADRequestHandler",
                "SQLite format 3",
                "extensions",
                "[fonts]",
                "DEFAULT_CONFIG",
            ]
            leaked = any(tok in body for tok in leaks)
            if status in (400, 403, 404) and not leaked:
                blocked_count += 1
                print(f"  [PASS] Traversal blocked (HTTP {status}): {payload}")
            else:
                print(f"  [FAIL] Traversal VULNERABILITY! Status: {status}, Leaked: {leaked}, URL: {url}")
                raise AssertionError(f"Path traversal succeeded on {payload}: {status}")

        assert blocked_count == len(traversal_payloads)
        results["path_traversal_defense"] = f"PASS ({blocked_count}/{len(traversal_payloads)} blocked)"

        # --- 2B: Script Injection / Stored XSS ---
        session = requests.Session()
        # Login as patient
        login_r = session.post(
            f"{server.base_url}/api/login",
            json={"username": "patient.durand", "password": "demo123"},
            timeout=5,
        )
        assert login_r.status_code == 200, "Patient login failed"

        xss_payload = "<script>alert('XSS-ATTACK')</script><img src=x onerror=alert(1)>"
        report_r = session.post(
            f"{server.base_url}/api/reports",
            json={
                "patient_id": 1,
                "symptom_id": "nausea",
                "symptom_category": "gastrointestinal",
                "severity_score": 2,
                "notes": xss_payload,
            },
            timeout=5,
        )
        assert report_r.status_code == 201, f"Report submission failed: {report_r.text}"
        rep_id = report_r.json().get("report_id")

        # Fetch reports and timeline
        list_r = session.get(f"{server.base_url}/api/reports?patient_id=1", timeout=5)
        assert list_r.status_code == 200
        assert "application/json" in list_r.headers.get("Content-Type", "")
        # Payload must be safely encapsulated as raw JSON string data
        reports_list = list_r.json().get("reports", [])
        stored_rep = next((r for r in reports_list if r["id"] == rep_id), None)
        assert stored_rep is not None, "Submitted report not found"
        assert stored_rep["notes"] == xss_payload, "XSS string corrupted or altered"
        results["stored_xss_json_isolation"] = "PASS"

        # --- 2C: HTTP Security Headers ---
        endpoints_to_check = [
            ("/", "Static Home"),
            ("/static/css/style.css", "Static CSS"),
            ("/api/whoami", "API JSON"),
            ("/api/patients", "Protected API"),
        ]

        for path, label in endpoints_to_check:
            resp = requests.get(f"{server.base_url}{path}", timeout=5)
            csp = resp.headers.get("Content-Security-Policy", "")
            nosniff = resp.headers.get("X-Content-Type-Options")
            xframe = resp.headers.get("X-Frame-Options")
            ref_pol = resp.headers.get("Referrer-Policy")

            assert "default-src 'self'" in csp, f"CSP missing on {label}: {csp}"
            assert "unsafe-inline" not in csp.split("script-src")[1].split(";")[0] if "script-src" in csp else True, f"unsafe-inline script allowed on {label}"
            assert nosniff == "nosniff", f"X-Content-Type-Options missing on {label}"
            assert xframe == "DENY", f"X-Frame-Options missing on {label}"
            assert ref_pol == "strict-origin-when-cross-origin", f"Referrer-Policy missing on {label}"
            print(f"  [PASS] Security headers verified on {label} ({path})")

        results["security_headers_hardened"] = "PASS"
        print("[+] Probe 2: All security boundary checks PASSED!")
    finally:
        server.stop()

    return results


# =====================================================================
# PROBE 3: Brute-Force Password Spraying & 5-Failure Account Lockout
# =====================================================================
def test_probe_3_brute_force_lockout() -> dict:
    print("\n" + "=" * 70)
    print("PROBE 3: Brute-Force Defense & 5-Failure Account Lockout")
    print("=" * 70)
    results = {}

    server = ExeServerController()
    server.start()

    try:
        # Step 1: 5 consecutive failed login attempts on patient.durand
        for attempt in range(1, 6):
            r = requests.post(
                f"{server.base_url}/api/login",
                json={"username": "patient.durand", "password": f"wrong_guess_{attempt}"},
                timeout=5,
            )
            assert r.status_code == 401, f"Attempt {attempt} returned {r.status_code}, expected 401"
            print(f"  [*] Attempt {attempt}/5: Failed login correctly rejected with 401")

        # Step 2: 6th attempt with CORRECT password 'demo123' -> MUST BE REJECTED
        r6 = requests.post(
            f"{server.base_url}/api/login",
            json={"username": "patient.durand", "password": "demo123"},
            timeout=5,
        )
        assert r6.status_code == 401, f"Attempt 6 with VALID password succeeded unexpectedly! Code: {r6.status_code}"
        err_msg = r6.json().get("error", "").lower()
        assert "locked" in err_msg or "invalid" in err_msg, f"Unexpected error message: {err_msg}"
        print(f"  [PASS] Attempt 6 with VALID password rejected with 401: '{r6.json().get('error')}'")
        results["lockout_enforced_attempt_6"] = "PASS"

        # Step 3: Account Isolation Check (other user should NOT be locked out)
        r_doc = requests.post(
            f"{server.base_url}/api/login",
            json={"username": "dr.martin", "password": "demo123"},
            timeout=5,
        )
        assert r_doc.status_code == 200, f"Unaffected user dr.martin was locked out! Code: {r_doc.status_code}"
        assert r_doc.json().get("user", {}).get("role") == "oncologist"
        print("  [PASS] Unaffected user dr.martin logged in successfully (no global DoS)")
        results["account_isolation_no_dos"] = "PASS"

        # Step 4: Audit Log Verification
        # Login as admin to read audit log
        admin_session = requests.Session()
        admin_login = admin_session.post(
            f"{server.base_url}/api/login",
            json={"username": "admin", "password": "admin123"},
            timeout=5,
        )
        assert admin_login.status_code == 200, "Admin login failed"

        audit_r = admin_session.get(f"{server.base_url}/api/audit-log", timeout=5)
        assert audit_r.status_code == 200
        logs = audit_r.json().get("audit_log", [])
        failed_attempts = [
            l for l in logs if l.get("action") == "login_failed" and "patient.durand" in l.get("details", "")
        ]
        assert len(failed_attempts) >= 5, f"Expected >=5 failed login audit logs, got {len(failed_attempts)} in {logs}"
        print(f"  [PASS] Audit log recorded {len(failed_attempts)} failed login events for patient.durand")
        results["audit_log_security_events"] = f"PASS ({len(failed_attempts)} events)"

        # Step 5: Lockout Persistence Across Process Restart
        print("  [*] Testing lockout persistence across executable restart...")
        # Save db path, stop server, and restart targeting same db
        db_path = server.db_path
        port = server.port
        server.proc.terminate()
        server.proc.wait(timeout=5)
        server.proc = None

        # Relaunch pointing to same DB
        restarted_server = ExeServerController(port=port, db_path=db_path)
        restarted_server.start()
        try:
            r_restarted = requests.post(
                f"{restarted_server.base_url}/api/login",
                json={"username": "patient.durand", "password": "demo123"},
                timeout=5,
            )
            assert r_restarted.status_code == 401, "Lockout was lost after server restart!"
            print("  [PASS] Account lockout persisted across process restart (SQLite state durable)")
            results["lockout_persistence_across_restart"] = "PASS"
        finally:
            restarted_server.stop()

        print("[+] Probe 3: All brute-force defense checks PASSED!")
    finally:
        server.stop()

    return results


# =====================================================================
# PROBE 4: Concurrency Stress Testing (25 Concurrent Threads)
# =====================================================================
def test_probe_4_concurrency_stress() -> dict:
    print("\n" + "=" * 70)
    print("PROBE 4: Concurrent Request Resilience (25 Threads on Standalone Exe)")
    print("=" * 70)
    results = {}

    server = ExeServerController()
    server.start()

    try:
        # Pre-authenticate a pool of sessions
        num_write_workers = 15
        num_read_workers = 10
        total_workers = num_write_workers + num_read_workers

        print(f"[*] Firing {total_workers} concurrent requests against {server.base_url}...")

        def write_task(worker_id: int) -> tuple[int, bool, str]:
            session = requests.Session()
            login_r = session.post(
                f"{server.base_url}/api/login",
                json={"username": "patient.durand", "password": "demo123"},
                timeout=10,
            )
            if login_r.status_code != 200:
                return worker_id, False, f"Login failed: {login_r.status_code}"

            symptom_list = ["nausea", "vomiting", "fatigue", "fever", "diarrhea"]
            symptom = symptom_list[worker_id % len(symptom_list)]
            score = (worker_id % 3) + 1

            post_r = session.post(
                f"{server.base_url}/api/reports",
                json={
                    "patient_id": 1,
                    "symptom_id": symptom,
                    "symptom_category": "gastrointestinal" if symptom != "fever" else "constitutional",
                    "severity_score": score,
                    "notes": f"Concurrent stress report from worker #{worker_id}",
                },
                timeout=10,
            )
            if post_r.status_code != 201:
                return worker_id, False, f"Report post failed: {post_r.status_code} - {post_r.text}"
            return worker_id, True, str(post_r.json().get("report_id"))

        def read_task(worker_id: int) -> tuple[int, bool, str]:
            session = requests.Session()
            session.post(
                f"{server.base_url}/api/login",
                json={"username": "dr.martin", "password": "demo123"},
                timeout=10,
            )
            endpoints = ["/api/reports?patient_id=1", "/api/timeline", "/api/patients", "/api/whoami"]
            ep = endpoints[worker_id % len(endpoints)]
            get_r = session.get(f"{server.base_url}{ep}", timeout=10)
            if get_r.status_code != 200:
                return worker_id, False, f"Read {ep} failed: {get_r.status_code}"
            return worker_id, True, f"OK ({len(get_r.text)} bytes)"

        start_time = time.time()
        with concurrent.futures.ThreadPoolExecutor(max_workers=25) as executor:
            write_futures = [executor.submit(write_task, i) for i in range(num_write_workers)]
            read_futures = [executor.submit(read_task, i) for i in range(num_read_workers)]

            write_results = [f.result() for f in write_futures]
            read_results = [f.result() for f in read_futures]

        elapsed = time.time() - start_time
        print(f"[+] All {total_workers} requests completed in {elapsed:.2f} seconds")

        # Analyze write results
        write_failures = [r for r in write_results if not r[1]]
        write_report_ids = [r[2] for r in write_results if r[1]]
        assert len(write_failures) == 0, f"Write failures encountered: {write_failures}"
        assert len(set(write_report_ids)) == num_write_workers, "Duplicate report IDs returned"
        print(f"  [PASS] {num_write_workers} concurrent writes succeeded with 0 lock errors")

        # Analyze read results
        read_failures = [r for r in read_results if not r[1]]
        assert len(read_failures) == 0, f"Read failures encountered: {read_failures}"
        print(f"  [PASS] {num_read_workers} concurrent reads succeeded with HTTP 200")

        # Direct DB validation
        wal_pragma = server.query_db("PRAGMA journal_mode")
        assert wal_pragma[0]["journal_mode"].lower() == "wal", f"DB is not in WAL mode: {wal_pragma}"
        print(f"  [PASS] SQLite engine confirmed operating in WAL mode: {wal_pragma[0]['journal_mode']}")

        db_reports = server.query_db(
            "SELECT count(*) as cnt FROM toxicity_reports WHERE notes LIKE 'Concurrent stress report from worker%'"
        )
        assert db_reports[0]["cnt"] == num_write_workers, f"Expected {num_write_workers} reports, found {db_reports[0]['cnt']}"
        print(f"  [PASS] Verified {db_reports[0]['cnt']} committed records in SQLite database")

        results["concurrency_25_threads"] = f"PASS ({total_workers}/{total_workers} ok in {elapsed:.2f}s)"
        results["sqlite_wal_mode"] = "PASS"
        print("[+] Probe 4: Concurrency stress test PASSED!")
    finally:
        server.stop()

    return results


# =====================================================================
# PROBE 5: Port Conflict & Dynamic Port Binding Resilience
# =====================================================================
def test_probe_5_port_conflict_resilience() -> dict:
    print("\n" + "=" * 70)
    print("PROBE 5: Port Conflict & Dynamic Port Binding Resilience")
    print("=" * 70)
    results = {}

    # Test 5A: Binding when specified port is occupied
    conflict_port = allocate_free_port()
    occupying_sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
    occupying_sock.bind(("127.0.0.1", conflict_port))
    occupying_sock.listen(1)
    print(f"[*] Artificially occupied port {conflict_port} with raw TCP socket")

    temp_dir = tempfile.mkdtemp(prefix="had_port_test_")
    proc = None
    try:
        cmd = [str(EXE_PATH), "--port", str(conflict_port)]
        proc = subprocess.Popen(
            cmd,
            cwd=temp_dir,
            stdout=subprocess.PIPE,
            stderr=subprocess.PIPE,
            text=True,
        )
        try:
            stdout, stderr = proc.communicate(timeout=10)
        except subprocess.TimeoutExpired:
            subprocess.run(["taskkill", "/F", "/T", "/PID", str(proc.pid)], capture_output=True, check=False)
            stdout, stderr = proc.communicate()
            raise AssertionError(f"Process hung on port conflict instead of exiting! STDOUT: {stdout}, STDERR: {stderr}")

        exit_code = proc.returncode
        print(f"  [*] Executable exited with code: {exit_code}")
        print(f"  [*] STDERR snippet:\n{stderr.strip()}")

        # Assert clean rejection: non-zero exit code, socket error logged
        assert exit_code != 0, "Process should have failed to bind to occupied port"
        assert (
            "10048" in stderr
            or "10013" in stderr
            or "address already in use" in stderr.lower()
            or "error while attempting to bind" in stderr.lower()
            or "permissionerror" in stderr.lower()
            or "oserror" in stderr.lower()
        ), f"Expected port collision OSError, got:\n{stderr}"
        print("  [PASS] Port collision cleanly caught; process terminated non-zero with socket error")
        results["port_collision_clean_termination"] = "PASS"
    finally:
        occupying_sock.close()
        shutil.rmtree(temp_dir, ignore_errors=True)

    time.sleep(1.0)

    # Test 5B: Clean binding when port is released or free
    free_port = allocate_free_port()
    server = ExeServerController(port=free_port)
    try:
        server.start(timeout=20.0)
        r = requests.get(f"{server.base_url}/api/whoami", timeout=5)
        assert r.status_code in (200, 401)
        print(f"  [PASS] Clean startup and dynamic binding on port {free_port}")
        results["dynamic_port_binding"] = "PASS"
    finally:
        server.stop()

    print("[+] Probe 5: Port conflict & binding resilience PASSED!")
    return results


def main():
    print("=" * 70)
    print("      HAD DIGITAL MVP - EMPIRICAL GATE 2 CHALLENGE HARNESS")
    print(f"Target Binary: {EXE_PATH}")
    print(f"File Size:     {EXE_PATH.stat().st_size:,} bytes")
    print(f"Timestamp:     {time.strftime('%Y-%m-%dT%H:%M:%SZ', time.gmtime())}")
    print("=" * 70)

    all_results = {}
    try:
        all_results.update(test_probe_1_standalone_packaging())
        all_results.update(test_probe_2_security_boundaries())
        all_results.update(test_probe_3_brute_force_lockout())
        all_results.update(test_probe_4_concurrency_stress())
        all_results.update(test_probe_5_port_conflict_resilience())
    except Exception as e:
        print(f"\n[!] PROBE EXECUTION ENCOUNTERED AN UNHANDLED EXCEPTION: {e}")
        import traceback
        traceback.print_exc()
        sys.exit(1)

    print("\n" + "=" * 70)
    print("                    FINAL RESULTS SUMMARY")
    print("=" * 70)
    all_passed = True
    for test_name, status in all_results.items():
        print(f"  {test_name:<40} : {status}")
        if not status.startswith("PASS"):
            all_passed = False

    print("=" * 70)
    if all_passed:
        print("VERDICT: >>> APPROVE <<< - ALL 5 GATES RIGOROUSLY SATISFIED")
        sys.exit(0)
    else:
        print("VERDICT: >>> REJECT <<< - FAILURES DETECTED")
        sys.exit(2)


if __name__ == "__main__":
    main()
