from __future__ import annotations

import json
import logging
import os
import uuid
from collections.abc import Iterable, Mapping
from contextvars import ContextVar
from dataclasses import dataclass, field
from datetime import UTC, datetime, timedelta
from logging.handlers import RotatingFileHandler
from pathlib import Path
from time import perf_counter
from typing import Any


_LOG_CONTEXT: ContextVar[dict[str, Any] | None] = ContextVar("log_context", default=None)
_AUDIT_EMIT_DEPTH: ContextVar[int] = ContextVar("audit_emit_depth", default=0)
_AUDIT_REPOSITORY: MongoAuditRepository | None = None
_AUDIT_LOGGER_INSTANCE: AuditLogger | None = None
_LOGGING_CONFIGURED = False

_SENSITIVE_KEYS = {
    "access_token",
    "accessToken",
    "analisis",
    "analisis_html",
    "analysis_structured",
    "api_key",
    "authorization",
    "contenido_pdf",
    "cookie",
    "cookies",
    "descripcion",
    "diagnosticos_html",
    "gemini_api_key",
    "google_api_key",
    "groq_api_key",
    "historia_html",
    "html",
    "input",
    "jwt",
    "message",
    "messages",
    "nombre_paciente",
    "output",
    "password",
    "patient_name",
    "pdf_text",
    "prompt",
    "raw_text",
    "response",
    "respuesta",
    "text",
    "texto",
    "token",
}
_LOG_CONTEXT_KEYS = (
    "request_id",
    "trace_id",
    "username",
    "batch_id",
    "case_key",
    "file_id",
    "job_id",
    "document_type",
    "provider",
    "model",
)


def _env_flag(name: str, default: bool) -> bool:
    value = os.getenv(name)
    if value is None:
        return default
    return value.strip().lower() in {"1", "true", "yes", "on"}


def _now_utc() -> datetime:
    return datetime.now(UTC)


def _iso_now() -> str:
    return _now_utc().isoformat()


def _safe_int(value: Any, default: int) -> int:
    try:
        return int(value)
    except (TypeError, ValueError):
        return default


def generate_correlation_id() -> str:
    return uuid.uuid4().hex


def get_log_context() -> dict[str, Any]:
    return dict(_LOG_CONTEXT.get() or {})


def bind_log_context(**values: Any) -> dict[str, Any]:
    current = get_log_context()
    for key, value in values.items():
        if value is None:
            continue
        normalized = str(value).strip() if isinstance(value, str) else value
        if normalized == "" or normalized == [] or normalized == {} or normalized == ():
            continue
        current[key] = normalized
    _LOG_CONTEXT.set(current)
    return current


def set_log_context(values: Mapping[str, Any] | None) -> dict[str, Any]:
    normalized = dict(values or {})
    _LOG_CONTEXT.set(normalized)
    return normalized


def clear_log_context() -> None:
    _LOG_CONTEXT.set(None)


def _is_audit_emit_active() -> bool:
    return _AUDIT_EMIT_DEPTH.get() > 0


def _coerce_jsonable(value: Any) -> Any:
    if isinstance(value, (str, int, float, bool)) or value is None:
        return value
    if isinstance(value, datetime):
        return value.isoformat()
    if isinstance(value, Path):
        return str(value)
    return str(value)


def sanitize_for_audit(value: Any, *, depth: int = 0, max_items: int = 12) -> Any:
    if depth > 4:
        return "<truncated>"
    if isinstance(value, Mapping):
        sanitized: dict[str, Any] = {}
        for key, item in list(value.items())[:max_items]:
            if str(key).lower() in _SENSITIVE_KEYS:
                sanitized[str(key)] = "<redacted>"
                continue
            sanitized[str(key)] = sanitize_for_audit(item, depth=depth + 1, max_items=max_items)
        return sanitized
    if isinstance(value, (list, tuple, set)):
        items = list(value)[:max_items]
        return [sanitize_for_audit(item, depth=depth + 1, max_items=max_items) for item in items]
    if isinstance(value, str):
        compact = " ".join(value.split())
        if len(compact) > 180:
            return f"{compact[:177]}..."
        return compact
    return _coerce_jsonable(value)


def summarize_query(value: Any) -> Any:
    return sanitize_for_audit(value)


def _payload_counts(payload: Any) -> dict[str, int]:
    if isinstance(payload, Mapping):
        return {"field_count": len(payload)}
    if isinstance(payload, (list, tuple, set)):
        return {"item_count": len(payload)}
    return {}


@dataclass(slots=True)
class AuditEvent:
    event_type: str
    action: str
    outcome: str
    service: str
    timestamp: str = field(default_factory=_iso_now)
    level: str = "INFO"
    context: dict[str, Any] = field(default_factory=dict)
    metrics: dict[str, Any] = field(default_factory=dict)
    resource: dict[str, Any] = field(default_factory=dict)
    error: dict[str, Any] = field(default_factory=dict)

    def to_record(self) -> dict[str, Any]:
        context = {**get_log_context(), **sanitize_for_audit(self.context)}
        metrics = sanitize_for_audit(self.metrics)
        resource = sanitize_for_audit(self.resource)
        error = sanitize_for_audit(self.error)
        record = {
            "timestamp": self.timestamp,
            "level": self.level,
            "event_type": self.event_type,
            "action": self.action,
            "outcome": self.outcome,
            "service": self.service,
            "request_id": context.get("request_id", ""),
            "trace_id": context.get("trace_id", ""),
            "username": context.get("username", ""),
            "batch_id": context.get("batch_id", ""),
            "case_key": context.get("case_key", ""),
            "file_id": context.get("file_id", ""),
            "job_id": context.get("job_id", ""),
            "document_type": context.get("document_type", ""),
            "provider": context.get("provider", ""),
            "model": context.get("model", ""),
            "collection": resource.get("collection", ""),
            "duration_ms": metrics.get("duration_ms"),
            "error_code": error.get("code", ""),
            "error_class": error.get("class", ""),
            "context": context,
            "metrics": metrics,
            "resource": resource,
            "error": error,
        }
        return record


class ContextFilter(logging.Filter):
    def filter(self, record: logging.LogRecord) -> bool:
        context = get_log_context()
        for key in _LOG_CONTEXT_KEYS:
            setattr(record, key, context.get(key, ""))
        return True


class HumanConsoleFormatter(logging.Formatter):
    def format(self, record: logging.LogRecord) -> str:
        if not hasattr(record, "request_id"):
            ContextFilter().filter(record)
        timestamp = datetime.fromtimestamp(record.created, tz=UTC).isoformat()
        context = []
        if record.trace_id:
            context.append(f"trace={record.trace_id}")
        if record.request_id:
            context.append(f"req={record.request_id}")
        if record.username:
            context.append(f"user={record.username}")
        suffix = f" [{' '.join(context)}]" if context else ""
        return f"{timestamp} {record.levelname} {record.name}{suffix} {record.getMessage()}"


class JsonAuditFormatter(logging.Formatter):
    def format(self, record: logging.LogRecord) -> str:
        if hasattr(record, "audit_event"):
            payload = dict(record.audit_event)
        else:
            payload = {
                "timestamp": datetime.fromtimestamp(record.created, tz=UTC).isoformat(),
                "level": record.levelname,
                "event_type": "app.log",
                "action": record.name,
                "outcome": "error" if record.levelno >= logging.ERROR else "success",
                "service": record.name,
                "message": record.getMessage(),
                "context": get_log_context(),
            }
        return json.dumps(payload, ensure_ascii=False, default=_coerce_jsonable)


def configure_logging() -> None:
    global _LOGGING_CONFIGURED
    if _LOGGING_CONFIGURED:
        return

    audit_enabled = _env_flag("AUDIT_LOG_ENABLED", True)
    audit_level_name = os.getenv("AUDIT_LOG_LEVEL", "INFO").upper()
    app_level_name = os.getenv("APP_LOG_LEVEL", audit_level_name).upper()
    audit_log_file = os.getenv("AUDIT_LOG_FILE", "logs/audit.jsonl")
    root_level = getattr(logging, app_level_name, logging.INFO)
    audit_level = getattr(logging, audit_level_name, logging.INFO)

    console_handler = logging.StreamHandler()
    console_handler.setLevel(root_level)
    console_handler.setFormatter(HumanConsoleFormatter())
    console_handler.addFilter(ContextFilter())

    root_logger = logging.getLogger()
    root_logger.handlers.clear()
    root_logger.setLevel(root_level)
    root_logger.addHandler(console_handler)

    audit_logger = logging.getLogger("audit")
    audit_logger.handlers.clear()
    audit_logger.propagate = False
    audit_logger.setLevel(audit_level)
    audit_logger.addHandler(console_handler)

    if audit_enabled:
        log_path = Path(audit_log_file)
        log_path.parent.mkdir(parents=True, exist_ok=True)
        file_handler = RotatingFileHandler(
            log_path,
            maxBytes=_safe_int(os.getenv("AUDIT_LOG_MAX_BYTES"), 10_485_760),
            backupCount=_safe_int(os.getenv("AUDIT_LOG_BACKUP_COUNT"), 10),
            encoding="utf-8",
        )
        file_handler.setLevel(audit_level)
        file_handler.setFormatter(JsonAuditFormatter())
        file_handler.addFilter(ContextFilter())
        audit_logger.addHandler(file_handler)

    logging.getLogger("uvicorn.access").handlers.clear()
    logging.getLogger("uvicorn.access").propagate = True
    _LOGGING_CONFIGURED = True


def get_logger(name: str) -> logging.Logger:
    configure_logging()
    return logging.getLogger(name)


class MongoAuditRepository:
    def __init__(self, collection: Any, *, retention_days: int = 180) -> None:
        self.collection = collection
        self.retention_days = retention_days
        self._indexes_ready = False
        self._logger = logging.getLogger(__name__)

    def ensure_indexes(self) -> None:
        if self._indexes_ready:
            return
        try:
            self.collection.create_index([("timestamp", 1)])
            self.collection.create_index(
                [("expires_at", 1)],
                expireAfterSeconds=0,
                name="audit_events_ttl",
            )
        except Exception:
            self._logger.exception("Failed to ensure audit indexes")
        self._indexes_ready = True

    def persist(self, event: dict[str, Any]) -> None:
        try:
            payload = dict(event)
            payload["expires_at"] = _now_utc() + timedelta(days=max(1, int(self.retention_days)))
            self.collection.insert_one(payload)
        except Exception:
            self._logger.exception("Failed to persist audit event")


def set_audit_repository(repository: MongoAuditRepository | None) -> None:
    global _AUDIT_REPOSITORY
    _AUDIT_REPOSITORY = repository


class AuditLogger:
    def __init__(self) -> None:
        configure_logging()
        self._logger = logging.getLogger("audit")
        self._fallback_logger = logging.getLogger(__name__)

    def emit(self, event: AuditEvent) -> None:
        payload = event.to_record()
        token = _AUDIT_EMIT_DEPTH.set(_AUDIT_EMIT_DEPTH.get() + 1)
        try:
            self._logger.info(
                "%s %s",
                payload["event_type"],
                payload["action"],
                extra={"audit_event": payload},
            )
            if _AUDIT_REPOSITORY is not None:
                try:
                    _AUDIT_REPOSITORY.persist(payload)
                except Exception:
                    self._fallback_logger.exception("Audit repository write failed")
        finally:
            _AUDIT_EMIT_DEPTH.reset(token)

    def http_event(
        self,
        *,
        action: str,
        outcome: str,
        service: str,
        resource: Mapping[str, Any] | None = None,
        metrics: Mapping[str, Any] | None = None,
        error: Mapping[str, Any] | None = None,
    ) -> None:
        self.emit(
            AuditEvent(
                event_type="http.request_completed",
                action=action,
                outcome=outcome,
                service=service,
                resource=dict(resource or {}),
                metrics=dict(metrics or {}),
                error=dict(error or {}),
            )
        )

    def db_event(
        self,
        *,
        action: str,
        outcome: str,
        service: str,
        collection: str,
        metrics: Mapping[str, Any] | None = None,
        resource: Mapping[str, Any] | None = None,
        error: Mapping[str, Any] | None = None,
    ) -> None:
        if _is_audit_emit_active():
            return
        data = {"collection": collection, **dict(resource or {})}
        self.emit(
            AuditEvent(
                event_type="db.write" if action in {"insert_one", "insert_many", "update_one", "delete_one", "delete_many"} else "db.query",
                action=action,
                outcome=outcome,
                service=service,
                resource=data,
                metrics=dict(metrics or {}),
                error=dict(error or {}),
            )
        )

    def llm_event(
        self,
        *,
        action: str,
        outcome: str,
        service: str,
        provider: str,
        model: str,
        metrics: Mapping[str, Any] | None = None,
        resource: Mapping[str, Any] | None = None,
        error: Mapping[str, Any] | None = None,
    ) -> None:
        context = {"provider": provider, "model": model}
        self.emit(
            AuditEvent(
                event_type="llm.invoke",
                action=action,
                outcome=outcome,
                service=service,
                context=context,
                metrics=dict(metrics or {}),
                resource=dict(resource or {}),
                error=dict(error or {}),
            )
        )

    def rag_event(
        self,
        *,
        action: str,
        outcome: str,
        service: str,
        resource: Mapping[str, Any] | None = None,
        metrics: Mapping[str, Any] | None = None,
        error: Mapping[str, Any] | None = None,
    ) -> None:
        self.emit(
            AuditEvent(
                event_type="rag.search",
                action=action,
                outcome=outcome,
                service=service,
                resource=dict(resource or {}),
                metrics=dict(metrics or {}),
                error=dict(error or {}),
            )
        )

    def batch_event(
        self,
        *,
        action: str,
        outcome: str,
        service: str,
        resource: Mapping[str, Any] | None = None,
        metrics: Mapping[str, Any] | None = None,
        error: Mapping[str, Any] | None = None,
    ) -> None:
        self.emit(
            AuditEvent(
                event_type="batch.status_changed",
                action=action,
                outcome=outcome,
                service=service,
                resource=dict(resource or {}),
                metrics=dict(metrics or {}),
                error=dict(error or {}),
            )
        )

    def business_event(
        self,
        *,
        event_type: str,
        action: str,
        outcome: str,
        service: str,
        resource: Mapping[str, Any] | None = None,
        metrics: Mapping[str, Any] | None = None,
        error: Mapping[str, Any] | None = None,
    ) -> None:
        self.emit(
            AuditEvent(
                event_type=event_type,
                action=action,
                outcome=outcome,
                service=service,
                resource=dict(resource or {}),
                metrics=dict(metrics or {}),
                error=dict(error or {}),
            )
        )


def get_audit_logger() -> AuditLogger:
    global _AUDIT_LOGGER_INSTANCE
    if _AUDIT_LOGGER_INSTANCE is None:
        _AUDIT_LOGGER_INSTANCE = AuditLogger()
    return _AUDIT_LOGGER_INSTANCE


def _error_payload(exc: Exception | None) -> dict[str, Any]:
    if exc is None:
        return {}
    return {
        "class": exc.__class__.__name__,
        "message": str(exc),
    }


class InstrumentedMongoCursor:
    def __init__(
        self,
        *,
        cursor: Any,
        audit_logger: AuditLogger,
        service: str,
        collection_name: str,
        action: str,
        started_at: float,
        resource: Mapping[str, Any],
    ) -> None:
        self._cursor = cursor
        self._audit_logger = audit_logger
        self._service = service
        self._collection_name = collection_name
        self._action = action
        self._started_at = started_at
        self._resource = dict(resource)
        self._emitted = False

    def sort(self, *args: Any, **kwargs: Any) -> InstrumentedMongoCursor:
        self._cursor = self._cursor.sort(*args, **kwargs)
        return self

    def limit(self, *args: Any, **kwargs: Any) -> InstrumentedMongoCursor:
        self._cursor = self._cursor.limit(*args, **kwargs)
        return self

    def __iter__(self):
        count = 0
        try:
            for item in self._cursor:
                count += 1
                yield item
            self._emit("success", result_count=count)
        except Exception as exc:
            self._emit("error", error=exc, result_count=count)
            raise

    def _emit(self, outcome: str, *, error: Exception | None = None, result_count: int | None = None) -> None:
        if self._emitted:
            return
        metrics = {"duration_ms": round((perf_counter() - self._started_at) * 1000, 3)}
        if result_count is not None:
            metrics["result_count"] = result_count
        self._audit_logger.db_event(
            action=self._action,
            outcome=outcome,
            service=self._service,
            collection=self._collection_name,
            resource=self._resource,
            metrics=metrics,
            error=_error_payload(error),
        )
        self._emitted = True


class InstrumentedMongoCollection:
    def __init__(self, collection: Any, *, service: str) -> None:
        self._collection = collection
        self._service = service
        self._audit_logger = get_audit_logger()

    @property
    def name(self) -> str:
        return getattr(self._collection, "name", "")

    @property
    def database(self) -> Any:
        return InstrumentedMongoDatabase(self._collection.database, service=self._service)

    def __getattr__(self, name: str) -> Any:
        return getattr(self._collection, name)

    def _run(self, action: str, *, resource: Mapping[str, Any], fn) -> Any:
        started_at = perf_counter()
        try:
            result = fn()
            metrics = {"duration_ms": round((perf_counter() - started_at) * 1000, 3)}
            if hasattr(result, "matched_count"):
                metrics["matched_count"] = getattr(result, "matched_count", 0)
                metrics["modified_count"] = getattr(result, "modified_count", 0)
                if getattr(result, "upserted_id", None) is not None:
                    metrics["upserted_id"] = str(result.upserted_id)
            elif hasattr(result, "deleted_count"):
                metrics["deleted_count"] = getattr(result, "deleted_count", 0)
            elif hasattr(result, "inserted_id"):
                metrics["inserted_id"] = str(result.inserted_id)
            elif hasattr(result, "inserted_ids"):
                metrics["inserted_count"] = len(list(getattr(result, "inserted_ids", [])))
            elif isinstance(result, str):
                metrics["result"] = result
            elif result is not None and action == "find_one":
                metrics["result_found"] = True
            elif action == "find_one":
                metrics["result_found"] = False
            self._audit_logger.db_event(
                action=action,
                outcome="success",
                service=self._service,
                collection=self.name,
                resource=resource,
                metrics=metrics,
            )
            return result
        except Exception as exc:
            self._audit_logger.db_event(
                action=action,
                outcome="error",
                service=self._service,
                collection=self.name,
                resource=resource,
                metrics={"duration_ms": round((perf_counter() - started_at) * 1000, 3)},
                error=_error_payload(exc),
            )
            raise

    def find_one(self, *args: Any, **kwargs: Any) -> Any:
        query = args[0] if args else kwargs.get("filter")
        resource = {"query": summarize_query(query)}
        return self._run("find_one", resource=resource, fn=lambda: self._collection.find_one(*args, **kwargs))

    def find(self, *args: Any, **kwargs: Any) -> InstrumentedMongoCursor:
        query = args[0] if args else kwargs.get("filter")
        started_at = perf_counter()
        cursor = self._collection.find(*args, **kwargs)
        return InstrumentedMongoCursor(
            cursor=cursor,
            audit_logger=self._audit_logger,
            service=self._service,
            collection_name=self.name,
            action="find",
            started_at=started_at,
            resource={"query": summarize_query(query)},
        )

    def insert_one(self, document: Mapping[str, Any], *args: Any, **kwargs: Any) -> Any:
        resource = {"payload_summary": summarize_query(document), **_payload_counts(document)}
        return self._run("insert_one", resource=resource, fn=lambda: self._collection.insert_one(document, *args, **kwargs))

    def insert_many(self, documents: Iterable[Mapping[str, Any]], *args: Any, **kwargs: Any) -> Any:
        docs = list(documents)
        resource = {"payload_summary": summarize_query(docs), "item_count": len(docs)}
        return self._run("insert_many", resource=resource, fn=lambda: self._collection.insert_many(docs, *args, **kwargs))

    def update_one(self, filter_doc: Mapping[str, Any], update_doc: Mapping[str, Any], *args: Any, **kwargs: Any) -> Any:
        resource = {
            "query": summarize_query(filter_doc),
            "update_summary": summarize_query(update_doc),
            **_payload_counts(update_doc),
        }
        return self._run(
            "update_one",
            resource=resource,
            fn=lambda: self._collection.update_one(filter_doc, update_doc, *args, **kwargs),
        )

    def replace_one(
        self,
        filter_doc: Mapping[str, Any],
        replacement_doc: Mapping[str, Any],
        *args: Any,
        **kwargs: Any,
    ) -> Any:
        resource = {
            "query": summarize_query(filter_doc),
            "payload_summary": summarize_query(replacement_doc),
            **_payload_counts(replacement_doc),
        }
        return self._run(
            "replace_one",
            resource=resource,
            fn=lambda: self._collection.replace_one(filter_doc, replacement_doc, *args, **kwargs),
        )

    def delete_many(self, filter_doc: Mapping[str, Any], *args: Any, **kwargs: Any) -> Any:
        resource = {"query": summarize_query(filter_doc)}
        return self._run("delete_many", resource=resource, fn=lambda: self._collection.delete_many(filter_doc, *args, **kwargs))

    def delete_one(self, filter_doc: Mapping[str, Any], *args: Any, **kwargs: Any) -> Any:
        resource = {"query": summarize_query(filter_doc)}
        return self._run("delete_one", resource=resource, fn=lambda: self._collection.delete_one(filter_doc, *args, **kwargs))

    def aggregate(self, pipeline: list[dict[str, Any]], *args: Any, **kwargs: Any) -> InstrumentedMongoCursor:
        started_at = perf_counter()
        cursor = self._collection.aggregate(pipeline, *args, **kwargs)
        return InstrumentedMongoCursor(
            cursor=cursor,
            audit_logger=self._audit_logger,
            service=self._service,
            collection_name=self.name,
            action="aggregate",
            started_at=started_at,
            resource={"pipeline": summarize_query(pipeline)},
        )

    def create_index(self, keys: list[tuple[str, int]], *args: Any, **kwargs: Any) -> Any:
        resource = {"keys": summarize_query(keys), "options": summarize_query(kwargs)}
        return self._run("create_index", resource=resource, fn=lambda: self._collection.create_index(keys, *args, **kwargs))


class InstrumentedMongoDatabase:
    def __init__(self, database: Any, *, service: str) -> None:
        self._database = database
        self._service = service

    def __getitem__(self, name: str) -> InstrumentedMongoCollection:
        return InstrumentedMongoCollection(self._database[name], service=self._service)

    def __getattr__(self, name: str) -> InstrumentedMongoCollection:
        return InstrumentedMongoCollection(getattr(self._database, name), service=self._service)


class InstrumentedAsyncMongoCursor:
    def __init__(
        self,
        *,
        cursor: Any,
        audit_logger: AuditLogger,
        service: str,
        collection_name: str,
        started_at: float,
        resource: Mapping[str, Any],
    ) -> None:
        self._cursor = cursor
        self._audit_logger = audit_logger
        self._service = service
        self._collection_name = collection_name
        self._started_at = started_at
        self._resource = dict(resource)

    def sort(self, *args: Any, **kwargs: Any) -> InstrumentedAsyncMongoCursor:
        self._cursor = self._cursor.sort(*args, **kwargs)
        return self

    async def to_list(self, length: int | None = None) -> list[Any]:
        try:
            items = await self._cursor.to_list(length=length)
            self._audit_logger.db_event(
                action="find",
                outcome="success",
                service=self._service,
                collection=self._collection_name,
                resource=self._resource,
                metrics={
                    "duration_ms": round((perf_counter() - self._started_at) * 1000, 3),
                    "result_count": len(items),
                },
            )
            return items
        except Exception as exc:
            self._audit_logger.db_event(
                action="find",
                outcome="error",
                service=self._service,
                collection=self._collection_name,
                resource=self._resource,
                metrics={"duration_ms": round((perf_counter() - self._started_at) * 1000, 3)},
                error=_error_payload(exc),
            )
            raise


class InstrumentedAsyncMongoCollection:
    def __init__(self, collection: Any, *, service: str) -> None:
        self._collection = collection
        self._service = service
        self._audit_logger = get_audit_logger()

    @property
    def name(self) -> str:
        return getattr(self._collection, "name", "")

    def find(self, *args: Any, **kwargs: Any) -> InstrumentedAsyncMongoCursor:
        query = args[0] if args else kwargs.get("filter")
        return InstrumentedAsyncMongoCursor(
            cursor=self._collection.find(*args, **kwargs),
            audit_logger=self._audit_logger,
            service=self._service,
            collection_name=self.name,
            started_at=perf_counter(),
            resource={"query": summarize_query(query)},
        )

    async def find_one(self, *args: Any, **kwargs: Any) -> Any:
        query = args[0] if args else kwargs.get("filter")
        started_at = perf_counter()
        try:
            result = await self._collection.find_one(*args, **kwargs)
            self._audit_logger.db_event(
                action="find_one",
                outcome="success",
                service=self._service,
                collection=self.name,
                resource={"query": summarize_query(query)},
                metrics={
                    "duration_ms": round((perf_counter() - started_at) * 1000, 3),
                    "result_found": result is not None,
                },
            )
            return result
        except Exception as exc:
            self._audit_logger.db_event(
                action="find_one",
                outcome="error",
                service=self._service,
                collection=self.name,
                resource={"query": summarize_query(query)},
                metrics={"duration_ms": round((perf_counter() - started_at) * 1000, 3)},
                error=_error_payload(exc),
            )
            raise

    async def insert_one(self, document: Mapping[str, Any], *args: Any, **kwargs: Any) -> Any:
        started_at = perf_counter()
        try:
            result = await self._collection.insert_one(document, *args, **kwargs)
            self._audit_logger.db_event(
                action="insert_one",
                outcome="success",
                service=self._service,
                collection=self.name,
                resource={"payload_summary": summarize_query(document), **_payload_counts(document)},
                metrics={
                    "duration_ms": round((perf_counter() - started_at) * 1000, 3),
                    "inserted_id": str(getattr(result, "inserted_id", "")),
                },
            )
            return result
        except Exception as exc:
            self._audit_logger.db_event(
                action="insert_one",
                outcome="error",
                service=self._service,
                collection=self.name,
                resource={"payload_summary": summarize_query(document), **_payload_counts(document)},
                metrics={"duration_ms": round((perf_counter() - started_at) * 1000, 3)},
                error=_error_payload(exc),
            )
            raise

    async def update_one(self, filter_doc: Mapping[str, Any], update_doc: Mapping[str, Any], *args: Any, **kwargs: Any) -> Any:
        started_at = perf_counter()
        try:
            result = await self._collection.update_one(filter_doc, update_doc, *args, **kwargs)
            self._audit_logger.db_event(
                action="update_one",
                outcome="success",
                service=self._service,
                collection=self.name,
                resource={
                    "query": summarize_query(filter_doc),
                    "update_summary": summarize_query(update_doc),
                },
                metrics={
                    "duration_ms": round((perf_counter() - started_at) * 1000, 3),
                    "matched_count": getattr(result, "matched_count", 0),
                    "modified_count": getattr(result, "modified_count", 0),
                },
            )
            return result
        except Exception as exc:
            self._audit_logger.db_event(
                action="update_one",
                outcome="error",
                service=self._service,
                collection=self.name,
                resource={
                    "query": summarize_query(filter_doc),
                    "update_summary": summarize_query(update_doc),
                },
                metrics={"duration_ms": round((perf_counter() - started_at) * 1000, 3)},
                error=_error_payload(exc),
            )
            raise

    async def delete_one(self, filter_doc: Mapping[str, Any], *args: Any, **kwargs: Any) -> Any:
        started_at = perf_counter()
        try:
            result = await self._collection.delete_one(filter_doc, *args, **kwargs)
            self._audit_logger.db_event(
                action="delete_one",
                outcome="success",
                service=self._service,
                collection=self.name,
                resource={"query": summarize_query(filter_doc)},
                metrics={
                    "duration_ms": round((perf_counter() - started_at) * 1000, 3),
                    "deleted_count": getattr(result, "deleted_count", 0),
                },
            )
            return result
        except Exception as exc:
            self._audit_logger.db_event(
                action="delete_one",
                outcome="error",
                service=self._service,
                collection=self.name,
                resource={"query": summarize_query(filter_doc)},
                metrics={"duration_ms": round((perf_counter() - started_at) * 1000, 3)},
                error=_error_payload(exc),
            )
            raise

    def __getattr__(self, name: str) -> Any:
        return getattr(self._collection, name)


class InstrumentedAsyncMongoDatabase:
    def __init__(self, database: Any, *, service: str) -> None:
        self._database = database
        self._service = service

    def __getitem__(self, name: str) -> InstrumentedAsyncMongoCollection:
        return InstrumentedAsyncMongoCollection(self._database[name], service=self._service)

    def __getattr__(self, name: str) -> InstrumentedAsyncMongoCollection:
        return InstrumentedAsyncMongoCollection(getattr(self._database, name), service=self._service)


def _text_size(value: Any) -> int:
    if value is None:
        return 0
    if isinstance(value, str):
        return len(value)
    if isinstance(value, Mapping):
        return len(json.dumps(sanitize_for_audit(value), ensure_ascii=False, default=_coerce_jsonable))
    if isinstance(value, (list, tuple)):
        total = 0
        for item in value:
            total += _text_size(item)
        return total
    return len(str(value))


class _GroqCompletionsProxy:
    def __init__(self, client: Any, *, provider: str, service: str) -> None:
        self._client = client
        self._provider = provider
        self._service = service
        self._audit_logger = get_audit_logger()

    def create(self, *args: Any, **kwargs: Any) -> Any:
        started_at = perf_counter()
        model = str(kwargs.get("model") or "")
        messages = kwargs.get("messages") or []
        try:
            result = self._client.create(*args, **kwargs)
            content = ""
            try:
                content = result.choices[0].message.content or ""
            except Exception:
                content = ""
            self._audit_logger.llm_event(
                action="chat.completions.create",
                outcome="success",
                service=self._service,
                provider=self._provider,
                model=model,
                metrics={
                    "duration_ms": round((perf_counter() - started_at) * 1000, 3),
                    "input_chars": _text_size(messages),
                    "output_chars": _text_size(content),
                },
                resource={"operation": "chat_completion"},
            )
            return result
        except Exception as exc:
            self._audit_logger.llm_event(
                action="chat.completions.create",
                outcome="error",
                service=self._service,
                provider=self._provider,
                model=model,
                metrics={
                    "duration_ms": round((perf_counter() - started_at) * 1000, 3),
                    "input_chars": _text_size(messages),
                },
                resource={"operation": "chat_completion"},
                error=_error_payload(exc),
            )
            raise


class _GroqChatProxy:
    def __init__(self, chat: Any, *, provider: str, service: str) -> None:
        self.completions = _GroqCompletionsProxy(chat.completions, provider=provider, service=service)


class InstrumentedGroqClient:
    def __init__(self, client: Any, *, provider: str = "groq", service: str = "llm") -> None:
        self._client = client
        self.chat = _GroqChatProxy(client.chat, provider=provider, service=service)

    def __getattr__(self, name: str) -> Any:
        return getattr(self._client, name)


class InstrumentedGeminiModelAdapter:
    def __init__(self, adapter: Any, *, provider: str = "gemini", service: str = "llm") -> None:
        self._adapter = adapter
        self._provider = provider
        self._service = service
        self._audit_logger = get_audit_logger()

    def generate_content(self, contents: Any, generation_config: Any = None):
        started_at = perf_counter()
        model_name = str(getattr(self._adapter, "_model_name", ""))
        try:
            result = self._adapter.generate_content(contents, generation_config=generation_config)
            text = getattr(result, "text", "") or ""
            self._audit_logger.llm_event(
                action="generate_content",
                outcome="success",
                service=self._service,
                provider=self._provider,
                model=model_name,
                metrics={
                    "duration_ms": round((perf_counter() - started_at) * 1000, 3),
                    "input_chars": _text_size(contents),
                    "output_chars": _text_size(text),
                },
                resource={"operation": "generate_content"},
            )
            return result
        except Exception as exc:
            self._audit_logger.llm_event(
                action="generate_content",
                outcome="error",
                service=self._service,
                provider=self._provider,
                model=model_name,
                metrics={
                    "duration_ms": round((perf_counter() - started_at) * 1000, 3),
                    "input_chars": _text_size(contents),
                },
                resource={"operation": "generate_content"},
                error=_error_payload(exc),
            )
            raise

    def __getattr__(self, name: str) -> Any:
        return getattr(self._adapter, name)
