from __future__ import annotations

import uuid

from fastapi import APIRouter, Request
from fastapi.exception_handlers import request_validation_exception_handler
from fastapi.exceptions import RequestValidationError
from fastapi.responses import JSONResponse

from app.retrieval.errors import RetrievalServiceError
from app.retrieval.models import (
    RetrieveErrorBody,
    RetrieveErrorResponse,
    RetrieveRequest,
    RetrieveResponse,
)
from app.retrieval.retrieval_service import retrieve

router = APIRouter(tags=["Internal Retrieval"])

INTERNAL_RETRIEVE_PATH = "/internal/retrieve"


def _request_id_from_payload(payload: dict | None) -> str:
    if not payload:
        return f"req_{uuid.uuid4().hex[:12]}"
    value = payload.get("request_id")
    if isinstance(value, str) and value.strip():
        return value.strip()
    return f"req_{uuid.uuid4().hex[:12]}"


def _error_response(
    *,
    status_code: int,
    code: str,
    message: str,
    request_id: str,
) -> JSONResponse:
    body = RetrieveErrorResponse(
        error=RetrieveErrorBody(
            code=code,
            message=message,
            request_id=request_id,
        )
    )
    return JSONResponse(status_code=status_code, content=body.model_dump())


async def _safe_json_body(request: Request) -> dict | None:
    try:
        body = await request.json()
    except Exception:
        return None
    return body if isinstance(body, dict) else None


def _is_internal_retrieve_request(request: Request) -> bool:
    return request.url.path.rstrip("/") == INTERNAL_RETRIEVE_PATH


@router.post(
    "/internal/retrieve",
    response_model=RetrieveResponse,
    response_model_exclude_none=True,
)
async def internal_retrieve(request: RetrieveRequest) -> RetrieveResponse:
    return await retrieve(request)


def register_internal_retrieve_exception_handlers(app) -> None:
    """Attach contract-shaped errors for POST /internal/retrieve."""

    @app.exception_handler(RetrievalServiceError)
    async def _handle_retrieval_service_error(
        request: Request,
        exc: RetrievalServiceError,
    ) -> JSONResponse:
        if not _is_internal_retrieve_request(request):
            raise exc

        request_id = _request_id_from_payload(await _safe_json_body(request))
        return _error_response(
            status_code=exc.status_code,
            code=exc.code,
            message=exc.message,
            request_id=request_id,
        )

    @app.exception_handler(RequestValidationError)
    async def _handle_request_validation_error(
        request: Request,
        exc: RequestValidationError,
    ) -> JSONResponse:
        if not _is_internal_retrieve_request(request):
            return await request_validation_exception_handler(request, exc)

        request_id = _request_id_from_payload(await _safe_json_body(request))
        return _error_response(
            status_code=422,
            code="validation_error",
            message="; ".join(_format_validation_errors(exc.errors())),
            request_id=request_id,
        )


def _format_validation_errors(errors: list[dict]) -> list[str]:
    messages: list[str] = []
    for error in errors:
        location = ".".join(str(part) for part in error.get("loc", ()))
        message = error.get("msg", "invalid value")
        messages.append(f"{location}: {message}" if location else message)
    return messages
