"""RAG pipeline debug logging — enable with RAG_DEBUG=true."""

from __future__ import annotations

import logging
from typing import Any, Iterable

from langchain_core.documents import Document

from app.config import get_settings

logger = logging.getLogger("app.rag.kb")


def rag_debug_enabled() -> bool:
    return bool(get_settings().rag_debug)


def log_info(msg: str, *args: Any) -> None:
    logger.info(msg, *args)


def log_warn(msg: str, *args: Any) -> None:
    logger.warning(msg, *args)


def log_error(msg: str, *args: Any) -> None:
    logger.error(msg, *args)


def print_rag_debug(
    *,
    query: str,
    docs: list[Document],
    scores: list[float | None] | None = None,
    prompt: str | None = None,
    reason_if_empty: str | None = None,
) -> None:
    if not rag_debug_enabled():
        return

    scores = scores or []
    lines = [
        "",
        "======== RAG DEBUG ========",
        "",
        "User Query:",
        query,
        "",
        "Retrieved Documents:",
    ]
    if not docs:
        lines.append(f"(none) {reason_if_empty or ''}".rstrip())
    else:
        for i, doc in enumerate(docs, start=1):
            meta = doc.metadata or {}
            score = scores[i - 1] if i - 1 < len(scores) else None
            preview = (doc.page_content or "").replace("\n", " ")[:240]
            lines.append(
                f"[{i}] score={score} page={meta.get('page')} "
                f"file={meta.get('source_dataset')} id={meta.get('record_id')}"
            )
            lines.append(f"    {preview}")

    lines.extend(["", "Similarity Scores:"])
    if scores:
        lines.append(", ".join(str(s) for s in scores))
    else:
        lines.append("(not available)")

    lines.extend(["", "Chunk IDs:"])
    lines.append(
        ", ".join(str((d.metadata or {}).get("record_id") or "?") for d in docs) or "(none)"
    )

    lines.extend(["", "Source Files:"])
    sources = sorted(
        {
            str((d.metadata or {}).get("source_dataset") or "?")
            for d in docs
        }
    )
    lines.append(", ".join(sources) or "(none)")

    if prompt is not None:
        lines.extend(["", "Prompt Sent to LLM:", prompt[:4000], ("...(truncated)" if len(prompt) > 4000 else "")])

    if reason_if_empty and not docs:
        lines.extend(["", "Empty Retrieval Reason:", reason_if_empty])

    lines.extend(["", "===========================", ""])
    block = "\n".join(lines)
    print(block, flush=True)
    logger.info(block)


def summarize_chunks(chunks: Iterable[str]) -> dict[str, Any]:
    sizes = [len(c) for c in chunks]
    if not sizes:
        return {"count": 0, "avg": 0, "min": 0, "max": 0}
    return {
        "count": len(sizes),
        "avg": round(sum(sizes) / len(sizes), 1),
        "min": min(sizes),
        "max": max(sizes),
    }
