Spaces:
Running
Running
| """FastAPI entry point for the conversational agent.""" | |
| import io | |
| import json | |
| import os | |
| import sys | |
| import traceback | |
| from pathlib import Path | |
| from urllib.parse import parse_qs | |
| from fastapi import FastAPI, HTTPException, Request, WebSocket, WebSocketDisconnect | |
| from fastapi.middleware.cors import CORSMiddleware | |
| from fastapi.responses import FileResponse, StreamingResponse | |
| from pydantic import BaseModel, Field | |
| from starlette.concurrency import run_in_threadpool | |
| from chat_history import ( | |
| add_message, | |
| create_conversation, | |
| ensure_conversation_owner, | |
| format_history, | |
| list_conversations, | |
| list_messages, | |
| list_users, | |
| ) | |
| from hf_qwen_client import ( | |
| DEFAULT_MODEL, | |
| MODEL_ALIASES, | |
| generate_response, | |
| generate_response_stream, | |
| ) | |
| from rag_context import build_context_prompt | |
| def _configure_utf8_stream(stream): | |
| if stream is None: | |
| return stream | |
| try: | |
| stream.reconfigure(encoding="utf-8", errors="backslashreplace") | |
| return stream | |
| except (AttributeError, ValueError, OSError): | |
| buffer = getattr(stream, "buffer", None) | |
| if buffer is not None: | |
| return io.TextIOWrapper(buffer, encoding="utf-8", errors="backslashreplace") | |
| return stream | |
| sys.stdout = _configure_utf8_stream(sys.stdout) | |
| sys.stderr = _configure_utf8_stream(sys.stderr) | |
| os.environ["PYTHONIOENCODING"] = "utf-8" | |
| app = FastAPI(title="Agent API") | |
| FRONTEND_PATH = Path(__file__).parent / "static" / "chat.html" | |
| app.add_middleware( | |
| CORSMiddleware, | |
| allow_origins=["*"], | |
| allow_credentials=False, | |
| allow_methods=["*"], | |
| allow_headers=["*"], | |
| ) | |
| async def print_incoming_request(request, call_next): | |
| print( | |
| { | |
| "method": request.method, | |
| "url": str(request.url), | |
| "content_type": request.headers.get("content-type"), | |
| "content_length": request.headers.get("content-length"), | |
| }, | |
| flush=True, | |
| ) | |
| return await call_next(request) | |
| class ChatRequest(BaseModel): | |
| text: str = Field(..., min_length=1) | |
| user_id: int = Field(..., ge=1) | |
| conversation_id: int | None = Field(default=None, ge=1) | |
| model: str = DEFAULT_MODEL | |
| context_k: int = Field(default=4, ge=1, le=24) | |
| class ChatResponse(BaseModel): | |
| model: str | |
| response: str | |
| conversation_id: int | |
| def generate_agent_response( | |
| text: str, | |
| model: str = DEFAULT_MODEL, | |
| context_k: int = 4, | |
| conversation_history: str = "", | |
| ) -> ChatResponse: | |
| selected_model = MODEL_ALIASES.get(model, model) | |
| prompt = build_context_prompt( | |
| text, | |
| k=context_k, | |
| conversation_history=conversation_history, | |
| ) | |
| response = generate_response(prompt, model=selected_model) | |
| return ChatResponse(model=selected_model, response=response, conversation_id=0) | |
| def read_root(): | |
| return FileResponse(FRONTEND_PATH) | |
| def chat_ui(): | |
| return FileResponse(FRONTEND_PATH) | |
| def health_check(): | |
| return {"status": "ok"} | |
| def chat_users(): | |
| return list_users() | |
| def chat_conversations(user_id: int): | |
| return list_conversations(user_id) | |
| def chat_messages(conversation_id: int, user_id: int): | |
| try: | |
| return list_messages(conversation_id, user_id) | |
| except ValueError as exc: | |
| raise HTTPException(status_code=404, detail=str(exc)) from exc | |
| def _chat(request: ChatRequest) -> ChatResponse: | |
| try: | |
| if request.conversation_id is None: | |
| conversation = create_conversation(request.user_id, request.text.strip()[:80]) | |
| conversation_id = conversation["id"] | |
| else: | |
| conversation_id = request.conversation_id | |
| ensure_conversation_owner(conversation_id, request.user_id) | |
| history = list_messages(conversation_id, request.user_id, limit=12) | |
| add_message(conversation_id, "user", request.text) | |
| selected_model = MODEL_ALIASES.get(request.model, request.model) | |
| prompt = build_context_prompt( | |
| request.text, | |
| k=request.context_k, | |
| conversation_history=format_history(history), | |
| ) | |
| response = generate_response(prompt, model=selected_model) | |
| add_message( | |
| conversation_id, | |
| "assistant", | |
| response, | |
| {"model": selected_model, "context_k": request.context_k}, | |
| ) | |
| return ChatResponse( | |
| model=selected_model, | |
| response=response, | |
| conversation_id=conversation_id, | |
| ) | |
| except Exception: | |
| traceback.print_exc() | |
| raise | |
| async def _read_chat_request(http_request: Request) -> ChatRequest: | |
| content_type = http_request.headers.get("content-type", "").lower() | |
| raw_body = await http_request.body() | |
| stripped_body = raw_body.lstrip() | |
| if "application/json" in content_type or stripped_body.startswith(b"{"): | |
| try: | |
| payload = json.loads(raw_body) | |
| except (UnicodeDecodeError, json.JSONDecodeError) as exc: | |
| raise HTTPException(status_code=422, detail="Invalid JSON body") from exc | |
| elif "application/x-www-form-urlencoded" in content_type: | |
| values = parse_qs( | |
| raw_body.decode("utf-8"), | |
| keep_blank_values=True, | |
| ) | |
| payload = {key: items[-1] for key, items in values.items()} | |
| else: | |
| raise HTTPException(status_code=415, detail="Use JSON or form-urlencoded") | |
| if set(payload) == {"data"}: | |
| try: | |
| nested_payload = json.loads(payload["data"]) | |
| if isinstance(nested_payload, dict): | |
| payload = nested_payload | |
| except (TypeError, json.JSONDecodeError): | |
| pass | |
| aliases = { | |
| "message": "text", | |
| "prompt": "text", | |
| "query": "text", | |
| "userId": "user_id", | |
| "conversationId": "conversation_id", | |
| "contextK": "context_k", | |
| } | |
| for source, target in aliases.items(): | |
| if target not in payload and source in payload: | |
| payload[target] = payload.pop(source) | |
| for optional_field in ("conversation_id", "model", "context_k"): | |
| value = payload.get(optional_field) | |
| if isinstance(value, str) and value.strip().lower() in { | |
| "", | |
| "none", | |
| "null", | |
| "undefined", | |
| }: | |
| payload.pop(optional_field) | |
| try: | |
| return ChatRequest(**payload) | |
| except Exception as exc: | |
| print( | |
| { | |
| "chat_validation_error": str(exc), | |
| "received_field_count": len(payload), | |
| "recognized_fields": sorted( | |
| key | |
| for key in payload | |
| if key | |
| in { | |
| "text", | |
| "user_id", | |
| "conversation_id", | |
| "model", | |
| "context_k", | |
| } | |
| ), | |
| }, | |
| flush=True, | |
| ) | |
| raise HTTPException(status_code=422, detail=str(exc)) from exc | |
| async def chat(http_request: Request): | |
| request = await _read_chat_request(http_request) | |
| return await run_in_threadpool(_chat, request) | |
| def chat_stream(request: ChatRequest): | |
| """Entrega eventos NDJSON: metadata, delta, done o error.""" | |
| def event_stream(): | |
| try: | |
| if request.conversation_id is None: | |
| conversation = create_conversation( | |
| request.user_id, request.text.strip()[:80] | |
| ) | |
| conversation_id = conversation["id"] | |
| else: | |
| conversation_id = request.conversation_id | |
| ensure_conversation_owner(conversation_id, request.user_id) | |
| history = list_messages(conversation_id, request.user_id, limit=12) | |
| add_message(conversation_id, "user", request.text) | |
| selected_model = MODEL_ALIASES.get(request.model, request.model) | |
| yield _ndjson_event( | |
| "metadata", | |
| conversation_id=conversation_id, | |
| model=selected_model, | |
| ) | |
| yield _ndjson_event("status", text="Buscando contexto ASTM...") | |
| prompt = build_context_prompt( | |
| request.text, | |
| k=request.context_k, | |
| conversation_history=format_history(history), | |
| ) | |
| yield _ndjson_event("status", text="Generando respuesta...") | |
| print( | |
| f"Streaming generation started | conversation={conversation_id}", | |
| flush=True, | |
| ) | |
| response_parts = [] | |
| for delta in generate_response_stream(prompt, model=selected_model): | |
| if not response_parts: | |
| print( | |
| f"First streamed token | conversation={conversation_id}", | |
| flush=True, | |
| ) | |
| response_parts.append(delta) | |
| yield _ndjson_event("delta", text=delta) | |
| response = "".join(response_parts).strip() | |
| add_message( | |
| conversation_id, | |
| "assistant", | |
| response, | |
| {"model": selected_model, "context_k": request.context_k}, | |
| ) | |
| print( | |
| f"Streaming generation completed | conversation={conversation_id} " | |
| f"| chars={len(response)}", | |
| flush=True, | |
| ) | |
| yield _ndjson_event("done", response=response) | |
| except Exception as exc: | |
| traceback.print_exc() | |
| yield _ndjson_event("error", detail=str(exc)) | |
| return StreamingResponse( | |
| event_stream(), | |
| media_type="text/event-stream", | |
| headers={ | |
| "Cache-Control": "no-cache, no-transform", | |
| "Content-Encoding": "identity", | |
| "X-Accel-Buffering": "no", | |
| }, | |
| ) | |
| def _ndjson_event(event_type: str, **payload) -> str: | |
| return json.dumps({"type": event_type, **payload}, ensure_ascii=False) + "\n" | |
| async def websocket_chat(websocket: WebSocket): | |
| await websocket.accept() | |
| try: | |
| while True: | |
| data = await websocket.receive_text() | |
| if not data.strip(): | |
| await websocket.send_json({"error": "Message text is required."}) | |
| continue | |
| try: | |
| result = await run_in_threadpool(generate_agent_response, data) | |
| await websocket.send_json(result.model_dump()) | |
| except Exception as exc: | |
| traceback.print_exc() | |
| await websocket.send_json({"error": f"Hugging Face inference failed: {exc}"}) | |
| except WebSocketDisconnect: | |
| print("Client disconnected from WS", flush=True) | |
| import os | |
| from fastapi import Depends, Header | |
| from typing import Any, Dict | |
| from report_agent import generate_agentic_report | |
| INTERNAL_API_KEY = os.getenv("INTERNAL_API_KEY", "alberti-internal-secret") | |
| async def verify_token(x_service_token: str = Header(...)): | |
| if x_service_token != INTERNAL_API_KEY: | |
| raise HTTPException(status_code=403, detail="Invalid internal service token") | |
| return x_service_token | |
| class AgentRequest(BaseModel): | |
| report_id: int | |
| muestra_id: int | |
| user_id: int | |
| raw_data: Dict[str, Any] | |
| async def generate_report(req: AgentRequest, token: str = Depends(verify_token)): | |
| try: | |
| enriched_data = generate_agentic_report( | |
| report_id=req.report_id, | |
| muestra_id=req.muestra_id, | |
| user_id=req.user_id, | |
| raw_data=req.raw_data | |
| ) | |
| return enriched_data | |
| except Exception as e: | |
| raise HTTPException(status_code=500, detail=f"Agent analysis failed: {str(e)}") | |