"""FastAPI sub-app for document upload, approval, and presigned URL endpoints (Airflow 3.x plugin).

POST /upload — receives multipart file uploads from BFF services,
validates HMAC headers, writes to SeaweedFS, creates DB record (pending approval).

POST /{document_id}/approve — approve a document for extraction (triggers DAG).
POST /{document_id}/reject — reject a document (deletes S3 file).

GET /{document_id}/download-url — presigned S3 URL with attachment disposition.
GET /{document_id}/preview-url — presigned S3 URL with inline disposition for PDFs.

Mounted on the Airflow api-server via AirflowPlugin.fastapi_apps at
url_prefix="/api/v1/documents". Full path: /api/v1/documents/upload

NOTE: This module runs inside the Airflow api-server process which uses
/home/airflow/.local/bin/python. Do NOT import from src.*.
"""

import logging
import mimetypes
import os
import tempfile
import time
from datetime import datetime, timezone
from pathlib import Path
from typing import Optional
from urllib.parse import quote
from uuid import UUID, uuid4

from botocore.exceptions import ClientError
from fastapi import FastAPI, File, Request, UploadFile
from pydantic import BaseModel, Field
from sqlalchemy import text

from api.config import get_max_filename_length, get_max_upload_bytes, get_presigned_url_ttl, get_upload_bucket
from api.metrics import approve_total, delete_total, record_outcome, reject_total, upload_duration_seconds, upload_total
from api.dag import trigger_dag
from api.db import get_engine
from api.middleware import (
    CorrelationIdMiddleware,
    HMACVerificationMiddleware,
    RequestIdMiddleware,
)
from api.milvus import delete_chunks_by_document_id
from api.responses import error_response, success_response
from api.rls import rls_connection
from api.s3 import (
    delete_object,
    delete_objects,
    ensure_bucket,
    generate_presigned_url,
    head_object,
    list_objects_v2,
    put_object,
    validate_s3_path,
)
from api.validation import (
    PasswordProtectedError,
    check_password_protected,
    sanitize_filename,
    validate_extension,
    validate_magic_bytes,
    validate_markdown_utf8,
)

log = logging.getLogger(__name__)

upload_app = FastAPI(title="TextIQ Upload API")

# Middleware registration — Starlette executes last-added first.
# Execution order: RequestId → CorrelationId → HMAC → routes.
upload_app.add_middleware(HMACVerificationMiddleware)
upload_app.add_middleware(CorrelationIdMiddleware)
upload_app.add_middleware(RequestIdMiddleware)

DAG_ID = "document_extraction_dag"


# ---------------------------------------------------------------------------
# Pydantic models for approval endpoints
# ---------------------------------------------------------------------------


class ApproveRequest(BaseModel):
    note: Optional[str] = Field(None, max_length=500)


class RejectRequest(BaseModel):
    rejection_reason: str = Field(..., min_length=10, max_length=500)



@upload_app.post("/upload")
def upload(request: Request, file: UploadFile = File(...)):
    """Handle document upload: auth -> validate -> S3 -> DB -> DAG trigger.

    NOTE: This is a sync def (not async) so FastAPI runs it in a threadpool,
    which is required because boto3 and SQLAlchemy 1.4 are blocking.
    """
    temp_path = None
    s3_uploaded = False
    s3_key = None
    bucket = get_upload_bucket()
    max_size = get_max_upload_bytes()
    max_filename_len = get_max_filename_length()
    _t0 = time.perf_counter()

    try:
        # 1. Auth context (set by HMACVerificationMiddleware)
        ctx = request.state.hmac_ctx
        tenant_id = ctx["tenant_id"]
        node_id = ctx["node_id"]

        # 2. File extraction
        if not file or not file.filename:
            resp = error_response("MISSING_FILE", "No file provided")
            outcome = record_outcome(upload_total, resp.status_code, tenant_id=tenant_id)
            upload_duration_seconds.labels(outcome=outcome).observe(time.perf_counter() - _t0)
            return resp

        filename = file.filename

        # 3. Validation (fail-fast)
        try:
            validate_extension(filename)
        except ValueError as e:
            resp = error_response("INVALID_EXTENSION", str(e))
            outcome = record_outcome(upload_total, resp.status_code, tenant_id=tenant_id)
            upload_duration_seconds.labels(outcome=outcome).observe(time.perf_counter() - _t0)
            return resp

        ext = Path(filename).suffix.lower()
        is_markdown = ext == ".md"

        # Read file content (sync read — we're in a threadpool)
        content = file.file.read()
        size = len(content)

        if size > max_size:
            resp = error_response("FILE_TOO_LARGE", f"File size {size} exceeds {max_size} bytes")
            outcome = record_outcome(upload_total, resp.status_code, tenant_id=tenant_id)
            upload_duration_seconds.labels(outcome=outcome).observe(time.perf_counter() - _t0)
            return resp

        if is_markdown:
            try:
                validate_markdown_utf8(content)
            except ValueError as e:
                resp = error_response("INVALID_FORMAT", str(e))
                outcome = record_outcome(upload_total, resp.status_code, tenant_id=tenant_id)
                upload_duration_seconds.labels(outcome=outcome).observe(time.perf_counter() - _t0)
                return resp
        else:
            try:
                validate_magic_bytes(content)
            except ValueError as e:
                resp = error_response("INVALID_FORMAT", str(e))
                outcome = record_outcome(upload_total, resp.status_code, tenant_id=tenant_id)
                upload_duration_seconds.labels(outcome=outcome).observe(time.perf_counter() - _t0)
                return resp

            # Save to temp file for password detection
            suffix = Path(filename).suffix
            with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as tmp:
                temp_path = tmp.name
                tmp.write(content)

            try:
                check_password_protected(temp_path)
            except PasswordProtectedError:
                resp = error_response("PASSWORD_PROTECTED", "Document is password-protected")
                outcome = record_outcome(upload_total, resp.status_code, tenant_id=tenant_id)
                upload_duration_seconds.labels(outcome=outcome).observe(time.perf_counter() - _t0)
                return resp

        # 4. Sanitize filename
        safe_filename = sanitize_filename(filename)
        if not safe_filename:
            resp = error_response("INVALID_FILENAME", "Filename is empty after sanitization")
            outcome = record_outcome(upload_total, resp.status_code, tenant_id=tenant_id)
            upload_duration_seconds.labels(outcome=outcome).observe(time.perf_counter() - _t0)
            return resp
        if len(safe_filename) > max_filename_len:
            resp = error_response("FILENAME_TOO_LONG", f"Filename exceeds {max_filename_len} characters")
            outcome = record_outcome(upload_total, resp.status_code, tenant_id=tenant_id)
            upload_duration_seconds.labels(outcome=outcome).observe(time.perf_counter() - _t0)
            return resp

        # 5. Generate IDs and S3 key
        document_id = uuid4()
        s3_key = f"{tenant_id}/{node_id}/{document_id}/{safe_filename}"

        # 6. S3 path validation
        validate_s3_path(tenant_id, node_id, s3_key)

        # 7. S3 upload
        ensure_bucket(bucket)
        put_object(bucket, s3_key, content)
        s3_uploaded = True

        # 8. DB insert (no DAG trigger — approval gate)
        row_status = "pending"
        row_approval_status = "pending"
        engine = get_engine()
        conn = engine.connect()
        txn = conn.begin()
        try:
            now = datetime.now(timezone.utc)
            # Set RLS tenant context for row-level security policy
            conn.execute(text("SET LOCAL app.tenant_id = :tid"), {"tid": str(tenant_id)})
            conn.execute(
                text(
                    "INSERT INTO documents (id, filename, file_path, file_size_bytes, status, "
                    "tenant_id, node_id, created_at, approval_status) "
                    "VALUES (:id, :filename, :file_path, :file_size_bytes, :status, "
                    ":tenant_id, :node_id, :created_at, :approval_status)"
                ),
                {
                    "id": str(document_id),
                    "filename": safe_filename,
                    "file_path": s3_key,
                    "file_size_bytes": size,
                    "status": row_status,
                    "tenant_id": str(tenant_id),
                    "node_id": str(node_id),
                    "created_at": now,
                    "approval_status": row_approval_status,
                },
            )

            txn.commit()
        except Exception:
            txn.rollback()
            if s3_uploaded:
                try:
                    delete_object(bucket, s3_key)
                except Exception as cleanup_err:
                    log.error("Failed to cleanup S3 object %s: %s", s3_key, cleanup_err)
            raise
        finally:
            conn.close()

        # 10. Success response
        log.info(
            "Upload successful: document_id=%s, key=%s, correlation_id=%s, request_id=%s",
            document_id,
            s3_key,
            getattr(request.state, "correlation_id", None),
            getattr(request.state, "request_id", None),
        )
        resp = success_response(
            {"document_id": str(document_id), "filename": safe_filename, "status": row_status},
            status_code=202,
        )
        outcome = record_outcome(upload_total, resp.status_code, tenant_id=tenant_id)
        upload_duration_seconds.labels(outcome=outcome).observe(time.perf_counter() - _t0)
        return resp

    except PermissionError as e:
        log.error("S3 path validation failed: %s", e)
        resp = error_response("ACCESS_DENIED", str(e), status_code=403)
        _tid = request.state.hmac_ctx["tenant_id"]
        outcome = record_outcome(upload_total, resp.status_code, tenant_id=_tid)
        upload_duration_seconds.labels(outcome=outcome).observe(time.perf_counter() - _t0)
        return resp
    except Exception as e:
        log.exception("Upload failed: %s", e)
        resp = error_response("INTERNAL_ERROR", "Upload failed", status_code=500)
        _tid = request.state.hmac_ctx["tenant_id"]
        outcome = record_outcome(upload_total, resp.status_code, tenant_id=_tid)
        upload_duration_seconds.labels(outcome=outcome).observe(time.perf_counter() - _t0)
        return resp
    finally:
        # 11. Cleanup temp file
        if temp_path and os.path.exists(temp_path):
            os.unlink(temp_path)


def _delete_s3_folder(tenant_id: str, node_id: str, document_id: str) -> None:
    """Delete all S3 objects under {tenant}/{node}/{doc_id}/ and the folder itself in both buckets."""
    prefix = f"{tenant_id}/{node_id}/{document_id}/"
    for bucket in (get_upload_bucket(), os.getenv("IMAGE_BUCKET", "textiq-images")):
        try:
            resp = list_objects_v2(bucket, prefix)
            keys = [obj["Key"] for obj in resp.get("Contents", [])]
            if keys:
                delete_objects(bucket, {"Objects": [{"Key": k} for k in keys]})
            # Delete the directory marker itself (SeaweedFS keeps empty dirs)
            delete_object(bucket, prefix)
        except Exception as e:
            log.warning("S3 cleanup failed for bucket %s, document %s: %s", bucket, document_id, e)


class _EarlyReturn(Exception):
    """Control-flow exception to exit RLS block with a response."""

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


@upload_app.delete("/{document_id}")
def delete_document(document_id: str, request: Request):
    """Hard-delete a document and all associated data.

    Order: verify existence → Milvus vectors → PG chunks → PG domain terms
    → PG document row → S3 files.
    """
    try:
        # 1. Validate UUID format
        try:
            doc_uuid = UUID(document_id)
        except ValueError:
            resp = error_response("INVALID_ID", "document_id is not a valid UUID", status_code=422)
            record_outcome(delete_total, resp.status_code, tenant_id=request.state.hmac_ctx["tenant_id"])
            return resp

        # 2. Auth context
        ctx = request.state.hmac_ctx
        tenant_id = ctx["tenant_id"]
        node_id = ctx["node_id"]

        # 3. Verify document exists and belongs to tenant (read-only check)
        with rls_connection(tenant_id) as conn:
            row = conn.execute(
                text(
                    "SELECT id, file_path FROM documents "
                    "WHERE id = :doc_id AND tenant_id = :tid AND node_id = :nid"
                ),
                {"doc_id": str(doc_uuid), "tid": tenant_id, "nid": node_id},
            ).fetchone()

        if not row:
            resp = error_response("NOT_FOUND", "Document not found", status_code=404)
            record_outcome(delete_total, resp.status_code, tenant_id=tenant_id)
            return resp

        # 4. Delete Milvus vectors (fail fast if unavailable)
        try:
            delete_chunks_by_document_id(str(doc_uuid), tenant_id, node_id)
        except Exception as e:
            log.error("Milvus deletion failed, aborting: %s", e)
            resp = error_response(
                "MILVUS_ERROR",
                "Failed to delete vectors from Milvus, no data was modified",
                status_code=500,
            )
            record_outcome(delete_total, resp.status_code, tenant_id=tenant_id)
            return resp

        # 5. Delete document row under RLS (CASCADE handles chunks, domain terms,
        #    extracted_records, symbol_dictionaries)
        with rls_connection(tenant_id) as conn:
            conn.execute(
                text("DELETE FROM documents WHERE id = :doc_id"),
                {"doc_id": str(doc_uuid)},
            )

        # 6. S3 cleanup (best-effort) — delete {doc_id}/ folder in both buckets
        _delete_s3_folder(tenant_id, node_id, str(doc_uuid))

        log.info("Document deleted: document_id=%s", doc_uuid)
        resp = success_response({"document_id": str(doc_uuid)})
        record_outcome(delete_total, resp.status_code, tenant_id=tenant_id)
        return resp

    except Exception as e:
        log.exception("Delete failed: %s", e)
        resp = error_response("INTERNAL_ERROR", "Delete failed", status_code=500)
        record_outcome(delete_total, resp.status_code, tenant_id=request.state.hmac_ctx["tenant_id"])
        return resp


# ---------------------------------------------------------------------------
# Approval endpoints
# ---------------------------------------------------------------------------


@upload_app.post("/{document_id}/approve")
def approve_document(document_id: str, body: ApproveRequest, request: Request):
    """Approve a document for extraction.

    Updates approval columns and triggers the extraction DAG when the
    document's processing status is still 'pending'.
    """
    try:
        try:
            doc_uuid = UUID(document_id)
        except ValueError:
            resp = error_response("INVALID_ID", "document_id is not a valid UUID", status_code=422)
            record_outcome(approve_total, resp.status_code, tenant_id=request.state.hmac_ctx["tenant_id"])
            return resp

        ctx = request.state.hmac_ctx
        tenant_id = ctx["tenant_id"]
        node_id = ctx["node_id"]
        user_id = ctx["user_id"]

        now = datetime.now(timezone.utc)

        with rls_connection(tenant_id) as conn:
            row = conn.execute(
                text(
                    "SELECT id, status, file_path FROM documents "
                    "WHERE id = :doc_id AND tenant_id = :tid AND node_id = :nid "
                    "FOR UPDATE"
                ),
                {"doc_id": str(doc_uuid), "tid": tenant_id, "nid": node_id},
            ).fetchone()

            if not row:
                raise _EarlyReturn(
                    error_response("NOT_FOUND", "Document not found", status_code=404)
                )

            doc_status = row[1]
            file_path = row[2]

            # Update all 5 approval columns atomically
            conn.execute(
                text(
                    "UPDATE documents SET "
                    "approval_status = 'approved', "
                    "reviewed_by = :user_id, "
                    "reviewed_at = :now, "
                    "rejection_reason = NULL, "
                    "reviewed_note = :note "
                    "WHERE id = :doc_id AND tenant_id = :tid AND node_id = :nid"
                ),
                {
                    "user_id": user_id,
                    "now": now,
                    "note": body.note,
                    "doc_id": str(doc_uuid),
                    "tid": tenant_id,
                    "nid": node_id,
                },
            )

        # Trigger DAG only when document hasn't been processed yet
        if doc_status == "pending":
            bucket = get_upload_bucket()
            conf = {
                "document_id": str(doc_uuid),
                "bucket": bucket,
                "object_key": file_path,
                "tenant_id": tenant_id,
                "node_id": node_id,
                "correlation_id": getattr(request.state, "correlation_id", None),
                "request_id": getattr(request.state, "request_id", None),
            }
            trigger_dag(DAG_ID, conf=conf)

        log.info("Document approved: document_id=%s, dag_triggered=%s", doc_uuid, doc_status == "pending")
        resp = success_response({
            "document_id": str(doc_uuid),
            "approval_status": "approved",
            "dag_triggered": doc_status == "pending",
        })
        record_outcome(approve_total, resp.status_code, tenant_id=tenant_id)
        return resp

    except _EarlyReturn as e:
        record_outcome(approve_total, e.response.status_code, tenant_id=request.state.hmac_ctx["tenant_id"])
        return e.response
    except Exception as e:
        log.exception("Approve failed: %s", e)
        resp = error_response("INTERNAL_ERROR", "Approve failed", status_code=500)
        record_outcome(approve_total, resp.status_code, tenant_id=request.state.hmac_ctx["tenant_id"])
        return resp


@upload_app.post("/{document_id}/reject")
def reject_document(document_id: str, body: RejectRequest, request: Request):
    """Reject a document.

    Updates approval columns, and deletes the S3 file when the document's
    processing status is still 'pending' (file hasn't been extracted yet).
    """
    try:
        try:
            doc_uuid = UUID(document_id)
        except ValueError:
            resp = error_response("INVALID_ID", "document_id is not a valid UUID", status_code=422)
            record_outcome(reject_total, resp.status_code, tenant_id=request.state.hmac_ctx["tenant_id"])
            return resp

        ctx = request.state.hmac_ctx
        tenant_id = ctx["tenant_id"]
        node_id = ctx["node_id"]
        user_id = ctx["user_id"]

        now = datetime.now(timezone.utc)

        with rls_connection(tenant_id) as conn:
            row = conn.execute(
                text(
                    "SELECT id, status FROM documents "
                    "WHERE id = :doc_id AND tenant_id = :tid AND node_id = :nid "
                    "FOR UPDATE"
                ),
                {"doc_id": str(doc_uuid), "tid": tenant_id, "nid": node_id},
            ).fetchone()

            if not row:
                raise _EarlyReturn(
                    error_response("NOT_FOUND", "Document not found", status_code=404)
                )

            doc_status = row[1]

            # Update all 5 approval columns atomically
            conn.execute(
                text(
                    "UPDATE documents SET "
                    "approval_status = 'rejected', "
                    "reviewed_by = :user_id, "
                    "reviewed_at = :now, "
                    "rejection_reason = :reason, "
                    "reviewed_note = NULL "
                    "WHERE id = :doc_id AND tenant_id = :tid AND node_id = :nid"
                ),
                {
                    "user_id": user_id,
                    "now": now,
                    "reason": body.rejection_reason,
                    "doc_id": str(doc_uuid),
                    "tid": tenant_id,
                    "nid": node_id,
                },
            )

        # Only delete S3 files if document hasn't been extracted yet
        if doc_status == "pending":
            _delete_s3_folder(tenant_id, node_id, str(doc_uuid))

        log.info("Document rejected: document_id=%s, s3_deleted=%s", doc_uuid, doc_status == "pending")
        resp = success_response({
            "document_id": str(doc_uuid),
            "approval_status": "rejected",
            "s3_deleted": doc_status == "pending",
        })
        record_outcome(reject_total, resp.status_code, tenant_id=tenant_id)
        return resp

    except _EarlyReturn as e:
        record_outcome(reject_total, e.response.status_code, tenant_id=request.state.hmac_ctx["tenant_id"])
        return e.response
    except Exception as e:
        log.exception("Reject failed: %s", e)
        resp = error_response("INTERNAL_ERROR", "Reject failed", status_code=500)
        record_outcome(reject_total, resp.status_code, tenant_id=request.state.hmac_ctx["tenant_id"])
        return resp


# ---------------------------------------------------------------------------
# Presigned URL endpoints
# ---------------------------------------------------------------------------

PREVIEWABLE_CONTENT_TYPES = {"application/pdf"}


def _infer_content_type(filename: str) -> str:
    """Infer MIME type from filename extension, defaulting to application/octet-stream."""
    content_type, _ = mimetypes.guess_type(filename)
    return content_type or "application/octet-stream"


def _generate_presigned_response(document_id: str, request: Request, inline_preference: bool):
    """Shared logic for download-url and preview-url endpoints.

    Validates UUID, extracts HMAC context, queries document with RLS,
    verifies S3 object existence, and generates a presigned URL.
    """
    try:
        try:
            doc_uuid = UUID(document_id)
        except ValueError:
            return error_response("INVALID_ID", "document_id is not a valid UUID", status_code=422)

        ctx = request.state.hmac_ctx
        tenant_id = ctx["tenant_id"]
        node_id = ctx["node_id"]

        # Query document scoped by tenant + node under RLS
        with rls_connection(tenant_id) as conn:
            row = conn.execute(
                text(
                    "SELECT id, filename, file_path, file_size_bytes "
                    "FROM documents "
                    "WHERE id = :doc_id AND tenant_id = :tid AND node_id = :nid"
                ),
                {"doc_id": str(doc_uuid), "tid": tenant_id, "nid": node_id},
            ).fetchone()

        if not row:
            return error_response("NOT_FOUND", "Document not found", status_code=404)

        filename = row._mapping["filename"]
        file_path = row._mapping["file_path"]
        file_size_bytes = row._mapping["file_size_bytes"]

        # Validate S3 key matches tenant/node scope
        validate_s3_path(tenant_id, node_id, file_path)

        bucket = get_upload_bucket()

        # Verify object exists in S3 before generating presigned URL
        try:
            head_object(bucket, file_path)
        except ClientError as e:
            if e.response["Error"]["Code"] in ("404", "NoSuchKey"):
                return error_response("FILE_NOT_FOUND", "File not found in storage", status_code=404)
            raise

        content_type = _infer_content_type(filename)
        inline = inline_preference and content_type in PREVIEWABLE_CONTENT_TYPES

        # Build Content-Disposition header value (RFC 5987 filename* for Unicode safety)
        encoded_filename = quote(filename, safe="")
        if inline:
            disposition = f"inline; filename*=UTF-8''{encoded_filename}"
        else:
            disposition = f"attachment; filename*=UTF-8''{encoded_filename}"

        ttl = get_presigned_url_ttl()

        presigned_url = generate_presigned_url(
            "get_object",
            {
                "Bucket": bucket,
                "Key": file_path,
                "ResponseContentDisposition": disposition,
                "ResponseContentType": content_type,
            },
            ttl,
        )

        return success_response({
            "url": presigned_url,
            "filename": filename,
            "content_type": content_type,
            "file_size_bytes": file_size_bytes,
            "expires_in": ttl,
            "inline": inline,
        })

    except PermissionError as e:
        log.error("S3 path validation failed: %s", e)
        return error_response("ACCESS_DENIED", str(e), status_code=403)
    except Exception as e:
        log.exception("Presigned URL generation failed: %s", e)
        return error_response("INTERNAL_ERROR", "Failed to generate presigned URL", status_code=500)


@upload_app.get("/{document_id}/download-url")
def download_url(document_id: str, request: Request):
    """Generate a presigned S3 URL for downloading a document (attachment disposition)."""
    return _generate_presigned_response(document_id, request, inline_preference=False)


@upload_app.get("/{document_id}/preview-url")
def preview_url(document_id: str, request: Request):
    """Generate a presigned S3 URL for previewing a document (inline for PDFs, attachment fallback)."""
    return _generate_presigned_response(document_id, request, inline_preference=True)
