"""Request-local buffering for diagnostics that must survive request rollback.

Admin API requests run inside ``ATOMIC_REQUESTS``.  A 4xx/5xx response is
deliberately marked for rollback by ``AdminApiBase`` so that a partial domain
mutation cannot commit without its integration outbox rows.  Database-backed
request/error logs written inside that transaction would be rolled back too.

This module keeps only a request-local, sanitized copy until response middleware
can persist it after the request transaction has exited.
"""

from __future__ import annotations

import datetime
import decimal
import hashlib
import json
import uuid
from collections.abc import Mapping
from contextvars import ContextVar
from dataclasses import dataclass, field
from typing import Any

from django.core.files.uploadedfile import UploadedFile


SENSITIVE_KEY_PARTS = (
    "authorization",
    "cookie",
    "credential",
    "password",
    "secret",
    "signature",
    "token",
)
MAX_IDENTIFIER_LENGTH = 255


def _safe_identifier(value: Any) -> str | None:
    if value in (None, ""):
        return None
    return str(value).replace("\r", "").replace("\n", "")[:MAX_IDENTIFIER_LENGTH]


def _sensitive_key(value: Any) -> bool:
    normalized = str(value).strip().lower().replace("-", "_")
    return any(part in normalized for part in SENSITIVE_KEY_PARTS)


def sanitize_diagnostic_value(value: Any) -> Any:
    """Return JSON-safe diagnostics without credential-bearing fields."""

    if isinstance(value, UploadedFile):
        return {
            "name": value.name,
            "size": value.size,
            "content_type": getattr(value, "content_type", None),
        }
    if isinstance(value, Mapping):
        return {
            str(key): (
                "[REDACTED]"
                if _sensitive_key(key)
                else sanitize_diagnostic_value(item)
            )
            for key, item in value.items()
        }
    if isinstance(value, (list, tuple, set)):
        return [sanitize_diagnostic_value(item) for item in value]
    if isinstance(value, (datetime.date, datetime.datetime, datetime.time)):
        return value.isoformat()
    if isinstance(value, decimal.Decimal):
        return float(value)
    if value is None or isinstance(value, (bool, int, float, str)):
        return value
    return str(value)


@dataclass
class BufferedCriticalError:
    user_id: int | None
    descr: str
    source: str
    trace: str
    data: Any = None

    def safe_payload(self) -> dict[str, Any]:
        return {
            "user_id": self.user_id,
            "descr": self.descr,
            "source": self.source,
            "trace": self.trace,
            "data": sanitize_diagnostic_value(self.data),
        }

    def sha256(self) -> str:
        serialized = json.dumps(
            self.safe_payload(), sort_keys=True, separators=(",", ":"), default=str
        )
        return hashlib.sha256(serialized.encode("utf-8")).hexdigest()


@dataclass
class AdminRequestDiagnostic:
    request_id: str | None
    correlation_id: str
    path: str
    ip: str | None
    user_id: int | None = None
    incoming_params: Any = None
    outgoing_params: Any = None
    status_code: int | None = None
    original_log_id: int | None = None
    errors: list[BufferedCriticalError] = field(default_factory=list)


_buffer: ContextVar[AdminRequestDiagnostic | None] = ContextVar(
    "timetabler_admin_request_diagnostic", default=None
)


def start_admin_diagnostic(
    *, request_id: Any, correlation_id: Any, path: str, ip: str | None
):
    safe_request_id = _safe_identifier(request_id)
    safe_correlation_id = _safe_identifier(correlation_id) or str(uuid.uuid4())
    diagnostic = AdminRequestDiagnostic(
        request_id=safe_request_id,
        correlation_id=safe_correlation_id,
        path=str(path)[:255],
        ip=_safe_identifier(ip),
    )
    return diagnostic, _buffer.set(diagnostic)


def reset_admin_diagnostic(token) -> None:
    _buffer.reset(token)


def current_admin_diagnostic() -> AdminRequestDiagnostic | None:
    return _buffer.get()


def current_diagnostic_identifiers() -> tuple[str | None, str | None]:
    diagnostic = current_admin_diagnostic()
    if diagnostic is None:
        return None, None
    return diagnostic.request_id, diagnostic.correlation_id


def record_admin_request(
    *, user_id: int | None, incoming_params: Any, path: str, ip: str | None
) -> None:
    diagnostic = current_admin_diagnostic()
    if diagnostic is None:
        return
    diagnostic.user_id = user_id
    diagnostic.incoming_params = sanitize_diagnostic_value(incoming_params)
    diagnostic.path = str(path)[:255]
    diagnostic.ip = _safe_identifier(ip)


def record_admin_log_id(log_id: int | None) -> None:
    diagnostic = current_admin_diagnostic()
    if diagnostic is not None:
        diagnostic.original_log_id = log_id


def record_admin_response(*, outgoing_params: Any, status_code: int) -> None:
    diagnostic = current_admin_diagnostic()
    if diagnostic is None:
        return
    diagnostic.outgoing_params = sanitize_diagnostic_value(outgoing_params)
    diagnostic.status_code = int(status_code)


def buffer_critical_error(
    *, user_id: int | None, descr: Any, source: Any, trace: Any, data: Any
) -> BufferedCriticalError | None:
    diagnostic = current_admin_diagnostic()
    if diagnostic is None:
        return None
    error = BufferedCriticalError(
        user_id=user_id,
        descr=str(descr or ""),
        source=str(source or "")[:255],
        trace=str(trace or ""),
        data=sanitize_diagnostic_value(data),
    )
    diagnostic.errors.append(error)
    return error
