prazy1208/text2sql
0
1"""2Table metadata retrieval for the Table Agent (Step 5a).3 4Loads per-schema table metadata from metadata_store. Candidate **text** for the Table Agent5is schema + table name + table description only (column details are left to a later agent).6FAISS shortlist (N > threshold) still uses the existing table index built from full metadata.7If the schema has more than TABLE_SHORTLIST_THRESHOLD tables, runs FAISS similarity search8with the same embedding model as the index build (all-MiniLM-L6-v2).9"""10 11from __future__ import annotations12 13import json14import logging15from typing import Any16 17from backend.config import (18 FAISS_INDEX_DIR,19 FAISS_INDEX_NAMES,20 METADATA_STORE_DIR,21 USE_CASE_TO_SCHEMA,22)23# Reuse one SentenceTransformer instance with business-rules retrieval (same model name).24from backend.services.business_rules_retrieval import _get_embedding_model25 26logger = logging.getLogger(__name__)27 28TABLE_SHORTLIST_THRESHOLD = 1029DEFAULT_TOP_K = 1030 31_table_index_cache: dict[str, tuple[Any, list[dict]]] = {}32 33 34def table_metadata_summary_for_agent(m: dict) -> dict:35 """36 Table-level fields only for the Table Agent (no columns; another agent will use column metadata).37 """38 return {39 "schema_name": m["schema_name"],40 "table_name": m["table_name"],41 "table_description": m.get("table_description") or "",42 }43 44 45def metadata_entry_to_text(m: dict) -> str:46 """One table → short structured text (schema, table, description only)."""47 desc = m.get("table_description") or "(no description)"48 return "\n".join(49 [50 f"Schema: {m['schema_name']}",51 f"Table: {m['table_name']}",52 f"Description: {desc}",53 ]54 )55 56 57def load_table_metadata_list(schema_name: str) -> list[dict]:58 """59 Load metadata_store/{schema_name}_metadata.json (list of table dicts, same order as FAISS).60 """61 path = METADATA_STORE_DIR / f"{schema_name}_metadata.json"62 if not path.exists():63 raise FileNotFoundError(64 f"Table metadata not found: {path}. Run build_vector_store.py for this schema."65 )66 with open(path, "r", encoding="utf-8") as f:67 data = json.load(f)68 if not isinstance(data, list):69 raise ValueError(f"Expected list in {path}, got {type(data).__name__}")70 return data71 72 73def _build_retrieval_query(rephrased_question: str, keywords: list[str] | None) -> str:74 parts: list[str] = []75 if rephrased_question and rephrased_question.strip():76 parts.append(rephrased_question.strip())77 if keywords:78 for k in keywords:79 if k and str(k).strip():80 parts.append(str(k).strip())81 return " ".join(parts) if parts else ""82 83 84def _load_table_faiss_and_metadata(schema_name: str) -> tuple[Any, list[dict]]:85 """FAISS index + table metadata list (cached). Vectors align with metadata list indices."""86 if schema_name in _table_index_cache:87 return _table_index_cache[schema_name]88 89 index_filename = FAISS_INDEX_NAMES.get(schema_name)90 if not index_filename:91 raise ValueError(f"No FAISS index mapping for schema {schema_name!r}")92 93 index_path = FAISS_INDEX_DIR / index_filename94 if not index_path.exists():95 raise FileNotFoundError(96 f"Table FAISS index not found: {index_path}. Run build_vector_store.py."97 )98 99 metadata_list = load_table_metadata_list(schema_name)100 101 import faiss102 103 index = faiss.read_index(str(index_path))104 n_meta = len(metadata_list)105 n_index = index.ntotal106 if n_meta != n_index:107 logger.warning(108 "Metadata count (%d) != FAISS ntotal (%d) for %s; mapping by index may be wrong.",109 n_meta,110 n_index,111 schema_name,112 )113 114 _table_index_cache[schema_name] = (index, metadata_list)115 logger.debug(116 "Loaded table FAISS index for %s (%d vectors, %d metadata rows)",117 schema_name,118 n_index,119 n_meta,120 )121 return index, metadata_list122 123 124def shortlist_candidate_tables(125 use_case: str,126 rephrased_question: str,127 keywords: list[str] | None = None,128 top_k: int = DEFAULT_TOP_K,129) -> list[dict]:130 """131 Return candidate table metadata dicts for the Table Agent.132 133 - Resolves schema from use_case via USE_CASE_TO_SCHEMA.134 - If N <= TABLE_SHORTLIST_THRESHOLD: returns all tables (full metadata dicts).135 - If N > TABLE_SHORTLIST_THRESHOLD: embeds query from rephrased_question + keywords,136 FAISS search with k = min(top_k, N), returns metadata rows at those indices.137 138 Each dict has schema_name, table_name, table_description only (no columns).139 """140 schema_name = USE_CASE_TO_SCHEMA.get(use_case)141 if not schema_name:142 logger.warning("Unknown use_case %r, cannot shortlist table metadata", use_case)143 return []144 145 metadata_list = load_table_metadata_list(schema_name)146 n = len(metadata_list)147 if n == 0:148 return []149 150 if n <= TABLE_SHORTLIST_THRESHOLD:151 logger.debug(152 "Schema %s has %d tables (<= %d); returning all as candidates.",153 schema_name,154 n,155 TABLE_SHORTLIST_THRESHOLD,156 )157 return [table_metadata_summary_for_agent(m) for m in metadata_list]158 159 query_text = _build_retrieval_query(rephrased_question, keywords)160 if not query_text.strip():161 logger.warning(162 "No query text for FAISS shortlist (%s); returning first %d tables.",163 schema_name,164 min(top_k, n),165 )166 return [167 table_metadata_summary_for_agent(m)168 for m in metadata_list[: min(top_k, n)]169 ]170 171 index, meta_aligned = _load_table_faiss_and_metadata(schema_name)172 if index.ntotal == 0:173 return []174 175 k = min(top_k, n, index.ntotal)176 import numpy as np177 178 model = _get_embedding_model()179 query_embedding = model.encode([query_text], show_progress_bar=False)180 query_vec = np.asarray(query_embedding, dtype=np.float32)181 distances, indices = index.search(query_vec, k)182 idx_list = indices[0].tolist()183 184 seen: set[int] = set()185 result: list[dict] = []186 for idx in idx_list:187 if idx < 0:188 continue189 if idx in seen:190 continue191 seen.add(idx)192 if idx >= len(meta_aligned):193 continue194 result.append(table_metadata_summary_for_agent(meta_aligned[idx]))195 196 return result197 198 199def candidate_tables_as_texts(candidates: list[dict]) -> list[str]:200 """Structured text per candidate table (for LLM context)."""201 return [metadata_entry_to_text(m) for m in candidates]202 