import asyncio
import logging
import tempfile
import uuid
from pathlib import Path
from typing import Annotated

from fastapi import Depends, FastAPI, File, Form, HTTPException, Query, UploadFile
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import FileResponse
from fastapi.staticfiles import StaticFiles
from pydantic import BaseModel, Field

from app import database
from app.auth import ADMIN_PASSWORD, ADMIN_USERNAME, create_access_token, verify_admin_token
from app.config import BASE_DIR, get_settings
from app.models.schemas import (
    AuditItem,
    ChatRequest,
    ChatResponse,
    ConversationDetail,
    ConversationMessage,
    ConversationSummary,
    FeedbackItem,
    FeedbackRequest,
    FeedbackResponse,
    HealthResponse,
    PdfConfirmRequest,
    PdfJobResponse,
    PdfRow,
    ReviewRequest,
)
from app.pdf import (
    calculate_job,
    clear_session_job,
    confirm_job,
    create_job_from_upload,
    get_job,
    ingest_pdf_for_chat,
)
from app.rag.chain import get_rag_service
from app.saas import saas_router
from app.saas.db import init_saas_tables

settings = get_settings()
UPLOAD_DIR = BASE_DIR / "data" / "pdf_uploads"

app = FastAPI(
    title="Accountants SaaS Chatbot",
    description="Multi-tenant RAG chatbot agent for accounting and finance (Ollama + ChromaDB)",
    version="2.0.0",
)

origins = [o.strip() for o in settings.cors_origins.split(",") if o.strip()]
app.add_middleware(
    CORSMiddleware,
    allow_origins=origins if origins != ["*"] else ["*"],
    allow_credentials=True,
    allow_methods=["*"],
    allow_headers=["*"],
)

app.include_router(saas_router)

STATIC_DIR = Path(__file__).resolve().parent.parent / "static"


@app.on_event("startup")
async def startup() -> None:
    logging.basicConfig(
        level=logging.DEBUG if settings.rag_debug else logging.INFO,
        format="%(asctime)s %(levelname)s [%(name)s] %(message)s",
    )
    logging.getLogger("app.rag.kb").setLevel(
        logging.DEBUG if settings.rag_debug else logging.INFO
    )
    if settings.rag_debug:
        logging.getLogger("app.rag.kb").info("RAG_DEBUG enabled")

    database.init_db()
    init_saas_tables()
    try:
        from app.saas.db import ensure_admin_account

        ensure_admin_account()
    except Exception:
        logging.getLogger(__name__).warning("Failed to ensure admin account", exc_info=True)
    try:
        database.warm_pg_pool()
    except Exception:
        logging.getLogger(__name__).warning("Postgres pool warm-up failed", exc_info=True)
    # Privacy mode: remove any previously saved user PDF uploads / job metas
    if not settings.persist_user_data and UPLOAD_DIR.exists():
        for path in UPLOAD_DIR.iterdir():
            if path.is_file() and path.suffix.lower() in {".pdf", ".json"}:
                try:
                    path.unlink()
                except OSError:
                    pass
    service = get_rag_service()
    asyncio.create_task(asyncio.to_thread(service.warmup))

    def _warm_statements() -> None:
        try:
            from app.saas.kb_answer import warm_local_statement_caches

            n = warm_local_statement_caches()
            logging.getLogger(__name__).info("Warmed %s statement PDF cache(s)", n)
        except Exception:
            logging.getLogger(__name__).warning("Statement cache warm-up failed", exc_info=True)

    asyncio.create_task(asyncio.to_thread(_warm_statements))


# ── Health ────────────────────────────────────────────────────────────────────

@app.get("/health", response_model=HealthResponse, tags=["system"])
async def health() -> HealthResponse:
    service = get_rag_service()
    return HealthResponse(
        status="ok",
        ollama_reachable=await service.ollama_reachable(),
        ollama_model=settings.ollama_model,
        vector_store_ready=service.vector_store_ready(),
        document_count=service.document_count(),
    )


# ── Chat ──────────────────────────────────────────────────────────────────────

@app.post("/chat", response_model=ChatResponse, tags=["chat"])
async def chat(request: ChatRequest) -> ChatResponse:
    service = get_rag_service()

    if not service.vector_store_ready():
        raise HTTPException(
            status_code=503,
            detail="Vector store is empty. Run: python scripts/ingest.py",
        )

    if not await service.ollama_reachable():
        raise HTTPException(
            status_code=503,
            detail=f"Ollama is not reachable at {settings.ollama_base_url}. "
                   f"Start Ollama and pull {settings.ollama_model}.",
        )

    try:
        return await asyncio.to_thread(service.chat, request.question, request.session_id)
    except RuntimeError as exc:
        raise HTTPException(status_code=503, detail=str(exc)) from exc


@app.post("/chat/upload-pdf", response_model=ChatResponse, tags=["chat"])
async def chat_upload_pdf(
    file: UploadFile = File(...),
    session_id: str | None = Form(default=None),
    message: str | None = Form(default=None),
) -> ChatResponse:
    """Upload a PDF into the chat session: extract, calculate, and answer-ready.

    Optional ``message`` is answered in the same turn using the uploaded statement.
    By default the PDF is processed from a temp file and deleted immediately (no disk retention).
    """
    if not file.filename or not file.filename.lower().endswith(".pdf"):
        raise HTTPException(status_code=400, detail="Please upload a .pdf file")

    raw = await file.read()
    if not raw:
        raise HTTPException(status_code=400, detail="Empty file")
    if len(raw) > 20 * 1024 * 1024:
        raise HTTPException(status_code=400, detail="PDF too large (max 20 MB)")

    sid = session_id or str(uuid.uuid4())
    saved: Path | None = None
    try:
        if settings.persist_user_data:
            UPLOAD_DIR.mkdir(parents=True, exist_ok=True)
            saved = UPLOAD_DIR / f"{uuid.uuid4()}_{file.filename}"
            saved.write_bytes(raw)
        else:
            tmp = tempfile.NamedTemporaryFile(suffix=".pdf", delete=False)
            try:
                tmp.write(raw)
                tmp.flush()
                saved = Path(tmp.name)
            finally:
                tmp.close()

        job = await asyncio.to_thread(ingest_pdf_for_chat, saved, file.filename, sid)
    except Exception as exc:
        raise HTTPException(status_code=500, detail=f"PDF ingest failed: {exc}") from exc
    finally:
        # Always drop the temp/upload file after extract unless durable mode is on
        if saved is not None and not settings.persist_user_data:
            try:
                saved.unlink(missing_ok=True)
            except OSError:
                pass

    user_note = (message or "").strip()
    summary = (job.get("calculation") or {}).get("summary")
    has_content = bool(
        job.get("rows")
        or job.get("calculation")
        or job.get("document_text")
        or job.get("text_preview")
    )

    if not has_content:
        answer = job.get("error") or f"Could not read content from {file.filename}."
        log_id = database.log_query(
            question=f"[uploaded PDF] {file.filename}",
            answer=answer,
            citations_json="[]",
            duration_ms=0,
            session_id=sid,
        )
        return ChatResponse(
            answer=answer,
            citations=[],
            session_id=sid,
            log_id=log_id or None,
            mode="pdf",
            pdf_job_id=job.get("job_id"),
            pdf_summary=summary,
        )

    # Always answer with the model reading the PDF — never a canned "Loaded…" message
    question = user_note or "What is this PDF about? Summarize what the document contains."
    service = get_rag_service()
    try:
        resp = await asyncio.to_thread(service.chat, question, sid)
        return ChatResponse(
            answer=resp.answer,
            citations=resp.citations,
            session_id=sid,
            log_id=resp.log_id,
            mode="pdf",
            pdf_job_id=job.get("job_id"),
            pdf_summary=summary,
        )
    except RuntimeError as exc:
        raise HTTPException(status_code=503, detail=str(exc)) from exc


@app.delete("/chat/pdf", tags=["chat"])
async def chat_clear_pdf(session_id: str = Query(...)) -> dict:
    """Detach the active uploaded PDF from a chat session so normal RAG resumes."""
    clear_session_job(session_id)
    return {"ok": True, "session_id": session_id}


# ── Conversations (history sidebar) ───────────────────────────────────────────

@app.get("/conversations", response_model=list[ConversationSummary], tags=["chat"])
async def conversations(limit: int = Query(default=40, ge=1, le=100)) -> list[ConversationSummary]:
    rows = database.list_conversations(limit=limit)
    return [ConversationSummary(**r) for r in rows]


@app.get("/conversations/{session_id}", response_model=ConversationDetail, tags=["chat"])
async def conversation_detail(session_id: str) -> ConversationDetail:
    messages = database.get_session_messages(session_id)
    if not messages:
        raise HTTPException(status_code=404, detail="Conversation not found")
    title = "Chat"
    convos = database.list_conversations(limit=200)
    for c in convos:
        if c["id"] == session_id:
            title = c.get("title") or title
            break
    else:
        title = _title_fallback(messages[0]["question"])
    return ConversationDetail(
        id=session_id,
        title=title,
        messages=[ConversationMessage(**m) for m in messages],
    )


@app.delete("/conversations/{session_id}", tags=["chat"])
async def conversation_delete(session_id: str) -> dict:
    database.delete_conversation(session_id)
    return {"ok": True, "session_id": session_id}


def _title_fallback(question: str) -> str:
    text = (question or "").strip()
    if text.startswith("[uploaded PDF]"):
        text = text.replace("[uploaded PDF]", "").strip() or "PDF upload"
    text = " ".join(text.split())
    return (text[:45] + "…") if len(text) > 48 else (text or "Chat")


# ── Feedback ──────────────────────────────────────────────────────────────────

@app.post("/feedback", response_model=FeedbackResponse, tags=["chat"])
async def submit_feedback(body: FeedbackRequest) -> FeedbackResponse:
    """Save thumbs up/down into SQLite for later model training / preference tuning."""
    feedback_id = database.save_feedback(
        log_id=body.log_id,
        session_id=body.session_id,
        question=body.question,
        answer=body.answer,
        rating=body.rating,
        correction=body.correction,
        message_id=body.message_id,
        mode=body.mode,
    )
    return FeedbackResponse(
        feedback_id=feedback_id,
        message="Thank you for your feedback!",
        rating=body.rating,
    )

class FeedbackCorrectionRequest(BaseModel):
    correction: str | None = Field(default=None, max_length=4000)


@app.patch("/feedback/{feedback_id}", tags=["chat"])
async def patch_feedback_correction(
    feedback_id: int,
    body: FeedbackCorrectionRequest,
) -> dict:
    database.update_feedback_correction(feedback_id, body.correction)
    return {"ok": True, "feedback_id": feedback_id}

# ── Admin auth ────────────────────────────────────────────────────────────────

class LoginRequest(BaseModel):
    username: str
    password: str


class TokenResponse(BaseModel):
    access_token: str
    token_type: str = "bearer"


@app.post("/admin/login", response_model=TokenResponse, tags=["admin"])
async def admin_login(body: LoginRequest) -> TokenResponse:
    if body.username != ADMIN_USERNAME or body.password != ADMIN_PASSWORD:
        raise HTTPException(status_code=401, detail="Invalid credentials")
    token = create_access_token({"sub": body.username})
    return TokenResponse(access_token=token)


# ── Admin — feedback management ───────────────────────────────────────────────

@app.get("/admin/feedback", response_model=list[FeedbackItem], tags=["admin"])
async def list_feedback(
    _admin: Annotated[str, Depends(verify_admin_token)],
    rating: str | None = Query(default=None, pattern="^(up|down)$"),
    reviewed: int | None = Query(default=None, ge=0, le=2),
    org_id: str | None = Query(default=None),
    agent_id: str | None = Query(default=None),
    limit: int = Query(default=50, ge=1, le=500),
    offset: int = Query(default=0, ge=0),
) -> list[FeedbackItem]:
    rows = database.get_feedback(
        limit=limit,
        offset=offset,
        rating=rating,
        reviewed=reviewed,
        org_id=org_id,
        agent_id=agent_id,
    )
    return [FeedbackItem(**r) for r in rows]


@app.get("/admin/feedback/training", tags=["admin"])
async def export_training_feedback(
    _admin: Annotated[str, Depends(verify_admin_token)],
    org_id: str | None = Query(default=None),
    agent_id: str | None = Query(default=None),
    limit: int = Query(default=5000, ge=1, le=20000),
) -> dict:
    """Export rated Q/A pairs for preference tuning / fine-tuning."""
    rows = database.get_training_pairs(org_id=org_id, agent_id=agent_id, limit=limit)
    pairs = [
        {
            "id": r["id"],
            "question": r["question"],
            "answer": r["answer"],
            "rating": r["rating"],
            "correction": r.get("correction"),
            "mode": r.get("mode"),
            "model_name": r.get("model_name"),
            "org_id": r.get("org_id"),
            "agent_id": r.get("agent_id"),
            "created_at": r.get("created_at"),
        }
        for r in rows
    ]
    return {
        "count": len(pairs),
        "up": sum(1 for p in pairs if p["rating"] == "up"),
        "down": sum(1 for p in pairs if p["rating"] == "down"),
        "pairs": pairs,
    }

@app.patch("/admin/feedback/{feedback_id}", tags=["admin"])
async def review_feedback(
    feedback_id: int,
    body: ReviewRequest,
    _admin: Annotated[str, Depends(verify_admin_token)],
) -> dict:
    database.update_feedback_review(feedback_id, body.reviewed)
    labels = {0: "pending", 1: "approved", 2: "dismissed"}
    return {"feedback_id": feedback_id, "reviewed": labels[body.reviewed]}


# ── Admin — audit log ─────────────────────────────────────────────────────────

@app.get("/admin/audit", response_model=list[AuditItem], tags=["admin"])
async def audit_log(
    _admin: Annotated[str, Depends(verify_admin_token)],
    limit: int = Query(default=50, ge=1, le=500),
    offset: int = Query(default=0, ge=0),
) -> list[AuditItem]:
    rows = database.get_audit_log(limit=limit, offset=offset)
    return [AuditItem(**r) for r in rows]


# ── PDF calculation (extract → confirm → calculate) ───────────────────────────

def _job_response(job: dict, message: str = "") -> PdfJobResponse:
    return PdfJobResponse(
        job_id=job["job_id"],
        filename=job["filename"],
        source_type=job["source_type"],
        trusted=job["trusted"],
        needs_confirm=job["needs_confirm"],
        confirmed=job.get("confirmed", False),
        status=job["status"],
        rows=[PdfRow(**r) for r in job.get("rows") or []],
        text_preview=job.get("text_preview") or "",
        warnings=job.get("warnings") or [],
        validation=job.get("validation") or {},
        ocr_used=job.get("ocr_used", False),
        calculation=job.get("calculation"),
        message=message,
    )


@app.post("/pdf/extract", response_model=PdfJobResponse, tags=["pdf"])
async def pdf_extract(file: UploadFile = File(...)) -> PdfJobResponse:
    if not file.filename or not file.filename.lower().endswith(".pdf"):
        raise HTTPException(status_code=400, detail="Please upload a .pdf file")

    raw = await file.read()
    if not raw:
        raise HTTPException(status_code=400, detail="Empty file")
    if len(raw) > 20 * 1024 * 1024:
        raise HTTPException(status_code=400, detail="PDF too large (max 20 MB)")

    saved: Path | None = None
    try:
        if settings.persist_user_data:
            UPLOAD_DIR.mkdir(parents=True, exist_ok=True)
            saved = UPLOAD_DIR / f"{uuid.uuid4()}_{file.filename}"
            saved.write_bytes(raw)
        else:
            tmp = tempfile.NamedTemporaryFile(suffix=".pdf", delete=False)
            try:
                tmp.write(raw)
                tmp.flush()
                saved = Path(tmp.name)
            finally:
                tmp.close()

        job = await asyncio.to_thread(create_job_from_upload, saved, file.filename)
    except Exception as exc:
        raise HTTPException(status_code=500, detail=f"Extraction failed: {exc}") from exc
    finally:
        if saved is not None and not settings.persist_user_data:
            try:
                saved.unlink(missing_ok=True)
            except OSError:
                pass

    msg = (
        "Digital PDF extracted and validated — ready to calculate."
        if job["trusted"] and not job["needs_confirm"]
        else "Review and confirm amounts before calculating (required for OCR / unverified data)."
    )
    return _job_response(job, message=msg)


@app.get("/pdf/{job_id}", response_model=PdfJobResponse, tags=["pdf"])
async def pdf_get(job_id: str) -> PdfJobResponse:
    job = get_job(job_id)
    if not job:
        raise HTTPException(status_code=404, detail="Job not found")
    return _job_response(job)


@app.post("/pdf/{job_id}/confirm", response_model=PdfJobResponse, tags=["pdf"])
async def pdf_confirm(job_id: str, body: PdfConfirmRequest) -> PdfJobResponse:
    try:
        job = confirm_job(job_id, [r.model_dump() for r in body.rows])
    except KeyError:
        raise HTTPException(status_code=404, detail="Job not found") from None
    msg = (
        "Amounts confirmed. You can calculate now."
        if job.get("trusted")
        else "Confirmation saved but validation failed — fix highlighted issues."
    )
    return _job_response(job, message=msg)


@app.post("/pdf/{job_id}/calculate", response_model=PdfJobResponse, tags=["pdf"])
async def pdf_calculate(job_id: str) -> PdfJobResponse:
    try:
        job = await asyncio.to_thread(calculate_job, job_id)
    except KeyError:
        raise HTTPException(status_code=404, detail="Job not found") from None
    except PermissionError as exc:
        raise HTTPException(status_code=403, detail=str(exc)) from exc
    except ValueError as exc:
        raise HTTPException(status_code=400, detail=str(exc)) from exc
    return _job_response(job, message="Calculation complete (deterministic totals).")


# ── Static UI ─────────────────────────────────────────────────────────────────

STUDIO_DIR = STATIC_DIR / "studio"


def _studio_index() -> FileResponse:
    index = STUDIO_DIR / "index.html"
    if not index.exists():
        raise HTTPException(
            status_code=404,
            detail="Studio UI not built. Run: cd frontend && npm install && npm run build",
        )
    return FileResponse(index)


@app.get("/admin", tags=["ui"])
async def admin_ui():
    page = STATIC_DIR / "admin.html"
    if not page.exists():
        raise HTTPException(status_code=404, detail="Admin UI not found")
    return FileResponse(page)


@app.get("/pdf", tags=["ui"])
async def pdf_ui():
    page = STATIC_DIR / "pdf.html"
    if not page.exists():
        raise HTTPException(status_code=404, detail="PDF UI not found")
    return FileResponse(page)


@app.get("/", tags=["ui"])
async def chat_ui():
    return _studio_index()


@app.get("/app", tags=["ui"])
async def app_ui():
    return _studio_index()


@app.get("/chat-classic", tags=["ui"])
async def classic_chat_ui():
    index = STATIC_DIR / "chat.html"
    if not index.exists():
        raise HTTPException(status_code=404, detail="Chat UI not found")
    return FileResponse(index)


@app.get("/legacy-studio", tags=["ui"])
async def legacy_saas_ui():
    page = STATIC_DIR / "saas.html"
    if not page.exists():
        raise HTTPException(status_code=404, detail="Legacy SaaS UI not found")
    return FileResponse(page)


# Built Vite assets live under /studio/* (index references /studio/assets/...)
if STUDIO_DIR.exists():
    app.mount("/studio", StaticFiles(directory=str(STUDIO_DIR), html=True), name="studio")
