"""
Local Chroma knowledge-base service for the OpenAI Realtime voice lab.

Run (from this directory, inside .venv):
  uvicorn app:app --host 127.0.0.1 --port 8100
"""

from __future__ import annotations

import json
import os
from pathlib import Path
from typing import Any

from dotenv import load_dotenv
from fastapi import FastAPI, HTTPException
from pydantic import BaseModel, Field

load_dotenv()

ROOT = Path(__file__).resolve().parent
SEED_PATH = ROOT / "seed_docs.json"
CHROMA_PATH = os.getenv("CHROMA_PATH", str(ROOT / "chroma_data"))
COLLECTION_NAME = os.getenv("CHROMA_COLLECTION", "realtime_voice_kb")
EMBEDDING_MODEL = os.getenv("EMBEDDING_MODEL", "text-embedding-3-small")

app = FastAPI(title="Realtime Voice KB", version="0.1.0")

_collection = None


def _require_api_key() -> str:
    key = os.getenv("OPENAI_API_KEY", "").strip()
    if not key:
        raise RuntimeError(
            "OPENAI_API_KEY is not set. Copy .env.example to .env and set your key."
        )
    return key


def get_collection():
    """Lazy-init persistent Chroma collection with OpenAI embeddings."""
    global _collection
    if _collection is not None:
        return _collection

    import chromadb
    from chromadb.utils.embedding_functions import OpenAIEmbeddingFunction

    api_key = _require_api_key()
    Path(CHROMA_PATH).mkdir(parents=True, exist_ok=True)

    client = chromadb.PersistentClient(path=CHROMA_PATH)
    embedding_fn = OpenAIEmbeddingFunction(
        api_key=api_key,
        model_name=EMBEDDING_MODEL,
    )
    _collection = client.get_or_create_collection(
        name=COLLECTION_NAME,
        embedding_function=embedding_fn,
        metadata={"hnsw:space": "cosine"},
    )
    return _collection


def load_seed_docs() -> list[dict[str, Any]]:
    with SEED_PATH.open(encoding="utf-8") as f:
        docs = json.load(f)
    if not isinstance(docs, list) or not docs:
        raise RuntimeError(f"No seed docs found in {SEED_PATH}")
    return docs


def seed_collection(*, force: bool = False) -> dict[str, Any]:
    collection = get_collection()
    count = collection.count()
    if count > 0 and not force:
        return {
            "seeded": False,
            "reason": "collection_already_has_documents",
            "count": count,
        }

    docs = load_seed_docs()
    if force and count > 0:
        # Recreate cleanly by deleting then re-adding IDs.
        existing = collection.get()
        if existing and existing.get("ids"):
            collection.delete(ids=existing["ids"])

    ids = [str(d["id"]) for d in docs]
    documents = [str(d["text"]) for d in docs]
    metadatas = [
        {
            "source": str(d.get("source", d["id"])),
            "title": str(d.get("title", d["id"])),
        }
        for d in docs
    ]

    collection.upsert(ids=ids, documents=documents, metadatas=metadatas)
    return {"seeded": True, "count": collection.count(), "ids": ids}


class SearchRequest(BaseModel):
    query: str = Field(..., min_length=1)
    k: int = Field(default=3, ge=1, le=10)


class Snippet(BaseModel):
    text: str
    source: str
    title: str
    score: float | None = None


class SearchResponse(BaseModel):
    query: str
    found: bool
    snippets: list[Snippet]


class IngestDoc(BaseModel):
    id: str
    text: str
    source: str | None = None
    title: str | None = None


class IngestRequest(BaseModel):
    documents: list[IngestDoc]


@app.on_event("startup")
def on_startup() -> None:
    # Soft-seed so first boot is ready for the voice lab.
    try:
        result = seed_collection(force=False)
        print(f"[realtime-kb] startup seed: {result}")
    except Exception as exc:  # noqa: BLE001 — surface clearly in logs
        print(f"[realtime-kb] startup seed skipped: {exc}")


@app.get("/health")
def health() -> dict[str, Any]:
    try:
        collection = get_collection()
        return {
            "ok": True,
            "collection": COLLECTION_NAME,
            "count": collection.count(),
            "chroma_path": CHROMA_PATH,
            "embedding_model": EMBEDDING_MODEL,
        }
    except Exception as exc:  # noqa: BLE001
        return {"ok": False, "error": str(exc)}


@app.post("/seed")
def seed(force: bool = False) -> dict[str, Any]:
    try:
        return seed_collection(force=force)
    except Exception as exc:  # noqa: BLE001
        raise HTTPException(status_code=500, detail=str(exc)) from exc


@app.post("/search", response_model=SearchResponse)
def search(body: SearchRequest) -> SearchResponse:
    query = body.query.strip()
    if not query:
        raise HTTPException(status_code=400, detail="query is required")

    try:
        collection = get_collection()
        if collection.count() == 0:
            seed_collection(force=False)

        result = collection.query(
            query_texts=[query],
            n_results=min(body.k, max(collection.count(), 1)),
            include=["documents", "metadatas", "distances"],
        )
    except Exception as exc:  # noqa: BLE001
        raise HTTPException(status_code=500, detail=str(exc)) from exc

    documents = (result.get("documents") or [[]])[0]
    metadatas = (result.get("metadatas") or [[]])[0]
    distances = (result.get("distances") or [[]])[0]

    snippets: list[Snippet] = []
    for i, text in enumerate(documents):
        if not text:
            continue
        meta = metadatas[i] if i < len(metadatas) and metadatas[i] else {}
        distance = distances[i] if i < len(distances) else None
        # Cosine distance → rough similarity score for the voice model.
        score = None if distance is None else round(max(0.0, 1.0 - float(distance)), 4)
        snippets.append(
            Snippet(
                text=str(text),
                source=str(meta.get("source", "unknown")),
                title=str(meta.get("title", "untitled")),
                score=score,
            )
        )

    # Drop weak matches (cosine distance high → low score).
    snippets = [s for s in snippets if s.score is None or s.score >= 0.25]

    return SearchResponse(query=query, found=len(snippets) > 0, snippets=snippets)


@app.post("/ingest")
def ingest(body: IngestRequest) -> dict[str, Any]:
    if not body.documents:
        raise HTTPException(status_code=400, detail="documents is required")

    try:
        collection = get_collection()
        ids = [d.id for d in body.documents]
        documents = [d.text for d in body.documents]
        metadatas = [
            {
                "source": d.source or d.id,
                "title": d.title or d.id,
            }
            for d in body.documents
        ]
        collection.upsert(ids=ids, documents=documents, metadatas=metadatas)
        return {"ingested": len(ids), "count": collection.count(), "ids": ids}
    except Exception as exc:  # noqa: BLE001
        raise HTTPException(status_code=500, detail=str(exc)) from exc


if __name__ == "__main__":
    import uvicorn

    host = os.getenv("HOST", "127.0.0.1")
    port = int(os.getenv("PORT", "8100"))
    uvicorn.run("app:app", host=host, port=port, reload=True)
