"""프로비저닝 시스템 테스트 — OperationLog 모델, ProvenanceRecorder, API 엔드포인트."""

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 fpdf import FPDF
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker

from app.core.config import settings
from app.core.database import Base, engine, SessionLocal
from app.core.version_registry import MODULE_VERSIONS, get_module_info
from app.main import app
from app.models.project import OperationLog
from app.modules.provenance import ProvenanceRecorder, OperationContext
from tests._safety_guards import safe_drop_all, safe_rmtree


# -- Fixtures --

def make_test_pdf(text="Test screenplay content") -> bytes:
    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


@pytest.fixture()
def project_db():
    """Create an in-memory DB for unit tests."""
    test_engine = create_engine("sqlite:///:memory:", echo=False)
    Base.metadata.create_all(test_engine)
    Session = sessionmaker(bind=test_engine)
    session = Session()
    yield session
    session.close()


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": "Provenance Test Proj"})
    assert resp.status_code == 200
    return resp.json()["id"]


def _insert_operation(project_id: str, **overrides) -> str:
    """Directly insert an OperationLog into the DB for testing."""
    db = SessionLocal()
    try:
        op_id = overrides.pop("id", str(uuid.uuid4()))
        now = datetime.now(timezone.utc).isoformat()
        defaults = {
            "id": op_id,
            "project_id": project_id,
            "operation_type": "entity_extraction",
            "episode_id": None,
            "module_name": "entity_extractor",
            "module_version": "1.0.0",
            "prompt_name": "entity_extraction/v6",
            "prompt_version": "v6",
            "prompt_hash": None,
            "input_summary": '{"fulltext_chars": 5000}',
            "output_summary": '{"entities": 10}',
            "status": "success",
            "error_message": None,
            "duration_ms": 1500,
            "token_usage": "{}",
            "metadata_json": "{}",
            "created_at": now,
        }
        defaults.update(overrides)
        op = OperationLog(**defaults)
        db.add(op)
        db.commit()
        return op_id
    finally:
        db.close()


# -- Test: OperationLog model creation --

class TestOperationLogModel:
    def test_operation_log_table_created(self, project_db):
        """OperationLog table is created in DB."""
        table_names = project_db.bind.dialect.get_table_names(project_db.bind.connect())
        assert "operation_log" in table_names

    def test_operation_log_insert(self, project_db):
        """Insert an OperationLog row and retrieve it."""
        op = OperationLog(
            id=str(uuid.uuid4()),
            project_id="test-project",
            operation_type="entity_extraction",
            episode_id="ep-001",
            module_name="entity_extractor",
            module_version="1.0.0",
            prompt_name="entity_extraction/v5",
            prompt_version="v5",
            prompt_hash="abc123",
            input_summary='{"chars": 5000}',
            output_summary='{"entities": 10}',
            status="success",
            duration_ms=1500,
            token_usage='{"input_tokens": 500}',
            metadata_json="{}",
            created_at=datetime.now(timezone.utc).isoformat(),
        )
        project_db.add(op)
        project_db.flush()

        found = project_db.query(OperationLog).filter(OperationLog.id == op.id).first()
        assert found is not None
        assert found.operation_type == "entity_extraction"
        assert found.module_name == "entity_extractor"
        assert found.status == "success"
        assert found.duration_ms == 1500

    def test_operation_log_fields(self, project_db):
        """All expected columns exist on OperationLog."""
        columns = {c.name for c in OperationLog.__table__.columns}
        expected = {
            "id", "project_id", "operation_type", "episode_id", "module_name", "module_version",
            "prompt_name", "prompt_version", "prompt_hash",
            "input_summary", "output_summary", "status", "error_message",
            "duration_ms", "token_usage", "metadata_json", "created_at",
        }
        assert expected.issubset(columns)


# -- Test: ProvenanceRecorder context manager --

class TestProvenanceRecorder:
    def test_start_operation_creates_log(self, project_db):
        """start_operation context manager creates an OperationLog on success."""
        recorder = ProvenanceRecorder(project_db, "test-project")
        with recorder.start_operation("entity_extraction", "entity_extractor", episode_id="ep-001") as op:
            op.set_input({"fulltext_chars": 5000})
            op.set_output({"entities": 10})

        logs = project_db.query(OperationLog).all()
        assert len(logs) == 1
        assert logs[0].operation_type == "entity_extraction"
        assert logs[0].status == "success"
        assert logs[0].duration_ms >= 0
        assert json.loads(logs[0].input_summary)["fulltext_chars"] == 5000

    def test_context_manager_records_error(self, project_db):
        """Context manager records error on exception."""
        recorder = ProvenanceRecorder(project_db, "test-project")
        with pytest.raises(ValueError):
            with recorder.start_operation("image_generation", "scene_image_generator") as op:
                op.set_input({"still_id": "s-001"})
                raise ValueError("test error")

        logs = project_db.query(OperationLog).all()
        assert len(logs) == 1
        assert logs[0].status == "error"
        assert "test error" in logs[0].error_message

    def test_explicit_fail(self, project_db):
        """Explicit fail() call records error status."""
        recorder = ProvenanceRecorder(project_db, "test-project")
        with recorder.start_operation("image_generation", "reference_image_generator") as op:
            op.set_input({"entity_id": "e-001"})
            op.fail("moderation_blocked: SAFETY")

        logs = project_db.query(OperationLog).all()
        assert len(logs) == 1
        assert logs[0].status == "error"
        assert "moderation_blocked" in logs[0].error_message

    def test_token_usage_recorded(self, project_db):
        """Token usage is recorded."""
        recorder = ProvenanceRecorder(project_db, "test-project")
        with recorder.start_operation("entity_extraction", "entity_extractor") as op:
            op.set_token_usage(input_tokens=500, output_tokens=1000, cost_usd=0.05)

        logs = project_db.query(OperationLog).all()
        usage = json.loads(logs[0].token_usage)
        assert usage["input_tokens"] == 500
        assert usage["output_tokens"] == 1000
        assert usage["cost_usd"] == 0.05

    def test_module_version_from_registry(self, project_db):
        """Module version comes from MODULE_VERSIONS."""
        recorder = ProvenanceRecorder(project_db, "test-project")
        with recorder.start_operation("entity_extraction", "entity_extractor") as op:
            pass

        logs = project_db.query(OperationLog).all()
        assert logs[0].module_version == MODULE_VERSIONS["entity_extractor"]

    def test_prompt_hash_computation(self):
        """compute_prompt_hash returns None for nonexistent file."""
        result = ProvenanceRecorder.compute_prompt_hash("/nonexistent/path.md")
        assert result is None

    def test_get_operations_filter(self, project_db):
        """get_operations filters by operation_type."""
        recorder = ProvenanceRecorder(project_db, "test-project")
        with recorder.start_operation("entity_extraction", "entity_extractor") as op:
            pass
        with recorder.start_operation("image_generation", "scene_image_generator") as op:
            pass

        all_ops = recorder.get_operations()
        assert len(all_ops) == 2

        filtered = recorder.get_operations(operation_type="entity_extraction")
        assert len(filtered) == 1
        assert filtered[0].operation_type == "entity_extraction"

    def test_get_summary(self, project_db):
        """get_summary returns aggregated stats."""
        recorder = ProvenanceRecorder(project_db, "test-project")
        with recorder.start_operation("entity_extraction", "entity_extractor") as op:
            pass
        with recorder.start_operation("entity_extraction", "entity_extractor") as op:
            op.fail("test error")
        with recorder.start_operation("image_generation", "scene_image_generator") as op:
            pass

        summary = recorder.get_summary()
        assert len(summary) == 2

        ee_summary = next(s for s in summary if s["operation_type"] == "entity_extraction")
        assert ee_summary["total"] == 2
        assert ee_summary["success"] == 1
        assert ee_summary["error"] == 1

        ig_summary = next(s for s in summary if s["operation_type"] == "image_generation")
        assert ig_summary["total"] == 1
        assert ig_summary["success"] == 1

    def test_get_operation_by_id(self, project_db):
        """get_operation returns a single operation by ID."""
        recorder = ProvenanceRecorder(project_db, "test-project")
        with recorder.start_operation("entity_extraction", "entity_extractor") as op:
            ctx_id = op.operation_id

        found = recorder.get_operation(ctx_id)
        assert found is not None
        assert found.id == ctx_id

    def test_get_operation_not_found(self, project_db):
        """get_operation returns None for nonexistent ID."""
        recorder = ProvenanceRecorder(project_db, "test-project")
        found = recorder.get_operation("nonexistent-id")
        assert found is None


# -- Test: version_registry get_module_info --

class TestVersionRegistry:
    def test_get_module_info_known(self):
        info = get_module_info("entity_extractor")
        # Wave3 (2026-05-21, e2e-review-fix-v1): registry synced to 1.5.0 +
        # entity_extractor_v2/v12 (active stem hygiene). 이전: C8 1.4.0 / v11.
        assert info["version"] == "1.5.0"
        assert info["prompt_dependency"] == "entity_extractor_v2/v12"
        assert info["updated_at"] is not None

    def test_get_module_info_unknown(self):
        info = get_module_info("nonexistent_module")
        assert info["version"] == "0.0.0"
        assert info["prompt_dependency"] is None

    def test_provenance_in_registry(self):
        assert "provenance" in MODULE_VERSIONS
        assert MODULE_VERSIONS["provenance"] == "1.0.0"


# -- Test: operations API endpoints --

class TestOperationsAPI:
    def test_list_operations_empty(self, client: TestClient):
        project_id = _create_project(client)
        resp = client.get(f"/api/v1/projects/{project_id}/operations")
        assert resp.status_code == 200
        assert resp.json() == []

    def test_list_operations_with_data(self, client: TestClient):
        project_id = _create_project(client)
        _insert_operation(project_id)
        _insert_operation(project_id, operation_type="image_generation", module_name="scene_image_generator")

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

    def test_list_operations_filter_by_type(self, client: TestClient):
        project_id = _create_project(client)
        _insert_operation(project_id, operation_type="entity_extraction")
        _insert_operation(project_id, operation_type="image_generation", module_name="scene_image_generator")

        resp = client.get(f"/api/v1/projects/{project_id}/operations?operation_type=entity_extraction")
        assert resp.status_code == 200
        data = resp.json()
        assert len(data) == 1
        assert data[0]["operation_type"] == "entity_extraction"

    def test_list_operations_pagination(self, client: TestClient):
        project_id = _create_project(client)
        for i in range(5):
            _insert_operation(project_id)

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

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

    def test_get_operation_detail(self, client: TestClient):
        project_id = _create_project(client)
        op_id = _insert_operation(project_id, prompt_hash="sha256abc", input_summary='{"test": true}')

        resp = client.get(f"/api/v1/projects/{project_id}/operations/{op_id}")
        assert resp.status_code == 200
        data = resp.json()
        assert data["id"] == op_id
        assert data["prompt_hash"] == "sha256abc"
        assert data["input_summary"] == '{"test": true}'

    def test_get_operation_not_found(self, client: TestClient):
        project_id = _create_project(client)
        resp = client.get(f"/api/v1/projects/{project_id}/operations/nonexistent-id")
        assert resp.status_code == 404
        assert resp.json()["error"]["code"] == "operation.not_found"

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

    def test_operations_summary_with_data(self, client: TestClient):
        project_id = _create_project(client)
        _insert_operation(project_id, operation_type="entity_extraction", status="success", duration_ms=1000)
        _insert_operation(project_id, operation_type="entity_extraction", status="error", duration_ms=500, error_message="fail")
        _insert_operation(
            project_id, operation_type="image_generation",
            module_name="scene_image_generator", status="success", duration_ms=2000,
        )

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

        ee = next(s for s in data if s["operation_type"] == "entity_extraction")
        assert ee["total"] == 2
        assert ee["success"] == 1
        assert ee["error"] == 1
        assert ee["avg_duration_ms"] == 750.0

    def test_operations_response_schema(self, client: TestClient):
        project_id = _create_project(client)
        _insert_operation(project_id)

        resp = client.get(f"/api/v1/projects/{project_id}/operations")
        assert resp.status_code == 200
        item = resp.json()[0]
        assert "id" in item
        assert "operation_type" in item
        assert "module_name" in item
        assert "module_version" in item
        assert "status" in item
        assert "created_at" in item
