| """ |
| 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__) |
|
|
| |
| |
| 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() |
| |
| |
| 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] |
| |
| |
| schema = self.geo_engine.get_table_schemas_for_tables(tables) |
| |
| |
| if len(self._schema_cache) >= self._schema_cache_max_size: |
| |
| 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.""" |
|
|
| |
| |
| |
|
|
| 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})} |
|
|
| |
| 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()) |
| ) |
| |
| |
| |
| 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": 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)} |
|
|
| |
| |
| 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 |
| |
| |
| |
| |
| yield {"event": "status", "data": dumps_safe({"status": "🧠 Analyzing query..."})} |
| |
| |
| semantic_task = asyncio.create_task( |
| asyncio.to_thread( |
| self.semantic_search.search_table_names, |
| query, 15, allowed_datasets |
| ) |
| ) |
| |
| |
| if allowed_datasets: |
| candidate_summaries = self.catalog.get_summaries_for_tables(allowed_datasets) |
| else: |
| candidate_summaries = self.catalog.get_all_table_summaries() |
| |
| |
| 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}") |
|
|
| |
| semantic_tables = await semantic_task |
| |
| |
| |
| |
| |
| if intent == "GENERAL_CHAT": |
| |
| |
| relevant_tables_chat = semantic_tables |
| |
| |
| user_layers = self.geo_engine.get_user_layers() |
| if user_layers: |
| relevant_tables_chat.extend(user_layers) |
| |
| |
| 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 [])] |
| |
| |
| 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 |
| |
| |
| |
| |
| |
| if intent in ["DATA_QUERY", "MAP_REQUEST", "STAT_QUERY"]: |
| |
| |
| |
| |
| |
| |
| include_map = True |
|
|
|
|
| |
| 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 |
| |
| |
| 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 |
| |
| |
| |
| |
| |
| 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 |
| |
| |
| 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": ""})} |
|
|
| |
| |
| |
|
|
| 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. |
| """ |
| |
| relevant_tables = list(set(llm_selected_tables + semantic_tables[:5])) |
| |
| |
| user_layers = self.geo_engine.get_user_layers() |
| if user_layers: |
| relevant_tables.extend(user_layers) |
| relevant_tables = list(set(relevant_tables)) |
| |
| |
| if allowed_datasets: |
| relevant_tables = [t for t in relevant_tables if t in allowed_datasets or t in user_layers] |
| |
| |
| 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) |
| |
| |
| table_schema = self._get_cached_schema(feature_tables) if feature_tables else self.geo_engine.get_table_schemas() |
| |
| |
| 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}") |
| |
| |
| 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 |
| |
| |
| 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"})} |
|
|
|
|
| |
| |
| |
| yield {"event": "status", "data": dumps_safe({"status": "💬 Generating response..."})} |
| |
| |
| data_summary = ResponseFormatter.generate_data_summary(features) |
| citations = ResponseFormatter.generate_citations(relevant_tables, sql) |
|
|
| |
| |
| 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) |
| |
| |
| |
| 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) |
| ) |
|
|
| |
| 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: |
| 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}") |
| |
| |
| 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) |
|
|
| |
| 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 |
| })} |
|
|
| |
| |
| async for ev in self._emit_followups(query, explanation_text, allowed_datasets, history): |
| yield ev |
|
|
| |
| |
| |
|
|
| 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..."})} |
| |
| |
| 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}" |
| ) |
| |
| |
| if allowed_datasets: |
| relevant_tables = [t for t in relevant_tables if t in allowed_datasets] |
| |
| |
| 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.") |
| |
| |
| 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(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}" |
| |
| |
| 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}") |
| |
| |
| 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"})} |
|
|
| |
| |
| 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": [] |
| })} |
|
|
| |
| |
| |
|
|
| 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: |
| 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() |
|
|
| |
| |
| |
| 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.""" |
|
|
| |
| |
| |
|
|
| 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.""" |
| |
| 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()) |
| |
| |
| 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: |
| |
| 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 |
| |
| |
| 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) |
| })} |
| |
| |
| 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) |
| |
| |
| 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)} |
| |
| |
| 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"]})} |
| |
| |
| 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") |
| |
| |
| 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) |
| })} |
|
|
| |
| |
| |
|
|
| 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" |
| } |
| |
| |
| 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] |
| |
| |
| 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() |
| |
| |
| 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" |
| } |
| |
| |
| 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" |
| } |
| |
| |
| 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 |