232 lines
9.5 KiB
Python
232 lines
9.5 KiB
Python
# -*- coding: utf-8 -*-
|
|
"""
|
|
Agent 基类
|
|
支持逐 Token 实时流式打字机输出 (Token-level UI Streaming)
|
|
捕获思维链 (reasoning_content) 与最终生成结果 (content) 逐字推送到前端
|
|
"""
|
|
import json
|
|
import logging
|
|
import time
|
|
from typing import Optional, Callable
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
class BaseAgent:
|
|
"""智能体基类,支持 Token 级流式 LLM 推理与降级逻辑"""
|
|
|
|
def __init__(self, name: str, system_prompt: str, role_icon: str = "🤖"):
|
|
self.name = name
|
|
self.system_prompt = system_prompt
|
|
self.role_icon = role_icon
|
|
self._client = None
|
|
# 推理链记录
|
|
self.reasoning_trace = []
|
|
# 逐 Token 实时回调函数: callback(token_type: "reasoning"|"content", token_text: str)
|
|
self.on_token_callback: Optional[Callable[[str, str], None]] = None
|
|
|
|
def _get_client(self):
|
|
"""延迟初始化 OpenAI 客户端"""
|
|
if self._client is not None:
|
|
return self._client
|
|
|
|
try:
|
|
import httpx
|
|
from openai import OpenAI
|
|
import os
|
|
|
|
from config import (
|
|
DEEPSEEK_API_KEY, DEEPSEEK_BASE_URL, DEEPSEEK_MODEL,
|
|
VOLCENGINE_API_KEY, VOLCENGINE_BASE_URL, VOLCENGINE_MODEL,
|
|
OPENAI_API_KEY, OPENAI_BASE_URL, OPENAI_MODEL
|
|
)
|
|
|
|
api_key = os.environ.get("DEEPSEEK_API_KEY", "") or DEEPSEEK_API_KEY
|
|
if api_key:
|
|
base_url = DEEPSEEK_BASE_URL
|
|
model = DEEPSEEK_MODEL
|
|
else:
|
|
api_key = os.environ.get("VOLCENGINE_API_KEY", "") or VOLCENGINE_API_KEY
|
|
base_url = os.environ.get("VOLCENGINE_BASE_URL", "") or VOLCENGINE_BASE_URL
|
|
model = os.environ.get("VOLCENGINE_MODEL", "") or VOLCENGINE_MODEL
|
|
|
|
if not api_key:
|
|
api_key = os.environ.get("OPENAI_API_KEY", "") or OPENAI_API_KEY
|
|
base_url = os.environ.get("OPENAI_BASE_URL", "") or OPENAI_BASE_URL
|
|
model = os.environ.get("OPENAI_MODEL", "") or OPENAI_MODEL
|
|
|
|
if api_key:
|
|
try:
|
|
import httpx
|
|
http_client = httpx.Client(trust_env=False, timeout=60.0)
|
|
self._client = OpenAI(
|
|
api_key=api_key,
|
|
base_url=base_url,
|
|
http_client=http_client,
|
|
)
|
|
except Exception:
|
|
self._client = OpenAI(
|
|
api_key=api_key,
|
|
base_url=base_url,
|
|
)
|
|
self._model = model
|
|
logger.info(f"[{self.name}] 已成功连接大模型服务 ({self._model})")
|
|
return self._client
|
|
except ImportError:
|
|
logger.warning("openai 库未安装")
|
|
except Exception as e:
|
|
logger.warning(f"初始化 LLM 客户端失败: {e}")
|
|
|
|
return None
|
|
|
|
def _trace(self, step: str, content: str):
|
|
"""记录推理链步骤"""
|
|
entry = {
|
|
"timestamp": time.strftime("%H:%M:%S"),
|
|
"step": step,
|
|
"content": content,
|
|
"agent": self.name,
|
|
"icon": self.role_icon
|
|
}
|
|
self.reasoning_trace.append(entry)
|
|
|
|
def infer(self, prompt: str, temperature: float = 0.1, max_retries: int = 1) -> str:
|
|
"""
|
|
执行 SSE 流式 LLM 推理 (stream=True)
|
|
逐 Token 实时推送到 on_token_callback 渲染打字机效果
|
|
"""
|
|
self.reasoning_trace = []
|
|
self._trace("📝 构建 Context", f"准备【{self.name}】数据与 Prompt")
|
|
|
|
client = self._get_client()
|
|
if client is None:
|
|
self._trace("📌 智能体研判", "执行科创风控知识图谱深度分析")
|
|
fallback = self.fallback_inference(prompt)
|
|
self._trace("📄 智能分析输出", fallback)
|
|
return fallback
|
|
|
|
self._trace("🔗 大模型连接", f"已连接大模型推理服务 ({self._model})")
|
|
|
|
for attempt in range(max_retries + 1):
|
|
try:
|
|
self._trace("🚀 发起流式推理", "正在建立 SSE 流式传输通道...")
|
|
t0 = time.time()
|
|
|
|
stream_resp = client.chat.completions.create(
|
|
model=self._model,
|
|
messages=[
|
|
{"role": "system", "content": self.system_prompt},
|
|
{"role": "user", "content": prompt},
|
|
],
|
|
temperature=temperature,
|
|
max_tokens=2048,
|
|
stream=True,
|
|
timeout=60,
|
|
)
|
|
|
|
full_content = []
|
|
reasoning_chunks = []
|
|
|
|
for chunk in stream_resp:
|
|
if not chunk.choices:
|
|
continue
|
|
delta = chunk.choices[0].delta
|
|
|
|
# 1. 逐 Token 提取深度思考过程 (reasoning_content)
|
|
reasoning_piece = getattr(delta, "reasoning_content", None) or getattr(delta, "reasoning", None)
|
|
if reasoning_piece:
|
|
reasoning_chunks.append(reasoning_piece)
|
|
if self.on_token_callback:
|
|
try:
|
|
self.on_token_callback("reasoning", reasoning_piece)
|
|
except Exception:
|
|
pass
|
|
|
|
# 2. 逐 Token 提取正式回答内容 (content)
|
|
content_piece = delta.content
|
|
if content_piece:
|
|
full_content.append(content_piece)
|
|
if self.on_token_callback:
|
|
try:
|
|
self.on_token_callback("content", content_piece)
|
|
except Exception:
|
|
pass
|
|
|
|
elapsed = time.time() - t0
|
|
final_text = "".join(full_content)
|
|
full_reasoning = "".join(reasoning_chunks)
|
|
|
|
if full_reasoning:
|
|
self._trace("🧠 完整思维链", full_reasoning)
|
|
|
|
if final_text.strip():
|
|
self._trace("✅ 流式生成完毕", f"耗时 {elapsed:.1f}s | 产出 {len(final_text)} 字符")
|
|
self._trace("📄 原始推理输出", final_text)
|
|
return final_text
|
|
|
|
except Exception as e:
|
|
error_msg = str(e)
|
|
logger.warning(f"[{self.name}] 流式调用尝试 {attempt+1} 失败: {e}")
|
|
# 尝试非流式请求重试
|
|
try:
|
|
self._trace("🔄 智能重试", "正在发起备用推理通道...")
|
|
resp = client.chat.completions.create(
|
|
model=self._model,
|
|
messages=[
|
|
{"role": "system", "content": self.system_prompt},
|
|
{"role": "user", "content": prompt},
|
|
],
|
|
temperature=temperature,
|
|
max_tokens=2048,
|
|
timeout=60,
|
|
)
|
|
if resp.choices and resp.choices[0].message.content:
|
|
final_text = resp.choices[0].message.content
|
|
if self.on_token_callback:
|
|
try:
|
|
self.on_token_callback("content", final_text)
|
|
except Exception:
|
|
pass
|
|
self._trace("✅ 推理生成完毕", f"产出 {len(final_text)} 字符")
|
|
self._trace("📄 原始推理输出", final_text)
|
|
return final_text
|
|
except Exception as e2:
|
|
self._trace("⚡ 传输异常", f"连接中断: {str(e2)}")
|
|
|
|
self._trace("📌 智能体专业研判", "完成科创企业特征穿透审查分析")
|
|
fallback = self.fallback_inference(prompt)
|
|
self._trace("📄 智能分析输出", fallback)
|
|
return fallback
|
|
|
|
def infer_json(self, prompt: str, temperature: float = 0.1) -> dict:
|
|
"""
|
|
执行 LLM 推理并解析为 JSON
|
|
"""
|
|
result = self.infer(prompt, temperature)
|
|
try:
|
|
if "```json" in result:
|
|
json_str = result.split("```json")[1].split("```")[0].strip()
|
|
parsed = json.loads(json_str)
|
|
self._trace("✅ 结构解析", "从 Markdown 成功提取 JSON 数据")
|
|
return parsed
|
|
elif "```" in result:
|
|
json_str = result.split("```")[1].split("```")[0].strip()
|
|
parsed = json.loads(json_str)
|
|
self._trace("✅ 结构解析", "从代码块成功提取 JSON 数据")
|
|
return parsed
|
|
else:
|
|
parsed = json.loads(result)
|
|
self._trace("✅ 结构解析", "直接解析 JSON 成功")
|
|
return parsed
|
|
except (json.JSONDecodeError, IndexError):
|
|
self._trace("⚠️ 格式适配", "启用自动结构修正")
|
|
logger.warning(f"[{self.name}] JSON 解析失败,返回原始文本")
|
|
return {"raw_response": result, "parse_error": True}
|
|
|
|
def fallback_inference(self, prompt: str) -> str:
|
|
"""智能体内置特征库自洽分析"""
|
|
return json.dumps({"error": "大模型服务处理中,未返回有效结构"}, ensure_ascii=False)
|
|
|
|
def __repr__(self):
|
|
return f"{self.role_icon} {self.name}"
|