"""DMS AI — API views.

Endpoints
---------
OCR:
  GET  /dms/documents/{doc_id}/versions/{version_id}/ocr/
      List OCR results for a document version.
  POST /dms/documents/{doc_id}/versions/{version_id}/ocr/
      Submit a new OCR job (creates OCRResult in PENDING state).

  GET  /dms/documents/{doc_id}/versions/{version_id}/ocr/{ocr_id}/
      Retrieve a specific OCR result.
  POST /dms/documents/{doc_id}/versions/{version_id}/ocr/{ocr_id}/record/
      Record OCR results (used by provider adapters/Celery tasks).
  POST /dms/documents/{doc_id}/versions/{version_id}/ocr/{ocr_id}/fail/
      Record an OCR failure.

AI Classification:
  GET  /dms/documents/{doc_id}/versions/{version_id}/classifications/
      List AI classifications for a document version.
  POST /dms/documents/{doc_id}/versions/{version_id}/classifications/
      Submit a new AI classification job.

  GET  /dms/documents/{doc_id}/versions/{version_id}/classifications/{cls_id}/
      Retrieve a specific AI classification.
  PUT  /dms/documents/{doc_id}/versions/{version_id}/classifications/{cls_id}/record/
      Record classification result (used by provider adapters/Celery tasks).

Extracted Entities:
  GET  /dms/documents/{doc_id}/versions/{version_id}/entities/
      List extracted entities for a document version.
  POST /dms/documents/{doc_id}/versions/{version_id}/entities/
      Bulk-add extracted entities.
  DELETE /dms/documents/{doc_id}/versions/{version_id}/entities/
      Clear (soft-delete) all extracted entities for a version.
"""

from __future__ import annotations

from rest_framework import status
from rest_framework.decorators import api_view, permission_classes
from rest_framework.permissions import IsAuthenticated
from rest_framework.request import Request
from rest_framework.response import Response

from simorgh.apps.dms.ai import queries, services
from simorgh.apps.dms.ai.iam_permissions import (
    PERM_AI_SUBMIT,
    PERM_AI_VIEW,
    PERM_ENTITIES_VIEW,
    PERM_OCR_SUBMIT,
    PERM_OCR_VIEW,
)
from simorgh.apps.dms.ai.serializers import (
    AIClassificationRecordSerializer,
    AIClassificationSerializer,
    AIClassificationSubmitSerializer,
    ExtractedEntityBulkSerializer,
    ExtractedEntitySerializer,
    OCRFailSerializer,
    OCRRecordResultSerializer,
    OCRResultSerializer,
    OCRSubmitSerializer,
)
from simorgh.apps.dms.ai.services import AIClassificationError, OCRError
from simorgh.apps.dms.common.exceptions import AssetNotFound


# ---------------------------------------------------------------------------
# Shared helpers
# ---------------------------------------------------------------------------

def _tenant(request: Request):
    tenant = getattr(request, "tenant", None)
    if tenant is None:
        from rest_framework.exceptions import PermissionDenied
        raise PermissionDenied("Tenant required.")
    return tenant


def _org_node_id(request: Request) -> int:
    from simorgh.apps.memberships.models import Membership, MembershipStatus

    t = _tenant(request)
    m = (
        Membership.objects.filter(
            tenant=t, users=request.user, status=MembershipStatus.ACTIVE
        )
        .order_by("id")
        .first()
    )
    if m is not None:
        return m.organization_node_id
    if getattr(request.user, "is_superuser", False):
        from simorgh.apps.organizations.models import OrganizationNode

        root = OrganizationNode.objects.filter(tenant=t).order_by("id").first()
        if root is not None:
            return root.pk
    from rest_framework.exceptions import PermissionDenied
    raise PermissionDenied("No active membership for this tenant.")


def _require_perm(request: Request, codename: str) -> None:
    from rest_framework.exceptions import PermissionDenied
    from simorgh.core.context import current_request_context

    ctx = current_request_context()
    if ctx.is_superuser:
        return
    if codename not in (ctx.permissions or set()):
        raise PermissionDenied(f"Permission required: {codename}")


def _get_document_version(request: Request, doc_id: str, version_id: str):
    """Return the DocumentVersion ensuring it belongs to the request tenant."""
    from simorgh.apps.dms.documents.queries import get_document, get_version

    tenant = _tenant(request)
    document = get_document(tenant.pk, doc_id)
    return get_version(document, version_id)


# ---------------------------------------------------------------------------
# OCR views
# ---------------------------------------------------------------------------

@api_view(["GET", "POST"])
@permission_classes([IsAuthenticated])
def ocr_result_list(request: Request, doc_id: str, version_id: str) -> Response:
    """
    GET  — list OCR results for a document version.
    POST — submit a new OCR job.
    """
    if request.method == "GET":
        _require_perm(request, PERM_OCR_VIEW)
        tenant = _tenant(request)
        qs = queries.list_ocr_results(tenant.pk, version_id)
        return Response(OCRResultSerializer(qs, many=True).data)

    # POST — submit job
    _require_perm(request, PERM_OCR_SUBMIT)
    ser = OCRSubmitSerializer(data=request.data)
    ser.is_valid(raise_exception=True)

    try:
        version = _get_document_version(request, doc_id, version_id)
    except AssetNotFound as exc:
        return Response({"detail": str(exc)}, status=status.HTTP_404_NOT_FOUND)

    tenant = _tenant(request)
    result = services.submit_ocr(
        tenant_id=tenant.pk,
        organization_node_id=_org_node_id(request),
        document_version_id=version.pk,
        provider=ser.validated_data.get("provider", "default"),
        language_code=ser.validated_data.get("language_code", ""),
        submitted_by=request.user,
    )
    return Response(OCRResultSerializer(result).data, status=status.HTTP_201_CREATED)


@api_view(["GET"])
@permission_classes([IsAuthenticated])
def ocr_result_detail(
    request: Request, doc_id: str, version_id: str, ocr_id: str
) -> Response:
    """Retrieve a specific OCR result."""
    _require_perm(request, PERM_OCR_VIEW)
    tenant = _tenant(request)
    try:
        result = queries.get_ocr_result(tenant.pk, ocr_id)
    except AssetNotFound as exc:
        return Response({"detail": str(exc)}, status=status.HTTP_404_NOT_FOUND)
    return Response(OCRResultSerializer(result).data)


@api_view(["POST"])
@permission_classes([IsAuthenticated])
def ocr_record_result(
    request: Request, doc_id: str, version_id: str, ocr_id: str
) -> Response:
    """Record OCR processing results. Called by provider adapters / Celery tasks."""
    _require_perm(request, PERM_OCR_SUBMIT)
    tenant = _tenant(request)
    try:
        ocr_result = queries.get_ocr_result(tenant.pk, ocr_id)
    except AssetNotFound as exc:
        return Response({"detail": str(exc)}, status=status.HTTP_404_NOT_FOUND)

    ser = OCRRecordResultSerializer(data=request.data)
    ser.is_valid(raise_exception=True)
    d = ser.validated_data
    try:
        updated = services.record_ocr_result(
            ocr_result=ocr_result,
            full_text=d.get("full_text", ""),
            page_count=d.get("page_count", 0),
            confidence_score=d.get("confidence_score"),
            language_code=d.get("language_code", ""),
            processing_metadata=d.get("processing_metadata"),
            updated_by=request.user,
        )
    except OCRError as exc:
        return Response({"detail": str(exc)}, status=status.HTTP_400_BAD_REQUEST)
    return Response(OCRResultSerializer(updated).data)


@api_view(["POST"])
@permission_classes([IsAuthenticated])
def ocr_fail(
    request: Request, doc_id: str, version_id: str, ocr_id: str
) -> Response:
    """Record an OCR failure."""
    _require_perm(request, PERM_OCR_SUBMIT)
    tenant = _tenant(request)
    try:
        ocr_result = queries.get_ocr_result(tenant.pk, ocr_id)
    except AssetNotFound as exc:
        return Response({"detail": str(exc)}, status=status.HTTP_404_NOT_FOUND)

    ser = OCRFailSerializer(data=request.data)
    ser.is_valid(raise_exception=True)
    try:
        updated = services.fail_ocr(
            ocr_result=ocr_result,
            error_message=ser.validated_data["error_message"],
            updated_by=request.user,
        )
    except OCRError as exc:
        return Response({"detail": str(exc)}, status=status.HTTP_400_BAD_REQUEST)
    return Response(OCRResultSerializer(updated).data)


# ---------------------------------------------------------------------------
# AI Classification views
# ---------------------------------------------------------------------------

@api_view(["GET", "POST"])
@permission_classes([IsAuthenticated])
def classification_list(
    request: Request, doc_id: str, version_id: str
) -> Response:
    """
    GET  — list AI classifications for a document version.
    POST — submit a new AI classification job.
    """
    if request.method == "GET":
        _require_perm(request, PERM_AI_VIEW)
        tenant = _tenant(request)
        classification_type = request.query_params.get("type")
        qs = queries.list_ai_classifications(tenant.pk, version_id, classification_type)
        return Response(AIClassificationSerializer(qs, many=True).data)

    # POST — submit classification
    _require_perm(request, PERM_AI_SUBMIT)
    ser = AIClassificationSubmitSerializer(data=request.data)
    ser.is_valid(raise_exception=True)

    try:
        version = _get_document_version(request, doc_id, version_id)
    except AssetNotFound as exc:
        return Response({"detail": str(exc)}, status=status.HTTP_404_NOT_FOUND)

    tenant = _tenant(request)
    classification = services.submit_ai_classification(
        tenant_id=tenant.pk,
        organization_node_id=_org_node_id(request),
        document_version_id=version.pk,
        classification_type=ser.validated_data.get("classification_type", "document_type"),
        provider=ser.validated_data.get("provider", "default"),
        created_by=request.user,
    )
    return Response(
        AIClassificationSerializer(classification).data,
        status=status.HTTP_201_CREATED,
    )


@api_view(["GET"])
@permission_classes([IsAuthenticated])
def classification_detail(
    request: Request, doc_id: str, version_id: str, cls_id: str
) -> Response:
    """Retrieve a specific AI classification."""
    _require_perm(request, PERM_AI_VIEW)
    tenant = _tenant(request)
    try:
        classification = queries.get_ai_classification(tenant.pk, cls_id)
    except AssetNotFound as exc:
        return Response({"detail": str(exc)}, status=status.HTTP_404_NOT_FOUND)
    return Response(AIClassificationSerializer(classification).data)


@api_view(["PUT"])
@permission_classes([IsAuthenticated])
def classification_record(
    request: Request, doc_id: str, version_id: str, cls_id: str
) -> Response:
    """Record AI classification results. Called by provider adapters / Celery tasks."""
    _require_perm(request, PERM_AI_SUBMIT)
    tenant = _tenant(request)
    try:
        classification = queries.get_ai_classification(tenant.pk, cls_id)
    except AssetNotFound as exc:
        return Response({"detail": str(exc)}, status=status.HTTP_404_NOT_FOUND)

    ser = AIClassificationRecordSerializer(data=request.data)
    ser.is_valid(raise_exception=True)
    d = ser.validated_data
    updated = services.record_ai_classification(
        classification=classification,
        label=d.get("label", ""),
        confidence_score=d.get("confidence_score"),
        summary=d.get("summary", ""),
        tags=d.get("tags"),
        processing_metadata=d.get("processing_metadata"),
        updated_by=request.user,
    )
    return Response(AIClassificationSerializer(updated).data)


# ---------------------------------------------------------------------------
# Extracted entity views
# ---------------------------------------------------------------------------

@api_view(["GET", "POST", "DELETE"])
@permission_classes([IsAuthenticated])
def entity_list(request: Request, doc_id: str, version_id: str) -> Response:
    """
    GET    — list extracted entities for a document version.
    POST   — bulk-add extracted entities.
    DELETE — clear (soft-delete) all extracted entities for a version.
    """
    if request.method == "GET":
        _require_perm(request, PERM_ENTITIES_VIEW)
        tenant = _tenant(request)
        entity_type = request.query_params.get("entity_type")
        qs = queries.list_extracted_entities(tenant.pk, version_id, entity_type)
        return Response(ExtractedEntitySerializer(qs, many=True).data)

    if request.method == "DELETE":
        _require_perm(request, PERM_OCR_SUBMIT)  # reuse OCR submit for destructive clear
        try:
            version = _get_document_version(request, doc_id, version_id)
        except AssetNotFound as exc:
            return Response({"detail": str(exc)}, status=status.HTTP_404_NOT_FOUND)
        tenant = _tenant(request)
        count = services.clear_extracted_entities(
            tenant_id=tenant.pk,
            document_version_id=version.pk,
            updated_by=request.user,
        )
        return Response({"cleared": count})

    # POST — bulk add entities
    _require_perm(request, PERM_OCR_SUBMIT)
    ser = ExtractedEntityBulkSerializer(data=request.data)
    ser.is_valid(raise_exception=True)

    try:
        version = _get_document_version(request, doc_id, version_id)
    except AssetNotFound as exc:
        return Response({"detail": str(exc)}, status=status.HTTP_404_NOT_FOUND)

    tenant = _tenant(request)
    ocr_result = None
    ocr_id = ser.validated_data.get("ocr_result_id")
    if ocr_id:
        try:
            ocr_result = queries.get_ocr_result(tenant.pk, str(ocr_id))
        except AssetNotFound as exc:
            return Response({"detail": str(exc)}, status=status.HTTP_404_NOT_FOUND)

    entities = services.add_extracted_entities(
        tenant_id=tenant.pk,
        organization_node_id=_org_node_id(request),
        document_version_id=version.pk,
        entities_data=ser.validated_data["entities"],
        ocr_result=ocr_result,
        created_by=request.user,
    )
    return Response(
        ExtractedEntitySerializer(entities, many=True).data,
        status=status.HTTP_201_CREATED,
    )
