Enterprise-AI-Copilot / agents /sql_agent.py
Punit1's picture
Initial commit
939c0c0
Raw
History Blame Contribute Delete
6.72 kB
"""
SQL Intelligence Specialist Agent Node
=======================================
Translates natural language queries into executable SQL commands.
Enforces strict security guardrails (SELECT only; blocks DDL/DML mutations).
Formats tabular SQL results for consumption by the Synthesizer node.
"""
from __future__ import annotations
import re
import structlog
from typing import Any, Dict, List, Tuple
from sqlalchemy import inspect, text
from app.core.config import settings
from app.core.database import AsyncSessionLocal, engine
from app.services.llm_gateway import llm_gateway
from agents.state import CopilotState
logger = structlog.get_logger(__name__)
# Forbidden SQL keywords that alter database state
FORBIDDEN_KEYWORDS = {
"INSERT", "UPDATE", "DELETE", "DROP", "ALTER", "TRUNCATE",
"CREATE", "GRANT", "REVOKE", "EXEC", "EXECUTE", "MERGE"
}
TEXT_TO_SQL_PROMPT = """You are a Database Expert AI Agent.
Your job is to translate the user's natural language request into a valid, read-only SQL query.
Database Dialect: {dialect}
Database Schema:
{schema_description}
Rules:
1. Generate ONLY a single SELECT query. Never attempt to INSERT, UPDATE, DELETE, DROP, or ALTER.
2. Return ONLY the raw SQL statement inside ```sql ... ``` block or plain text without explanations.
3. Use proper joins and table aliases if referencing multiple tables.
4. Restrict query results with LIMIT 50 if returning many rows.
"""
def is_safe_sql(query_sql: str) -> Tuple[bool, str]:
"""
Validate that the generated SQL is strictly read-only (SELECT).
"""
cleaned = re.sub(r"--.*?\n|/\*.*?\*/", "", query_sql, flags=re.DOTALL).strip()
uppercase_words = set(re.findall(r"\b[A-Z]+\b", cleaned.upper()))
forbidden_found = uppercase_words.intersection(FORBIDDEN_KEYWORDS)
if forbidden_found:
return False, f"Security Violation: Query contains forbidden statement(s) {forbidden_found}"
if not cleaned.upper().startswith("SELECT") and not cleaned.upper().startswith("WITH"):
return False, "Security Violation: Only SELECT queries are permitted"
return True, "Safe"
async def get_database_schema() -> str:
"""Extract schema tables and columns for LLM prompt context."""
schema_info = []
try:
async with engine.connect() as conn:
def _inspect_schema(connection):
inspector = inspect(connection)
tables = inspector.get_table_names()
res = []
for table in tables[:10]: # inspect first 10 tables
columns = inspector.get_columns(table)
col_str = ", ".join([f"{c['name']} ({c['type']})" for c in columns])
res.append(f"Table '{table}': {col_str}")
return "\n".join(res)
schema_info = await conn.run_sync(_inspect_schema)
except Exception as e:
logger.warning("Could not introspect database schema", error=str(e))
schema_info = "Tables: users (id, email, full_name, role), tenants (id, name, slug), documents (id, filename, doc_type, status, chunk_count)"
return schema_info or "No tables found"
async def execute_sql_query(query_sql: str) -> Tuple[List[str], List[Dict[str, Any]]]:
"""Execute safe read-only SQL query and return (columns, rows)."""
async with AsyncSessionLocal() as session:
result = await session.execute(text(query_sql))
columns = list(result.keys())
rows = [dict(zip(columns, row)) for row in result.fetchall()]
return columns, rows
async def sql_node(state: CopilotState) -> CopilotState:
"""
SQL Specialist Node.
Generates, validates, and executes read-only SQL.
"""
query = state.get("query", "")
logger.info("SQL Agent executing", query=query)
try:
# 1. Introspect Schema
schema_str = await get_database_schema()
dialect = "SQLite" if settings.use_sqlite else "PostgreSQL"
# 2. Text-to-SQL LLM Prompt
prompt = f"User Request: {query}\n\nGenerated SQL:"
system_prompt = TEXT_TO_SQL_PROMPT.format(dialect=dialect, schema_description=schema_str)
response = await llm_gateway.generate(
prompt=prompt,
system_prompt=system_prompt,
temperature=0.0,
)
sql_text = response.content.strip()
# Extract SQL from markdown code block if present
if "```" in sql_text:
sql_text = sql_text.split("```")[1]
if sql_text.lower().startswith("sql"):
sql_text = sql_text[3:].strip()
logger.info("Generated SQL statement", sql=sql_text)
# 3. Security Validation
safe, reason = is_safe_sql(sql_text)
if not safe:
logger.warning("SQL execution blocked by security guardrail", reason=reason)
state["error"] = reason
state["retrieved_chunks"] = [{
"document_name": "SQL Guardrail",
"text": f"⚠️ Query blocked: {reason}. Only read-only SELECT queries are allowed."
}]
return state
# 4. Execute Query
columns, rows = await execute_sql_query(sql_text)
logger.info("SQL Query executed successfully", row_count=len(rows))
# Format tabular output into markdown text for Synthesizer
formatted_table = ""
if rows:
headers = " | ".join(columns)
formatted_table += f"| {headers} |\n"
formatted_table += f"| {' | '.join(['---'] * len(columns))} |\n"
for row in rows[:20]:
formatted_table += f"| {' | '.join(str(row.get(c, '')) for c in columns)} |\n"
else:
formatted_table = "Query returned 0 rows."
# Update state
state["retrieved_chunks"] = [{
"document_id": "sql_result",
"document_name": f"SQL Query ({sql_text})",
"text": f"Executed SQL: `{sql_text}`\n\nResults:\n{formatted_table}",
"score": 1.0,
"doc_type": "sql"
}]
outputs = state.get("agent_outputs", [])
outputs.append({
"agent_name": "sql",
"content": f"Executed SQL query: `{sql_text}`. Returned {len(rows)} rows.",
"metadata": {"sql": sql_text, "row_count": len(rows)}
})
state["agent_outputs"] = outputs
state["active_agent"] = "sql"
except Exception as e:
logger.error("SQL Agent execution error", error=str(e))
state["error"] = f"SQL Error: {str(e)}"
state["retrieved_chunks"] = [{
"document_name": "SQL Error",
"text": f"Failed to execute SQL query: {str(e)}"
}]
return state