perch / backend /api /endpoints /chat.py
GerardCB's picture
Fix 405 when saving a drawn layer in the deployed build
05d7afb verified
Raw
History Blame Contribute Delete
2.92 kB
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())