296 lines
11 KiB
Python
296 lines
11 KiB
Python
#!/usr/bin/env python3
|
||
"""Invoke the configured AG-UI sentiment gateway and normalize its SSE stream.
|
||
|
||
This script intentionally uses only Python's standard library, so an installed
|
||
skill does not need to install an extra dependency. It accepts the external
|
||
AG-UI stream and prints one normalized JSON object per line (NDJSON) as soon as
|
||
each event arrives, which makes it suitable for a caller that renders a live
|
||
stream. ``--format markdown`` is a convenience mode for ordinary agent calls.
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import argparse
|
||
import base64
|
||
import json
|
||
import os
|
||
import ssl
|
||
import sys
|
||
import uuid
|
||
from collections.abc import Iterable, Iterator
|
||
from pathlib import Path
|
||
from typing import Any
|
||
from urllib.error import HTTPError, URLError
|
||
from urllib.request import HTTPSHandler, Request, build_opener
|
||
|
||
|
||
DEFAULT_CONFIG_PATH = Path(__file__).resolve().parents[1] / "config" / "sentiment-agent.json"
|
||
VALID_ROLES = {"user", "assistant"}
|
||
|
||
|
||
class SentimentAgentError(RuntimeError):
|
||
"""A safe, user-facing error from the external sentiment gateway."""
|
||
|
||
|
||
def load_config(path: Path) -> dict[str, Any]:
|
||
"""Read and validate the skill's separate gateway configuration file."""
|
||
|
||
try:
|
||
raw = json.loads(path.read_text(encoding="utf-8"))
|
||
except FileNotFoundError as exc:
|
||
raise SentimentAgentError(f"未找到技能配置文件:{path}") from exc
|
||
except json.JSONDecodeError as exc:
|
||
raise SentimentAgentError(f"技能配置文件不是合法 JSON:{path}") from exc
|
||
if not isinstance(raw, dict):
|
||
raise SentimentAgentError("技能配置文件根节点必须是 JSON 对象。")
|
||
|
||
api_url = str(raw.get("api_url") or "").strip()
|
||
if not api_url.startswith(("https://", "http://")):
|
||
raise SentimentAgentError("请先在 config/sentiment-agent.json 配置有效的 api_url。")
|
||
return raw
|
||
|
||
|
||
def _read_auth_value(auth: dict[str, Any], key: str) -> str:
|
||
"""Read a credential, allowing the deployment environment to override JSON."""
|
||
|
||
env_name = str(auth.get(f"{key}_env") or "").strip()
|
||
if env_name:
|
||
env_value = os.getenv(env_name)
|
||
if env_value:
|
||
return env_value
|
||
return str(auth.get(key) or "")
|
||
|
||
|
||
def build_headers(config: dict[str, Any]) -> dict[str, str]:
|
||
"""Create upstream request headers without ever printing credentials."""
|
||
|
||
headers = {
|
||
"Content-Type": "application/json",
|
||
"Accept": "text/event-stream",
|
||
}
|
||
auth = config.get("auth")
|
||
if not isinstance(auth, dict):
|
||
return headers
|
||
username = _read_auth_value(auth, "username")
|
||
password = _read_auth_value(auth, "password")
|
||
if username and password:
|
||
token = base64.b64encode(f"{username}:{password}".encode("utf-8")).decode("ascii")
|
||
headers["Authorization"] = f"Basic {token}"
|
||
return headers
|
||
|
||
|
||
def load_history(path: Path | None) -> list[dict[str, str]]:
|
||
"""Read a caller-supplied AG-UI history file and validate its safe subset."""
|
||
|
||
if path is None:
|
||
return []
|
||
try:
|
||
raw = json.loads(path.read_text(encoding="utf-8"))
|
||
except FileNotFoundError as exc:
|
||
raise SentimentAgentError(f"未找到历史消息文件:{path}") from exc
|
||
except json.JSONDecodeError as exc:
|
||
raise SentimentAgentError(f"历史消息文件不是合法 JSON:{path}") from exc
|
||
if not isinstance(raw, list):
|
||
raise SentimentAgentError("历史消息文件必须是 JSON 数组。")
|
||
|
||
messages: list[dict[str, str]] = []
|
||
for index, item in enumerate(raw, start=1):
|
||
if not isinstance(item, dict):
|
||
raise SentimentAgentError(f"历史消息第 {index} 项必须是对象。")
|
||
role = str(item.get("role") or "")
|
||
content = str(item.get("content") or "").strip()
|
||
if role not in VALID_ROLES or not content:
|
||
raise SentimentAgentError(f"历史消息第 {index} 项需要非空 role(user/assistant) 和 content。")
|
||
message_id = str(item.get("id") or f"history-{index}-{uuid.uuid4()}")
|
||
messages.append({"role": role, "id": message_id, "content": content})
|
||
return messages
|
||
|
||
|
||
def build_payload(
|
||
*,
|
||
thread_id: str,
|
||
run_id: str,
|
||
question: str,
|
||
history: list[dict[str, str]],
|
||
) -> dict[str, Any]:
|
||
"""Build the AG-UI request body documented by the external gateway."""
|
||
|
||
text = question.strip()
|
||
if not text:
|
||
raise SentimentAgentError("问题不能为空。")
|
||
return {
|
||
"threadId": thread_id,
|
||
"runId": run_id,
|
||
"state": {},
|
||
"tools": [],
|
||
"context": [],
|
||
"forwardedProps": {},
|
||
"messages": [
|
||
*history,
|
||
{"role": "user", "id": f"message-{uuid.uuid4()}", "content": text},
|
||
],
|
||
}
|
||
|
||
|
||
def parse_sse_payloads(lines: Iterable[bytes]) -> Iterator[dict[str, Any]]:
|
||
"""Yield JSON payloads from an SSE byte stream, including an unterminated last frame."""
|
||
|
||
data_lines: list[str] = []
|
||
|
||
def flush() -> dict[str, Any] | None:
|
||
if not data_lines:
|
||
return None
|
||
joined = "\n".join(data_lines)
|
||
data_lines.clear()
|
||
try:
|
||
parsed = json.loads(joined)
|
||
except json.JSONDecodeError:
|
||
return None
|
||
return parsed if isinstance(parsed, dict) else None
|
||
|
||
for raw_line in lines:
|
||
line = raw_line.decode("utf-8", errors="replace").rstrip("\r\n")
|
||
if not line:
|
||
payload = flush()
|
||
if payload is not None:
|
||
yield payload
|
||
continue
|
||
if line.startswith("data:"):
|
||
data_lines.append(line[5:].lstrip(" "))
|
||
|
||
payload = flush()
|
||
if payload is not None:
|
||
yield payload
|
||
|
||
|
||
def normalize_agui_event(payload: dict[str, Any], *, thread_id: str, run_id: str) -> dict[str, Any] | None:
|
||
"""Convert an upstream AG-UI event to the stable, UI-friendly event contract."""
|
||
|
||
event_type = str(payload.get("type") or "")
|
||
message_id = str(payload.get("message_id") or payload.get("messageId") or "")
|
||
delta = payload.get("delta")
|
||
text_delta = delta if isinstance(delta, str) else ""
|
||
|
||
if event_type == "RUN_STARTED":
|
||
return {"type": "run_started", "threadId": thread_id, "runId": run_id}
|
||
if event_type == "THINKING_TEXT_MESSAGE_START":
|
||
return {"type": "thinking_start", "messageId": message_id}
|
||
if event_type == "THINKING_TEXT_MESSAGE_CONTENT":
|
||
return {"type": "thinking_delta", "messageId": message_id, "delta": text_delta}
|
||
if event_type == "THINKING_TEXT_MESSAGE_END":
|
||
return {"type": "thinking_end", "messageId": message_id}
|
||
if event_type == "TEXT_MESSAGE_START":
|
||
return {"type": "answer_start", "messageId": message_id}
|
||
if event_type == "TEXT_MESSAGE_CONTENT":
|
||
return {"type": "answer_delta", "messageId": message_id, "delta": text_delta}
|
||
if event_type == "TEXT_MESSAGE_END":
|
||
return {"type": "answer_end", "messageId": message_id}
|
||
if event_type == "RUN_FINISHED":
|
||
return {"type": "run_finished", "threadId": thread_id, "runId": run_id}
|
||
if event_type == "RUN_ERROR":
|
||
message = str(payload.get("message") or "舆情分析服务返回错误。")
|
||
raise SentimentAgentError(message)
|
||
return None
|
||
|
||
|
||
def stream_events(config: dict[str, Any], payload: dict[str, Any]) -> Iterator[dict[str, Any]]:
|
||
"""Call the configured gateway and yield normalized events as soon as they arrive."""
|
||
|
||
verify_tls = bool(config.get("verify_tls", True))
|
||
timeout_seconds = float(config.get("timeout_seconds", 300))
|
||
if timeout_seconds <= 0:
|
||
raise SentimentAgentError("timeout_seconds 必须大于 0。")
|
||
|
||
request = Request(
|
||
str(config["api_url"]),
|
||
data=json.dumps(payload, ensure_ascii=False).encode("utf-8"),
|
||
headers=build_headers(config),
|
||
method="POST",
|
||
)
|
||
ssl_context = ssl.create_default_context() if verify_tls else ssl._create_unverified_context()
|
||
opener = build_opener(HTTPSHandler(context=ssl_context))
|
||
try:
|
||
with opener.open(request, timeout=timeout_seconds) as response:
|
||
for upstream in parse_sse_payloads(response):
|
||
normalized = normalize_agui_event(
|
||
upstream,
|
||
thread_id=str(payload["threadId"]),
|
||
run_id=str(payload["runId"]),
|
||
)
|
||
if normalized is not None:
|
||
yield normalized
|
||
except HTTPError as exc:
|
||
raise SentimentAgentError(f"舆情分析服务请求失败(HTTP {exc.code})。") from exc
|
||
except URLError as exc:
|
||
raise SentimentAgentError("无法连接舆情分析服务,请检查接口地址和网络。") from exc
|
||
except TimeoutError as exc:
|
||
raise SentimentAgentError("舆情分析服务响应超时。") from exc
|
||
|
||
|
||
def write_ndjson(events: Iterable[dict[str, Any]]) -> None:
|
||
"""Write every normalized event immediately for real-time consumers."""
|
||
|
||
reasoning = ""
|
||
answer = ""
|
||
for event in events:
|
||
if event["type"] == "thinking_delta":
|
||
reasoning += str(event.get("delta") or "")
|
||
elif event["type"] == "answer_delta":
|
||
answer += str(event.get("delta") or "")
|
||
print(json.dumps(event, ensure_ascii=False), flush=True)
|
||
print(json.dumps({"type": "done", "reasoning": reasoning, "answer": answer}, ensure_ascii=False), flush=True)
|
||
|
||
|
||
def write_markdown(events: Iterable[dict[str, Any]]) -> None:
|
||
"""Collect the stream for a normal chat-friendly Markdown result."""
|
||
|
||
reasoning = ""
|
||
answer = ""
|
||
for event in events:
|
||
if event["type"] == "thinking_delta":
|
||
reasoning += str(event.get("delta") or "")
|
||
elif event["type"] == "answer_delta":
|
||
answer += str(event.get("delta") or "")
|
||
if reasoning:
|
||
print("<details><summary>分析过程</summary>\n\n" + reasoning + "\n\n</details>\n")
|
||
print(answer or "舆情分析服务未返回正文。")
|
||
|
||
|
||
def parse_args() -> argparse.Namespace:
|
||
parser = argparse.ArgumentParser(description="调用舆情分析 AG-UI 流式网关。")
|
||
parser.add_argument("--question", required=True, help="本轮需要分析的问题。")
|
||
parser.add_argument("--thread-id", default=f"sentiment-{uuid.uuid4()}", help="会话标识,默认自动生成。")
|
||
parser.add_argument("--run-id", default=f"run-{uuid.uuid4()}", help="运行标识,默认自动生成。")
|
||
parser.add_argument("--messages-file", type=Path, help="可选的历史消息 JSON 文件。")
|
||
parser.add_argument("--config", type=Path, default=DEFAULT_CONFIG_PATH, help="网关配置 JSON 文件。")
|
||
parser.add_argument("--format", choices=("ndjson", "markdown"), default="ndjson", help="输出格式,默认 ndjson。")
|
||
return parser.parse_args()
|
||
|
||
|
||
def main() -> int:
|
||
args = parse_args()
|
||
try:
|
||
config = load_config(args.config)
|
||
payload = build_payload(
|
||
thread_id=args.thread_id,
|
||
run_id=args.run_id,
|
||
question=args.question,
|
||
history=load_history(args.messages_file),
|
||
)
|
||
events = stream_events(config, payload)
|
||
if args.format == "markdown":
|
||
write_markdown(events)
|
||
else:
|
||
write_ndjson(events)
|
||
return 0
|
||
except SentimentAgentError as exc:
|
||
if args.format == "ndjson":
|
||
print(json.dumps({"type": "error", "message": str(exc)}, ensure_ascii=False), flush=True)
|
||
else:
|
||
print(f"舆情分析调用失败:{exc}", file=sys.stderr)
|
||
return 1
|
||
|
||
|
||
if __name__ == "__main__":
|
||
raise SystemExit(main())
|