CoolFace
Apppublic

prazy1208/text2sql

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
table_metadata_retrieval.py202 linesDownload Raw Back to services
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