agent-api / main.py
github-actions[bot]
Sync GitHub snapshot to Hugging Face
9a1014e
Raw
History Blame Contribute Delete
12 kB
"""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=["*"],
)
@app.middleware("http")
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)
@app.get("/")
def read_root():
return FileResponse(FRONTEND_PATH)
@app.get("/ui", include_in_schema=False)
def chat_ui():
return FileResponse(FRONTEND_PATH)
@app.get("/health")
def health_check():
return {"status": "ok"}
@app.get("/chat/users")
def chat_users():
return list_users()
@app.get("/chat/conversations")
def chat_conversations(user_id: int):
return list_conversations(user_id)
@app.get("/chat/conversations/{conversation_id}/messages")
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
@app.post("/chat", response_model=ChatResponse)
async def chat(http_request: Request):
request = await _read_chat_request(http_request)
return await run_in_threadpool(_chat, request)
@app.post("/chat/stream")
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"
@app.websocket("/ws/chat")
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]
@app.post("/generate-report")
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)}")