from typing import Literal, List, Generator, Optional import json import threading from fastapi import APIRouter, HTTPException, Request from fastapi.responses import StreamingResponse from pydantic import BaseModel, Field from backend.services.chat import ChatService # Import the service class from backend.services.event_logger import log_error from backend.services.llm_override import ProviderOverride from postgres_client import pg router = APIRouter(prefix="/api", tags=["chat"]) # Pydantic models for request/response validation class ChatMessage(BaseModel): role: Literal["system", "user", "assistant"] content: str = Field(..., min_length=1) class ChatRequest(BaseModel): messages: List[ChatMessage] top_k: int = Field(default=5, ge=1, le=10) focus_ids: List[str] | None = Field(default=None, description="Optional ids from previously shared results.") provider_override: ProviderOverride | None = Field(default=None, description="Optional user-supplied provider.") use_editor: bool = Field(default=False, description="Run the editor pass (fact-check + language polish) before returning.") quick: bool = Field(default=False, description="Skip the planner/Researcher pre-pass and let the orchestrator answer directly. Faster, less thorough.") session_id: str | None = Field(default=None, description="Browser session id; used to group opt-in TEST eval logs.") class ChatSource(BaseModel): _id: str chunk_index: int heading: str | None debateurl: str | None snippet: str speaker: str | None = None party: str | None = None intressent_id: str | None = None date: str | None = None class PersonRef(BaseModel): intressent_id: str name: str party: str class AttributionWarning(BaseModel): paragraph_idx: int name: str party: str cited_ns: List[int] = Field(default_factory=list) reason: str class ChatResponse(BaseModel): answer: str sources: List[ChatSource] persons: List[PersonRef] = Field(default_factory=list) tables: List[dict] = Field(default_factory=list) focus_ids: List[str] = Field(default_factory=list) attribution_warnings: List[AttributionWarning] = Field(default_factory=list) # Instantiate the chat service once (can be reused for all requests) chat_service = ChatService() @router.post("/chat", response_model=ChatResponse) def chat_endpoint(payload: ChatRequest, request: Request) -> ChatResponse: """ Handles chat requests from the frontend. Uses ChatService to generate a response. Args: payload (ChatRequest): The chat history and parameters from the frontend. Returns: ChatResponse: The assistant's answer and a list of sources. """ # Convert Pydantic models to dicts for the service messages = [msg.model_dump() for msg in payload.messages] try: result = chat_service.get_chat_response( messages=messages, top_k=payload.top_k, focus_ids=payload.focus_ids or [], provider_override=payload.provider_override, use_editor=payload.use_editor, quick=payload.quick, session_id=payload.session_id or request.headers.get("X-Session-Id"), ) except ValueError as e: raise HTTPException(status_code=400, detail=str(e)) except Exception as e: import traceback print("UNHANDLED ERROR in chat_endpoint:", e) traceback.print_exc() log_error("http_500", e, route="/api/chat") raise HTTPException(status_code=500, detail=f"Internal server error: {e}") raw_answer = result.get("answer", "") if not isinstance(raw_answer, str): raw_answer = str(raw_answer) raw_sources = result.get("sources", []) if not isinstance(raw_sources, list): raw_sources = [] sources = [ ChatSource( _id=src.get("_id", ""), chunk_index=src.get("chunk_index", 0), heading=src.get("heading"), debateurl=src.get("debateurl"), snippet=src.get("snippet", ""), speaker=src.get("speaker"), party=src.get("party"), intressent_id=src.get("intressent_id"), date=src.get("date"), ) for src in raw_sources ] persons = [ PersonRef(**p) for p in result.get("persons", []) if isinstance(p, dict) and "intressent_id" in p and "name" in p ] warnings = [ AttributionWarning(**w) for w in result.get("attribution_warnings", []) if isinstance(w, dict) and "paragraph_idx" in w and "name" in w ] return ChatResponse( answer=raw_answer, sources=sources, persons=persons, tables=result.get("tables", []), focus_ids=result.get("focus_ids", []), attribution_warnings=warnings, ) @router.post("/chat/stream") def chat_stream_endpoint(payload: ChatRequest, request: Request) -> StreamingResponse: """ SSE endpoint: streams tool-call progress events followed by the final answer. Each event is a line of the form: data: \\n\\n Event types: "tool_call", "status", "answer", "error". Using streaming avoids Cloudflare's 100-second proxy timeout for long-running queries. """ messages = [msg.model_dump() for msg in payload.messages] session_id = payload.session_id or request.headers.get("X-Session-Id") def generate() -> Generator[str, None, None]: # SSE comment keepalive: nginx and Cloudflare buffer connections until # data arrives. By sending a ": keepalive\n\n" comment every 5 seconds # we force buffer flushes so progress hints reach the browser in real-time. import queue as _queue import time as _time done_event = threading.Event() pipe: _queue.Queue[str] = _queue.Queue() def produce() -> None: try: for event in chat_service.stream_chat_response( messages=messages, top_k=payload.top_k, focus_ids=payload.focus_ids or [], provider_override=payload.provider_override, use_editor=payload.use_editor, quick=payload.quick, session_id=session_id, ): pipe.put(f"data: {json.dumps(event, ensure_ascii=False)}\n\n") if event.get("type") in ("answer", "error"): break except Exception as exc: import traceback traceback.print_exc() log_error("sse_stream_exception", exc, route="/api/chat/stream") pipe.put(f"data: {json.dumps({'type': 'error', 'message': str(exc)})}\n\n") finally: done_event.set() threading.Thread(target=produce, daemon=True).start() while not done_event.is_set() or not pipe.empty(): try: chunk = pipe.get(timeout=5) yield chunk except _queue.Empty: # No data for 5 s — send an SSE comment to keep the connection alive # and force proxies/CDNs to flush their buffers. if not done_event.is_set(): yield ": keepalive\n\n" return StreamingResponse( generate(), media_type="text/event-stream", headers={ "Cache-Control": "no-cache", "X-Accel-Buffering": "no", # Disable nginx buffering for SSE }, ) # ── Provider list + model list proxy ───────────────────────────────────────── @router.get("/providers") def list_providers() -> dict: """Return all configured providers from providers.yaml (no secrets included).""" from backend.services.provider_registry import list_providers as _list return { "providers": [ {"id": p.id, "name": p.name, "user_api_key": p.user_api_key} for p in _list() ] } @router.get("/providers/{provider_id}/models") def list_provider_models(provider_id: str, request: Request) -> dict: """ Proxy GET {provider_base_url}/models with the user-supplied API key. The key must be passed in the X-Provider-Key header — it is never logged or stored. """ import requests as _requests from backend.services.provider_registry import get_provider provider = get_provider(provider_id) if provider is None: raise HTTPException(status_code=404, detail=f"Unknown provider: {provider_id!r}") api_key = request.headers.get("X-Provider-Key", "").strip() if not api_key: raise HTTPException(status_code=400, detail="X-Provider-Key header is required") base = provider.base_url.rstrip("/") # OpenRouter supports a query param to filter by tool-capable models directly. params = {} if "openrouter.ai" in base: params["supported_parameters"] = "tools" try: resp = _requests.get( f"{base}/models", headers={"Authorization": f"Bearer {api_key}"}, params=params, timeout=10, ) except Exception as exc: raise HTTPException(status_code=502, detail=f"Could not reach provider: {exc}") if not resp.ok: raise HTTPException(status_code=resp.status_code, detail=f"Provider error: {resp.text[:200]}") data = resp.json() raw_models = [m for m in data.get("data", []) if isinstance(m, dict) and "id" in m] # Filter to tool-capable models. Each provider exposes this differently. if "berget.ai" in base: # Berget exposes capabilities.function_calling per model. raw_models = [m for m in raw_models if m.get("capabilities", {}).get("function_calling")] elif "openai.com" in base: # OpenAI /v1/models includes embeddings, TTS, image models — keep only text generation. _SKIP = ("whisper", "tts", "dall-e", "embedding", "moderation", "babbage", "davinci", "o1-mini", "realtime") raw_models = [m for m in raw_models if not any(s in m["id"] for s in _SKIP)] # OpenRouter already filtered by ?supported_parameters=tools above. models = [m["id"] for m in raw_models] return {"models": sorted(models)} # ── MP (person) endpoints ────────────────────────────────────────────────────── class Uppdrag(BaseModel): typ: str | None = None organ_kod: str | None = None roll_kod: str | None = None status: str | None = None uppgift: str | None = None from_: str | None = Field(None, alias="from") tom: str | None = None model_config = {"populate_by_name": True} class PersonDetail(BaseModel): intressent_id: str namn: str tilltalsnamn: str | None = None efternamn: str | None = None parti: str | None = None valkrets: str | None = None status: str | None = None bild_url_192: str | None = None fodd_ar: str | None = None uppdrag: list[Uppdrag] | None = None @router.get("/person/{intressent_id}", response_model=PersonDetail) def get_person(intressent_id: str) -> PersonDetail: """Fetch basic person data for a Riksdag member.""" rows = pg.execute( """SELECT intressent_id, namn, tilltalsnamn, efternamn, parti, valkrets, status, bild_url_192, fodd_ar, personuppdrag FROM people WHERE intressent_id = %s""", (intressent_id,) ) if not rows: raise HTTPException(status_code=404, detail="Person not found") row = dict(rows[0]) raw = row.pop("personuppdrag", None) or {} uppdrag_list = raw.get("uppdrag", []) if isinstance(raw, dict) else [] return PersonDetail(**row, uppdrag=[Uppdrag(**u) for u in uppdrag_list]) class MpChatRequest(BaseModel): messages: List[ChatMessage] intressent_id: str initial_talk_id: Optional[str] = None provider_override: ProviderOverride | None = Field(default=None, description="Optional user-supplied provider.") @router.post("/chat/mp/stream") def mp_chat_stream_endpoint(payload: MpChatRequest) -> StreamingResponse: """ SSE streaming endpoint for chatting with an MP persona. The LLM role-plays as the specified person, grounded in their actual speeches. """ from backend.services.mp_chat import MpChatService try: mp_service = MpChatService( intressent_id=payload.intressent_id, initial_talk_id=payload.initial_talk_id, provider_override=payload.provider_override, ) except ValueError as e: # Same exception type covers "no such person" and "no such provider"; # only the former should read as a missing resource. status = 400 if "provider" in str(e).lower() else 404 raise HTTPException(status_code=status, detail=str(e)) messages = [msg.model_dump() for msg in payload.messages] def generate() -> Generator[str, None, None]: import queue as _queue done_event = threading.Event() pipe: _queue.Queue[str] = _queue.Queue() def produce() -> None: try: for event in mp_service.stream_chat_response(messages=messages): pipe.put(f"data: {json.dumps(event, ensure_ascii=False)}\n\n") if event.get("type") in ("answer", "error"): break except Exception as exc: import traceback traceback.print_exc() log_error("sse_stream_exception", exc, route="/api/chat/mp/stream") pipe.put(f"data: {json.dumps({'type': 'error', 'message': str(exc)})}\n\n") finally: done_event.set() threading.Thread(target=produce, daemon=True).start() while not done_event.is_set() or not pipe.empty(): try: chunk = pipe.get(timeout=5) yield chunk except _queue.Empty: if not done_event.is_set(): yield ": keepalive\n\n" return StreamingResponse( generate(), media_type="text/event-stream", headers={ "Cache-Control": "no-cache", "X-Accel-Buffering": "no", }, )