Cognitive-rag / backend / routes_chat.py
routes_chat.py
Raw
from __future__ import annotations

import threading
from typing import Any
import importlib
import re
from difflib import SequenceMatcher

from fastapi import APIRouter, HTTPException, Request
from pydantic import BaseModel

from backend import config
from backend.logging_audit import audit_event, resolve_user, summarize_text
from backend.run_report import write_prompt_reports


router = APIRouter(tags=["chat"])

_TURN_LOCK = threading.Lock()
_TURN_BY_USER: dict[str, int] = {}
_SCORE_FIELDS = (
    "score",
    "dense_score",
    "sparse_score",
    "rrf_score",
    "cross_encoder_score",
    "original_score",
)
_GLOBAL_MEMORY_SESSION_ID = "shared_memory"
_DEFAULT_EMBEDDING_MODEL = "nvidia/llama-3.2-nv-embedqa-1b-v2"
_FALLBACK_ANSWER_PHRASES = (
    "provided documents do not contain this information",
    "provided context does not contain this information",
    "provided evidence is insufficient",
    "do not contain this information",
    "i could not find",
    "there is no information about",
    "not enough information",
)


class ChatRequest(BaseModel):
    query: str
    debug: bool = False
    use_multihop: bool = False


def _best_doc_id(chunk: dict[str, Any]) -> str | None:
    value = chunk.get("doc_id")
    if value:
        return value
    value = chunk.get("document_name")
    if value:
        return value
    value = chunk.get("source_path")
    if value:
        return value
    return None


def _chunk_to_source(chunk: dict[str, Any]) -> dict[str, Any]:
    source: dict[str, Any] = {}
    doc_id = _best_doc_id(chunk)
    if doc_id:
        source["doc_id"] = doc_id
    if chunk.get("chunk_id"):
        source["chunk_id"] = chunk["chunk_id"]
    if chunk.get("text"):
        source["text"] = chunk["text"]

    for field in ("modality", "page", "section", "timestamp", "source_path"):
        value = chunk.get(field)
        if value is not None and value != "":
            source[field] = value
    for field in _SCORE_FIELDS:
        value = chunk.get(field)
        if value is not None and value != "":
            source[field] = value
    return source


def _is_memory_origin_chunk(chunk: dict[str, Any]) -> bool:
    doc_id = str(chunk.get("doc_id") or "").strip().lower()
    section = str(chunk.get("section") or "").strip().lower()
    chunk_id = str(chunk.get("chunk_id") or "").strip().lower()
    source_path = str(chunk.get("source_path") or "").strip().lower()
    text_value = str(chunk.get("text") or "")

    if doc_id in {"shared_memory", "conversation_memory"}:
        return True
    if section == "conversation_memory":
        return True
    if chunk_id.startswith("mem_"):
        return True
    if "conversation_memory" in source_path:
        return True
    if "query:" in text_value.lower() and "answer:" in text_value.lower() and (
        doc_id == "shared_memory" or section == "conversation_memory"
    ):
        return True
    return False


def _filter_user_facing_evidence_chunks(chunks: list[dict[str, Any]]) -> list[dict[str, Any]]:
    filtered: list[dict[str, Any]] = []
    seen_ids: set[str] = set()
    seen_text: set[str] = set()

    for chunk in chunks:
        if _is_memory_origin_chunk(chunk):
            continue
        text_value = str(chunk.get("text") or "").strip()
        chunk_id = str(chunk.get("chunk_id") or "").strip()
        if not text_value or not chunk_id:
            continue
        norm_text = _normalize_text(text_value)
        if chunk_id in seen_ids or (norm_text and norm_text in seen_text):
            continue
        seen_ids.add(chunk_id)
        if norm_text:
            seen_text.add(norm_text)
        filtered.append(chunk)
    return filtered


def _best_source_identifier_for_evidence(chunk: dict[str, Any]) -> str:
    source_path = str(chunk.get("source_path") or "").strip()
    if source_path:
        normalized = source_path.replace("\\", "/").strip("/")
        if normalized:
            return normalized.split("/")[-1]
    document_name = str(chunk.get("document_name") or "").strip()
    if document_name:
        return document_name
    doc_id = str(chunk.get("doc_id") or "").strip()
    if doc_id:
        return doc_id
    chunk_id = str(chunk.get("chunk_id") or "").strip()
    return chunk_id or "unknown_source"


def _format_evidence_label(chunk: dict[str, Any]) -> str:
    label_parts = [_best_source_identifier_for_evidence(chunk)]
    page = chunk.get("page")
    section = chunk.get("section")
    timestamp = chunk.get("timestamp")
    if page not in (None, ""):
        label_parts.append(f"page {page}")
    if section not in (None, ""):
        label_parts.append(f"section {section}")
    if timestamp not in (None, ""):
        label_parts.append(f"timestamp {timestamp}")
    return ", ".join(str(part) for part in label_parts if str(part).strip())


def _append_evidence_section(answer: str, evidence_chunks: list[dict[str, Any]]) -> str:
    base = (answer or "").strip()
    if not evidence_chunks:
        return base
    lines = [base, "", "Evidence used:"]
    for idx, chunk in enumerate(evidence_chunks, 1):
        text_value = str(chunk.get("text") or "").strip()
        if not text_value:
            continue
        label = _format_evidence_label(chunk)
        lines.append(f"{idx}. {text_value} [{label}]")
    return "\n".join(lines).strip()


def _select_context_chunks_by_evidence(
    context_chunks: list[dict[str, Any]],
    evidence_texts: list[str],
) -> list[dict[str, Any]]:
    if not evidence_texts:
        return context_chunks

    by_text: dict[str, list[int]] = {}
    for idx, row in enumerate(context_chunks):
        text = (row.get("text") or "").strip()
        by_text.setdefault(text, []).append(idx)

    selected_indices: list[int] = []
    used_indices: set[int] = set()

    for evidence in evidence_texts:
        key = (evidence or "").strip()
        candidates = by_text.get(key, [])
        for idx in candidates:
            if idx not in used_indices:
                selected_indices.append(idx)
                used_indices.add(idx)
                break

    if not selected_indices:
        return context_chunks
    return [context_chunks[i] for i in selected_indices]


def _next_turn(user_n: str) -> int:
    with _TURN_LOCK:
        current = _TURN_BY_USER.get(user_n, 0) + 1
        _TURN_BY_USER[user_n] = current
        return current


def _build_memory_payload(query: str, answer: str) -> str:
    return f'Query:\n    "{query}"\n\nAnswer:\n    {answer}\n'


def _resolve_embedding_model(vectordb_config: Any) -> str:
    model_from_config = getattr(vectordb_config, "EMBEDDING_MODEL", None)
    if isinstance(model_from_config, str) and model_from_config.strip():
        return model_from_config.strip()
    model_from_env = config.env_value("EMBEDDING_MODEL", "").strip()
    if model_from_env:
        return model_from_env
    return _DEFAULT_EMBEDDING_MODEL


def _extract_answer_section(raw_output: str) -> str:
    text = (raw_output or "").replace("\r\n", "\n").replace("\r", "\n").strip()
    if not text:
        return ""
    match = re.search(r"(?is)\bAnswer\s*:\s*", text)
    if not match:
        return ""
    tail = text[match.end():]
    boundary = re.search(r"(?im)^\s*(?:Evidence|Query)\s*:\s*", tail)
    if boundary:
        tail = tail[: boundary.start()]
    return tail.strip()


def _cleanup_final_answer(answer: str) -> str:
    text = (answer or "").replace("\r\n", "\n").replace("\r", "\n").strip()
    if not text:
        return ""

    text = re.sub(r"(?im)^\s*Answer\s*:\s*", "", text)
    text = re.sub(r"(?im)^\s*(?:Query|Evidence)\s*:\s*.*$", "", text)
    text = re.sub(r"\n{3,}", "\n\n", text)

    cleaned_lines: list[str] = []
    previous_norm = ""
    seen_bullets: set[str] = set()

    for raw_line in text.split("\n"):
        if re.search(r"(?i)\[(?:shared_memory|conversation_memory)\]", raw_line):
            continue
        if re.search(r"(?i)^\s*(?:shared_memory|conversation_memory)\s*:", raw_line):
            continue
        line = raw_line.strip()
        if not line:
            if cleaned_lines and cleaned_lines[-1]:
                cleaned_lines.append("")
            previous_norm = ""
            continue
        if re.fullmatch(r"(?i)(?:query|answer|evidence)\s*:?", line):
            continue
        line = re.sub(r"(?i)^\s*(?:shared_memory|conversation_memory)\s*:?\s*", "", line).strip()
        if not line:
            continue

        bullet = re.match(r"^\s*(?:[-*•]|\d+[.)])\s+(.+)$", line)
        if bullet:
            bullet_norm = _normalize_text(bullet.group(1))
            if bullet_norm and bullet_norm in seen_bullets:
                continue
            if bullet_norm:
                seen_bullets.add(bullet_norm)

        norm = _normalize_text(line)
        if previous_norm and norm == previous_norm:
            continue
        cleaned_lines.append(line)
        previous_norm = norm

    while cleaned_lines and not cleaned_lines[-1]:
        cleaned_lines.pop()
    result = "\n".join(cleaned_lines).strip()

    final_lines = [ln.strip() for ln in result.split("\n") if ln.strip()]
    if len(final_lines) >= 2:
        last_line = final_lines[-1]
        prev_line = final_lines[-2]
        if len(last_line) >= 12 and len(last_line) < len(prev_line) and prev_line.lower().startswith(last_line.lower()):
            final_lines = final_lines[:-1]
            result = "\n".join(final_lines).strip()

    return result


def _derive_canonical_answer(raw_output: str, memory_module: Any) -> str:
    answer = _extract_answer_section(raw_output)
    if not answer and memory_module is not None:
        try:
            parsed = memory_module.parse_LLM_output(raw_output)
        except Exception:
            parsed = None
        if isinstance(parsed, dict):
            parsed_answer = (parsed.get("answer") or "").strip()
            if parsed_answer:
                answer = parsed_answer
    if not answer:
        answer = (raw_output or "").strip()
    return _cleanup_final_answer(answer)


def _source_path(chunk: dict[str, Any]) -> str | None:
    value = chunk.get("source_path")
    if value:
        return value
    source = chunk.get("source")
    if isinstance(source, dict):
        nested = source.get("path") or source.get("source_path")
        if nested:
            return nested
    return None


def _normalize_chunk(chunk: dict[str, Any], *, include_content: bool) -> dict[str, Any]:
    text_value = chunk.get("text") or chunk.get("content") or ""
    row: dict[str, Any] = {
        "chunk_id": chunk.get("chunk_id"),
        "doc_id": chunk.get("doc_id"),
        "document_name": chunk.get("document_name"),
        "text": text_value,
        "modality": chunk.get("modality"),
        "page": chunk.get("page"),
        "section": chunk.get("section"),
        "timestamp": chunk.get("timestamp"),
        "source_path": _source_path(chunk),
    }
    for score_key in _SCORE_FIELDS:
        value = chunk.get(score_key)
        if value is not None and value != "":
            row[score_key] = value
    if "score" not in row and row.get("cross_encoder_score") is not None:
        row["score"] = row["cross_encoder_score"]
    if include_content:
        row["content"] = text_value
    return row


def _annotate_retrieval_channels(
    chunks: list[dict[str, Any]],
    dense_ids: set[str],
    sparse_ids: set[str],
) -> list[dict[str, Any]]:
    annotated: list[dict[str, Any]] = []
    for chunk in chunks:
        chunk_id = chunk.get("chunk_id")
        channels: list[str] = []
        if chunk_id in dense_ids:
            channels.append("dense")
        if chunk_id in sparse_ids:
            channels.append("sparse")
        out = dict(chunk)
        if channels:
            out["retrieval_channels"] = channels
        annotated.append(out)
    return annotated


def _merge_unique_chunks(
    first: list[dict[str, Any]],
    second: list[dict[str, Any]],
) -> list[dict[str, Any]]:
    merged: list[dict[str, Any]] = []
    seen: set[str] = set()

    for chunk in first + second:
        chunk_id = str(chunk.get("chunk_id") or "")
        if not chunk_id:
            continue
        if chunk_id in seen:
            continue
        seen.add(chunk_id)
        merged.append(chunk)
    return merged


def _normalize_text(value: str | None) -> str:
    return " ".join((value or "").strip().lower().split())


def _dedupe_memory_chunks(chunks: list[dict[str, Any]]) -> tuple[list[dict[str, Any]], int]:
    deduped: list[dict[str, Any]] = []
    seen: set[str] = set()
    dropped = 0
    for chunk in chunks:
        key = _normalize_text(chunk.get("text"))
        if not key:
            key = str(chunk.get("chunk_id") or "")
        if not key:
            continue
        if key in seen:
            dropped += 1
            continue
        seen.add(key)
        deduped.append(chunk)
    return deduped, dropped


def _combine_and_rerank_for_generation(
    *,
    query: str,
    retriever_module: Any,
    vectordb_config: Any,
    document_chunks: list[dict[str, Any]],
    memory_chunks: list[dict[str, Any]],
) -> list[dict[str, Any]]:
    candidates: list[dict[str, Any]] = []
    seen: set[str] = set()

    for chunk in document_chunks + memory_chunks:
        chunk_id = str(chunk.get("chunk_id") or "")
        if not chunk_id or chunk_id in seen:
            continue
        seen.add(chunk_id)
        candidates.append(chunk)

    if not candidates:
        return []

    reranker_enabled = bool(getattr(vectordb_config, "RERANKER_ENABLED", False))
    reranker_top_n = config.env_int(
        "RERANKER_TOP_N",
        int(getattr(vectordb_config, "RERANKER_TOP_N", len(candidates))),
    )
    reranker_top_n = max(1, min(reranker_top_n, len(candidates)))

    if reranker_enabled and query:
        try:
            return retriever_module.rerank(query, candidates, top_n=reranker_top_n)
        except Exception:
            pass

    return sorted(
        candidates,
        key=lambda x: x.get("cross_encoder_score", x.get("score", 0)),
        reverse=True,
    )[:reranker_top_n]


def _build_retrieval_section(
    *,
    query: str,
    strategy: str,
    context_chunks: list[dict[str, Any]],
    retrieved_topk_chunks: list[dict[str, Any]],
    dense_chunks: list[dict[str, Any]],
    sparse_chunks: list[dict[str, Any]],
    related_chunks: list[dict[str, Any]],
    memory_chunks: list[dict[str, Any]],
    memory_raw_count: int,
    memory_dropped_count: int,
) -> dict[str, Any]:
    dense_ids = {
        str(item.get("chunk_id"))
        for item in dense_chunks
        if item.get("chunk_id") is not None
    }
    sparse_ids = {
        str(item.get("chunk_id"))
        for item in sparse_chunks
        if item.get("chunk_id") is not None
    }
    return {
        "query": query,
        "strategy": strategy,
        "context_sent_to_generator": _annotate_retrieval_channels(
            [_normalize_chunk(chunk, include_content=False) for chunk in context_chunks],
            dense_ids,
            sparse_ids,
        ),
        "retrieved_topk": _annotate_retrieval_channels(
            [_normalize_chunk(chunk, include_content=False) for chunk in retrieved_topk_chunks],
            dense_ids,
            sparse_ids,
        ),
        "dense_topk": _annotate_retrieval_channels(
            [_normalize_chunk(chunk, include_content=False) for chunk in dense_chunks],
            dense_ids,
            sparse_ids,
        ),
        "sparse_topk": _annotate_retrieval_channels(
            [_normalize_chunk(chunk, include_content=False) for chunk in sparse_chunks],
            dense_ids,
            sparse_ids,
        ),
        "related_chunks": [
            _normalize_chunk(chunk, include_content=False) for chunk in related_chunks
        ],
        "memory_chunks": [
            _normalize_chunk(chunk, include_content=False) for chunk in memory_chunks
        ],
        "memory_raw_count": int(memory_raw_count),
        "memory_dropped_count": int(memory_dropped_count),
    }


def _as_float(value: str, default: float) -> float:
    try:
        return float(value)
    except (TypeError, ValueError):
        return default


def _extract_query_matches(response: Any) -> list[dict[str, Any]]:
    if response is None:
        return []
    if isinstance(response, dict):
        matches = response.get("matches", [])
        return matches if isinstance(matches, list) else []
    matches_attr = getattr(response, "matches", None)
    if isinstance(matches_attr, list):
        out: list[dict[str, Any]] = []
        for item in matches_attr:
            if isinstance(item, dict):
                out.append(item)
                continue
            score = getattr(item, "score", None)
            item_id = getattr(item, "id", None)
            out.append({"id": item_id, "score": score})
        return out
    return []


def _as_int(value: str, default: int) -> int:
    try:
        return int(value)
    except (TypeError, ValueError):
        return default


def _best_retrieval_score(chunks: list[dict[str, Any]]) -> float:
    if not chunks:
        return 0.0
    best = 0.0
    for chunk in chunks:
        score = chunk.get("cross_encoder_score")
        if score is None:
            score = chunk.get("score")
        if score is None:
            score = chunk.get("rrf_score")
        try:
            value = float(score)
        except (TypeError, ValueError):
            continue
        if value > best:
            best = value
    return best


def _word_count(text: str) -> int:
    return len(re.findall(r"[A-Za-z0-9]+", text or ""))


def _contains_fallback_language(text: str) -> bool:
    lowered = (text or "").strip().lower()
    if not lowered:
        return True
    return any(phrase in lowered for phrase in _FALLBACK_ANSWER_PHRASES)


def _sentence_count(text: str) -> int:
    parts = [p for p in re.split(r"[.!?]+", text or "") if p.strip()]
    return len(parts)


def _is_definition_like(query: str, answer: str) -> bool:
    q = (query or "").strip().lower()
    starts = ("what is ", "who is ", "define ", "what are ", "who are ")
    if not any(q.startswith(s) for s in starts):
        return False
    return _word_count(answer) <= 70 and _sentence_count(answer) <= 2


def _lexical_similarity(a: str, b: str) -> float:
    left = _normalize_text(a)
    right = _normalize_text(b)
    if not left or not right:
        return 0.0
    return float(SequenceMatcher(None, left, right).ratio())


def _best_lexical_similarity(text: str, candidates: list[str]) -> tuple[float, str]:
    best_score = 0.0
    best_text = ""
    for candidate in candidates:
        score = _lexical_similarity(text, candidate)
        if score > best_score:
            best_score = score
            best_text = candidate
    return best_score, best_text


def _filter_memory_chunks_for_quality(
    memory_chunks: list[dict[str, Any]],
) -> tuple[list[dict[str, Any]], list[dict[str, Any]]]:
    kept: list[dict[str, Any]] = []
    filtered: list[dict[str, Any]] = []
    for chunk in memory_chunks:
        text_value = (chunk.get("text") or "").strip()
        reasons: list[str] = []
        if _contains_fallback_language(text_value):
            reasons.append("fallback_or_negative_memory")
        if _word_count(text_value) < 12:
            reasons.append("low_information_memory")
        if reasons:
            filtered.append(
                {
                    "chunk_id": chunk.get("chunk_id"),
                    "reasons": reasons,
                    "score": chunk.get("score"),
                }
            )
            continue
        kept.append(chunk)
    return kept, filtered


@router.post("/chat")
async def chat(request: Request, payload: ChatRequest) -> dict[str, Any]:
    query = (payload.query or "").strip()
    if not query:
        raise HTTPException(status_code=400, detail="query is required")

    ip = request.client.host if request.client else "unknown"
    user_n = resolve_user(ip)

    modules = config.get_stage_modules()
    embed_query_module = modules["embedding_query"]
    retriever = modules["vectordb_retriever"]
    generation = modules["generation"]
    memory = modules["memory"]
    pinecone = modules["vectordb_pinecone"]
    vectordb_config = modules["vectordb_config"]
    embedding_model = _resolve_embedding_model(vectordb_config)

    try:
        embed_client = embed_query_module.load_client(
            embedding_model,
            config.env_value("NVIDIA_API_KEY", ""),
            "NONE",
        )
        query_embedding = embed_client.embed_query(query)
    except Exception as exc:
        raise HTTPException(status_code=500, detail=f"Embedding failed: {exc}") from exc

    try:
        memory_namespace = config.env_value(
            "CONVERSATION_MEMORY_NAMESPACE",
            getattr(
                getattr(memory, "memory_config", object()),
                "CONVERSATION_MEMORY_NAMESPACE",
                vectordb_config.CONVERSATION_MEMORY_NAMESPACE,
            ),
        )
        memory_top_k = _as_int(config.env_value("MEMORY_TOP_K", "8"), 8)
        memory_min_score = _as_float(config.env_value("MEMORY_MIN_SCORE", "0.0"), 0.0)
        memory_boost = _as_float(config.env_value("MEMORY_BOOST", "0.5"), 0.5)
        context = retriever.get_unified_context_for_llm(
            query_embedding=query_embedding,
            query_text=query,
            strategy="hybrid",
            doc_top_k=config.env_int("TOP_K", 30),
            memory_top_k=memory_top_k,
            memory_min_score=memory_min_score,
            memory_boost=memory_boost,
            doc_namespace=config.env_value("DOC_NAMESPACE", ""),
            memory_namespace=memory_namespace,
            include_related=True,
            export_json=False,
        )
    except Exception as exc:
        raise HTTPException(status_code=500, detail=f"Hybrid retrieval failed: {exc}") from exc

    doc_namespace = config.env_value("DOC_NAMESPACE", "")
    dense_raw_chunks: list[dict[str, Any]] = []
    sparse_raw_chunks: list[dict[str, Any]] = []
    try:
        dense_result = retriever.retrieve_dense(
            query_embedding=query_embedding,
            top_k=config.env_int("TOP_K", 30),
            namespace=doc_namespace,
        )
        dense_raw_chunks = dense_result.get("retrieved_chunks", [])
    except Exception:
        dense_raw_chunks = []
    try:
        sparse_result = retriever.retrieve_sparse(
            query_text=query,
            top_k=config.env_int("TOP_K", 30),
            namespace=doc_namespace,
        )
        sparse_raw_chunks = sparse_result.get("retrieved_chunks", [])
    except Exception:
        sparse_raw_chunks = []

    document_chunks = context.get("document_chunks", [])
    related_chunks = context.get("related_chunks", [])
    memory_chunks = context.get("memory_chunks", [])
    memory_chunks, memory_duplicates_removed = _dedupe_memory_chunks(memory_chunks)
    memory_chunks, memory_quality_filtered = _filter_memory_chunks_for_quality(memory_chunks)
    memory_quality_filtered_count = len(memory_quality_filtered)
    memory_raw_count = int(context.get("memory_raw_count", len(memory_chunks)))
    memory_dropped_count = int(context.get("memory_dropped_count", 0))
    generation_candidates = _combine_and_rerank_for_generation(
        query=query,
        retriever_module=retriever,
        vectordb_config=modules["vectordb_config"],
        document_chunks=document_chunks,
        memory_chunks=memory_chunks,
    )
    original_document_chunks = list(document_chunks)
    original_related_chunks = list(related_chunks)
    original_memory_chunks = list(memory_chunks)
    original_generation_candidates = list(generation_candidates)
    original_memory_raw_count = memory_raw_count
    original_memory_dropped_count = memory_dropped_count
    first_dense_ids = {
        str(item.get("chunk_id"))
        for item in dense_raw_chunks
        if item.get("chunk_id") is not None
    }
    first_sparse_ids = {
        str(item.get("chunk_id"))
        for item in sparse_raw_chunks
        if item.get("chunk_id") is not None
    }
    first_pass_context_chunks = _annotate_retrieval_channels([
        _normalize_chunk(chunk, include_content=True) for chunk in original_generation_candidates
    ], first_dense_ids, first_sparse_ids)
    first_pass_generation_output = ""
    first_pass_parsed: dict[str, Any] | None = None
    first_pass_answer_text = ""
    first_pass_evidence_texts: list[str] = []
    try:
        first_pass_generation_output = generation.generate_answer(
            {"query": query, "retrieved_results": first_pass_context_chunks}
        )
        first_pass_parsed = (
            memory.parse_LLM_output(first_pass_generation_output)
            if first_pass_generation_output
            else None
        )
        first_pass_answer_text = _derive_canonical_answer(first_pass_generation_output, memory)
        first_pass_evidence_texts = (
            first_pass_parsed.get("evidence", [])
            if isinstance(first_pass_parsed, dict) and isinstance(first_pass_parsed.get("evidence"), list)
            else []
        )
    except Exception as exc:
        raise HTTPException(status_code=500, detail=f"Generation failed: {exc}") from exc

    multi_hop_retrieval_sections: list[dict[str, Any]] = []
    multihop_used = False
    multihop_sub_query: str | None = None
    multi_hop_reason_list: list[str] = []
    multi_hop_reason = ""
    multihop_error: str | None = None
    multihop_validation: dict[str, Any] | None = None

    if payload.use_multihop:
        try:
            multihop_module = importlib.import_module("multihop_agent")
            MultiHopAgent = getattr(multihop_module, "MultiHopAgent")
            MultiHopConfig = getattr(multihop_module, "MultiHopConfig")

            sufficiency_threshold = _as_float(
                config.env_value("MULTIHOP_SUFFICIENCY_THRESHOLD", "0.75"),
                0.75,
            )
            min_context_chunks = config.env_int("MULTIHOP_MIN_CONTEXT_CHUNKS", 3)

            mh_config = MultiHopConfig(
                max_hops=2,
                per_hop_top_k=config.env_int("TOP_K", 30),
                total_doc_budget=config.env_int("TOP_K", 30) * 2,
                sufficiency_threshold=sufficiency_threshold,
                time_limit_sec=60,
                strategy="hybrid",
                include_related=True,
                memory_top_k=config.env_int("MEMORY_TOP_K", 3),
            )
            mh_agent = MultiHopAgent(
                embed_fn=lambda text: embed_client.embed_query(text),
                config=mh_config,
                bedrock_region=config.env_value("AWS_DEFAULT_REGION", "ca-central-1"),
                output_dir=str(config.RUN_REPORT_PROMPT_DIR),
            )

            evidence_items = [
                mh_agent._build_evidence_item(chunk)  # noqa: SLF001
                for chunk in generation_candidates
                if chunk.get("text")
            ]
            multihop_validation = mh_agent.validate(query, evidence_items)
            if not first_pass_evidence_texts:
                multi_hop_reason_list.append("first pass has missing evidence")
            answer_lower = first_pass_answer_text.lower()
            if any(phrase in answer_lower for phrase in _FALLBACK_ANSWER_PHRASES):
                multi_hop_reason_list.append("first pass answer indicates insufficiency")
            unique_docs = len({
                chunk.get("doc_id")
                for chunk in original_generation_candidates
                if chunk.get("doc_id")
            })
            if len(original_generation_candidates) < min_context_chunks:
                multi_hop_reason_list.append("retrieved chunks had weak coverage")
            if unique_docs < 2:
                multi_hop_reason_list.append("retrieved chunks lacked document diversity")
            min_top_score = _as_float(config.env_value("MULTIHOP_MIN_TOP_SCORE", "1.0"), 1.0)
            if _best_retrieval_score(original_generation_candidates) < min_top_score:
                multi_hop_reason_list.append("retrieval top score was weak")
            if float(multihop_validation.get("combined_score", 0.0)) < sufficiency_threshold:
                multi_hop_reason_list.append("answer confidence insufficient")

            if multi_hop_reason_list:
                multi_hop_reason = "; ".join(multi_hop_reason_list)
                evidence_summary = mh_agent._summarize_evidence(evidence_items)  # noqa: SLF001
                decomposition = mh_agent.decompose(query, evidence_summary)
                proposed = (decomposition.get("sub_question") or "").strip()
                if decomposition.get("action") == "ask" and proposed:
                    multihop_sub_query = proposed
                else:
                    multihop_sub_query = (
                        f"{query}. Focus on missing details and supporting evidence."
                    )

                second_query_embedding = embed_client.embed_query(multihop_sub_query)
                second_context = retriever.get_unified_context_for_llm(
                    query_embedding=second_query_embedding,
                    query_text=multihop_sub_query,
                    strategy="hybrid",
                    doc_top_k=config.env_int("TOP_K", 30),
                    memory_top_k=memory_top_k,
                    memory_min_score=memory_min_score,
                    memory_boost=memory_boost,
                    doc_namespace=doc_namespace,
                    memory_namespace=memory_namespace,
                    include_related=True,
                    export_json=False,
                )

                second_documents = second_context.get("document_chunks", [])
                second_related = second_context.get("related_chunks", [])
                second_memory = second_context.get("memory_chunks", [])
                second_memory, _ = _dedupe_memory_chunks(second_memory)
                second_memory, second_memory_quality_filtered = _filter_memory_chunks_for_quality(second_memory)
                memory_quality_filtered_count += len(second_memory_quality_filtered)
                second_generation_candidates = _combine_and_rerank_for_generation(
                    query=multihop_sub_query,
                    retriever_module=retriever,
                    vectordb_config=modules["vectordb_config"],
                    document_chunks=second_documents,
                    memory_chunks=second_memory,
                )
                second_dense_chunks: list[dict[str, Any]] = []
                second_sparse_chunks: list[dict[str, Any]] = []
                try:
                    second_dense_result = retriever.retrieve_dense(
                        query_embedding=second_query_embedding,
                        top_k=config.env_int("TOP_K", 30),
                        namespace=doc_namespace,
                    )
                    second_dense_chunks = second_dense_result.get("retrieved_chunks", [])
                except Exception:
                    second_dense_chunks = []
                try:
                    second_sparse_result = retriever.retrieve_sparse(
                        query_text=multihop_sub_query,
                        top_k=config.env_int("TOP_K", 30),
                        namespace=doc_namespace,
                    )
                    second_sparse_chunks = second_sparse_result.get("retrieved_chunks", [])
                except Exception:
                    second_sparse_chunks = []
                multi_hop_retrieval_sections.append(
                    _build_retrieval_section(
                        query=multihop_sub_query,
                        strategy="hybrid",
                        context_chunks=second_generation_candidates,
                        retrieved_topk_chunks=second_documents,
                        dense_chunks=second_dense_chunks,
                        sparse_chunks=second_sparse_chunks,
                        related_chunks=second_related,
                        memory_chunks=second_memory,
                        memory_raw_count=int(second_context.get("memory_raw_count", len(second_memory))),
                        memory_dropped_count=int(second_context.get("memory_dropped_count", 0)),
                    )
                )

                document_chunks = _merge_unique_chunks(document_chunks, second_documents)
                related_chunks = _merge_unique_chunks(related_chunks, second_related)
                memory_chunks = _merge_unique_chunks(memory_chunks, second_memory)
                memory_chunks, memory_duplicates_removed_second = _dedupe_memory_chunks(memory_chunks)
                memory_duplicates_removed += memory_duplicates_removed_second

                generation_candidates = _combine_and_rerank_for_generation(
                    query=query,
                    retriever_module=retriever,
                    vectordb_config=modules["vectordb_config"],
                    document_chunks=document_chunks,
                    memory_chunks=memory_chunks,
                )
                multihop_used = True
        except Exception as exc:
            multihop_error = str(exc)

    memory_chunk_ids = {
        str(item.get("chunk_id"))
        for item in memory_chunks
        if item.get("chunk_id") is not None
    }
    dense_ids = {
        str(item.get("chunk_id"))
        for item in dense_raw_chunks
        if item.get("chunk_id") is not None
    }
    sparse_ids = {
        str(item.get("chunk_id"))
        for item in sparse_raw_chunks
        if item.get("chunk_id") is not None
    }

    llm_context_chunks = _annotate_retrieval_channels([
        _normalize_chunk(chunk, include_content=True) for chunk in generation_candidates
    ], dense_ids, sparse_ids)
    memory_in_context_count = sum(
        1
        for chunk in llm_context_chunks
        if str(chunk.get("chunk_id")) in memory_chunk_ids
    )
    memory_raw_count = int(context.get("memory_raw_count", len(memory_chunks)))
    memory_dropped_count = int(context.get("memory_dropped_count", 0))
    retrieved_topk = _annotate_retrieval_channels([
        _normalize_chunk(chunk, include_content=False) for chunk in document_chunks
    ], dense_ids, sparse_ids)
    dense_topk = _annotate_retrieval_channels([
        _normalize_chunk(chunk, include_content=False) for chunk in dense_raw_chunks
    ], dense_ids, sparse_ids)
    sparse_topk = _annotate_retrieval_channels([
        _normalize_chunk(chunk, include_content=False) for chunk in sparse_raw_chunks
    ], dense_ids, sparse_ids)
    related_chunks_report = [
        _normalize_chunk(chunk, include_content=False) for chunk in related_chunks
    ]
    memory_chunks_report = [
        _normalize_chunk(chunk, include_content=False) for chunk in memory_chunks
    ]

    if multihop_used:
        try:
            generation_input = {"query": query, "retrieved_results": llm_context_chunks}
            generation_output = generation.generate_answer(generation_input)
        except Exception as exc:
            raise HTTPException(status_code=500, detail=f"Generation failed: {exc}") from exc

        parsed = memory.parse_LLM_output(generation_output) if generation_output else None
        answer_text = _derive_canonical_answer(generation_output, memory)
        evidence_texts = (
            parsed.get("evidence", [])
            if isinstance(parsed, dict) and isinstance(parsed.get("evidence"), list)
            else []
        )
    else:
        generation_output = first_pass_generation_output
        parsed = first_pass_parsed
        answer_text = first_pass_answer_text
        evidence_texts = list(first_pass_evidence_texts)

    selected_chunks = _select_context_chunks_by_evidence(llm_context_chunks, evidence_texts)
    used_chunks = selected_chunks or llm_context_chunks
    user_evidence_chunks = _filter_user_facing_evidence_chunks(used_chunks)
    if not user_evidence_chunks:
        user_evidence_chunks = _filter_user_facing_evidence_chunks(llm_context_chunks)

    sources = []
    for chunk in user_evidence_chunks:
        source_item = _chunk_to_source(chunk)
        if source_item.get("chunk_id") and source_item.get("text"):
            sources.append(source_item)

    if not sources:
        for chunk in _filter_user_facing_evidence_chunks(llm_context_chunks):
            source_item = _chunk_to_source(chunk)
            if source_item.get("chunk_id") and source_item.get("text"):
                sources.append(source_item)
    answer_for_user = _append_evidence_section(answer_text, user_evidence_chunks)

    memory_result: dict[str, Any] | None = None
    memory_error: str | None = None
    memory_duplicate_skipped = False
    memory_duplicate_score: float | None = None
    memory_write_attempted = False
    memory_write_stored = False
    memory_write_rejected = False
    memory_write_reasons: list[str] = []
    memory_similarity_scores: dict[str, Any] = {}
    try:
        index = pinecone.get_client().get_index()
        doc_store = pinecone.get_client()
        if hasattr(memory, "memory_config"):
            try:
                memory.memory_config.CONVERSATION_MEMORY_NAMESPACE = memory_namespace
            except Exception:
                pass
        memory_write_attempted = True
        raw_memory_payload = _build_memory_payload(query, answer_text)
        semantic_reject_threshold = _as_float(
            config.env_value("MEMORY_SEMANTIC_REJECT_THRESHOLD", "0.50"),
            0.50,
        )
        lexical_reject_threshold = _as_float(
            config.env_value("MEMORY_LEXICAL_REJECT_THRESHOLD", "0.90"),
            0.90,
        )
        dominance_threshold = _as_float(
            config.env_value("MEMORY_SINGLE_CHUNK_DOMINANCE_THRESHOLD", "0.72"),
            0.72,
        )
        min_words_required = _as_int(
            config.env_value("MEMORY_MIN_ANSWER_WORDS", "24"),
            24,
        )
        min_evidence_required = _as_int(
            config.env_value("MEMORY_MIN_EVIDENCE_COUNT", "2"),
            2,
        )
        min_used_chunks_required = _as_int(
            config.env_value("MEMORY_MIN_USED_CHUNKS", "2"),
            2,
        )

        if _contains_fallback_language(answer_text):
            memory_write_reasons.append("fallback_or_no_answer_response")
        if _word_count(answer_text) < min_words_required:
            memory_write_reasons.append("low_information_answer")
        if _is_definition_like(query, answer_text):
            memory_write_reasons.append("definition_style_answer")
        if len(evidence_texts) < min_evidence_required:
            memory_write_reasons.append("insufficient_evidence_count")
        if len(used_chunks) < min_used_chunks_required:
            memory_write_reasons.append("insufficient_used_chunks")

        distinct_docs = {
            chunk.get("doc_id")
            for chunk in used_chunks
            if chunk.get("doc_id")
        }
        distinct_sections = {
            f"{chunk.get('doc_id')}::{chunk.get('section')}"
            for chunk in used_chunks
            if chunk.get("section")
        }
        if len(distinct_docs) <= 1 and len(distinct_sections) <= 1:
            memory_write_reasons.append("no_multi_section_synthesis")

        answer_embedding = embed_client.embed_query(answer_text)
        doc_sim_results = doc_store.retrieve_similar(
            answer_embedding,
            top_k=5,
            namespace=doc_namespace,
        ).get("retrieved_chunks", [])
        memory_sim_results = doc_store.retrieve_similar(
            answer_embedding,
            top_k=5,
            namespace=memory_namespace,
        ).get("retrieved_chunks", [])

        best_doc_semantic = _best_retrieval_score(doc_sim_results)
        best_memory_semantic = _best_retrieval_score(memory_sim_results)

        if best_doc_semantic >= semantic_reject_threshold:
            memory_write_reasons.append("too_similar_to_document_chunk")
        if best_memory_semantic >= semantic_reject_threshold:
            memory_write_reasons.append("too_similar_to_existing_memory")

        doc_texts = [str(item.get("text") or "") for item in doc_sim_results if item.get("text")]
        memory_texts = [str(item.get("text") or "") for item in memory_sim_results if item.get("text")]
        used_texts = [str(item.get("text") or "") for item in used_chunks if item.get("text")]

        best_doc_lexical, _ = _best_lexical_similarity(answer_text, doc_texts)
        best_memory_lexical, _ = _best_lexical_similarity(answer_text, memory_texts)
        best_used_lexical, _ = _best_lexical_similarity(answer_text, used_texts)

        if best_doc_lexical >= lexical_reject_threshold:
            memory_write_reasons.append("lexical_duplicate_of_document_chunk")
        if best_memory_lexical >= lexical_reject_threshold:
            memory_write_reasons.append("lexical_duplicate_of_memory")
        if best_used_lexical >= dominance_threshold:
            memory_write_reasons.append("single_chunk_dominance")

        memory_similarity_scores = {
            "semantic_threshold": semantic_reject_threshold,
            "lexical_threshold": lexical_reject_threshold,
            "single_chunk_dominance_threshold": dominance_threshold,
            "best_document_semantic_similarity": best_doc_semantic,
            "best_memory_semantic_similarity": best_memory_semantic,
            "best_document_lexical_similarity": best_doc_lexical,
            "best_memory_lexical_similarity": best_memory_lexical,
            "best_used_chunk_lexical_similarity": best_used_lexical,
        }

        dedup_threshold = _as_float(config.env_value("MEMORY_DEDUP_THRESHOLD", "0.90"), 0.90)
        duplicate_check = index.query(
            vector=answer_embedding,
            top_k=1,
            namespace=memory_namespace,
            include_metadata=False,
        )
        matches = _extract_query_matches(duplicate_check)
        if matches:
            top_score = matches[0].get("score")
            if top_score is not None:
                memory_duplicate_score = float(top_score)
                memory_similarity_scores["duplicate_check_similarity"] = memory_duplicate_score
                if memory_duplicate_score >= dedup_threshold:
                    memory_duplicate_skipped = True
                    memory_write_reasons.append("duplicate_threshold_reached")
                    memory_result = {
                        "skipped_duplicate": True,
                        "duplicate_score": memory_duplicate_score,
                        "threshold": dedup_threshold,
                    }

        if memory_write_reasons:
            memory_write_rejected = True
            memory_result = memory_result or {
                "rejected": True,
                "reasons": memory_write_reasons,
            }
        elif not memory_duplicate_skipped:
            if hasattr(memory, "store_memory"):
                memory_result = await memory.store_memory(
                    raw_memory_payload,
                    _GLOBAL_MEMORY_SESSION_ID,
                    index,
                    _next_turn(_GLOBAL_MEMORY_SESSION_ID),
                )
            else:
                memory_result = memory.store_memory_sync(
                    raw_memory_payload,
                    index,
                    session_id=_GLOBAL_MEMORY_SESSION_ID,
                    turn_number=_next_turn(_GLOBAL_MEMORY_SESSION_ID),
                )
            memory_write_stored = bool(memory_result)
        else:
            memory_write_rejected = True
    except Exception as exc:
        memory_error = str(exc)
    if memory_write_attempted and not memory_write_stored and not memory_write_rejected and not memory_error:
        memory_write_rejected = True
        memory_write_reasons.append("memory_store_returned_empty")
    if memory_duplicate_skipped and "duplicate_threshold_reached" not in memory_write_reasons:
        memory_write_reasons.append("duplicate_threshold_reached")
    memory_write_debug = {
        "attempted": memory_write_attempted,
        "stored": memory_write_stored,
        "rejected": memory_write_rejected,
        "rejection_reasons": memory_write_reasons,
        "similarity_scores": memory_similarity_scores,
        "duplicate_skipped": memory_duplicate_skipped,
        "duplicate_score": memory_duplicate_score,
        "error": memory_error,
    }

    audit_event(
        "chat_request",
        {
            "query": query,
            "multihop_enabled": payload.use_multihop,
            "multihop_used": multihop_used,
            "multi_hop_triggered": multihop_used,
            "multi_hop_reason": multi_hop_reason,
            "multi_hop_reason_list": multi_hop_reason_list,
            "multihop_sub_query": multihop_sub_query,
            "multihop_validation": multihop_validation,
            "multihop_error": multihop_error,
            "multi_hop_retrievals": multi_hop_retrieval_sections,
            "sources": sources,
            "retrieval_stats": context.get("stats") or {
                "doc_count": context.get("doc_count", 0),
                "memory_count": context.get("memory_count", 0),
                "related_count": context.get("related_count", 0),
            },
            "memory_raw_count": memory_raw_count,
            "memory_dropped_count": memory_dropped_count,
            "memory_quality_filtered_count": memory_quality_filtered_count,
            "memory_in_context_count": memory_in_context_count,
            "memory_duplicates_removed_from_retrieval": memory_duplicates_removed,
            "memory_quality_filtered": memory_quality_filtered[:10],
            "context": summarize_text(str(context)),
            "answer": answer_for_user,
            "memory_stored": memory_write_stored,
            "memory_duplicate_skipped": memory_duplicate_skipped,
            "memory_duplicate_score": memory_duplicate_score,
            "memory_write": memory_write_debug,
            "memory_namespace": memory_namespace,
            "memory_session_id": _GLOBAL_MEMORY_SESSION_ID,
            "memory_error": memory_error,
        },
        user_n=user_n,
        ip=ip,
    )

    prompt_report = write_prompt_reports(
        query=query,
        strategy="hybrid+multihop_failsafe" if multihop_used else "hybrid",
        context_sent_to_generator=llm_context_chunks,
        retrieved_topk=retrieved_topk,
        dense_topk=dense_topk,
        sparse_topk=sparse_topk,
        related_chunks=related_chunks_report,
        memory_chunks=memory_chunks_report,
        memory_raw_count=memory_raw_count,
        memory_dropped_count=memory_dropped_count,
        answer=answer_text,
        used_chunks=used_chunks,
        evidence_block=evidence_texts or None,
        original_query_retrieval=_build_retrieval_section(
            query=query,
            strategy="hybrid",
            context_chunks=original_generation_candidates,
            retrieved_topk_chunks=original_document_chunks,
            dense_chunks=dense_raw_chunks,
            sparse_chunks=sparse_raw_chunks,
            related_chunks=original_related_chunks,
            memory_chunks=original_memory_chunks,
            memory_raw_count=original_memory_raw_count,
            memory_dropped_count=original_memory_dropped_count,
        ),
        multi_hop_retrievals=multi_hop_retrieval_sections or None,
        multi_hop_triggered=multihop_used,
        multi_hop_reason=multi_hop_reason or None,
        memory_write_debug=memory_write_debug,
        memory_retrieval_filter={
            "filtered_count": memory_quality_filtered_count,
            "filtered_entries": memory_quality_filtered[:20],
            "dedup_removed_count": memory_duplicates_removed,
        },
    )

    response: dict[str, Any] = {"answer": answer_for_user, "sources": sources}
    if payload.debug:
        response["debug"] = {
            "strategy": "hybrid+multihop_failsafe" if multihop_used else "hybrid",
            "multihop_enabled": payload.use_multihop,
            "multihop_used": multihop_used,
            "multi_hop_triggered": multihop_used,
            "multi_hop_reason": multi_hop_reason,
            "multi_hop_reason_list": multi_hop_reason_list,
            "multihop_sub_query": multihop_sub_query,
            "multihop_validation": multihop_validation,
            "multihop_error": multihop_error,
            "multi_hop_retrievals": multi_hop_retrieval_sections,
            "doc_count": context.get("doc_count", 0),
            "memory_count": context.get("memory_count", 0),
            "memory_raw_count": memory_raw_count,
            "memory_dropped_count": memory_dropped_count,
            "memory_top_k": memory_top_k,
            "memory_min_score": memory_min_score,
            "memory_boost": memory_boost,
            "related_count": context.get("related_count", 0),
            "evidence_count": len(evidence_texts),
            "llm_context_count": len(llm_context_chunks),
            "memory_in_context_count": memory_in_context_count,
            "memory_duplicates_removed_from_retrieval": memory_duplicates_removed,
            "memory_quality_filtered_count": memory_quality_filtered_count,
            "memory_quality_filtered": memory_quality_filtered[:20],
            "sources_count": len(sources),
            "retrieved_topk": retrieved_topk,
            "dense_topk": dense_topk,
            "sparse_topk": sparse_topk,
            "related_chunks": related_chunks_report,
            "memory_chunks": memory_chunks_report,
            "context_sent_to_generator": llm_context_chunks,
            "used_chunks": used_chunks,
            "evidence_block": evidence_texts,
            "memory_stored": memory_write_stored,
            "memory_duplicate_skipped": memory_duplicate_skipped,
            "memory_duplicate_score": memory_duplicate_score,
            "memory_write_attempted": memory_write_attempted,
            "memory_write_rejected": memory_write_rejected,
            "memory_write_reasons": memory_write_reasons,
            "memory_similarity_scores": memory_similarity_scores,
            "memory_write": memory_write_debug,
            "memory_namespace": memory_namespace,
            "memory_session_id": _GLOBAL_MEMORY_SESSION_ID,
            "memory_error": memory_error,
            "run_report": prompt_report,
        }
    return response