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