"""Persist failed admin diagnostics after the ATOMIC_REQUESTS rollback."""

from __future__ import annotations

import logging
import traceback

from django.db import connection, transaction

from api.services.diagnostics import (
    buffer_critical_error,
    current_admin_diagnostic,
    record_admin_response,
    reset_admin_diagnostic,
    sanitize_diagnostic_value,
    start_admin_diagnostic,
)


logger = logging.getLogger("api.admin.diagnostics")


def _client_ip(request) -> str | None:
    forwarded = request.META.get("HTTP_X_FORWARDED_FOR")
    if forwarded:
        return forwarded.split(",", 1)[0].strip()
    return request.META.get("REMOTE_ADDR")


def _response_payload(response, status_code: int):
    data = getattr(response, "data", None)
    if data is None:
        return {"code": status_code}
    return sanitize_diagnostic_value(data)


def persist_failed_admin_diagnostics(diagnostic) -> None:
    """Write diagnostics in a transaction distinct from the failed request."""

    if connection.in_atomic_block:
        raise RuntimeError(
            "refusing to persist failed admin diagnostics inside an atomic block"
        )

    from api.models import ErrorLog, IncomingApiAdmin

    with transaction.atomic():
        log = None
        if diagnostic.original_log_id is not None:
            log = IncomingApiAdmin.objects.filter(
                pk=diagnostic.original_log_id
            ).first()
        if log is None:
            log = IncomingApiAdmin.objects.create(
                user_id=diagnostic.user_id,
                incoming_params=diagnostic.incoming_params,
                outgoing_params=diagnostic.outgoing_params,
                url=diagnostic.path,
                ip=diagnostic.ip,
                scode=diagnostic.status_code,
                request_id=diagnostic.request_id,
                correlation_id=diagnostic.correlation_id,
            )
        else:
            log.user_id = diagnostic.user_id
            log.incoming_params = diagnostic.incoming_params
            log.outgoing_params = diagnostic.outgoing_params
            log.url = diagnostic.path
            log.ip = diagnostic.ip
            log.scode = diagnostic.status_code
            log.request_id = diagnostic.request_id
            log.correlation_id = diagnostic.correlation_id
            log.save(
                update_fields=[
                    "user_id",
                    "incoming_params",
                    "outgoing_params",
                    "url",
                    "ip",
                    "scode",
                    "request_id",
                    "correlation_id",
                    "updated_at",
                ]
            )

        ErrorLog.objects.bulk_create(
            [
                ErrorLog(
                    user_id=error.user_id,
                    descr=error.descr,
                    url=error.source,
                    trace=error.trace,
                    data=error.data,
                    request_id=diagnostic.request_id,
                    correlation_id=diagnostic.correlation_id,
                )
                for error in diagnostic.errors
            ]
        )


class AdminDiagnosticMiddleware:
    """Buffer admin diagnostics and persist failures after view atomicity ends."""

    def __init__(self, get_response):
        self.get_response = get_response

    def __call__(self, request):
        if not request.path.startswith("/api/admin/"):
            return self.get_response(request)

        request_id = request.headers.get("X-Request-ID")
        correlation_id = request.headers.get("X-Correlation-ID") or request_id
        diagnostic, token = start_admin_diagnostic(
            request_id=request_id,
            correlation_id=correlation_id,
            path=request.path,
            ip=_client_ip(request),
        )
        request._admin_diagnostic_correlation_id = diagnostic.correlation_id
        try:
            response = self.get_response(request)
            status_code = int(response.status_code)
            if not response.has_header("X-Correlation-ID"):
                response["X-Correlation-ID"] = diagnostic.correlation_id
            if diagnostic.request_id and not response.has_header("X-Request-ID"):
                response["X-Request-ID"] = diagnostic.request_id
            if diagnostic.status_code is None:
                record_admin_response(
                    outgoing_params=_response_payload(response, status_code),
                    status_code=status_code,
                )
            if status_code >= 400:
                try:
                    persist_failed_admin_diagnostics(diagnostic)
                except Exception:
                    logger.exception(
                        "failed_admin_diagnostic_persistence",
                        extra={
                            "request_id": diagnostic.request_id,
                            "correlation_id": diagnostic.correlation_id,
                            "path": diagnostic.path,
                            "status_code": diagnostic.status_code,
                            "error_sha256": [
                                error.sha256() for error in diagnostic.errors
                            ],
                        },
                    )
            return response
        except Exception as error:
            if not diagnostic.errors:
                buffer_critical_error(
                    user_id=diagnostic.user_id,
                    descr=str(error),
                    source=(
                        error.__traceback__.tb_frame.f_code.co_filename
                        if error.__traceback__
                        else ""
                    ),
                    trace="".join(
                        traceback.format_exception(type(error), error, error.__traceback__)
                    ),
                    data=None,
                )
            record_admin_response(outgoing_params={"code": 500}, status_code=500)
            try:
                persist_failed_admin_diagnostics(diagnostic)
            except Exception:
                logger.exception(
                    "failed_admin_diagnostic_persistence",
                    extra={
                        "request_id": diagnostic.request_id,
                        "correlation_id": diagnostic.correlation_id,
                        "path": diagnostic.path,
                        "status_code": 500,
                        "error_sha256": [
                            item.sha256() for item in diagnostic.errors
                        ],
                    },
                )
            raise
        finally:
            reset_admin_diagnostic(token)

    def process_exception(self, request, exception):
        """Capture exceptions Django exposes through its exception hook."""

        diagnostic = current_admin_diagnostic()
        if diagnostic is None or diagnostic.errors:
            return None
        buffer_critical_error(
            user_id=diagnostic.user_id,
            descr=str(exception),
            source=(
                exception.__traceback__.tb_frame.f_code.co_filename
                if exception.__traceback__
                else ""
            ),
            trace="".join(
                traceback.format_exception(
                    type(exception), exception, exception.__traceback__
                )
            ),
            data=None,
        )
        return None
