Team Ai
Apppublic

prazy1208/text2sql

sourceHugging Faceupdated 5mo agoView on Hugging Face
0likes
table_agent.py310 linesDownload Raw Back to agents
1"""2Table Agent (Stage 2 — table selection).3Inputs: use_case, rephrased_question, keywords (from Intent Agent).4Outputs: selected_tables as fully-qualified names schema.table_name validated against shortlist ∪ FK relationship tables.5 6Flow: shortlist candidate tables (metadata + optional FAISS) → LLM picks subset → validate against allowed FQNs (candidates plus any schema.table that appears in FK relationships).7"""8 9from __future__ import annotations10 11import json12import logging13import re14 15from backend.services.llm_client import chat_completion16from backend.services.table_metadata_retrieval import (17    DEFAULT_TOP_K,18    candidate_tables_as_texts,19    shortlist_candidate_tables,20)21 22logger = logging.getLogger(__name__)23 24# ---------------------------------------------------------------------------25# Table selection prompt (placeholders: rephrased_question, keywords_block,26# relationships_block, numbered_candidates)27# ---------------------------------------------------------------------------28TABLE_AGENT_PROMPT = """You are the Table Agent in a Natural Language to SQL system.29 30Your task is to select ALL database tables required to correctly answer the user's analytical question.31 32You are given:331. The analytical question342. Relevant keywords353. Known foreign-key relationships between tables364. A list of candidate tables (may be incomplete due to pre-filtering)37 38--------------------------------------------------39ANALYTICAL QUESTION40--------------------------------------------------41{rephrased_question}42 43--------------------------------------------------44KEYWORDS45--------------------------------------------------46{keywords_block}47 48--------------------------------------------------49KNOWN FOREIGN KEY RELATIONSHIPS50--------------------------------------------------51{relationships_block}52 53Each relationship is formatted as:54schema.table.column -> schema.table.column55 56--------------------------------------------------57CANDIDATE TABLES58--------------------------------------------------59{numbered_candidates}60 61--------------------------------------------------62SELECTION PRINCIPLES63--------------------------------------------------64 65- Select ALL tables necessary to correctly answer the question.66- Do NOT omit required tables, even if multiple tables are needed.67- Prefer correctness over minimality.68 69- Identify entities mentioned in the question (e.g., customers, accounts, patients, transactions, products, visits).70- Ensure each entity is represented by a selected table if required.71 72--------------------------------------------------73RELATIONSHIP-AWARE EXPANSION (CRITICAL)74--------------------------------------------------75 76- Use the provided relationships to understand how tables are connected.77- If two selected tables are NOT directly related, you MUST include the intermediate table(s) that connect them.78 79- If a required table is missing from the candidate list but is necessary to complete a relationship path:80  → You MUST include it using the relationships provided.81 82- Think in terms of join paths:83  table A → table B → table C84 85- NEVER assume direct relationships if they are not supported by the given relationships.86 87- Only add tables that exist in the relationship list (do NOT invent tables).88 89--------------------------------------------------90EXAMPLES91--------------------------------------------------92 93Finance Example:94Question: "Which accounts have the highest number of transactions?"95-> Selected Tables: ["finance_schema.accounts", "finance_schema.transactions"]96 97Retail Example:98Question: "Which customers bought which products?"99-> Selected Tables: ["retail_schema.customers", "retail_schema.orders", "retail_schema.products"]100 101Healthcare Example:102Question: "Which patients had which diagnoses during their visits?"103-> Selected Tables: ["healthcare_schema.patients", "healthcare_schema.visits", "healthcare_schema.diagnoses"]104 105--------------------------------------------------106STRICT RULES107--------------------------------------------------108 109- Use exact table identifiers: schema_name.table_name110- Do NOT invent or modify table names111- Do NOT include column names112- Do NOT include explanations or extra text113 114- Do NOT include unrelated tables115- Do NOT include all tables by default116 117- Tables may be selected from:118  1. Candidate list119  2. Relationship expansion (ONLY if required for joins)120 121- If the query is ambiguous and no table can be confidently selected, return an empty list122 123--------------------------------------------------124OUTPUT FORMAT125--------------------------------------------------126 127Return ONLY valid JSON:128 129{{130  "selected_tables": ["schema_name.table_name"]131}}132 133OR134 135{{136  "selected_tables": ["schema_name.table1", "schema_name.table2"]137}}138 139OR140 141{{142  "selected_tables": []143}}144 145--------------------------------------------------146FINAL VALIDATION147--------------------------------------------------148 149Before responding, ensure:150 151- All entities in the question are covered152- All selected tables are connected via valid relationship paths153- No required intermediate table is missing154- No unrelated table is included155- Output is valid JSON only"""156 157 158def _format_keywords(keywords: list[str] | None) -> str:159    if not keywords:160        return "(none)"161    lines = [f"- {k}" for k in keywords if k and str(k).strip()]162    return "\n".join(lines) if lines else "(none)"163 164 165def _format_relationships_block(relationships: list[dict] | None) -> str:166    if not relationships:167        return "(none)"168    lines: list[str] = []169    for r in relationships:170        rt = (r.get("relationship_text") or "").strip()171        if rt:172            lines.append(f"- {rt}")173    return "\n".join(lines) if lines else "(none)"174 175 176def _candidate_fqn(c: dict) -> str:177    """Fully-qualified table name as used in prompts and validation."""178    return f"{c['schema_name']}.{c['table_name']}"179 180 181def _table_fqns_from_relationships(relationships: list[dict] | None) -> set[str]:182    """Unique schema.table names appearing as FK source or target in relationship rows."""183    if not relationships:184        return set()185    out: set[str] = set()186    for r in relationships:187        sn = (r.get("schema_name") or "").strip()188        st = (r.get("source_table") or "").strip()189        ts = (r.get("target_schema") or "").strip()190        tt = (r.get("target_table") or "").strip()191        if sn and st:192            out.add(f"{sn}.{st}")193        if ts and tt:194            out.add(f"{ts}.{tt}")195    return out196 197 198def _build_numbered_candidates_text(candidates: list[dict]) -> str:199    """Numbered list of table descriptions + explicit FQN for each row."""200    texts = candidate_tables_as_texts(candidates)201    blocks = []202    for i, (c, text) in enumerate(zip(candidates, texts), start=1):203        fqn = _candidate_fqn(c)204        blocks.append(f"### Candidate {i} — `{fqn}`\n{text}")205    return "\n\n".join(blocks)206 207 208def _parse_table_agent_response(response_text: str) -> list[str] | None:209    """210    Parse LLM response; return list of selected table strings or None on failure.211    """212    text = response_text.strip()213    if "```json" in text:214        text = re.sub(r"^.*?```json\s*", "", text, flags=re.DOTALL)215    if "```" in text:216        text = re.sub(r"\s*```.*$", "", text, flags=re.DOTALL)217    text = text.strip()218    try:219        data = json.loads(text)220    except (json.JSONDecodeError, TypeError) as e:221        logger.warning("Failed to parse Table Agent JSON: %s", e)222        return None223    raw = data.get("selected_tables")224    if raw is None:225        return []226    if not isinstance(raw, list):227        return None228    out: list[str] = []229    for x in raw:230        if x is None:231            continue232        s = str(x).strip()233        if s:234            out.append(s)235    return out236 237 238def _validate_selected_tables(selected: list[str], allowed: set[str]) -> list[str]:239    """240    Keep only entries that exactly match an allowed FQN (shortlist candidate and/or241    schema.table from FK relationship rows). Preserves order, deduplicates.242    """243    seen: set[str] = set()244    result: list[str] = []245    for name in selected:246        if name not in allowed:247            logger.info(248                "Dropping invalid table selection (not in candidates or relationships): %r",249                name,250            )251            continue252        if name in seen:253            continue254        seen.add(name)255        result.append(name)256    return result257 258 259def run_table_agent(260    use_case: str,261    rephrased_question: str,262    keywords: list[str] | None = None,263    top_k: int = DEFAULT_TOP_K,264    *,265    relationships: list[dict] | None = None,266) -> dict:267    """268    Select tables needed for the question using shortlist + LLM + validation.269 270    Returns:271        { "selected_tables": list[str] }  # FQNs schema.table_name; allowed = shortlist ∪ FK endpoints272    """273    rq = (rephrased_question or "").strip()274    candidates = shortlist_candidate_tables(275        use_case,276        rq,277        keywords,278        top_k=top_k,279    )280 281    if not candidates:282        logger.warning("Table Agent: no candidate tables for use_case=%r", use_case)283        return {"selected_tables": []}284 285    allowed = {_candidate_fqn(c) for c in candidates}286    allowed |= _table_fqns_from_relationships(relationships)287    numbered = _build_numbered_candidates_text(candidates)288    user_content = TABLE_AGENT_PROMPT.format(289        rephrased_question=rq or "(empty)",290        keywords_block=_format_keywords(keywords),291        relationships_block=_format_relationships_block(relationships),292        numbered_candidates=numbered,293    )294 295    messages = [{"role": "user", "content": user_content}]296 297    try:298        response = chat_completion(messages)299    except Exception as e:300        logger.warning("Table Agent LLM call failed: %s", e)301        return {"selected_tables": []}302 303    parsed = _parse_table_agent_response(response)304    if parsed is None:305        logger.warning("Table Agent: could not parse LLM response.")306        return {"selected_tables": []}307 308    selected = _validate_selected_tables(parsed, allowed)309    return {"selected_tables": selected}310