"""T2I 콘텐츠 모더레이션, 프롬프트 수정, 생성 추적 테스트."""

import io
import json
import shutil
import uuid
from datetime import datetime, timezone
from pathlib import Path

import pytest
from fastapi.testclient import TestClient

from app.core.config import settings
from app.core.database import Base, engine, SessionLocal
from app.main import app
from app.models.project import GenerationTrace, ImageAsset
from app.modules.generation_tracker import GenerationTracker
from app.modules.llm.gemini_image_client import ModerationError
from app.modules.prompt_sanitizer import SANITIZE_SCHEMA, STRATEGIES, ATTEMPT_TO_STRATEGY
from tests._safety_guards import safe_drop_all, safe_rmtree


def _make_test_pdf(text="Test content") -> bytes:
    from fpdf import FPDF
    pdf = FPDF()
    pdf.add_page()
    pdf.set_font("Helvetica", size=12)
    pdf.cell(200, 10, text=text)
    return pdf.output()


@pytest.fixture(autouse=True)
def _setup_db():
    Base.metadata.create_all(engine)
    with TestClient(app):
        pass
    yield
    safe_drop_all(engine, Base.metadata)
    proj_dir = Path(settings.projects_dir)
    if proj_dir.exists():
        safe_rmtree(proj_dir)


@pytest.fixture()
def client():
    with TestClient(app, raise_server_exceptions=False) as c:
        yield c


def _admin_login(client: TestClient):
    resp = client.post("/api/v1/auth/login", json={"username": "admin", "password": "admin123"})
    assert resp.status_code == 200


def _create_project(client: TestClient) -> str:
    _admin_login(client)
    resp = client.post("/api/v1/projects/", json={"name": "Moderation Test Proj"})
    assert resp.status_code == 200
    return resp.json()["id"]


def _insert_image_asset(project_id: str, **overrides) -> str:
    """Directly insert an ImageAsset into the DB for testing."""
    db = SessionLocal()
    try:
        image_id = overrides.pop("id", str(uuid.uuid4()))
        now = datetime.now(timezone.utc).isoformat()
        defaults = {
            "id": image_id,
            "project_id": project_id,
            "asset_type": "reference",
            "entity_id": None,
            "still_id": None,
            "episode_id": None,
            "file_path": "test_image.png",
            "prompt_used": "test prompt",
            "generation_model": "test-model",
            "width": 800,
            "height": 600,
            "status": "generated",
            "review_notes": "",
            "created_at": now,
        }
        defaults.update(overrides)
        asset = ImageAsset(**defaults)
        db.add(asset)
        db.commit()
        return defaults["id"]
    finally:
        db.close()


def _insert_trace(project_id: str, **overrides) -> str:
    """Directly insert a GenerationTrace into the DB."""
    db = SessionLocal()
    try:
        trace_id = overrides.pop("id", str(uuid.uuid4()))
        now = datetime.now(timezone.utc).isoformat()
        defaults = {
            "id": trace_id,
            "project_id": project_id,
            "image_asset_id": None,
            "still_id": None,
            "entity_id": None,
            "attempt_number": 1,
            "prompt_used": "test prompt",
            "prompt_version": "original",
            "model_name": "test-model",
            "status": "success",
            "block_reason": None,
            "block_categories": "[]",
            "response_time_ms": 1000,
            "sanitizer_feedback": None,
            "created_at": now,
        }
        defaults.update(overrides)
        trace = GenerationTrace(**defaults)
        db.add(trace)
        db.commit()
        return defaults["id"]
    finally:
        db.close()


# ------------------------------------------------------------------
# Test: ModerationError is raised correctly
# ------------------------------------------------------------------

class TestModerationError:
    def test_moderation_error_attributes(self):
        err = ModerationError(
            block_reason="SAFETY",
            block_categories=["HARM_CATEGORY_VIOLENCE"],
            raw_response={"promptFeedback": {"blockReason": "SAFETY"}},
        )
        assert err.block_reason == "SAFETY"
        assert err.block_categories == ["HARM_CATEGORY_VIOLENCE"]
        assert err.raw_response == {"promptFeedback": {"blockReason": "SAFETY"}}
        assert "Content moderation blocked" in str(err)

    def test_moderation_error_is_exception(self):
        with pytest.raises(ModerationError) as exc_info:
            raise ModerationError(
                block_reason="HARM",
                block_categories=["HARM_CATEGORY_DANGEROUS"],
                raw_response={},
            )
        assert exc_info.value.block_reason == "HARM"

    def test_check_moderation_block_prompt_feedback(self):
        from app.modules.llm.gemini_image_client import GeminiImageClient
        payload = {
            "promptFeedback": {
                "blockReason": "SAFETY",
                "safetyRatings": [
                    {"category": "HARM_CATEGORY_VIOLENCE", "probability": "HIGH"},
                ],
            },
        }
        with pytest.raises(ModerationError) as exc_info:
            GeminiImageClient._check_moderation_block(payload)
        assert exc_info.value.block_reason == "SAFETY"
        assert "HARM_CATEGORY_VIOLENCE" in exc_info.value.block_categories

    def test_check_moderation_block_safety_finish(self):
        from app.modules.llm.gemini_image_client import GeminiImageClient
        payload = {
            "candidates": [
                {
                    "finishReason": "SAFETY",
                    "safetyRatings": [
                        {"category": "HARM_CATEGORY_SEXUALLY_EXPLICIT", "probability": "MEDIUM"},
                    ],
                }
            ],
        }
        with pytest.raises(ModerationError) as exc_info:
            GeminiImageClient._check_moderation_block(payload)
        assert exc_info.value.block_reason == "SAFETY"

    def test_check_moderation_block_no_block(self):
        from app.modules.llm.gemini_image_client import GeminiImageClient
        payload = {
            "candidates": [
                {
                    "finishReason": "STOP",
                    "content": {"parts": [{"text": "ok"}]},
                    "safetyRatings": [
                        {"category": "HARM_CATEGORY_VIOLENCE", "probability": "NEGLIGIBLE"},
                    ],
                }
            ],
        }
        GeminiImageClient._check_moderation_block(payload)


# ------------------------------------------------------------------
# Test: PromptSanitizer schema
# ------------------------------------------------------------------

class TestPromptSanitizerSchema:
    def test_sanitize_schema_structure(self):
        assert SANITIZE_SCHEMA["type"] == "object"
        required = SANITIZE_SCHEMA["required"]
        assert "sanitized_prompt" in required
        assert "changes" in required
        assert "strategy" in required
        props = SANITIZE_SCHEMA["properties"]
        assert props["sanitized_prompt"]["type"] == "string"
        assert props["changes"]["type"] == "string"
        assert props["strategy"]["type"] == "string"

    def test_sanitize_prompts_exist(self):
        # problems.md #14 후속: PROMPTS_DIR 직접 export 제거 → prompt_loader 통합.
        from app.modules.prompt_loader import get_effective_source

        for stem in ("sanitize_system", "sanitize_user"):
            src = get_effective_source("prompt_sanitizer", stem)
            assert src["candidates"]["file"] is not None, (
                f"prompts/_base/prompt_sanitizer/<v>/{stem}.md 누락"
            )

    def test_sanitize_system_prompt_content(self):
        from app.modules.prompt_sanitizer import _load_prompt

        content = _load_prompt("sanitize_system.md")
        assert "T2I" in content
        assert "film_previs" in content
        assert "movie_poster" in content
        assert "aftermath" in content
        assert "포토리얼리스틱" in content

    def test_sanitize_user_prompt_has_placeholders(self):
        from app.modules.prompt_sanitizer import _load_prompt

        content = _load_prompt("sanitize_user.md")
        assert "{original_prompt}" in content
        assert "{block_reason}" in content
        assert "{block_categories}" in content
        assert "{attempt}" in content
        assert "{strategy_name}" in content
        assert "{strategy_description}" in content
        assert "{strategy_prefix}" in content

    def test_strategies_defined(self):
        assert "film_previs" in STRATEGIES
        assert "movie_poster" in STRATEGIES
        assert "aftermath" in STRATEGIES
        assert len(STRATEGIES) == 3
        for key, strategy in STRATEGIES.items():
            assert "name" in strategy
            assert "description" in strategy
            assert "prefix" in strategy

    def test_attempt_to_strategy_mapping(self):
        assert ATTEMPT_TO_STRATEGY[1] == "film_previs"
        assert ATTEMPT_TO_STRATEGY[2] == "movie_poster"
        assert ATTEMPT_TO_STRATEGY[3] == "aftermath"


# ------------------------------------------------------------------
# Test: GenerationTracker record and query
# ------------------------------------------------------------------

class TestGenerationTracker:
    def test_record_creates_trace(self, client: TestClient):
        project_id = _create_project(client)
        db = SessionLocal()
        try:
            tracker = GenerationTracker(db, project_id)
            trace = tracker.record(
                still_id="still-1",
                attempt_number=1,
                prompt_used="test prompt",
                prompt_version="original",
                model_name="gemini-test",
                status="success",
                response_time_ms=500,
            )
            assert trace.id is not None
            assert trace.status == "success"
            assert trace.attempt_number == 1
            assert trace.prompt_version == "original"

            found = db.query(GenerationTrace).filter(GenerationTrace.id == trace.id).first()
            assert found is not None
            assert found.still_id == "still-1"
        finally:
            db.close()

    def test_record_blocked_trace(self, client: TestClient):
        project_id = _create_project(client)
        db = SessionLocal()
        try:
            tracker = GenerationTracker(db, project_id)
            trace = tracker.record(
                entity_id="entity-1",
                attempt_number=1,
                prompt_used="violent prompt",
                prompt_version="original",
                model_name="gemini-test",
                status="moderation_blocked",
                block_reason="SAFETY",
                block_categories=["HARM_CATEGORY_VIOLENCE"],
            )
            assert trace.status == "moderation_blocked"
            assert trace.block_reason == "SAFETY"
            categories = json.loads(trace.block_categories)
            assert "HARM_CATEGORY_VIOLENCE" in categories
        finally:
            db.close()

    def test_get_traces(self, client: TestClient):
        project_id = _create_project(client)
        db = SessionLocal()
        try:
            tracker = GenerationTracker(db, project_id)
            tracker.record(
                still_id="still-A",
                attempt_number=1,
                prompt_used="prompt A",
                prompt_version="original",
                status="success",
            )
            tracker.record(
                still_id="still-B",
                attempt_number=1,
                prompt_used="prompt B",
                prompt_version="original",
                status="success",
            )
            db.commit()

            traces_a = tracker.get_traces(still_id="still-A")
            assert len(traces_a) == 1
            assert traces_a[0]["still_id"] == "still-A"

            traces_b = tracker.get_traces(still_id="still-B")
            assert len(traces_b) == 1
        finally:
            db.close()

    def test_get_rejection_stats(self, client: TestClient):
        project_id = _create_project(client)
        db = SessionLocal()
        try:
            tracker = GenerationTracker(db, project_id)
            tracker.record(
                attempt_number=1,
                prompt_used="p1",
                prompt_version="original",
                status="success",
            )
            tracker.record(
                attempt_number=1,
                prompt_used="p2",
                prompt_version="original",
                status="moderation_blocked",
                block_reason="SAFETY",
            )
            tracker.record(
                attempt_number=2,
                prompt_used="p3",
                prompt_version="sanitized_v1",
                status="moderation_blocked",
                block_reason="SAFETY",
            )
            tracker.record(
                attempt_number=1,
                prompt_used="p4",
                prompt_version="original",
                status="error",
            )
            db.commit()

            stats = tracker.get_rejection_stats()
            assert stats["total"] == 4
            assert stats["success"] == 1
            assert stats["blocked"] == 2
            assert stats["error"] == 1
            assert stats["block_reasons"]["SAFETY"] == 2
        finally:
            db.close()

    def test_get_prompt_evolution(self, client: TestClient):
        project_id = _create_project(client)
        db = SessionLocal()
        try:
            tracker = GenerationTracker(db, project_id)
            tracker.record(
                still_id="still-X",
                attempt_number=1,
                prompt_used="original prompt",
                prompt_version="original",
                status="moderation_blocked",
                block_reason="SAFETY",
            )
            tracker.record(
                still_id="still-X",
                attempt_number=2,
                prompt_used="sanitized prompt v1",
                prompt_version="sanitized_v1",
                status="moderation_blocked",
                block_reason="SAFETY",
            )
            tracker.record(
                still_id="still-X",
                attempt_number=3,
                prompt_used="sanitized prompt v2",
                prompt_version="sanitized_v2",
                status="success",
            )
            db.commit()

            evolution = tracker.get_prompt_evolution("still-X")
            assert len(evolution) == 3
            assert evolution[0]["prompt_version"] == "original"
            assert evolution[1]["prompt_version"] == "sanitized_v1"
            assert evolution[2]["prompt_version"] == "sanitized_v2"
            assert evolution[2]["status"] == "success"
        finally:
            db.close()


# ------------------------------------------------------------------
# Test: API -- rejection stats endpoint
# ------------------------------------------------------------------

class TestTraceAPI:
    def test_rejection_stats_empty(self, client: TestClient):
        project_id = _create_project(client)
        resp = client.get(f"/api/v1/projects/{project_id}/generation-traces/stats")
        assert resp.status_code == 200
        data = resp.json()
        assert data["total"] == 0
        assert data["success"] == 0
        assert data["blocked"] == 0
        assert data["error"] == 0
        assert data["block_reasons"] == {}

    def test_rejection_stats_with_data(self, client: TestClient):
        project_id = _create_project(client)
        _insert_trace(project_id, status="success")
        _insert_trace(project_id, status="moderation_blocked", block_reason="SAFETY")
        _insert_trace(project_id, status="error")

        resp = client.get(f"/api/v1/projects/{project_id}/generation-traces/stats")
        assert resp.status_code == 200
        data = resp.json()
        assert data["total"] == 3
        assert data["success"] == 1
        assert data["blocked"] == 1
        assert data["error"] == 1
        assert data["block_reasons"]["SAFETY"] == 1

    def test_list_traces_empty(self, client: TestClient):
        project_id = _create_project(client)
        resp = client.get(f"/api/v1/projects/{project_id}/generation-traces")
        assert resp.status_code == 200
        assert resp.json() == []

    def test_list_traces_with_data(self, client: TestClient):
        project_id = _create_project(client)
        _insert_trace(project_id, status="success", prompt_used="test prompt 1")
        _insert_trace(project_id, status="moderation_blocked", prompt_used="blocked prompt")

        resp = client.get(f"/api/v1/projects/{project_id}/generation-traces")
        assert resp.status_code == 200
        data = resp.json()
        assert len(data) == 2

    def test_list_traces_pagination(self, client: TestClient):
        project_id = _create_project(client)
        for i in range(5):
            _insert_trace(project_id, status="success", prompt_used=f"prompt {i}")

        resp = client.get(f"/api/v1/projects/{project_id}/generation-traces?limit=2&offset=0")
        assert resp.status_code == 200
        assert len(resp.json()) == 2

        resp = client.get(f"/api/v1/projects/{project_id}/generation-traces?limit=2&offset=3")
        assert resp.status_code == 200
        assert len(resp.json()) == 2

    def test_image_traces(self, client: TestClient):
        project_id = _create_project(client)
        image_id = _insert_image_asset(project_id)
        _insert_trace(project_id, image_asset_id=image_id, status="success")
        _insert_trace(project_id, image_asset_id="other-id", status="success")

        resp = client.get(f"/api/v1/projects/{project_id}/images/{image_id}/traces")
        assert resp.status_code == 200
        data = resp.json()
        assert len(data) == 1
        assert data[0]["image_asset_id"] == image_id


# ------------------------------------------------------------------
# Test: GenerationTrace model in DB
# ------------------------------------------------------------------

class TestGenerationTraceModel:
    def test_trace_model_fields(self, client: TestClient):
        project_id = _create_project(client)
        db = SessionLocal()
        try:
            now = datetime.now(timezone.utc).isoformat()
            trace = GenerationTrace(
                id="trace-1",
                project_id=project_id,
                image_asset_id="asset-1",
                still_id="still-1",
                entity_id="entity-1",
                attempt_number=2,
                prompt_used="sanitized prompt",
                prompt_version="sanitized_v1",
                model_name="gemini-3.1-flash",
                status="success",
                block_reason=None,
                block_categories="[]",
                response_time_ms=1500,
                sanitizer_feedback="Removed violent keywords",
                created_at=now,
            )
            db.add(trace)
            db.commit()

            found = db.query(GenerationTrace).filter(GenerationTrace.id == "trace-1").first()
            assert found is not None
            assert found.image_asset_id == "asset-1"
            assert found.still_id == "still-1"
            assert found.entity_id == "entity-1"
            assert found.attempt_number == 2
            assert found.prompt_version == "sanitized_v1"
            assert found.model_name == "gemini-3.1-flash"
            assert found.response_time_ms == 1500
            assert found.sanitizer_feedback == "Removed violent keywords"
        finally:
            db.close()
