#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
=============================================================================
KIEU STORY — EMPIRICAL ADVERSARIAL TEST SUITE: M4 API & CONCURRENCY
=============================================================================
Agent: teamwork_preview_challenger (M4 Challenger 1)
Target: web_review/server.py

Empirical Stress-Testing Criteria:
1. /api/matrix/ep01:
   - Returns HTTP 200 with total_shots == 188
   - Structured across exactly 15 scenes (C01 to C15)
   - Every scene has exact shot breakdown matching screenplay contract
   - Scene 05 (Đạm Tiên) has exactly 27 shots
   - Closed-lips guard flags and audio guards present
2. Concurrency & State Machine Protection:
   - When gen_state["status"] == "generating": both POST /api/generate and POST /api/concat return HTTP 409 Conflict
   - When gen_state["status"] == "concatenating": both POST /api/generate and POST /api/concat return HTTP 409 Conflict
   - Rapid multi-request stress flood (50 requests) guarantees zero 500 errors and 100% 409 rejection
3. Request Validation Rigor:
   - POST /api/generate without prompt or shot_id returns HTTP 422 Unprocessable Entity
   - Empty JSON, empty string prompts, invalid types return HTTP 422
   - Valid payloads are accepted (HTTP 200, status "accepted")
4. Concat Endpoint Rigor:
   - POST /api/concat without scene_id or with blank whitespace returns HTTP 422
   - Valid ConcatRequest is accepted (HTTP 200)
5. Windows Encoding & Logging Stability:
   - log_msg() handles extensive Vietnamese diacritics without crashing
   - Emulated cp1252 / ASCII terminal encode errors trigger safe fallback without uncaught exceptions
   - Rolling buffer keeps exactly 30 entries
6. Query Filtering & Security Boundaries:
   - Filtering /api/shots by ?episode=ep01 returns 188 shots
   - Filtering /api/shots by ?scene=ep01_scene05 returns 27 shots
   - Directory traversal attempts on /api/episode/{filename} return HTTP 404
=============================================================================
"""

import os
import sys
import json
import time
import unittest
from pathlib import Path
from unittest.mock import patch, MagicMock
from typing import Dict, Any, List

# Ensure UTF-8 console output on Windows
if sys.platform == "win32":
    try:
        sys.stdout.reconfigure(encoding="utf-8")
        sys.stderr.reconfigure(encoding="utf-8")
    except Exception:
        pass

BASE_DIR = Path(__file__).resolve().parent.parent
WEB_REVIEW_DIR = BASE_DIR / "web_review"
PIPELINE_DIR = BASE_DIR / "05_Production_Pipeline"

for p in [str(WEB_REVIEW_DIR), str(PIPELINE_DIR), str(BASE_DIR)]:
    if p not in sys.path:
        sys.path.insert(0, p)

from starlette.testclient import TestClient
import server as srv


class TestAdversarialM4ApiConcurrency(unittest.TestCase):
    """Empirical adversarial test harness for web_review/server.py."""

    @classmethod
    def setUpClass(cls):
        cls.client = TestClient(srv.app)
        cls.results_json_path = BASE_DIR / "tests" / "adversarial_m4_api_concurrency_results.json"
        cls.telemetry: Dict[str, Any] = {
            "test_suite": "M4 API & Concurrency Stress Test",
            "timestamp": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()),
            "total_assertions": 0,
            "passed_assertions": 0,
            "metrics": {},
            "verdict": "PENDING"
        }

    @classmethod
    def tearDownClass(cls):
        cls.telemetry["verdict"] = "PASS" if cls.telemetry["passed_assertions"] == cls.telemetry["total_assertions"] else "FAIL"
        try:
            with open(cls.results_json_path, "w", encoding="utf-8") as f:
                json.dump(cls.telemetry, f, indent=2, ensure_ascii=False)
            print(f"\n[✓] Saved empirical telemetry to: {cls.results_json_path}")
        except Exception as e:
            print(f"[!] Warning: failed to save telemetry: {e}")

    def record_assertion(self, condition: bool, description: str):
        self.telemetry["total_assertions"] += 1
        if condition:
            self.telemetry["passed_assertions"] += 1
        self.assertTrue(condition, f"Assertion failed: {description}")

    # =========================================================================
    # 1. MATRIX EP01 EMPIRICAL TESTS
    # =========================================================================

    def test_01_matrix_ep01_total_shots_and_scenes(self):
        """Verify /api/matrix/ep01 returns exactly 188 shots structured across 15 scenes."""
        resp = self.client.get("/api/matrix/ep01")
        self.record_assertion(resp.status_code == 200, "GET /api/matrix/ep01 returns 200")
        
        data = resp.json()
        self.record_assertion(data.get("episode_id") == "ep01", "episode_id is ep01")
        self.record_assertion(data.get("total_shots") == 188, "total_shots is exactly 188")

        scenes = data.get("scenes", [])
        self.record_assertion(len(scenes) == 15, "Structured across exactly 15 scenes")

        # Hollywood screenplay breakdown specification
        expected_breakdown = {
            "ep01_scene01": 21,
            "ep01_scene02": 15,
            "ep01_scene03": 15,
            "ep01_scene04": 12,
            "ep01_scene05": 27,
            "ep01_scene06": 14,
            "ep01_scene07": 8,
            "ep01_scene08": 8,
            "ep01_scene09": 6,
            "ep01_scene10": 14,
            "ep01_scene11": 10,
            "ep01_scene12": 8,
            "ep01_scene13": 12,
            "ep01_scene14": 12,
            "ep01_scene15": 6
        }

        actual_breakdown = {}
        total_accumulated_shots = 0

        for sc in scenes:
            sc_id = sc["scene_id"]
            sc_shots = sc["shots"]
            shot_count = len(sc_shots)
            actual_breakdown[sc_id] = shot_count
            total_accumulated_shots += shot_count

            expected_cnt = expected_breakdown.get(sc_id)
            self.record_assertion(
                shot_count == expected_cnt,
                f"Scene {sc_id} shot count matches expected {expected_cnt} (actual: {shot_count})"
            )
            self.record_assertion(
                sc["shots_total"] == expected_cnt,
                f"Scene {sc_id} shots_total attribute matches expected count"
            )

        self.record_assertion(
            total_accumulated_shots == 188,
            f"Sum of shots across all 15 scenes == 188 (actual: {total_accumulated_shots})"
        )

        self.telemetry["metrics"]["matrix_ep01"] = {
            "total_shots": total_accumulated_shots,
            "scene_count": len(scenes),
            "breakdown": actual_breakdown
        }

    def test_02_matrix_shot_attributes_and_guards(self):
        """Stress-test detailed attributes of all 188 shots in the matrix."""
        resp = self.client.get("/api/matrix/ep01")
        data = resp.json()
        scenes = data.get("scenes", [])

        shots_with_closed_lips = 0
        all_shots_have_audio_guard = True
        all_shots_have_10s_duration = True
        all_shots_have_v1_target = True

        for sc in scenes:
            for s in sc["shots"]:
                shot_id = s["shot_id"]
                # Must match ep01_sceneXX_shotYY format
                self.record_assertion(
                    bool(shot_id.startswith("ep01_scene")),
                    f"Shot {shot_id} has valid ID prefix"
                )
                if s.get("duration_sec") != 10:
                    all_shots_have_10s_duration = False
                if not str(s.get("target_video_file", "")).endswith("_10s_v1.mp4"):
                    all_shots_have_v1_target = False
                if not s.get("audio_guard") or "KHÔNG sinh nhạc nền" not in s["audio_guard"]:
                    all_shots_have_audio_guard = False
                if s.get("has_lip_sync_guard"):
                    shots_with_closed_lips += 1

        self.record_assertion(all_shots_have_10s_duration, "All 188 shots have 10s duration specified")
        self.record_assertion(all_shots_have_v1_target, "All 188 shots target *_10s_v1.mp4 naming")
        self.record_assertion(all_shots_have_audio_guard, "All 188 shots have mandatory Audio Guard")
        self.record_assertion(shots_with_closed_lips > 0, "Identified shots with Closed Lips Guard active")

        self.telemetry["metrics"]["matrix_guards"] = {
            "shots_with_closed_lips": shots_with_closed_lips,
            "all_shots_have_audio_guard": all_shots_have_audio_guard,
            "all_shots_have_10s_duration": all_shots_have_10s_duration
        }

    def test_03_dam_tien_scene05_subsequences(self):
        """Empirically inspect reconstructed Đạm Tiên Scene 05 (27 shots)."""
        resp = self.client.get("/api/matrix/ep01")
        data = resp.json()
        sc05 = next((s for s in data["scenes"] if s["scene_id"] == "ep01_scene05"), None)
        self.record_assertion(sc05 is not None, "Scene 05 exists in matrix")
        self.record_assertion(len(sc05["shots"]) == 27, "Scene 05 has exactly 27 shots")

        # Verify shots are numbered 1 to 27 sequentially
        shot_ids = [s["shot_id"] for s in sc05["shots"]]
        expected_ids = [f"ep01_scene05_shot{i:02d}" for i in range(1, 28)]
        self.record_assertion(shot_ids == expected_ids, "Scene 05 shots are numbered 01 to 27 sequentially")

    # =========================================================================
    # 2. CONCURRENCY & HTTP 409 CONFLICT PROTECTION
    # =========================================================================

    def test_04_concurrency_generating_state_returns_409(self):
        """Verify both POST /api/generate and /api/concat return 409 when status is 'generating'."""
        orig_status = srv.gen_state["status"]
        try:
            srv.gen_state["status"] = "generating"

            # 1. POST /api/generate with valid prompt
            r_gen = self.client.post("/api/generate", json={"prompt": "Adversarial test prompt"})
            self.record_assertion(r_gen.status_code == 409, "POST /api/generate returns 409 when generating")
            self.record_assertion(r_gen.json().get("status") == "busy", "Status payload is 'busy'")

            # 2. POST /api/generate with valid shot_id
            r_gen_shot = self.client.post("/api/generate", json={"shot_id": "ep01_scene01_shot01"})
            self.record_assertion(r_gen_shot.status_code == 409, "POST /api/generate with shot_id returns 409 when generating")

            # 3. POST /api/concat with valid scene_id
            r_concat = self.client.post("/api/concat", json={"scene_id": "ep01_scene01"})
            self.record_assertion(r_concat.status_code == 409, "POST /api/concat returns 409 when generating")
            self.record_assertion(r_concat.json().get("status") == "busy", "Status payload is 'busy'")
        finally:
            srv.gen_state["status"] = orig_status

    def test_05_concurrency_concatenating_state_returns_409(self):
        """Verify both POST /api/generate and /api/concat return 409 when status is 'concatenating'."""
        orig_status = srv.gen_state["status"]
        try:
            srv.gen_state["status"] = "concatenating"

            # 1. POST /api/generate
            r_gen = self.client.post("/api/generate", json={"prompt": "Adversarial test prompt"})
            self.record_assertion(r_gen.status_code == 409, "POST /api/generate returns 409 when concatenating")
            self.record_assertion(r_gen.json().get("status") == "busy", "Status payload is 'busy'")

            # 2. POST /api/concat
            r_concat = self.client.post("/api/concat", json={"scene_id": "ep01_scene02"})
            self.record_assertion(r_concat.status_code == 409, "POST /api/concat returns 409 when concatenating")
            self.record_assertion(r_concat.json().get("status") == "busy", "Status payload is 'busy'")
        finally:
            srv.gen_state["status"] = orig_status

    def test_06_rapid_concurrency_stress_flood(self):
        """Stress-test: Send 50 rapid sequential requests during busy state to verify 100% 409 rate."""
        orig_status = srv.gen_state["status"]
        try:
            srv.gen_state["status"] = "generating"

            flood_count = 50
            responses_409 = 0
            unexpected_responses = []

            for i in range(flood_count):
                endpoint = "/api/generate" if i % 2 == 0 else "/api/concat"
                payload = {"prompt": f"flood_{i}"} if i % 2 == 0 else {"scene_id": "ep01_scene01"}
                r = self.client.post(endpoint, json=payload)
                if r.status_code == 409 and r.json().get("status") == "busy":
                    responses_409 += 1
                else:
                    unexpected_responses.append((endpoint, r.status_code, r.text))

            self.record_assertion(
                responses_409 == flood_count,
                f"100% of {flood_count} flood requests returned 409 Conflict (actual: {responses_409})"
            )
            self.record_assertion(len(unexpected_responses) == 0, "Zero unexpected response codes during flood")

            self.telemetry["metrics"]["concurrency_flood"] = {
                "total_requests": flood_count,
                "returned_409": responses_409,
                "rejection_rate_pct": (responses_409 / flood_count) * 100.0
            }
        finally:
            srv.gen_state["status"] = orig_status

    # =========================================================================
    # 3. GENERATION REQUEST VALIDATION RIGOR (HTTP 422)
    # =========================================================================

    def test_07_generate_validation_empty_and_invalid_payloads(self):
        """Verify invalid/empty generation requests return HTTP 422 Unprocessable Entity."""
        orig_status = srv.gen_state["status"]
        try:
            srv.gen_state["status"] = "idle"

            # Case 1: Empty JSON object {}
            r1 = self.client.post("/api/generate", json={})
            self.record_assertion(r1.status_code == 422, "Empty JSON returns 422")

            # Case 2: Only metadata without prompt or shot_id
            r2 = self.client.post("/api/generate", json={"scene_id": "ep01_scene01", "title": "Test"})
            self.record_assertion(r2.status_code == 422, "Missing prompt/shot_id returns 422")

            # Case 3: Empty string prompt without shot_id
            r3 = self.client.post("/api/generate", json={"prompt": ""})
            self.record_assertion(r3.status_code == 422, "Empty string prompt returns 422")

            # Case 4: None values for both prompt and shot_id
            r4 = self.client.post("/api/generate", json={"prompt": None, "shot_id": None})
            self.record_assertion(r4.status_code == 422, "None for prompt and shot_id returns 422")

            # Case 5: Invalid type for boolean field
            r5 = self.client.post("/api/generate", json={"prompt": "Valid prompt", "force": "not_a_bool"})
            self.record_assertion(r5.status_code == 422, "Invalid boolean field returns 422")

        finally:
            srv.gen_state["status"] = orig_status

    def test_08_generate_valid_payload_acceptance(self):
        """Verify valid generation requests are accepted (HTTP 200, status 'accepted')."""
        orig_status = srv.gen_state["status"]
        try:
            srv.gen_state["status"] = "idle"

            with patch("server.run_unified_generation") as mock_gen:
                # 1. Valid with prompt only
                r1 = self.client.post("/api/generate", json={"prompt": "Valid test prompt", "title": "Test"})
                self.record_assertion(r1.status_code == 200, "Valid prompt accepted with 200")
                self.record_assertion(r1.json().get("status") == "accepted", "Response status is 'accepted'")

                # 2. Valid with shot_id only
                r2 = self.client.post("/api/generate", json={"shot_id": "ep01_scene01_shot01"})
                self.record_assertion(r2.status_code == 200, "Valid shot_id accepted with 200")
                self.record_assertion(r2.json().get("status") == "accepted", "Response status is 'accepted'")

                # 3. Valid with both prompt and shot_id
                r3 = self.client.post("/api/generate", json={"shot_id": "ep01_scene01_shot02", "prompt": "Custom prompt"})
                self.record_assertion(r3.status_code == 200, "Valid dual prompt/shot accepted with 200")

        finally:
            srv.gen_state["status"] = orig_status

    # =========================================================================
    # 4. CONCAT REQUEST VALIDATION & EXECUTION
    # =========================================================================

    def test_09_concat_validation_invalid_payloads(self):
        """Verify invalid concat requests return HTTP 422."""
        orig_status = srv.gen_state["status"]
        try:
            srv.gen_state["status"] = "idle"

            # Case 1: Empty JSON
            r1 = self.client.post("/api/concat", json={})
            self.record_assertion(r1.status_code == 422, "Empty JSON concat returns 422")

            # Case 2: Missing scene_id
            r2 = self.client.post("/api/concat", json={"output_path": "out.mp4"})
            self.record_assertion(r2.status_code == 422, "Missing scene_id returns 422")

            # Case 3: Empty string scene_id
            r3 = self.client.post("/api/concat", json={"scene_id": ""})
            self.record_assertion(r3.status_code == 422, "Empty string scene_id returns 422")

            # Case 4: Whitespace-only scene_id
            r4 = self.client.post("/api/concat", json={"scene_id": "    "})
            self.record_assertion(r4.status_code == 422, "Whitespace-only scene_id returns 422")

            # Case 5: Invalid crossfade duration (string instead of float)
            r5 = self.client.post("/api/concat", json={"scene_id": "ep01_scene01", "crossfade_dur": "invalid_num"})
            self.record_assertion(r5.status_code == 422, "Invalid crossfade_dur returns 422")

        finally:
            srv.gen_state["status"] = orig_status

    def test_10_concat_valid_payload_acceptance(self):
        """Verify valid ConcatRequest is accepted when idle."""
        orig_status = srv.gen_state["status"]
        try:
            srv.gen_state["status"] = "idle"

            with patch("server.run_scene_concat") as mock_concat:
                r = self.client.post("/api/concat", json={
                    "scene_id": "ep01_scene01",
                    "output_path": "test_output.mp4",
                    "crossfade_dur": 1.5
                })
                self.record_assertion(r.status_code == 200, "Valid ConcatRequest returns 200")
                self.record_assertion(r.json().get("status") == "accepted", "Response status is 'accepted'")
                self.record_assertion("ep01_scene01" in r.json().get("message", ""), "Message confirms scene_id")

        finally:
            srv.gen_state["status"] = orig_status

    # =========================================================================
    # 5. WINDOWS ENCODING & VIETNAMESE DIACRITICS IN LOG_MSG()
    # =========================================================================

    def test_11_windows_encoding_vietnamese_diacritics(self):
        """Verify log_msg() does not crash on complex Vietnamese text."""
        test_phrases = [
            "Thập Ngũ Niên: Đoạn Trường Ký - Thúy Kiều & Thúy Vân",
            "Bước lần theo ngọn tiểu khê, lần xem phong cảnh có bề thanh thanh",
            "Rằng sao trong tiết thanh minh, mà đây hương khói vắng tanh thế mà",
            "Trăm năm trong cõi người ta, chữ tài chữ mệnh khéo là ghét nhau",
            "Tiết Thanh Minh tảo mộ, nấm mồ vô danh bên đường, Vương Quan thuật chuyện",
            "ÀÁÂÃÈÉÊÌÍÒÓÔÕÙÚÝàáâãèéêìíòóôõùúýĂăĐđĨĩŨũƠơƯưẠạẢảẤấẦầẨẩẪẫẬậẮắẰằẲẳẴẵẶặẸẹẺẻẼẽẾếỀềỂểỄễỆệỈỉỊịỌọỎỏỐốỒồỔổỖỗỘộỚớỜờỞởỠỡỢợỤụỦủỨứỪừỬửỮữỰựỲỳỴỵỶỷỸỹ"
        ]

        # Ensure no exception is raised
        for phrase in test_phrases:
            try:
                srv.log_msg(phrase)
                crashed = False
            except Exception as e:
                crashed = True
            self.record_assertion(not crashed, f"log_msg did not crash on Vietnamese text: '{phrase[:30]}...'")

        # Verify logs were appended to gen_state["logs"]
        latest_log = srv.gen_state["logs"][-1]
        self.record_assertion("ÀÁÂÃ" in latest_log, "Vietnamese characters correctly stored in gen_state logs buffer")

    def test_12_windows_cp1252_encoding_fallback_simulation(self):
        """Adversarially simulate UnicodeEncodeError (Windows cp1252) to verify ASCII fallback logic."""
        test_msg = "Kiều gặp Đạm Tiên: khóc thương số phận bạc mệnh"

        # Mock print to simulate terminal with cp1252 throwing UnicodeEncodeError
        call_count = {"normal": 0, "fallback": 0}

        def mock_failing_print(text):
            call_count["normal"] += 1
            if any(ord(c) > 127 for c in text):
                raise UnicodeEncodeError("cp1252", text, 0, 1, "character maps to <undefined>")

        with patch("builtins.print", side_effect=mock_failing_print):
            try:
                srv.log_msg(test_msg)
                crashed = False
            except Exception:
                crashed = True

            self.record_assertion(not crashed, "log_msg survived simulated cp1252 UnicodeEncodeError")

    def test_13_log_buffer_rolling_window_limit(self):
        """Verify gen_state['logs'] strictly caps at 30 items."""
        # Inject 45 log entries
        for i in range(45):
            srv.log_msg(f"Adversarial rolling log entry #{i:02d}")

        self.record_assertion(len(srv.gen_state["logs"]) == 30, "gen_state['logs'] strictly capped at 30 entries")
        self.record_assertion(
            srv.gen_state["logs"][-1].endswith("entry #44"),
            "Most recent log entry preserved at tail of buffer"
        )

    # =========================================================================
    # 6. QUERY FILTERING & SECURITY BOUNDARIES
    # =========================================================================

    def test_14_shots_query_filtering(self):
        """Verify query parameters on /api/shots."""
        # Filter by episode
        r_ep = self.client.get("/api/shots?episode=ep01")
        self.record_assertion(r_ep.status_code == 200, "GET /api/shots?episode=ep01 returns 200")
        shots_ep = r_ep.json().get("shots", {})
        self.record_assertion(len(shots_ep) == 188, f"Filtered episode=ep01 returns 188 shots (actual: {len(shots_ep)})")
        self.record_assertion(all(k.startswith("ep01") for k in shots_ep), "All returned shots start with ep01")

        # Filter by scene
        r_sc = self.client.get("/api/shots?scene=ep01_scene05")
        self.record_assertion(r_sc.status_code == 200, "GET /api/shots?scene=ep01_scene05 returns 200")
        shots_sc = r_sc.json().get("shots", {})
        self.record_assertion(len(shots_sc) == 27, f"Filtered scene=ep01_scene05 returns 27 shots (actual: {len(shots_sc)})")

        # Non-matching filter
        r_none = self.client.get("/api/shots?episode=nonexistent_ep99")
        self.record_assertion(len(r_none.json().get("shots", {})) == 0, "Non-matching episode returns 0 shots")

    def test_15_path_traversal_guards(self):
        """Verify path traversal attacks on /api/episode/{filename} return HTTP 404."""
        malicious_paths = [
            "../../windows/system32/cmd.exe",
            "../PROJECT.md",
            "..%2F..%2FPROJECT.md",
            "../../../boot.ini",
            "TAP_99_NONEXISTENT.md"
        ]

        for p in malicious_paths:
            r = self.client.get(f"/api/episode/{p}")
            self.record_assertion(
                r.status_code == 404,
                f"Path traversal / invalid episode '{p}' safely rejected with 404 (actual: {r.status_code})"
            )

        # Valid episode returns 200
        r_valid = self.client.get("/api/episode/TAP_01_XUAN_SAC_THE_NGUYEN_VA_GIONG_BAO_DOAN_TRUONG.md")
        self.record_assertion(r_valid.status_code == 200, "Valid episode TAP_01 returns 200")
        self.record_assertion("content" in r_valid.json(), "Valid episode returns markdown content")

    def test_16_banana_prompts_endpoint(self):
        """Verify /api/banana-prompts serves the 188 Start Frame prompts."""
        r = self.client.get("/api/banana-prompts")
        self.record_assertion(r.status_code == 200, "GET /api/banana-prompts returns 200")
        data = r.json()
        ep01_banana = [k for k in data.keys() if k.startswith("ep01_scene")]
        self.record_assertion(len(ep01_banana) == 188, f"Banana prompts catalog has 188 EP01 shots (actual: {len(ep01_banana)})")


if __name__ == "__main__":
    unittest.main()
