deerflow-code/offline-backend-20260512/backend/app/gateway/routers/sentiment_agent.py
2026-09-07 18:24:55 +08:00

221 lines
7.3 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""舆情分析虚拟智能体的 AG-UI SSE 反向代理。
前端对话页不直连外部网关(避免浏览器 CORS / 混合内容),改为把
``runtime-config.js`` 里的 ``apiUrl`` / Basic Auth 随 AG-UI 请求体传给本路由。
Gateway 关闭 TLS 校验后按传入地址转发上游 SSE,字节级透传给前端。
连接失败或上游非 2xx 时,把完整错误信息放进 HTTP ``detail``(流已开始则发
AG-UI ``RUN_ERROR`` 帧),前端原样展示。
Routes (prefix ``/api/sentiment-agent``):
POST /stream 透传上游 ``text/event-stream``
"""
from __future__ import annotations
import base64
import json
import logging
import ssl
from collections.abc import AsyncIterator, Mapping
from typing import Any
import httpx
from fastapi import APIRouter, HTTPException, Request
from fastapi.responses import StreamingResponse
from pydantic import BaseModel, ConfigDict, Field
from deerflow.config import get_app_config
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/api/sentiment-agent", tags=["sentiment-agent"])
_MAX_ERROR_BODY_CHARS = 32_768
class SentimentAgentStreamRequest(BaseModel):
"""前端传入的 AG-UI 体 + runtime-config 上游地址/凭据。"""
model_config = ConfigDict(extra="ignore")
apiUrl: str = Field(default="", min_length=0, max_length=2048)
authUsername: str = Field(default="", max_length=512)
authPassword: str = Field(default="", max_length=512)
threadId: str = Field(..., min_length=1, max_length=256)
runId: str = Field(..., min_length=1, max_length=256)
state: dict[str, Any] = Field(default_factory=dict)
tools: list[Any] = Field(default_factory=list)
context: list[Any] = Field(default_factory=list)
forwardedProps: dict[str, Any] = Field(default_factory=dict)
messages: list[Any] = Field(default_factory=list)
def _unverified_ssl_context() -> ssl.SSLContext:
"""Do not validate the upstream certificate (intranet / self-signed hosts)."""
ctx = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT)
ctx.check_hostname = False
ctx.verify_mode = ssl.CERT_NONE
return ctx
def _build_http_client(timeout_seconds: float) -> httpx.AsyncClient:
"""Upstream client: skip TLS cert checks, ignore env proxy, no redirects."""
return httpx.AsyncClient(
verify=_unverified_ssl_context(),
trust_env=False,
follow_redirects=False,
timeout=httpx.Timeout(timeout_seconds, connect=30.0),
)
def _format_exception(exc: BaseException) -> str:
parts = [f"{type(exc).__name__}: {exc}"]
seen = {id(exc)}
cause: BaseException | None = exc.__cause__ or exc.__context__
while cause is not None and id(cause) not in seen:
seen.add(id(cause))
parts.append(f"Caused by {type(cause).__name__}: {cause}")
cause = cause.__cause__ or cause.__context__
return " | ".join(parts)
def _decode_body(content: bytes | None) -> str:
if not content:
return ""
text = content.decode("utf-8", errors="replace").strip()
if len(text) > _MAX_ERROR_BODY_CHARS:
return text[:_MAX_ERROR_BODY_CHARS] + f"\n…(truncated, total {len(text)} chars)"
return text
def format_upstream_error(
*,
kind: str,
url: str,
status: int | None = None,
headers: Mapping[str, str] | None = None,
body: str = "",
exc: BaseException | None = None,
) -> str:
"""Assemble a complete, user-visible upstream failure string."""
lines = [f"舆情分析上游调用失败({kind})", f"url={url}"]
if status is not None:
lines.append(f"http_status={status}")
if headers:
location = headers.get("location") or headers.get("Location")
content_type = headers.get("content-type") or headers.get("Content-Type")
if location:
lines.append(f"location={location}")
if content_type:
lines.append(f"content_type={content_type}")
if body:
lines.append(f"response={body}")
if exc is not None:
lines.append(_format_exception(exc))
request_obj = getattr(exc, "request", None)
method = getattr(request_obj, "method", None)
url_obj = getattr(request_obj, "url", None)
if method and url_obj is not None:
lines.append(f"request={method} {url_obj}")
return "\n".join(lines)
def _upstream_headers(username: str, password: str) -> dict[str, str]:
headers = {
"Content-Type": "application/json",
"Accept": "text/event-stream",
}
if username and password:
token = base64.b64encode(f"{username}:{password}".encode()).decode("ascii")
headers["Authorization"] = f"Basic {token}"
return headers
def _agui_payload(body: SentimentAgentStreamRequest) -> dict[str, Any]:
return {
"threadId": body.threadId,
"runId": body.runId,
"state": body.state,
"tools": body.tools,
"context": body.context,
"forwardedProps": body.forwardedProps,
"messages": body.messages,
}
def _run_error_frame(message: str) -> bytes:
payload = json.dumps({"type": "RUN_ERROR", "message": message}, ensure_ascii=False)
return f"data: {payload}\n\n".encode()
@router.post("/stream")
async def stream_sentiment_agent(
request: Request,
body: SentimentAgentStreamRequest,
) -> StreamingResponse:
api_url = (body.apiUrl or "").strip()
if not api_url:
raise HTTPException(status_code=400, detail="缺少上游地址 apiUrl(runtime-config.js 的 VITE_SENTIMENT_AGENT_API_URL)。")
config = get_app_config().sentiment_agent
payload = _agui_payload(body)
headers = _upstream_headers(body.authUsername, body.authPassword)
client = _build_http_client(config.timeout_seconds)
upstream: httpx.Response | None = None
try:
upstream = await client.send(
client.build_request("POST", api_url, headers=headers, json=payload),
stream=True,
)
except httpx.RequestError as exc:
await client.aclose()
raise HTTPException(
status_code=502,
detail=format_upstream_error(kind="network", url=api_url, exc=exc),
) from exc
if upstream.status_code >= 400:
err_body = _decode_body(await upstream.aread())
err_headers = dict(upstream.headers)
status = upstream.status_code
await upstream.aclose()
await client.aclose()
raise HTTPException(
status_code=502,
detail=format_upstream_error(
kind="http",
url=api_url,
status=status,
headers=err_headers,
body=err_body,
),
)
async def pump() -> AsyncIterator[bytes]:
assert upstream is not None
try:
async for chunk in upstream.aiter_bytes():
if await request.is_disconnected():
break
if chunk:
yield chunk
except httpx.RequestError as exc:
yield _run_error_frame(format_upstream_error(kind="stream", url=api_url, exc=exc))
finally:
await upstream.aclose()
await client.aclose()
return StreamingResponse(
pump(),
media_type="text/event-stream",
headers={
"Cache-Control": "no-cache",
"Connection": "keep-alive",
"X-Accel-Buffering": "no",
},
)