import json import logging import traceback from typing import Any, Optional from fastapi import APIRouter from pydantic import BaseModel from sse_starlette.sse import EventSourceResponse from backend.services.executor import QueryExecutor logger = logging.getLogger(__name__) router = APIRouter() class MessageHistory(BaseModel): role: str # "user" or "assistant" content: str class ChatRequest(BaseModel): message: str history: list[MessageHistory] = [] allowed_datasets: Optional[list[str]] = None class ChartData(BaseModel): type: str # 'bar', 'line', 'pie', 'donut', 'histogram' title: Optional[str] = None data: list[dict] = [] xKey: Optional[str] = None yKey: Optional[str] = None series: Optional[list[dict]] = None stacked: Optional[bool] = None xAxisLabel: Optional[str] = None yAxisLabel: Optional[str] = None class ChatResponse(BaseModel): response: str sql_query: Optional[str] = None geojson: Optional[dict] = None data_citations: list[str] = [] intent: Optional[str] = None chart_data: Optional[ChartData] = None raw_data: list[dict[str, Any]] = [] @router.post("", response_model=ChatResponse) @router.post("/", response_model=ChatResponse, include_in_schema=False) async def chat(request: ChatRequest): """ Non-streaming chat endpoint. Routes to the appropriate handler based on detected intent. Prefer /stream for interactive use. """ executor = QueryExecutor() history = [{"role": h.role, "content": h.content} for h in request.history] result = await executor.process_query_with_context( query=request.message, history=history, allowed_datasets=request.allowed_datasets ) return ChatResponse( response=result.get("response", "I processed your request."), sql_query=result.get("sql_query"), geojson=result.get("geojson"), data_citations=result.get("data_citations", []), intent=result.get("intent"), chart_data=result.get("chart_data"), raw_data=result.get("raw_data") or [] ) @router.post("/stream") async def chat_stream(request: ChatRequest): """Streaming chat endpoint that returns Server-Sent Events (SSE).""" executor = QueryExecutor() history = [{"role": h.role, "content": h.content} for h in request.history] async def event_generator(): try: async for event in executor.process_query_stream( request.message, history, allowed_datasets=request.allowed_datasets ): yield event except Exception as e: logger.error(f"Stream error: {e}\n{traceback.format_exc()}") yield { "event": "chunk", "data": json.dumps({"type": "text", "content": f"\n\nError: {str(e)}"}) } return EventSourceResponse(event_generator())