from __future__ import annotations

import unittest
from types import SimpleNamespace
from unittest.mock import patch

from app.core.logging import (
    AuditEvent,
    AuditLogger,
    InstrumentedGeminiModelAdapter,
    InstrumentedGroqClient,
    InstrumentedMongoCollection,
    MongoAuditRepository,
    bind_log_context,
    clear_log_context,
    set_audit_repository,
)


class FakeAuditLogger:
    def __init__(self) -> None:
        self.events: list[dict] = []

    def db_event(self, **kwargs):
        self.events.append({"kind": "db", **kwargs})

    def llm_event(self, **kwargs):
        self.events.append({"kind": "llm", **kwargs})


class FakeInsertResult:
    inserted_id = "abc123"


class FakeCollection:
    name = "historias_analizadas"

    def __init__(self) -> None:
        self.documents: list[dict] = []

    def insert_one(self, document):
        self.last_document = document
        self.documents.append(document)
        return FakeInsertResult()


class FakeDatabase:
    def __init__(self) -> None:
        self.collections = {"audit_events": FakeCollection()}

    def __getitem__(self, name: str):
        return self.collections[name]


class FakeGroqCompletions:
    def create(self, *args, **kwargs):
        return SimpleNamespace(
            choices=[SimpleNamespace(message=SimpleNamespace(content="respuesta sintetica"))]
        )


class FakeGroqChat:
    def __init__(self) -> None:
        self.completions = FakeGroqCompletions()


class FakeGroqClient:
    def __init__(self) -> None:
        self.chat = FakeGroqChat()


class FakeGeminiAdapter:
    def __init__(self) -> None:
        self._model_name = "gemini-test"

    def generate_content(self, contents, generation_config=None):
        return SimpleNamespace(text="salida sintetica")


class AuditLoggingTest(unittest.TestCase):
    def tearDown(self) -> None:
        clear_log_context()

    def test_audit_event_includes_bound_context(self) -> None:
        bind_log_context(request_id="req-1", trace_id="trace-1", username="tester")

        event = AuditEvent(
            event_type="http.request_completed",
            action="GET /demo",
            outcome="success",
            service="test_service",
        )

        record = event.to_record()

        self.assertEqual(record["request_id"], "req-1")
        self.assertEqual(record["trace_id"], "trace-1")
        self.assertEqual(record["username"], "tester")

    def test_instrumented_mongo_collection_redacts_sensitive_payload(self) -> None:
        collection = InstrumentedMongoCollection(FakeCollection(), service="unit_test")
        fake_audit = FakeAuditLogger()
        collection._audit_logger = fake_audit

        collection.insert_one({"descripcion": "texto clinico sensible", "usuario": "tester"})

        self.assertEqual(len(fake_audit.events), 1)
        event = fake_audit.events[0]
        payload_summary = event["resource"]["payload_summary"]
        self.assertEqual(payload_summary["descripcion"], "<redacted>")
        self.assertEqual(payload_summary["usuario"], "tester")

    def test_build_audit_repository_uses_raw_audit_collection(self) -> None:
        runtime = SimpleNamespace(sync_database=FakeDatabase())

        repository = MongoAuditRepository(runtime.sync_database["audit_events"], retention_days=180)

        self.assertIsInstance(repository, MongoAuditRepository)
        self.assertIs(repository.collection, runtime.sync_database["audit_events"])
        self.assertNotIsInstance(repository.collection, InstrumentedMongoCollection)

    def test_audit_repository_write_does_not_reenter_db_audit_events(self) -> None:
        raw_collection = FakeCollection()
        instrumented_collection = InstrumentedMongoCollection(raw_collection, service="audit_repository")
        audit_logger = AuditLogger()
        instrumented_collection._audit_logger = audit_logger
        repository = MongoAuditRepository(instrumented_collection, retention_days=30)
        set_audit_repository(repository)

        try:
            audit_logger.business_event(
                event_type="rda.generated",
                action="run_case_rda_job",
                outcome="success",
                service="celery_case_rda",
                resource={"case_key": "CASE-001", "artifact_type": "patient"},
            )
        finally:
            set_audit_repository(None)

        self.assertEqual(len(raw_collection.documents), 1)
        self.assertEqual(raw_collection.documents[0]["event_type"], "rda.generated")

    def test_instrumented_groq_client_emits_metrics_without_prompt_content(self) -> None:
        fake_audit = FakeAuditLogger()
        with patch("app.core.logging.get_audit_logger", return_value=fake_audit):
            client = InstrumentedGroqClient(FakeGroqClient(), provider="groq", service="unit_test")
            result = client.chat.completions.create(
                model="groq-test",
                messages=[{"role": "user", "content": "contenido sensible"}],
            )

        self.assertEqual(result.choices[0].message.content, "respuesta sintetica")
        self.assertEqual(len(fake_audit.events), 1)
        event = fake_audit.events[0]
        self.assertEqual(event["provider"], "groq")
        self.assertEqual(event["model"], "groq-test")
        self.assertIn("input_chars", event["metrics"])
        self.assertIn("output_chars", event["metrics"])
        self.assertNotIn("messages", event)

    def test_instrumented_gemini_adapter_emits_metrics(self) -> None:
        fake_audit = FakeAuditLogger()
        with patch("app.core.logging.get_audit_logger", return_value=fake_audit):
            adapter = InstrumentedGeminiModelAdapter(
                FakeGeminiAdapter(),
                provider="gemini",
                service="unit_test",
            )
            result = adapter.generate_content("contenido sensible")

        self.assertEqual(result.text, "salida sintetica")
        self.assertEqual(len(fake_audit.events), 1)
        event = fake_audit.events[0]
        self.assertEqual(event["provider"], "gemini")
        self.assertEqual(event["model"], "gemini-test")
        self.assertIn("input_chars", event["metrics"])
        self.assertIn("output_chars", event["metrics"])


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