dgarzon/SQL-Explainer
0
1"""Core logic shared by the CLI and Streamlit versions of SQL Explainer Agent."""2 3from __future__ import annotations4 5import os6import textwrap7from functools import lru_cache8 9import torch10from transformers import AutoModelForCausalLM, AutoTokenizer11 12 13DEFAULT_MODEL = os.getenv("LOCAL_MODEL_ID", "Qwen/Qwen2.5-0.5B-Instruct")14DEFAULT_MAX_NEW_TOKENS = int(os.getenv("MAX_NEW_TOKENS", "320"))15MAX_SQL_CHARS = int(os.getenv("MAX_SQL_CHARS", "8000"))16FREE_LOCAL_MODELS = (17 "Qwen/Qwen2.5-0.5B-Instruct",18 "HuggingFaceTB/SmolLM2-360M-Instruct",19 "Qwen/Qwen2.5-1.5B-Instruct",20)21 22 23SYSTEM_PROMPT = """You are SQL Explainer Agent.24 25Your job is to explain SQL queries clearly for analytics or business users.26 27Follow these rules:28- Use practical, plain English. Avoid academic language.29- Do not invent the meaning of tables, columns, or business logic when it is not obvious.30- If something is unclear, say it clearly.31- If the SQL is incomplete, explain only what is visible.32- Be careful with row count changes and explain realistic causes.33- Keep the structure exactly as requested.34- Use the exact section titles shown below.35- Keep statements grounded in the visible SQL.36- Do not add a subject line, greeting, or sign-off in the email section.37- Do not assume a specific SQL dialect unless it is visible in the query.38- Do not claim case sensitivity, function behavior, or NULL behavior that depends on the SQL engine unless it is clearly visible or you label it as dialect-dependent.39- When discussing join duplicates, be cautious: grouping does not automatically remove duplicate impact on counts or sums.40 41Always answer using these sections and headings:421. Overview432. Tables / Data Sources443. Filters and Logic454. Why results may be lower or higher than expected465. Risks / Things to Check476. Short Email Explanation48 49Section guidance:50- Overview: explain what the query does in plain English.51- Tables / Data Sources: list the main tables, views, or CTEs that are visible in the SQL.52- Filters and Logic: explain joins, WHERE filters, GROUP BY, HAVING, DISTINCT, CASE, and any important transformations.53- Why results may be lower or higher than expected: explicitly include one bullet for each of these topics, even if the answer is "Not visible in this SQL" or "Not applicable based on the visible SQL": INNER JOIN vs LEFT JOIN, filters in WHERE, NULL handling, NOT IN behavior, DISTINCT, aggregations, duplicate rows caused by joins.54- Risks / Things to Check: mention ambiguities, missing conditions, risky assumptions, incomplete SQL, possible duplicate amplification, date logic, and edge cases.55- Short Email Explanation: write exactly one short professional paragraph in English that someone could send to a colleague. No bullets. No subject line. No greeting. No sign-off.56"""57 58 59def build_prompt(sql: str) -> str:60 """Build the final user prompt that combines instructions and the SQL query."""61 return textwrap.dedent(62 f"""\63 Explain the following SQL query for a non-technical analytics or business audience.64 65 Format requirements:66 - Use the six section titles exactly as written in the system instructions.67 - In section 4, include exactly seven bullets, one for each required risk topic.68 - In section 6, write exactly one paragraph and nothing else.69 70 SQL:71 ```sql72 {sql}73 ```74 """75 )76 77 78def validate_sql_input(sql: str) -> str:79 """Validate the SQL payload before sending it to the local model."""80 cleaned_sql = sql.strip()81 if not cleaned_sql:82 raise ValueError("No SQL query was provided.")83 84 if len(cleaned_sql) > MAX_SQL_CHARS:85 raise ValueError(86 f"SQL is too long for this demo ({len(cleaned_sql)} characters). "87 f"Keep it under {MAX_SQL_CHARS} characters."88 )89 90 return cleaned_sql91 92 93def get_device_label() -> str:94 """Return a short label for the active inference device."""95 return "GPU" if torch.cuda.is_available() else "CPU"96 97 98def get_free_local_models() -> tuple[str, ...]:99 """Return the curated list of safe local models for the free Space."""100 if DEFAULT_MODEL in FREE_LOCAL_MODELS:101 return FREE_LOCAL_MODELS102 return (DEFAULT_MODEL, *FREE_LOCAL_MODELS)103 104 105@lru_cache(maxsize=2)106def load_model_components(model_id: str) -> tuple[AutoTokenizer, AutoModelForCausalLM]:107 """Load and cache the tokenizer and model once per process."""108 tokenizer = AutoTokenizer.from_pretrained(model_id, use_fast=True)109 if tokenizer.pad_token_id is None:110 tokenizer.pad_token = tokenizer.eos_token111 112 model = AutoModelForCausalLM.from_pretrained(113 model_id,114 torch_dtype=torch.float32,115 low_cpu_mem_usage=True,116 )117 model.eval()118 119 cpu_threads = max(1, min(4, os.cpu_count() or 1))120 torch.set_num_threads(cpu_threads)121 122 return tokenizer, model123 124 125def call_local_model(126 sql: str,127 model_id: str,128 max_new_tokens: int = DEFAULT_MAX_NEW_TOKENS,129) -> str:130 """Generate the SQL explanation locally with Transformers."""131 sql = validate_sql_input(sql)132 tokenizer, model = load_model_components(model_id)133 134 messages = [135 {"role": "system", "content": SYSTEM_PROMPT},136 {"role": "user", "content": build_prompt(sql)},137 ]138 prompt = tokenizer.apply_chat_template(139 messages,140 tokenize=False,141 add_generation_prompt=True,142 )143 inputs = tokenizer(prompt, return_tensors="pt")144 145 with torch.inference_mode():146 output_ids = model.generate(147 **inputs,148 max_new_tokens=max_new_tokens,149 do_sample=False,150 pad_token_id=tokenizer.pad_token_id,151 eos_token_id=tokenizer.eos_token_id,152 )153 154 generated_ids = output_ids[0][inputs["input_ids"].shape[1] :]155 result = tokenizer.decode(generated_ids, skip_special_tokens=True).strip()156 if not result:157 raise RuntimeError("The local model returned an empty response.")158 159 return result160 