| 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 |
| content: str |
|
|
|
|
| class ChatRequest(BaseModel): |
| message: str |
| history: list[MessageHistory] = [] |
| allowed_datasets: Optional[list[str]] = None |
|
|
|
|
| class ChartData(BaseModel): |
| type: str |
| 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()) |
|
|