221 lines
7.3 KiB
Python
221 lines
7.3 KiB
Python
"""舆情分析虚拟智能体的 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",
|
||
},
|
||
)
|