prazy1208/text2sql
0
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 