Spaces:
Paused
Paused
File size: 18,016 Bytes
7758af3 3626208 0017280 fbd18c8 f95cdb7 fbd18c8 0017280 0bb6123 2317f6c 0bb6123 f9f7706 0bb6123 fbd18c8 0017280 fbd18c8 0017280 0bb6123 fbd18c8 0017280 0bb6123 0017280 9815e49 0bb6123 0017280 9815e49 0bb6123 0017280 2317f6c 1330e5c 2317f6c 0017280 f9f7706 0017280 f9f7706 0017280 9815e49 0017280 1330e5c f9f7706 0017280 9815e49 0017280 f9f7706 0017280 7758af3 f9f7706 7758af3 f9f7706 0017280 f9f7706 0017280 7758af3 0017280 7758af3 0017280 7758af3 0017280 7758af3 0017280 7758af3 d6a01c7 7758af3 0017280 9093d8f fbd18c8 f95cdb7 9093d8f b71f629 fbd18c8 f95cdb7 b71f629 f95cdb7 b71f629 6b689a5 7758af3 0017280 6b689a5 0017280 7758af3 6b689a5 7758af3 1330e5c 6b689a5 7758af3 6b689a5 d4a1a18 6b689a5 d4a1a18 b71f629 d4a1a18 fbd18c8 d4a1a18 6b689a5 7758af3 6b689a5 7758af3 6b689a5 7758af3 3d1d5b4 6b689a5 1330e5c 3d1d5b4 6b689a5 7758af3 3626208 7758af3 6b689a5 7758af3 6b689a5 7758af3 f97b9fe fbd18c8 f97b9fe 6b689a5 f97b9fe fbd18c8 0017280 7758af3 0017280 7758af3 0017280 f9f7706 3626208 f9f7706 7758af3 0017280 7758af3 6b689a5 0017280 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 | """Dual retriever: ChromaDB vector search + optional web search + SQL retrieval."""
import re
from typing import Iterable
from src.search import search, format_web_results_as_context
_collection_cache = {}
def clear_collection_cache():
"""Clear the cached ChromaDB collection so it's re-read after re-ingestion."""
_collection_cache.clear()
def _get_cached_collection(cfg: dict):
"""Get or create cached ChromaDB collection."""
from src.ingest import get_chroma_collection
db_path = cfg.get("paths", {}).get("vector_db", "chroma_db")
embed_provider = cfg.get("embeddings", {}).get("provider", "local")
cache_key = f"{db_path}:{embed_provider}"
if cache_key not in _collection_cache:
_collection_cache[cache_key] = get_chroma_collection(cfg)
return _collection_cache[cache_key]
NO_SOURCES_REFUSAL = (
"I don't have any information on this topic in my knowledge base. "
"No relevant local documents or web sources were found. "
"Please try a different question, or add relevant materials to "
"the knowledge_base/ folder and run ingestion again."
)
def retrieve_from_vectordb(query: str, cfg: dict) -> list[dict]:
"""Retrieve all relevant chunks from ChromaDB using a distance threshold.
Queries a large candidate pool, then filters to keep only chunks
with cosine distance below ``max_distance``. The ``top_k`` setting
acts as a hard cap to prevent overwhelming the LLM context.
When no chunks pass the distance filter, the KB meta overview chunk
is included as a fallback so that meta-questions about the knowledge
base can still be answered.
"""
retrieval_cfg = cfg.get("retrieval", {})
top_k = retrieval_cfg.get("top_k", 20)
max_distance = retrieval_cfg.get("max_distance", 0.55)
collection = _get_cached_collection(cfg)
total = collection.count()
if total == 0:
return []
# Query a large candidate pool, then filter by relevance
candidate_count = min(top_k, total)
results = collection.query(
query_texts=[query],
n_results=candidate_count,
)
chunks = []
documents = results.get("documents", [[]])[0]
distances = results.get("distances", [[]])[0]
metadatas = results.get("metadatas", [[]])[0]
for i in range(len(documents)):
distance = distances[i]
if distance > max_distance:
continue
chunks.append({
"text": documents[i],
"metadata": metadatas[i],
"distance": distance,
})
# Fallback: if no chunks passed the distance filter, include the KB
# meta overview chunk so meta-questions can still be answered.
if not chunks:
from src.kb_meta import META_CHUNK_ID
try:
meta_result = collection.get(
ids=[META_CHUNK_ID],
include=["documents", "metadatas"],
)
if meta_result.get("documents") and meta_result["documents"][0]:
chunks.append({
"text": meta_result["documents"][0],
"metadata": meta_result["metadatas"][0],
"distance": 1.0, # fallback β not a real distance match
})
except Exception:
pass
return chunks
def format_db_results_as_context(chunks: list[dict], max_chars: int = 0) -> str:
"""Format vector DB results with CHUNK-LOCAL IDs for citation anchoring.
Args:
chunks: Retrieved chunks sorted by distance (most relevant first).
max_chars: If > 0, stop adding chunks once total context exceeds
this character budget. Chunks already sorted by relevance
ensure the most important content is kept.
"""
if not chunks:
return ""
header = "=== Local Document Results (PRIMARY β always trust these over web sources) ===\n"
parts = [header]
total_chars = len(header)
for i, chunk in enumerate(chunks, 1):
meta = chunk.get("metadata", {})
dataset = meta.get("dataset", "")
dataset_label = f" [Dataset: {dataset}]" if dataset else ""
source_path = meta.get('source', 'unknown')
entry = (
f"[CHUNK-LOCAL-{i:03d}] From: {meta.get('source', 'unknown')}, "
f"Page/Section {meta.get('page', '?')}{dataset_label}\n"
f" Path: {source_path}\n"
f" {chunk.get('text', '')}\n"
)
if max_chars > 0 and total_chars + len(entry) > max_chars and i > 1:
break
parts.append(entry)
total_chars += len(entry)
return "\n".join(parts)
def build_combined_context(
db_results: list[dict], web_results: list[dict], sql_context: str = "",
max_context_chars: int = 0,
) -> str:
"""Build combined context string from local, SQL, and web results.
When *max_context_chars* > 0, the local-chunk context is truncated to
fit within the budget (after accounting for SQL and web context).
SQL and web context are never truncated; the budget applies to local chunks.
"""
# SQL context is precise and compact β always include in full.
# Web context is supplementary β include in full but account for its size.
# Allocate remaining budget to local chunks.
web_context = format_web_results_as_context(web_results)
local_budget = 0
if max_context_chars > 0:
reserved = len(sql_context) + len(web_context)
local_budget = max(max_context_chars - reserved, 2000)
db_context = format_db_results_as_context(db_results, max_chars=local_budget)
has_local = bool(db_results) or bool(sql_context)
has_web = bool(web_results)
if not has_local and not has_web:
return NO_SOURCES_REFUSAL
parts = []
if has_local and not has_web:
parts.append("ONLY use the local sources below. Do not add any outside knowledge.\n")
if has_local and has_web:
parts.append(
"IMPORTANT: Local document sources are the PRIMARY authority. "
"Web sources are SUPPLEMENTARY only. If any web source contradicts "
"a local document, trust the local document.\n"
)
if not has_local and has_web:
parts.append(
"NOTE: No local documents are indexed. All sources below come from "
"academic web search. You MUST include the URL link for every source cited. "
"State clearly that you have no curated local sources.\n"
)
if sql_context:
parts.append(sql_context)
if db_context:
parts.append(db_context)
if web_context:
parts.append(web_context)
return "\n\n".join(parts)
def _extract_table_from_query(sql_query: str) -> str:
"""Extract the first table name from a SQL query (best-effort)."""
match = re.search(r'\bFROM\s+["`]?(\w+)["`]?', sql_query, re.IGNORECASE)
return match.group(1) if match else ""
def _build_fallback_sql_query(query: str, cfg: dict) -> str | None:
"""Build a simple keyword search SQL query as a last-resort fallback.
Loads the SQL schema, finds TEXT columns in the first table, and builds
a LIKE-based search across those columns. Returns None if no schema
or no text columns exist.
"""
import json as _json
import os
from pathlib import Path
sql_db_dir = cfg.get("paths", {}).get("sql_db", "sql_db")
if not os.path.isabs(sql_db_dir):
project_root = Path(__file__).resolve().parent.parent
sql_db_dir = os.path.join(str(project_root), sql_db_dir)
schema_path = os.path.join(sql_db_dir, "sql_schemas.json")
try:
with open(schema_path) as f:
schema = _json.load(f)
except Exception:
return None
if not schema:
return None
max_rows = cfg.get("sql", {}).get("max_rows", 200)
# Extract meaningful search words (3+ chars, skip common stopwords)
stopwords = {"the", "and", "for", "are", "but", "not", "you", "all",
"can", "had", "her", "was", "one", "our", "out", "has",
"what", "how", "who", "which", "when", "where", "with",
"from", "that", "this", "than", "then", "they", "been"}
words = [w for w in re.split(r'\W+', query.lower()) if len(w) >= 3 and w not in stopwords]
if not words:
return None
# Try first table with TEXT columns
for table_name, info in schema.items():
text_cols = [
c.get("name") for c in info.get("columns", [])
if c.get("type", "").upper() == "TEXT" and c.get("name")
]
if not text_cols:
continue
conditions = []
for col in text_cols:
for word in words:
safe_word = word.replace("'", "''")
if not re.match(r'^\w+$', safe_word):
continue # Skip non-alphanumeric words
safe_word = safe_word.replace('%', '\\%').replace('_', '\\_')
conditions.append(f'"{col}" LIKE \'%{safe_word}%\' ESCAPE \'\\\'')
if conditions:
where_clause = " OR ".join(conditions)
return f'SELECT * FROM "{table_name}" WHERE {where_clause} LIMIT {max_rows}'
return None
def _try_alternate_columns(sql_query: str, cfg: dict) -> tuple[list[dict], str] | None:
"""Try phrase-level LIKE on other TEXT columns of the same table.
When ``WHERE "Country" LIKE '%South Korea%'`` fails because the value
is stored in a different column (e.g. ``Country_OLD``), this function
tries each TEXT column in the table until one matches.
Returns (rows, effective_query) or None if nothing found.
"""
import json as _json
import os
# Extract table name and search value from the query
table_name = _extract_table_from_query(sql_query)
if not table_name:
return None
# Extract the LIKE value or = value from the WHERE clause
value_match = re.search(
r"LIKE\s+'%([^%]+)%'|=\s*'([^']+)'",
sql_query, re.IGNORECASE,
)
if not value_match:
return None
search_value = value_match.group(1) or value_match.group(2)
if not search_value:
return None
# Extract the column currently being searched
col_match = re.search(
r'["`]?(\w+)["`]?\s*(?:LIKE|=)\s*',
sql_query, re.IGNORECASE,
)
current_col = col_match.group(1) if col_match else ""
# Load schema to find other TEXT columns
sql_db_dir = cfg.get("paths", {}).get("sql_db", "sql_db")
if not os.path.isabs(sql_db_dir):
from pathlib import Path
project_root = Path(__file__).resolve().parent.parent
sql_db_dir = os.path.join(str(project_root), sql_db_dir)
schema_path = os.path.join(sql_db_dir, "sql_schemas.json")
try:
with open(schema_path) as f:
schema = _json.load(f)
except Exception:
return None
table_info = schema.get(table_name)
if not table_info:
return None
text_cols = [
c.get("name") for c in table_info.get("columns", [])
if c.get("type", "").upper() == "TEXT"
and c.get("name")
and c.get("name") != current_col
]
from src.sql_retriever import execute_sql_query
max_rows = cfg.get("sql", {}).get("max_rows", 200)
safe_value = search_value.replace("'", "''")
# Strip dangerous characters for defense-in-depth
if not re.match(r"^[\w\s.,'-]+$", safe_value):
safe_value = re.sub(r"[;'\\\"]", "", safe_value)
safe_value = safe_value.replace('%', '\\%').replace('_', '\\_')
for col in text_cols:
alt_query = f'SELECT * FROM "{table_name}" WHERE "{col}" LIKE \'%{safe_value}%\' ESCAPE \'\\\' LIMIT {max_rows}'
rows = execute_sql_query(alt_query, cfg)
if rows:
return rows, alt_query
return None
def _run_sql_retrieval(sql_query: str, cfg: dict) -> tuple[list[dict], str, str]:
"""Execute SQL query and format results as context.
Tries the exact query first. If it returns zero rows, automatically
retries with fuzzy matching (``=`` β ``LIKE '%β¦%'``).
Returns (sql_rows, sql_context, match_type) where match_type is
``"exact"``, ``"fuzzy"``, or ``""`` (no results).
"""
import json as _json
import os
from src.sql_retriever import (
execute_sql_query, format_sql_results_as_context,
_lookup_source_file, make_fuzzy_query,
)
max_rows = cfg.get("sql", {}).get("max_rows", 200)
# Step 1: try exact query
sql_rows = execute_sql_query(sql_query, cfg)
effective_query = sql_query
match_type = "exact"
# Step 2: if no results, try phrase-level fuzzy (LIKE '%South Korea%')
if not sql_rows:
fuzzy = make_fuzzy_query(sql_query)
if fuzzy:
sql_rows = execute_sql_query(fuzzy, cfg)
if sql_rows:
effective_query = fuzzy
match_type = "fuzzy (phrase)"
# Step 2.5: if phrase fuzzy failed on original column, try other TEXT columns
if not sql_rows:
alt = _try_alternate_columns(fuzzy or sql_query, cfg)
if alt:
sql_rows, effective_query = alt
match_type = "fuzzy (alt column)"
# Step 3: if still no results, try word-level fuzzy (LIKE '%Korea%')
if not sql_rows:
word_fuzzy = make_fuzzy_query(sql_query, word_level=True)
if word_fuzzy and word_fuzzy != (fuzzy or ""):
sql_rows = execute_sql_query(word_fuzzy, cfg)
if sql_rows:
effective_query = word_fuzzy
match_type = "fuzzy (word)"
if not sql_rows:
return [], "", ""
sql_db_dir = cfg.get("paths", {}).get("sql_db", "sql_db")
if not os.path.isabs(sql_db_dir):
from pathlib import Path
project_root = Path(__file__).resolve().parent.parent
sql_db_dir = os.path.join(str(project_root), sql_db_dir)
schema_path = os.path.join(sql_db_dir, "sql_schemas.json")
schema = {}
if os.path.exists(schema_path):
try:
with open(schema_path) as f:
schema = _json.load(f)
except Exception:
pass
table_name = _extract_table_from_query(effective_query)
source_file = _lookup_source_file(table_name, schema)
table_info = schema.get(table_name)
sql_context = format_sql_results_as_context(
sql_rows, effective_query, source_file, table_info=table_info,
max_rows=max_rows,
)
return sql_rows, sql_context, match_type
def retrieve(
query: str,
cfg: dict,
route: str = "vector",
sql_query: str = None,
db_filter: Iterable = None
) -> dict:
"""Run retrieval with SQL/vector routing and fallback.
Returns dict with: context, db_results, web_results, sql_results, has_sources.
"""
sql_enabled = cfg.get("sql", {}).get("enabled", True)
sql_context = ""
sql_rows = []
sql_match_type = ""
# ββ SQL retrieval ββββββββββββββββββββββββββββββββββββββββββββββββ
if sql_enabled and route in ("sql", "both") and sql_query:
sql_rows, sql_context, sql_match_type = _run_sql_retrieval(sql_query, cfg)
# ββ Vector retrieval βββββββββββββββββββββββββββββββββββββββββββββ
db_results = []
if route in ("vector", "both") or (route == "sql" and not sql_rows):
db_results = retrieve_from_vectordb(query, cfg)
# ββ Fallback: try keyword-based SQL when primary path returned nothing β
# - route "vector"/"both": vector found nothing β try SQL keyword search
# - route "sql": LLM SQL query failed β try SQL keyword search as last resort
sql_fallback_needed = (
sql_enabled
and not sql_rows
and (
(route in ("vector", "both") and not db_results)
or (route == "sql" and not db_results)
)
)
if sql_fallback_needed:
if route == "sql":
# LLM-generated sql_query already failed; build a fresh keyword query
fallback_query = _build_fallback_sql_query(query, cfg)
if fallback_query:
sql_rows, sql_context, sql_match_type = _run_sql_retrieval(fallback_query, cfg)
else:
# route is "vector" or "both" β build fresh keyword query
# (if sql_query exists, it was already tried above and returned nothing)
fallback_sql = _build_fallback_sql_query(query, cfg)
if fallback_sql:
sql_rows, sql_context, sql_match_type = _run_sql_retrieval(fallback_sql, cfg)
# ββ Web search βββββββββββββββββββββββββββββββββββββββββββββββββββ
web_enabled = cfg.get("web_search", {}).get("enabled", False)
web_results = []
if web_enabled:
backend = cfg.get("web_search", {}).get("backend", "semantic_scholar")
max_results = cfg.get("web_search", {}).get("max_results", 5)
if not db_results and not sql_rows:
max_results *= 2
web_results = search(query, backend=backend, limit=max_results)
max_context_chars = cfg.get("retrieval", {}).get("max_context_chars", 12000)
if db_filter is not None:
import os
db_results = [r for r in db_results if os.path.basename(r['metadata']['source']) in db_filter]
combined_context = build_combined_context(
db_results, web_results, sql_context=sql_context,
max_context_chars=max_context_chars,
)
has_sources = bool(db_results) or bool(web_results) or bool(sql_rows)
return {
"context": combined_context,
"db_results": db_results,
"web_results": web_results,
"sql_results": sql_rows,
"sql_match_type": sql_match_type,
"has_sources": has_sources,
}
|