perch / backend /services /executor.py
GerardCB's picture
Perch: eBird Status and Trends, queried in plain language
969891d verified
Raw
History Blame Contribute Delete
51.3 kB
"""
Query Executor Service
Orchestrates the full natural-language-to-map query flow. The pipeline is
structured to minimize LLM round-trips:
1. Intent detection and table selection share a single LLM call.
2. Semantic search runs concurrently with that call.
3. Generated table schemas are cached across queries.
4. Layer naming and explanation generation run in parallel.
"""
from backend.core.llm_gateway import LLMGateway
from backend.core.geo_engine import get_geo_engine
from backend.services.response_formatter import ResponseFormatter
from backend.core.session_store import get_session_store
from backend.core.semantic_search import get_semantic_search
from backend.core.data_catalog import get_data_catalog
from backend.core.query_planner import get_query_planner
from backend.core.jsonutil import dumps_safe
from typing import List, Dict, Any, Optional
import os
import json
import datetime
import asyncio
import logging
logger = logging.getLogger(__name__)
# Session scope. The app is currently single-tenant, so all requests share one
# session; SessionStore is already keyed by id for when that changes.
DEFAULT_SESSION_ID = "default-session"
class QueryExecutor:
def __init__(self):
self.llm = LLMGateway()
self.geo_engine = get_geo_engine()
self.session_store = get_session_store()
self.semantic_search = get_semantic_search()
self.catalog = get_data_catalog()
self.query_planner = get_query_planner()
# Schema cache for optimization
self._schema_cache: Dict[str, str] = {}
self._schema_cache_max_size = 50
def _get_cached_schema(self, tables: List[str]) -> str:
"""Get schema with caching to avoid regeneration."""
cache_key = ",".join(sorted(tables))
if cache_key in self._schema_cache:
return self._schema_cache[cache_key]
# Generate schema
schema = self.geo_engine.get_table_schemas_for_tables(tables)
# Cache with LRU-style eviction
if len(self._schema_cache) >= self._schema_cache_max_size:
# Remove oldest entry
oldest_key = next(iter(self._schema_cache))
del self._schema_cache[oldest_key]
self._schema_cache[cache_key] = schema
return schema
def _build_chat_context(self, catalog_summary: str, query: str) -> str:
"""Wrap a general-chat question with the datasets currently in scope."""
return f"""Available geographic data:
{catalog_summary}
User question: {query}
Respond as Perch, the avian distribution intelligence assistant."""
# =========================================================================
# Main Streaming Entry Point
# =========================================================================
async def process_query_stream(self, query: str, history: List[Dict[str, str]], allowed_datasets: Optional[List[str]] = None):
"""
Stream a query to the user.
Routes through the tool-calling agent, which decides for itself what to
inspect and what to produce. Set PERCH_AGENT=legacy to fall back to the
original fixed pipeline (kept for comparison while the agent beds in).
"""
if os.getenv("PERCH_AGENT", "agent").lower() != "legacy":
async for event in self._process_with_agent(query, history, allowed_datasets):
yield event
return
async for event in self._process_legacy_pipeline(query, history, allowed_datasets):
yield event
async def _process_with_agent(
self, query: str, history: List[Dict[str, str]],
allowed_datasets: Optional[List[str]] = None,
):
"""Run the tool-calling agent, translating its events into the SSE contract."""
from backend.core.agent_loop import GeoAgent
from backend.core.agent_tools import AgentContext
from backend.core.agent_subagents import build_subagent_tool
session_id = DEFAULT_SESSION_ID
if not self.llm.client:
yield {"event": "result", "data": dumps_safe({
"response": "No API key configured, so I cannot answer questions.",
"sql_query": None, "geojson": None,
"data_citations": [], "chart_data": None, "raw_data": [],
})}
return
yield {"event": "status", "data": dumps_safe({"status": "🧠 Working on it..."})}
ctx = AgentContext(allowed_datasets=allowed_datasets or None)
agent = GeoAgent(
self.llm.client, self.llm.model,
extra_tools={"spawn_subagents": build_subagent_tool(
self.llm.client, self.llm.model, ctx
)},
)
answer = ""
pending_question = None
try:
async for event in agent.run(query, history, ctx):
kind = event.get("type")
if kind == "step":
yield {"event": "status", "data": dumps_safe({"status": event["text"]})}
elif kind == "thought":
yield {"event": "chunk", "data": dumps_safe(
{"type": "thought", "content": event["text"]})}
elif kind == "error":
answer = event.get("message", "Something went wrong.")
elif kind == "final":
answer = event.get("text", "")
pending_question = event.get("question")
except Exception as e:
logger.error(f"Agent run failed: {e}", exc_info=True)
answer = f"I hit an error working on that: {e}"
yield {"event": "chunk", "data": dumps_safe({"type": "text", "content": answer})}
# Register every layer so follow-up questions can reference them by name.
agent_layers = [g for g in ctx.layers if g and g.get("features")]
for geo in agent_layers:
try:
layer_id = geo.get("properties", {}).get("layer_id") or "agent"
table_name = self.geo_engine.register_layer(layer_id, geo)
self.session_store.add_layer(session_id, {
"id": layer_id,
"name": geo.get("properties", {}).get("layer_name", "Map Layer"),
"table_name": table_name,
"timestamp": datetime.datetime.now().isoformat(),
})
except Exception as e:
logger.warning(f"Failed to register agent layer: {e}")
geojson = agent_layers[0] if agent_layers else None
combined_sql = "\n\n".join(dict.fromkeys(ctx.sql_statements)) or None
citations = ResponseFormatter.generate_citations(
list(self.catalog.catalog.keys()), combined_sql or ""
) if combined_sql else []
referenced = ResponseFormatter._tables_referenced_in_sql(
combined_sql or "", list(self.catalog.catalog.keys())
)
# add_map_layer already credits each layer with the tables its own query
# read. Only fall back to the union for a layer that arrived without any,
# rather than overwriting accurate per-layer attribution with it.
for geo in agent_layers:
props = geo.setdefault("properties", {})
if not props.get("source_tables"):
props["source_tables"] = referenced
result: Dict[str, Any] = {
"response": answer,
"sql_query": combined_sql,
# `geojson` stays for compatibility; `geojson_layers` carries them all
# so a multi-species answer puts every layer on the map.
"geojson": geojson,
"geojson_layers": agent_layers,
"chart_data": ctx.chart_data,
"raw_data": ctx.raw_data,
"data_citations": citations,
}
if pending_question:
result["pending_question"] = pending_question
yield {"event": "result", "data": dumps_safe(result)}
# Skip follow-ups when we just asked the user something — they should
# answer the question, not be handed three new ones.
if not pending_question:
async for ev in self._emit_followups(query, answer, allowed_datasets, history):
yield ev
async def _process_legacy_pipeline(self, query: str, history: List[Dict[str, str]], allowed_datasets: Optional[List[str]] = None):
"""Original fixed pipeline: intent -> tables -> SQL -> execute -> explain."""
session_id = DEFAULT_SESSION_ID
# =====================================================================
# PHASE 1: Parallel Discovery (Semantic Search + Intent/Tables Detection)
# =====================================================================
yield {"event": "status", "data": dumps_safe({"status": "🧠 Analyzing query..."})}
# Start semantic search in background (doesn't need LLM)
semantic_task = asyncio.create_task(
asyncio.to_thread(
self.semantic_search.search_table_names,
query, 15, allowed_datasets
)
)
# Get catalog summaries for LLM (can use semantic results to filter, but we need some summaries)
if allowed_datasets:
candidate_summaries = self.catalog.get_summaries_for_tables(allowed_datasets)
else:
candidate_summaries = self.catalog.get_all_table_summaries()
# Combined intent + table detection in a single LLM call
detection_result = await self.llm.detect_intent_and_tables(query, candidate_summaries, history)
intent = detection_result["intent"]
llm_selected_tables = detection_result["tables"]
yield {"event": "intent", "data": dumps_safe({"intent": intent})}
logger.info(f"[Perch] Intent: {intent}, Tables: {llm_selected_tables}")
# Get semantic search results
semantic_tables = await semantic_task
# =====================================================================
# PHASE 2: Route by Intent
# =====================================================================
if intent == "GENERAL_CHAT":
# Enhance with context (match non-streaming logic)
# Use semantic search results already fetched
relevant_tables_chat = semantic_tables
# Add user layers
user_layers = self.geo_engine.get_user_layers()
if user_layers:
relevant_tables_chat.extend(user_layers)
# Filter
if allowed_datasets is not None:
relevant_tables_chat = [t for t in relevant_tables_chat if t in allowed_datasets or t in (user_layers or [])]
# Get summary
catalog_summary = "Catalog unavailable."
try:
if relevant_tables_chat:
catalog_summary = self.catalog.get_summaries_for_tables(relevant_tables_chat)
elif allowed_datasets is not None and len(allowed_datasets) == 0:
catalog_summary = "No datasets are currently selected."
else:
catalog_summary = self.catalog.get_all_table_summaries()
except Exception as e:
logger.warning(f"Failed to get catalog summary: {e}")
enhanced_query = self._build_chat_context(catalog_summary, query)
full_response = ""
async for chunk in self.llm.generate_response_stream(enhanced_query, history):
if chunk.get("type") == "content":
text = chunk.get("text", "")
full_response += text
yield {"event": "chunk", "data": dumps_safe({"type": "text", "content": text})}
elif chunk.get("type") == "thought":
yield {"event": "chunk", "data": dumps_safe({"type": "thought", "content": chunk.get("content")})}
yield {"event": "result", "data": dumps_safe({"response": full_response})}
return
# =====================================================================
# PHASE 3: Data/Map/Stat Queries
# =====================================================================
if intent in ["DATA_QUERY", "MAP_REQUEST", "STAT_QUERY"]:
# Always allow a map. Intent decides what to *emphasise*, not what the
# user is allowed to see: suppressing the map for STAT_QUERY meant a
# spatial question phrased as a comparison ("compare the breeding and
# wintering range") returned geometry that was never drawn. Whether a
# map actually appears is decided downstream by _has_geometry(), so a
# non-spatial aggregate still correctly produces charts only.
include_map = True
# Check query complexity for multi-step execution
complexity = self.query_planner.detect_complexity(query)
if complexity["is_complex"]:
yield {"event": "status", "data": dumps_safe({"status": "🔄 Complex query detected, planning steps..."})}
async for event in self._execute_multi_step_query(query, history, include_map, session_id, allowed_datasets):
yield event
return
# Simple query flow
async for event in self._handle_data_query_stream(
query, history, intent, include_map, session_id,
llm_selected_tables, semantic_tables, allowed_datasets
):
yield event
return
# =====================================================================
# PHASE 4: Spatial Operations
# =====================================================================
if intent == "SPATIAL_OP":
async for event in self._handle_spatial_op_stream(
query, history, session_id, llm_selected_tables, semantic_tables, allowed_datasets
):
yield event
return
# Fallback
yield {"event": "chunk", "data": dumps_safe({"type": "text", "content": "I'm not sure how to handle this query."})}
yield {"event": "result", "data": dumps_safe({"response": ""})}
# =========================================================================
# Data Query Handler (Streaming)
# =========================================================================
async def _handle_data_query_stream(
self,
query: str,
history: List[Dict[str, str]],
intent: str,
include_map: bool,
session_id: str,
llm_selected_tables: List[str],
semantic_tables: List[str],
allowed_datasets: Optional[List[str]]
):
"""
Optimized data query handling with parallel operations.
"""
# Merge LLM-selected and semantic tables, prioritizing LLM selection
relevant_tables = list(set(llm_selected_tables + semantic_tables[:5]))
# Add user layers
user_layers = self.geo_engine.get_user_layers()
if user_layers:
relevant_tables.extend(user_layers)
relevant_tables = list(set(relevant_tables))
# Filter by allowed_datasets
if allowed_datasets:
relevant_tables = [t for t in relevant_tables if t in allowed_datasets or t in user_layers]
# Load tables
if relevant_tables:
yield {"event": "status", "data": dumps_safe({"status": f"💾 Loading {len(relevant_tables)} tables..."})}
feature_tables = []
for table in relevant_tables:
if self.geo_engine.ensure_table_loaded(table):
feature_tables.append(table)
# Get schema (cached)
table_schema = self._get_cached_schema(feature_tables) if feature_tables else self.geo_engine.get_table_schemas()
# Generate SQL (streaming with thoughts)
yield {"event": "status", "data": dumps_safe({"status": "✍️ Writing SQL query..."})}
sql_buffer = ""
async for chunk in self.llm.stream_analytical_sql(query, table_schema, history):
if chunk["type"] == "thought":
yield {"event": "chunk", "data": dumps_safe({"type": "thought", "content": chunk["text"]})}
elif chunk["type"] == "content":
sql_buffer += chunk["text"]
sql = sql_buffer.replace("```sql", "").replace("```", "").strip()
logger.info(f"Generated SQL:\n{sql}")
# Check for DATA_UNAVAILABLE
if "DATA_UNAVAILABLE" in sql or sql.startswith("-- ERROR"):
yield {"event": "status", "data": dumps_safe({"status": "ℹ️ Data not available"})}
error_response = self._format_data_unavailable_response(sql)
yield {"event": "result", "data": dumps_safe({
"response": error_response,
"sql_query": sql,
"geojson": None,
"data_citations": [],
"chart_data": None,
"raw_data": []
})}
return
# Execute query
yield {"event": "status", "data": dumps_safe({"status": "⚡ Executing query..."})}
geojson, features, error_message = await self._execute_sql_with_retry(
sql, query, table_schema
)
if error_message:
yield {"event": "result", "data": dumps_safe({
"response": f"Query failed: {error_message}",
"sql_query": sql,
"geojson": None,
"data_citations": [],
"chart_data": None,
"raw_data": []
})}
return
yield {"event": "status", "data": dumps_safe({"status": f"✅ Found {len(features)} results"})}
# =====================================================================
# Parallel post-processing: layer name + explanation
# =====================================================================
yield {"event": "status", "data": dumps_safe({"status": "💬 Generating response..."})}
# Prepare data summary for explanation
data_summary = ResponseFormatter.generate_data_summary(features)
citations = ResponseFormatter.generate_citations(relevant_tables, sql)
# Record the source tables on the layer itself so the map popup can credit
# the datasets a feature came from.
if geojson is not None:
geojson.setdefault("properties", {})["source_tables"] = (
ResponseFormatter._tables_referenced_in_sql(sql, relevant_tables)
)
raw_data = ResponseFormatter.prepare_raw_data(features)
# Start parallel tasks. A map is only possible when the result actually
# carries geometry — a non-spatial table would produce an empty layer.
mappable = include_map and bool(features) and bool(geojson) and self._has_geometry(features)
if include_map and features and not mappable:
logger.info("Result has no geometry; returning stats only (no map layer).")
layer_task = None
if mappable:
layer_task = asyncio.create_task(self.llm.generate_layer_name(query, sql))
explanation_task = asyncio.create_task(
self.llm.generate_explanation(query, sql, data_summary, history, map_rendered=mappable)
)
# Wait for explanation (this is what we stream to user)
explanation_result = await explanation_task
explanation_text = explanation_result.get("explanation", "")
chart_config = explanation_result.get("chart_config")
# Stream explanation to user
yield {"event": "chunk", "data": dumps_safe({"type": "text", "content": explanation_text})}
# Wait for layer name and process map
if layer_task:
layer_info = await layer_task
layer_name_ai = layer_info.get("name", "Map Layer")
layer_emoji = layer_info.get("emoji", "📍")
point_style = layer_info.get("pointStyle")
color_by = layer_info.get("colorBy")
geojson, layer_id, layer_name = ResponseFormatter.format_geojson_layer(
query, geojson, features, layer_name_ai, layer_emoji, point_style, color_by=color_by
)
try:
table_name = self.geo_engine.register_layer(layer_id, geojson)
self.session_store.add_layer(session_id, {
"id": layer_id,
"name": layer_name,
"table_name": table_name,
"timestamp": datetime.datetime.now().isoformat()
})
except Exception as e:
logger.warning(f"Failed to register layer: {e}")
# Generate chart data
chart_data = ResponseFormatter.generate_chart_data(sql, features, query, chart_config)
if intent == "STAT_QUERY" and not chart_data and features:
chart_data = ResponseFormatter.generate_chart_data("GROUP BY forced", features, query, chart_config)
# Final result
yield {"event": "result", "data": dumps_safe({
"response": explanation_text,
"sql_query": sql,
"geojson": geojson if mappable else None,
"chart_data": chart_data,
"raw_data": raw_data,
"data_citations": citations
})}
# Best-effort follow-up suggestions (emitted after the result so they never
# delay the answer, and a failure here cannot break the response).
async for ev in self._emit_followups(query, explanation_text, allowed_datasets, history):
yield ev
# =========================================================================
# Spatial Operations Handler (Streaming)
# =========================================================================
async def _handle_spatial_op_stream(
self,
query: str,
history: List[Dict[str, str]],
session_id: str,
llm_selected_tables: List[str],
semantic_tables: List[str],
allowed_datasets: Optional[List[str]]
):
"""Handle spatial operations with optimized flow."""
yield {"event": "status", "data": dumps_safe({"status": "📐 Preparing spatial operation..."})}
# Merge tables
relevant_tables = list(set(llm_selected_tables + semantic_tables[:5]))
logger.info(
f"Spatial op table selection - LLM: {llm_selected_tables}, "
f"semantic: {semantic_tables[:5]}, merged: {relevant_tables}"
)
# Filter by allowed_datasets
if allowed_datasets:
relevant_tables = [t for t in relevant_tables if t in allowed_datasets]
# Load tables - add status message and debugging
if relevant_tables:
yield {"event": "status", "data": dumps_safe({"status": f"💾 Loading {len(relevant_tables)} tables..."})}
for table in relevant_tables:
loaded = self.geo_engine.ensure_table_loaded(table)
logger.info(f"Loaded table {table}: {loaded}")
else:
logger.warning("No relevant tables identified for spatial operation.")
# Get schema
base_table_schema = self._get_cached_schema(relevant_tables) if relevant_tables else self.geo_engine.get_table_schemas()
# Get session layers context
session_layers = self.session_store.get_layers(session_id)
user_layer_schemas = ""
if session_layers:
user_layer_names = [layer['table_name'] for layer in session_layers]
user_layer_schemas = self.geo_engine.get_table_schemas_for_tables(user_layer_names)
full_context = f"{base_table_schema}\n\n{user_layer_schemas}"
# Generate spatial SQL
yield {"event": "status", "data": dumps_safe({"status": "✍️ Writing spatial SQL..."})}
sql = await self.llm.generate_spatial_sql(query, full_context, history)
logger.info(f"Generated spatial SQL:\n{sql}")
# Execute
yield {"event": "status", "data": dumps_safe({"status": "⚙️ Processing geometry..."})}
geojson, features, error_message = await self._execute_sql_with_retry(
sql, query, full_context
)
if error_message:
yield {"event": "result", "data": dumps_safe({
"response": f"Spatial operation failed: {error_message}",
"sql_query": sql,
"geojson": None,
"data_citations": [],
"chart_data": None,
"raw_data": []
})}
return
yield {"event": "status", "data": dumps_safe({"status": f"✅ Result: {len(features)} features"})}
# Parallel post-processing
yield {"event": "status", "data": dumps_safe({"status": "💬 Generating response..."})}
layer_task = asyncio.create_task(self.llm.generate_layer_name(query, sql)) if features else None
data_summary = f"Spatial operation resulted in {len(features)} features."
explanation_task = asyncio.create_task(
self.llm.generate_explanation(query, sql, data_summary, history)
)
explanation_result = await explanation_task
explanation_text = explanation_result.get("explanation", "")
chart_config = explanation_result.get("chart_config")
yield {"event": "chunk", "data": dumps_safe({"type": "text", "content": explanation_text})}
if layer_task and features and geojson:
layer_info = await layer_task
geojson, layer_id, layer_name = ResponseFormatter.format_geojson_layer(
query, geojson, features,
layer_info.get("name", "Spatial Result"),
layer_info.get("emoji", "📐"),
layer_info.get("pointStyle"),
color_by=layer_info.get("colorBy")
)
try:
table_name = self.geo_engine.register_layer(layer_id, geojson)
self.session_store.add_layer(session_id, {
"id": layer_id,
"name": layer_name,
"table_name": table_name,
"timestamp": datetime.datetime.now().isoformat()
})
except Exception as e:
logger.warning(f"Failed to register spatial layer: {e}")
yield {"event": "result", "data": dumps_safe({
"response": explanation_text,
"sql_query": sql,
"geojson": geojson,
"chart_data": ResponseFormatter.generate_chart_data(sql, features, query, chart_config),
"raw_data": [],
"data_citations": []
})}
# =========================================================================
# Helper Methods
# =========================================================================
async def _emit_followups(
self, query: str, answer: str, allowed_datasets: Optional[List[str]],
history: Optional[List[Dict[str, str]]] = None,
):
"""Yield a 'suggestions' SSE event with follow-up questions, if any."""
try:
if allowed_datasets:
summary = self.catalog.get_summaries_for_tables(allowed_datasets)
else:
summary = self.catalog.get_all_table_summaries()
suggestions = await self.llm.suggest_followups(query, summary, answer, history)
if suggestions:
yield {"event": "suggestions", "data": dumps_safe({"suggestions": suggestions})}
except Exception as e:
logger.info(f"Follow-up emission skipped: {e}")
@staticmethod
def _has_geometry(features: List[Dict[str, Any]]) -> bool:
"""
True if any feature carries a geometry.
Non-spatial tables (attribute-only) produce rows whose geometry is None.
Registering those as a map layer creates a layer that renders nothing, so
callers use this to decide whether a map is even possible.
"""
return any(f.get("geometry") for f in (features or []))
async def _execute_sql_with_retry(self, sql: str, query: str, schema_context: str) -> tuple:
"""Execute SQL with one retry on failure."""
geojson = None
features = []
error_message = None
try:
geojson = self.geo_engine.execute_spatial_query(sql)
features = geojson.get("features", [])
except Exception as e:
error_message = str(e)
logger.warning(f"SQL execution error: {error_message}")
# Try to correct
try:
corrected_sql = await self.llm.correct_sql(query, sql, error_message, schema_context)
geojson = self.geo_engine.execute_spatial_query(corrected_sql)
features = geojson.get("features", [])
error_message = None
except Exception as e2:
error_message = f"Original: {error_message}, Correction failed: {str(e2)}"
return geojson, features, error_message
def _format_data_unavailable_response(self, sql: str) -> str:
"""Format a user-friendly response when data is unavailable."""
requested = "the requested data"
available = ""
for line in sql.split("\n"):
if "Requested:" in line:
requested = line.split("Requested:")[-1].strip()
elif "Available:" in line:
available = line.split("Available:")[-1].strip()
# The model echoes back whatever was in its schema context, which includes
# session layers (layer_ab12cd34). Those are transient results of earlier
# questions, so listing them as "available datasets" is just noise.
catalog_names = set(self.catalog.catalog.keys())
listed = [
t for t in (n.strip(" `") for n in available.replace(",", " ").split())
if t in catalog_names
]
if not listed:
listed = sorted(catalog_names)
shown = ", ".join(sorted(listed)[:12])
more = f" (+{len(listed) - 12} more)" if len(listed) > 12 else ""
return f"""I couldn't find data for **{requested}** in the current database.
**Available datasets include:** {shown}{more}
Try rephrasing, or ask "what data do you have?" for the full list."""
# =========================================================================
# Multi-Step Query Execution
# =========================================================================
async def _execute_multi_step_query(
self,
query: str,
history: List[Dict[str, str]],
include_map: bool,
session_id: str,
allowed_datasets: Optional[List[str]] = None
):
"""Execute complex queries by breaking into steps."""
# Get candidate tables
yield {"event": "status", "data": dumps_safe({"status": "📚 Discovering relevant datasets..."})}
candidate_tables = self.semantic_search.search_table_names(query, top_k=20, allowed_datasets=allowed_datasets)
if not candidate_tables and allowed_datasets is None:
candidate_tables = list(self.catalog.catalog.keys())
# Plan the query
yield {"event": "status", "data": dumps_safe({"status": "📋 Creating execution plan..."})}
plan = await self.query_planner.plan_query(query, candidate_tables, self.llm)
if not plan.is_complex or not plan.steps:
# Fallback to simple execution
yield {"event": "status", "data": dumps_safe({"status": "📚 Executing as simple query..."})}
candidate_summaries = self.catalog.get_summaries_for_tables(candidate_tables) if candidate_tables else self.catalog.get_summaries_for_tables(allowed_datasets or [])
relevant_tables = await self.llm.identify_relevant_tables(query, candidate_summaries)
if allowed_datasets:
relevant_tables = [t for t in relevant_tables if t in allowed_datasets]
for table in relevant_tables:
self.geo_engine.ensure_table_loaded(table)
table_schema = self._get_cached_schema(relevant_tables) if relevant_tables else self.geo_engine.get_table_schemas()
yield {"event": "status", "data": dumps_safe({"status": "✍️ Writing SQL query..."})}
sql = await self.llm.generate_analytical_sql(query, table_schema, history)
sql = sql.replace("```sql", "").replace("```", "").strip()
geojson, features, error_message = await self._execute_sql_with_retry(sql, query, table_schema)
if error_message:
yield {"event": "result", "data": dumps_safe({
"response": f"Query execution failed: {error_message}",
"sql_query": sql
})}
return
data_summary = ResponseFormatter.generate_data_summary(features)
explanation_result = await self.llm.generate_explanation(query, sql, data_summary, history)
explanation_text = explanation_result.get("explanation", "")
chart_config = explanation_result.get("chart_config")
yield {"event": "result", "data": dumps_safe({
"response": explanation_text,
"sql_query": sql,
"geojson": geojson if include_map and features else None,
"chart_data": ResponseFormatter.generate_chart_data(sql, features, query, chart_config),
"raw_data": ResponseFormatter.prepare_raw_data(features),
"data_citations": []
})}
return
# Show plan
step_descriptions = [f"Step {i+1}: {s.description}" for i, s in enumerate(plan.steps)]
yield {"event": "chunk", "data": dumps_safe({
"type": "thought",
"content": f"Planning multi-step execution:\n" + "\n".join(step_descriptions)
})}
# Load all needed tables
all_tables = set()
for step in plan.steps:
all_tables.update(step.tables_needed)
if all_tables:
yield {"event": "status", "data": dumps_safe({"status": f"💾 Loading {len(all_tables)} datasets..."})}
for table in all_tables:
self.geo_engine.ensure_table_loaded(table)
# Execute steps
intermediate_results = {}
all_features = []
all_sql = []
for group_idx, group in enumerate(plan.parallel_groups):
group_steps = [s for s in plan.steps if s.step_id in group]
yield {"event": "status", "data": dumps_safe({
"status": f"⚡ Executing step group {group_idx + 1}/{len(plan.parallel_groups)}..."
})}
for step in group_steps:
yield {"event": "status", "data": dumps_safe({"status": f"🔄 {step.description}..."})}
table_schema = self._get_cached_schema(list(all_tables)) if all_tables else self.geo_engine.get_table_schemas()
step_query = f"""Execute this step: {step.description}
Original user request: {query}
SQL Hint: {step.sql_template or 'None'}
Previous step results: {list(intermediate_results.keys())}"""
sql = await self.llm.generate_analytical_sql(step_query, table_schema, history)
sql = sql.replace("```sql", "").replace("```", "").strip()
if "DATA_UNAVAILABLE" in sql or sql.startswith("-- ERROR"):
intermediate_results[step.result_name] = {"features": [], "sql": sql}
continue
try:
geojson = self.geo_engine.execute_spatial_query(sql)
features = geojson.get("features", [])
intermediate_results[step.result_name] = {
"features": features,
"sql": sql,
"geojson": geojson
}
all_features.extend(features)
all_sql.append(f"-- {step.description}\n{sql}")
yield {"event": "status", "data": dumps_safe({"status": f"✅ Step got {len(features)} results"})}
except Exception as e:
logger.error(f"Step {step.step_id} failed: {e}")
try:
sql = await self.llm.correct_sql(step_query, sql, str(e), table_schema)
geojson = self.geo_engine.execute_spatial_query(sql)
features = geojson.get("features", [])
intermediate_results[step.result_name] = {
"features": features,
"sql": sql,
"geojson": geojson
}
all_features.extend(features)
all_sql.append(f"-- {step.description} (repaired)\n{sql}")
except Exception as e2:
intermediate_results[step.result_name] = {"features": [], "sql": sql, "error": str(e2)}
# Generate combined result
yield {"event": "status", "data": dumps_safe({"status": "💬 Generating combined analysis..."})}
result_summary = [f"{name}: {len(r.get('features', []))} records" for name, r in intermediate_results.items()]
combined_summary = f"Multi-step query completed.\nResults: {', '.join(result_summary)}\nCombination: {plan.final_combination_logic}"
explanation_buffer = ""
async for chunk in self.llm.stream_explanation(query, "\n\n".join(all_sql), combined_summary, history):
if chunk["type"] == "content":
explanation_buffer += chunk["text"]
yield {"event": "chunk", "data": dumps_safe({"type": "text", "content": chunk["text"]})}
# Find best geojson
best_geojson = None
best_features = []
for result in intermediate_results.values():
features = result.get("features", [])
if len(features) > len(best_features):
best_features = features
best_geojson = result.get("geojson")
# Generate layer
if include_map and best_features and best_geojson:
layer_info = await self.llm.generate_layer_name(query, all_sql[0] if all_sql else "")
best_geojson, layer_id, layer_name = ResponseFormatter.format_geojson_layer(
query, best_geojson, best_features,
layer_info.get("name", "Multi-Step Result"),
layer_info.get("emoji", "📊"),
layer_info.get("pointStyle"),
color_by=layer_info.get("colorBy")
)
try:
table_name = self.geo_engine.register_layer(layer_id, best_geojson)
self.session_store.add_layer(session_id, {
"id": layer_id,
"name": layer_name,
"table_name": table_name,
"timestamp": datetime.datetime.now().isoformat()
})
except Exception as e:
logger.warning(f"Failed to register multi-step layer: {e}")
yield {"event": "result", "data": dumps_safe({
"response": explanation_buffer,
"sql_query": "\n\n".join(all_sql),
"geojson": best_geojson if include_map and best_features else None,
"chart_data": ResponseFormatter.generate_chart_data("\n".join(all_sql), best_features, query),
"raw_data": ResponseFormatter.prepare_raw_data(best_features),
"data_citations": [],
"multi_step": True,
"steps_executed": len(plan.steps)
})}
# =========================================================================
# Non-Streaming Methods (Backward Compatibility)
# =========================================================================
async def process_query_with_context(self, query: str, history: List[Dict[str, str]], allowed_datasets: Optional[List[str]] = None) -> Dict[str, Any]:
"""Non-streaming query processing."""
intent = await self.llm.detect_intent(query, history)
if intent == "GENERAL_CHAT":
return await self._handle_general_chat(query, history, allowed_datasets)
elif intent in ["DATA_QUERY", "MAP_REQUEST"]:
return await self._handle_data_query(query, history, include_map=True, allowed_datasets=allowed_datasets)
elif intent == "SPATIAL_OP":
return await self._handle_spatial_op(query, history, allowed_datasets)
elif intent == "STAT_QUERY":
return await self._handle_stat_query(query, history, allowed_datasets)
else:
return await self._handle_general_chat(query, history, allowed_datasets)
async def _handle_general_chat(self, query: str, history: List[Dict[str, str]], allowed_datasets: Optional[List[str]] = None) -> Dict[str, Any]:
"""Handle general chat queries."""
try:
relevant_tables = self.semantic_search.search_table_names(query)
user_layers = self.geo_engine.get_user_layers()
if user_layers:
relevant_tables.extend(user_layers)
if allowed_datasets is not None:
relevant_tables = [t for t in relevant_tables if t in allowed_datasets or t in (user_layers or [])]
if relevant_tables:
catalog_summary = self.catalog.get_summaries_for_tables(relevant_tables)
elif allowed_datasets is not None and len(allowed_datasets) == 0:
catalog_summary = "No datasets are currently selected."
else:
catalog_summary = self.catalog.get_all_table_summaries()
except Exception as e:
logger.warning(f"Failed to get catalog summary: {e}")
catalog_summary = "Catalog unavailable."
enhanced_query = self._build_chat_context(catalog_summary, query)
response = await self.llm.generate_response(enhanced_query, history)
return {
"response": response,
"sql_query": None,
"geojson": None,
"data_citations": [],
"intent": "GENERAL_CHAT"
}
async def _handle_data_query(self, query: str, history: List[Dict[str, str]], include_map: bool = True, allowed_datasets: Optional[List[str]] = None) -> Dict[str, Any]:
"""Handle data queries (non-streaming)."""
if allowed_datasets is not None and len(allowed_datasets) == 0:
return {
"response": "No datasets selected. Please select at least one dataset.",
"sql_query": None,
"geojson": None,
"data_citations": [],
"intent": "DATA_QUERY"
}
# Get summaries
if allowed_datasets is not None:
summaries = self.catalog.get_summaries_for_tables(allowed_datasets)
else:
summaries = self.catalog.get_all_table_summaries()
relevant_tables = await self.llm.identify_relevant_tables(query, summaries)
if allowed_datasets:
relevant_tables = [t for t in relevant_tables if t in allowed_datasets]
# Load tables
feature_tables = []
for table in relevant_tables:
if self.geo_engine.ensure_table_loaded(table):
feature_tables.append(table)
table_schema = self._get_cached_schema(feature_tables) if feature_tables else self.geo_engine.get_table_schemas()
# Generate SQL
sql = await self.llm.generate_analytical_sql(query, table_schema, history)
if sql.startswith("-- Error"):
return {
"response": f"Could not generate query for: {query}",
"sql_query": sql,
"intent": "DATA_QUERY"
}
# Execute
geojson, features, error_message = await self._execute_sql_with_retry(sql, query, table_schema)
if error_message:
return {
"response": f"Query failed: {error_message}",
"sql_query": sql,
"intent": "DATA_QUERY"
}
# Post-process
citations = ResponseFormatter.generate_citations(relevant_tables, sql)
data_summary = ResponseFormatter.generate_data_summary(features)
explanation_result = await self.llm.generate_explanation(query, sql, data_summary, history)
explanation = explanation_result.get("explanation", "")
chart_config = explanation_result.get("chart_config")
if include_map and features:
layer_info = await self.llm.generate_layer_name(query, sql)
geojson, layer_id, layer_name = ResponseFormatter.format_geojson_layer(
query, geojson, features,
layer_info.get("name", "Map Layer"),
layer_info.get("emoji", "📍"),
layer_info.get("pointStyle"),
color_by=layer_info.get("colorBy")
)
try:
table_name = self.geo_engine.register_layer(layer_id, geojson)
self.session_store.add_layer(DEFAULT_SESSION_ID, {
"id": layer_id,
"name": layer_name,
"table_name": table_name,
"timestamp": datetime.datetime.now().isoformat()
})
except Exception as e:
logger.warning(f"Failed to register layer: {e}")
chart_data = ResponseFormatter.generate_chart_data(sql, features, query, chart_config)
raw_data = ResponseFormatter.prepare_raw_data(features)
return {
"response": explanation,
"sql_query": sql,
"geojson": geojson if include_map and features else None,
"data_citations": citations,
"chart_data": chart_data,
"raw_data": raw_data,
"intent": "DATA_QUERY" if not include_map else "MAP_REQUEST"
}
async def _handle_spatial_op(self, query: str, history: List[Dict[str, str]], allowed_datasets: Optional[List[str]] = None) -> Dict[str, Any]:
"""Handle spatial operations (non-streaming)."""
if allowed_datasets is not None and len(allowed_datasets) == 0:
return {
"response": "No datasets selected for spatial operations.",
"sql_query": None,
"geojson": None,
"data_citations": [],
"intent": "SPATIAL_OP"
}
if allowed_datasets is not None:
summaries = self.catalog.get_summaries_for_tables(allowed_datasets)
else:
summaries = self.catalog.get_all_table_summaries()
relevant_tables = await self.llm.identify_relevant_tables(query, summaries)
if allowed_datasets:
relevant_tables = [t for t in relevant_tables if t in allowed_datasets]
for table in relevant_tables:
self.geo_engine.ensure_table_loaded(table)
base_table_schema = self._get_cached_schema(relevant_tables) if relevant_tables else self.geo_engine.get_table_schemas()
session_layers = self.session_store.get_layers(DEFAULT_SESSION_ID)
user_layer_schemas = ""
if session_layers:
user_layer_schemas = "### User-Created Layers:\n"
for layer in session_layers:
user_layer_schemas += f"### Table: {layer['table_name']} ('{layer['name']}')\n"
user_layer_schemas += f"Columns: geom GEOMETRY, name TEXT\n\n"
full_context = f"{base_table_schema}\n\n{user_layer_schemas}"
sql = await self.llm.generate_spatial_sql(query, full_context, history)
geojson, features, error_message = await self._execute_sql_with_retry(sql, query, full_context)
if error_message:
return {
"response": f"Spatial operation failed: {error_message}",
"sql_query": sql,
"intent": "SPATIAL_OP"
}
if features:
layer_info = await self.llm.generate_layer_name(query, sql)
geojson, layer_id, layer_name = ResponseFormatter.format_geojson_layer(
query, geojson, features,
layer_info.get("name", "Spatial Result"),
layer_info.get("emoji", "📐"),
layer_info.get("pointStyle"),
color_by=layer_info.get("colorBy")
)
table_name = self.geo_engine.register_layer(layer_id, geojson)
self.session_store.add_layer(DEFAULT_SESSION_ID, {
"id": layer_id,
"name": layer_name,
"table_name": table_name,
"timestamp": datetime.datetime.now().isoformat()
})
data_summary = f"Spatial operation resulted in {len(features)} features."
explanation_result = await self.llm.generate_explanation(query, sql, data_summary, history)
explanation = explanation_result.get("explanation", "")
return {
"response": explanation,
"sql_query": sql,
"geojson": geojson,
"data_citations": [],
"intent": "SPATIAL_OP"
}
async def _handle_stat_query(self, query: str, history: List[Dict[str, str]], allowed_datasets: Optional[List[str]] = None) -> Dict[str, Any]:
"""Handle statistical queries."""
result = await self._handle_data_query(query, history, include_map=False, allowed_datasets=allowed_datasets)
result["intent"] = "STAT_QUERY"
if not result.get("chart_data") and result.get("raw_data"):
features_mock = [{"properties": d} for d in result["raw_data"]]
result["chart_data"] = ResponseFormatter.generate_chart_data(result.get("sql_query", ""), features_mock, query)
return result