mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-10-09 22:00:01 +08:00
✨ feat!(llm): 重构并升级大语言模型服务为全新 AI 智能体框架 (#2146)
* ✨ feat!(llm): 重构并升级大语言模型服务为全新 AI 智能体框架 - 【重构】将原 services/llm 重构并迁移至全新的 services/ai 架构,提供向下兼容垫片 - 【新增】引入 Agent、Team、Workflow 三大智能体与工作流编排范式 - 【新增】引入基于 RAG 的长期向量记忆与中期槽位记忆系统 - 【新增】引入基于 Docker 的安全代码执行沙箱环境 - 【新增】支持 MCP 协议,允许动态管理和调用 MCP 服务 - 【新增】引入输入输出安全合规护栏与自愈反思机制 - 【优化】重构并优化多厂商 API 适配器 (Gemini, OpenAI, DeepSeek, GLM 等) - 【优化】优化日志脱敏与 Token 预估机制 - 【移除】移除旧版 llm default 和 llm reset-key 命令,新增 llm mcp 管理命令 * 🔧 chore(deps): 更新项目依赖与配置 - 添加 mcp、jieba 和 aiodocker 依赖到配置文件及 requirements.txt - 在 pyright 配置中设置 reportMissingImports 为 none - 调整 .gitignore 中 resources 目录的忽略规则 * ♻️ refactor(tools): 重构工具终止机制并清理知识库日志输出 - 统一使用 `context.state["__end_run__"]` 替代 `EndRunResult` 控制任务结束 - 移除文件系统和向量知识库检索工具中 `ToolResult` 的 `.with_log` 调用 - 调整指令处理器(Directive)的返回值为 `tool_res.output` - 修复部分类型检查警告并优化联合类型判断语法 * ♻️ refactor(tools): 重构工具副作用指令与控制流熔断机制 - 引入 `DirectivePayload` 及 `ToolResult` 的子类以结构化表达工具副作用 - 移除通过 `context.state` 传递魔术变量的隐式控制流设计 - 重构 `DirectiveManager` 处理器接口,直接在处理器中修改 `AgentState` 并构建 `AgentRunResult` - 在 `StandardAgentExecutor` 中统一通过 `directive_manager` 调度工具返回的副作用指令 - 补全 `MessageBuilder` 中部分核心方法的文档注释 * 🐛 fix(sandbox): 修复 Docker 沙箱容器状态检测与会话清理逻辑 -【修复】修正 `is_alive` 中直接读取私有属性的问题,改用 `show()` 返回值 -【修复】解决 `execute_code` 中缓存的执行器与当前会话不一致的问题 -【优化】在清理工作区前增加容器存活检测,避免向已死容器发送请求 -【优化】创建容器时增加运行状态校验,若已停止则自动从缓存中移除并重建 -【优化】优化容器销毁和清理逻辑,静默处理容器不存在 (404) 的异常 * 📝 docs(core): 补充核心模块初始化方法的文档注释 * 🚨 auto fix by pre-commit hooks --------- Co-authored-by: webjoin111 <455457521@qq.com> Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
This commit is contained in:
co-authored by
webjoin111
pre-commit-ci[bot]
parent
bdc1374848
commit
80fc5b86a7
@@ -0,0 +1,178 @@
|
||||
from dataclasses import dataclass
|
||||
import hashlib
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
from zhenxun.utils.pydantic_compat import model_dump
|
||||
|
||||
|
||||
@dataclass
|
||||
class StablePrefixSnapshot:
|
||||
"""系统提示词与工具的稳定前缀快照"""
|
||||
|
||||
system_prompt: list[str]
|
||||
tools: list[Any]
|
||||
fingerprint: str
|
||||
|
||||
|
||||
class StablePrefix:
|
||||
"""
|
||||
一个冻结 of 系统前缀(系统提示词 + 工具)。
|
||||
通过比对特征指纹,在内容未改变时避免重新构建。
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self._snapshot: StablePrefixSnapshot | None = None
|
||||
self._version = 0
|
||||
|
||||
@property
|
||||
def fingerprint(self) -> str:
|
||||
return self._snapshot.fingerprint if self._snapshot else "<unbuilt>"
|
||||
|
||||
@property
|
||||
def version(self) -> int:
|
||||
return self._version
|
||||
|
||||
@property
|
||||
def built(self) -> bool:
|
||||
return self._snapshot is not None
|
||||
|
||||
def build(self, system_prompt: list[str], tools: list[Any]) -> bool:
|
||||
"""
|
||||
构建或重新构建前缀。
|
||||
返回 True 表示内容发生实质变化(缓存可能失效),False 表示使用旧快照。
|
||||
"""
|
||||
snapshot = self._take_snapshot(system_prompt, tools)
|
||||
if self._snapshot and self._snapshot.fingerprint == snapshot.fingerprint:
|
||||
return False
|
||||
self._snapshot = snapshot
|
||||
self._version += 1
|
||||
return True
|
||||
|
||||
def invalidate(self):
|
||||
self._snapshot = None
|
||||
|
||||
def to_context(self) -> tuple[list[str], list[Any]]:
|
||||
if not self._snapshot:
|
||||
raise RuntimeError("StablePrefix.to_context() called before build()")
|
||||
return self._snapshot.system_prompt, self._snapshot.tools
|
||||
|
||||
def _take_snapshot(
|
||||
self, system_prompt: list[str], tools: list[Any]
|
||||
) -> StablePrefixSnapshot:
|
||||
parsed_tools = []
|
||||
for t in tools:
|
||||
if hasattr(t, "name"):
|
||||
parsed_tools.append((t.name, getattr(t, "description", "")))
|
||||
elif isinstance(t, dict):
|
||||
parsed_tools.append((t.get("name", ""), t.get("description", "")))
|
||||
else:
|
||||
parsed_tools.append(str(t))
|
||||
|
||||
payload = {"s": system_prompt, "t": parsed_tools}
|
||||
json_str = json.dumps(payload, default=str, sort_keys=True)
|
||||
fingerprint = hashlib.md5(json_str.encode("utf-8")).hexdigest()[:8]
|
||||
return StablePrefixSnapshot(
|
||||
system_prompt=list(system_prompt),
|
||||
tools=list(tools),
|
||||
fingerprint=fingerprint,
|
||||
)
|
||||
|
||||
|
||||
class AppendOnlyLog:
|
||||
"""追加写入模式 of 对话日志管理器"""
|
||||
|
||||
def __init__(self):
|
||||
self._entries: list[Any] = []
|
||||
|
||||
@property
|
||||
def length(self) -> int:
|
||||
return len(self._entries)
|
||||
|
||||
def append(self, message: Any):
|
||||
self._entries.append(message)
|
||||
|
||||
def extend(self, messages: list[Any]):
|
||||
self._entries.extend(messages)
|
||||
|
||||
def clear(self):
|
||||
self._entries.clear()
|
||||
|
||||
def to_messages(self) -> list[Any]:
|
||||
"""返回浅拷贝 of 消息列表,防止外部意外修改"""
|
||||
return list(self._entries)
|
||||
|
||||
|
||||
class AppendOnlyContextManager:
|
||||
"""
|
||||
为大模型 Prefix Cache 深度定制 of 上下文管理器。
|
||||
将上下文拆分为绝对稳定 of Prefix (系统提示/工具) 和只增不减 of Log (对话历史)。
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.prefix = StablePrefix()
|
||||
self.log = AppendOnlyLog()
|
||||
self._last_sync_count = 0
|
||||
self._synced_digest = 0
|
||||
|
||||
def build(
|
||||
self, system_prompt: list[str], tools: list[Any]
|
||||
) -> tuple[list[str], list[Any], list[Any]]:
|
||||
"""装配并获取当前 of 完整上下文元组:(系统提示词, 历史消息, 工具)"""
|
||||
self.prefix.build(system_prompt, tools)
|
||||
sys_p, ts = self.prefix.to_context()
|
||||
return sys_p, self.log.to_messages(), ts
|
||||
|
||||
def sync_messages(self, normalized_messages: list[Any]):
|
||||
"""
|
||||
同步消息游标。
|
||||
通过滚动摘要算法(Rolling Digest)自动检测历史消息是否被就地篡改或截断。
|
||||
如果是,则自动重置基线;否则执行极速追加写入。
|
||||
"""
|
||||
if 0 < self._last_sync_count <= len(normalized_messages):
|
||||
synced_part = normalized_messages[: self._last_sync_count]
|
||||
if self._compute_digest(synced_part) != self._synced_digest:
|
||||
self.log.clear()
|
||||
self._last_sync_count = 0
|
||||
|
||||
if len(normalized_messages) < self._last_sync_count:
|
||||
self.log.clear()
|
||||
self._last_sync_count = 0
|
||||
|
||||
new_msgs = normalized_messages[self._last_sync_count :]
|
||||
for msg in new_msgs:
|
||||
self.log.append(msg)
|
||||
|
||||
self._last_sync_count = len(normalized_messages)
|
||||
self._synced_digest = self._compute_digest(normalized_messages)
|
||||
|
||||
def invalidate(self):
|
||||
"""使前缀快照失效(通常在模型发生变更时调用)"""
|
||||
self.prefix.invalidate()
|
||||
|
||||
def reset_sync_cursor(self):
|
||||
"""强制重置对话历史游标和日志"""
|
||||
self.log.clear()
|
||||
self._last_sync_count = 0
|
||||
self._synced_digest = 0
|
||||
|
||||
def _compute_digest(self, messages: list[Any]) -> int:
|
||||
"""核心:计算消息列表 of 指纹,用于识别内容篡改。包含 role 与 content。"""
|
||||
from zhenxun.services.ai.core.engine.context_renderer import ContextConverter
|
||||
|
||||
payloads = []
|
||||
flattened = ContextConverter.flatten_to_llm_messages(messages)
|
||||
for msg in flattened:
|
||||
try:
|
||||
d = model_dump(msg, include={"role", "content"})
|
||||
payloads.append(d)
|
||||
except Exception:
|
||||
payloads.append(str(msg))
|
||||
|
||||
def _default(obj):
|
||||
if isinstance(obj, bytes):
|
||||
return "<bytes>"
|
||||
return str(obj)
|
||||
|
||||
json_str = json.dumps(payloads, default=_default, sort_keys=True)
|
||||
return int(hashlib.md5(json_str.encode("utf-8")).hexdigest()[:8], 16)
|
||||
@@ -0,0 +1,52 @@
|
||||
from collections.abc import Sequence
|
||||
from typing import Any
|
||||
|
||||
from zhenxun.services.ai.core.messages import AgentEvent, AgentMessage, LLMMessage
|
||||
from zhenxun.services.log import logger
|
||||
|
||||
|
||||
class ContextConverter:
|
||||
"""
|
||||
上下文边界降维转换器。
|
||||
负责将内存中混合了 AgentEvent 与 LLMMessage 的业务事件流,
|
||||
安全拍平为底层大模型可读的原生 API 载体。
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def flatten_to_llm_messages(
|
||||
messages: Sequence[AgentMessage], context: Any | None = None
|
||||
) -> list[LLMMessage]:
|
||||
flattened: list[LLMMessage] = []
|
||||
|
||||
for msg in messages:
|
||||
if isinstance(msg, LLMMessage):
|
||||
flattened.append(msg)
|
||||
elif isinstance(msg, AgentEvent):
|
||||
try:
|
||||
res = msg.to_llm_message(context)
|
||||
if res is None:
|
||||
continue
|
||||
|
||||
if isinstance(res, str):
|
||||
flattened.append(LLMMessage.system(res))
|
||||
elif isinstance(res, LLMMessage):
|
||||
flattened.append(res)
|
||||
elif isinstance(res, list):
|
||||
flattened.extend(res)
|
||||
else:
|
||||
logger.warning(
|
||||
f"事件 {msg.__class__.__name__} 的 to_llm_message "
|
||||
f"返回了不支持的类型: {type(res)}"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"业务事件 [{msg.__class__.__name__}] "
|
||||
f"在降维渲染为大模型 Prompt 时发生崩溃: {e}\n"
|
||||
f"防呆拦截:请检查该事件 to_llm_message 方法的实现。"
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
f"ContextConverter 遇到未知类型的消息,已跳过: {type(msg)}"
|
||||
)
|
||||
|
||||
return flattened
|
||||
@@ -0,0 +1,214 @@
|
||||
import types
|
||||
from typing import Any, Generic, Union, cast, get_origin
|
||||
|
||||
import json_repair
|
||||
from nonebot.compat import type_validate_json
|
||||
from pydantic import BaseModel, Field, ValidationError, create_model
|
||||
|
||||
from zhenxun.services.ai.core.exceptions import (
|
||||
ControlFlowExit,
|
||||
ModelRetry,
|
||||
SchemaParseError,
|
||||
)
|
||||
from zhenxun.services.ai.core.models import ToolDefinition
|
||||
from zhenxun.services.ai.run.models import OutputDataT
|
||||
from zhenxun.services.ai.tools.core.tool import BaseTool
|
||||
from zhenxun.services.ai.tools.models import StructuredSubmissionResult, ToolResult
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.pydantic_compat import model_json_schema, model_validate
|
||||
|
||||
DEFAULT_IVR_TEMPLATE = (
|
||||
"### ❌ [输出内容或格式验证失败]\n"
|
||||
"你的上一次输出未能通过系统的校验与规则检查。请立即启动修正流程:\n\n"
|
||||
"**错误反馈报告:**\n"
|
||||
"> {error_msg}\n\n"
|
||||
"**修正要求:** 请结合反馈报告,"
|
||||
"仔细反思你的输出内容或格式,\n"
|
||||
"并重新生成正确的数据以满足所有的规则与规范。"
|
||||
)
|
||||
|
||||
|
||||
class BaseOutputProcessor(Generic[OutputDataT]):
|
||||
"""
|
||||
统一的结构化输出处理器。
|
||||
负责管理 Schema 生成、Prompt 约束注入以及最终的
|
||||
JSON 反序列化和业务校验。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
response_model: type[Any] | None = None,
|
||||
error_template: str | None = None,
|
||||
raw_schema: dict[str, Any] | None = None,
|
||||
):
|
||||
"""
|
||||
初始化结构化输出处理器。
|
||||
|
||||
参数:
|
||||
response_model: 期望的输出目标 Pydantic 模型类或 Union 类型,默认 None。
|
||||
error_template: 当 JSON 解析或模型验证失败时,反馈给大模型的 IVR 纠错提示词模板,默认 None。
|
||||
raw_schema: 显式传入的原始 JSON Schema 字典,
|
||||
如果不为 None 则跳过根据 response_model 生成,默认 None。
|
||||
""" # noqa: E501
|
||||
self.original_model = response_model
|
||||
self.error_template = error_template or DEFAULT_IVR_TEMPLATE
|
||||
self.raw_schema = raw_schema
|
||||
self.target_model = None
|
||||
self.is_union_wrapped = False
|
||||
|
||||
if response_model is not None:
|
||||
self.target_model, self.is_union_wrapped = self._create_union_wrapper(
|
||||
response_model
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _create_union_wrapper(union_type: Any) -> tuple[type[BaseModel], bool]:
|
||||
"""[私有方法] 如果是 Union 类型,
|
||||
动态构建带 kind 区分字段的模型"""
|
||||
origin = get_origin(union_type)
|
||||
union_types = [Union]
|
||||
if hasattr(types, "UnionType"):
|
||||
union_types.append(types.UnionType)
|
||||
|
||||
if origin not in union_types:
|
||||
return union_type, False
|
||||
|
||||
UnionWrapper = create_model(
|
||||
"UnionResponseWrapper",
|
||||
result=(
|
||||
union_type,
|
||||
Field(..., description="根据你的决策,输出对应的结构化数据"),
|
||||
),
|
||||
)
|
||||
return UnionWrapper, True
|
||||
|
||||
def get_json_schema(self) -> dict[str, Any]:
|
||||
"""提取目标模型的 JSON Schema"""
|
||||
if self.raw_schema is not None:
|
||||
return self.raw_schema
|
||||
if self.target_model is None:
|
||||
raise ValueError("未提供 response_model 或 raw_schema")
|
||||
try:
|
||||
return model_json_schema(self.target_model)
|
||||
except AttributeError:
|
||||
return self.target_model.schema()
|
||||
|
||||
def _parse_and_validate(self, text: str) -> Any:
|
||||
"""[私有方法] 执行带有容错修复的 JSON 解析与模型验证"""
|
||||
if self.raw_schema is not None:
|
||||
import json
|
||||
|
||||
try:
|
||||
return json.loads(text)
|
||||
except Exception:
|
||||
try:
|
||||
return json_repair.loads(text, skip_json_loads=True)
|
||||
except Exception as repair_error:
|
||||
raise SchemaParseError(f"JSON格式损坏: {repair_error}")
|
||||
if self.target_model is None:
|
||||
raise SchemaParseError("未提供 response_model 或 raw_schema")
|
||||
try:
|
||||
return type_validate_json(self.target_model, text)
|
||||
except (ValidationError, ValueError) as e:
|
||||
try:
|
||||
logger.warning(f"标准JSON解析失败,尝试使用json_repair修复: {e}")
|
||||
repaired_obj = json_repair.loads(text, skip_json_loads=True)
|
||||
return model_validate(self.target_model, repaired_obj)
|
||||
except Exception as repair_error:
|
||||
logger.error(
|
||||
f"LLM结构化输出校验最终失败: {repair_error}",
|
||||
e=repair_error,
|
||||
)
|
||||
raise SchemaParseError(
|
||||
f"JSON格式损坏或字段不匹配,未能通过Schema验证: {repair_error}"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"解析LLM结构化输出时发生未知错误: {e}", e=e)
|
||||
raise SchemaParseError(f"解析LLM的JSON输出时失败: {e}")
|
||||
|
||||
async def validate_and_parse(self, text: str, context: Any = None) -> OutputDataT:
|
||||
"""执行 JSON 解析与回调验证"""
|
||||
try:
|
||||
parsed_obj = self._parse_and_validate(text)
|
||||
|
||||
current_obj = parsed_obj
|
||||
|
||||
if getattr(self, "is_union_wrapped", False):
|
||||
current_obj = getattr(current_obj, "result")
|
||||
|
||||
final_obj = cast(OutputDataT, current_obj)
|
||||
|
||||
return final_obj
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
||||
|
||||
class SubmitFinalResultExecutable(BaseTool):
|
||||
"""
|
||||
动态生成的提交最终结果工具。
|
||||
用于将大模型的结构化输出拦截并终止 AgentExecutor 的循环。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
output_processor: BaseOutputProcessor,
|
||||
guardrails: list[Any] | None = None,
|
||||
):
|
||||
"""
|
||||
初始化提交最终结果的动态执行工具。
|
||||
|
||||
参数:
|
||||
output_processor: 绑定的结构化输出处理器,用于验证提交的最终结果。
|
||||
guardrails: 用于在结果输出前进行安全合规拦截的护栏中间件列表,默认 None。
|
||||
"""
|
||||
super().__init__(
|
||||
name="submit_final_result",
|
||||
description=(
|
||||
"当你完成所有必要的调查 and 思考后,"
|
||||
"必须且只能调用此工具来提交最终的结构化结果。"
|
||||
"提交后任务将立刻结束。"
|
||||
),
|
||||
)
|
||||
self.output_processor = output_processor
|
||||
self.guardrails = guardrails or []
|
||||
|
||||
async def get_definition(self, context: Any | None = None) -> ToolDefinition | None:
|
||||
if getattr(self, "_dynamic_def", None) is not None:
|
||||
return self._dynamic_def
|
||||
schema = self.output_processor.get_json_schema()
|
||||
return ToolDefinition(
|
||||
name=self.name,
|
||||
description=self.description,
|
||||
parameters=schema,
|
||||
)
|
||||
|
||||
async def execute(self, context: Any | None = None, **kwargs) -> ToolResult:
|
||||
parse_target = kwargs
|
||||
if isinstance(kwargs, dict):
|
||||
if "kwargs" in kwargs and len(kwargs) == 1:
|
||||
parse_target = kwargs["kwargs"]
|
||||
elif "result" in kwargs and len(kwargs) == 1:
|
||||
parse_target = kwargs["result"]
|
||||
|
||||
try:
|
||||
json_str = __import__("json").dumps(parse_target, ensure_ascii=False)
|
||||
final_obj = await self.output_processor.validate_and_parse(
|
||||
json_str, context=context
|
||||
)
|
||||
from zhenxun.services.ai.guardrails import GuardrailPipeline
|
||||
|
||||
pipeline = GuardrailPipeline(self.guardrails)
|
||||
json_str, final_obj = await pipeline.run_output_pipeline(
|
||||
json_str, final_obj, context
|
||||
)
|
||||
|
||||
return StructuredSubmissionResult(
|
||||
output="结构化数据已成功提交", parsed_obj=final_obj
|
||||
)
|
||||
except ControlFlowExit as e:
|
||||
raise e
|
||||
except ModelRetry as e:
|
||||
raise e
|
||||
except Exception as e:
|
||||
error_msg = f"系统捕获到解析异常:\n{e}"
|
||||
raise SchemaParseError(error_msg)
|
||||
@@ -0,0 +1,191 @@
|
||||
"""
|
||||
LLM Token 动态预估与上下文管理模块
|
||||
"""
|
||||
|
||||
from collections.abc import Sequence
|
||||
import math
|
||||
import re
|
||||
from typing import Any
|
||||
|
||||
from zhenxun.services.ai.core.messages import (
|
||||
AgentEvent,
|
||||
AgentMessage,
|
||||
AudioPart,
|
||||
FilePart,
|
||||
ImagePart,
|
||||
LLMMessage,
|
||||
TextPart,
|
||||
ThoughtPart,
|
||||
ToolCallPart,
|
||||
ToolMessage,
|
||||
ToolReturnPart,
|
||||
UsageInfo,
|
||||
VideoPart,
|
||||
)
|
||||
|
||||
|
||||
class TokenCounter:
|
||||
"""
|
||||
Token 计数器
|
||||
基于确定性规则,摆脱外部库依赖,提供绝对稳定的 Token 消耗预估基线。
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def _count_text(text: str) -> int:
|
||||
"""基于字符类型近似计算纯文本的 Token 消耗量。"""
|
||||
if not text:
|
||||
return 0
|
||||
cjk_chars = len(re.findall(r"[\u4e00-\u9fff\u3000-\u303f\uff00-\uffef]", text))
|
||||
ascii_chars = len(text) - cjk_chars
|
||||
return math.ceil(cjk_chars * 1.2 + ascii_chars * 0.3)
|
||||
|
||||
@staticmethod
|
||||
def _count_image(resolution_hint: str | None, model_name: str) -> int:
|
||||
"""根据分辨率策略和模型厂商计算单张图片的 Token 消耗量。"""
|
||||
if "gemini" in model_name.lower():
|
||||
res = (resolution_hint or "").upper()
|
||||
if "ULTRA_HIGH" in res:
|
||||
return 6192
|
||||
if "HIGH" in res:
|
||||
return 3096
|
||||
if "LOW" in res:
|
||||
return 258
|
||||
return 1032
|
||||
return 765
|
||||
|
||||
@classmethod
|
||||
def count_tools_schema(cls, obj: dict | list | str | Any) -> int:
|
||||
"""递归计算 JSON Schema 结构在被大模型作为工具时的 Token 开销。"""
|
||||
if isinstance(obj, dict):
|
||||
cost = len(obj.keys()) * 12
|
||||
for k, v in obj.items():
|
||||
if k == "description" and isinstance(v, str):
|
||||
cost += int(len(v) * 0.3)
|
||||
else:
|
||||
cost += cls.count_tools_schema(v)
|
||||
return cost
|
||||
elif isinstance(obj, list):
|
||||
return sum(cls.count_tools_schema(item) for item in obj)
|
||||
return 0
|
||||
|
||||
@classmethod
|
||||
def count_message(cls, msg: LLMMessage, model_name: str) -> int:
|
||||
"""累加计算单条包含多模态片段和工具调用的消息 Token 总数。"""
|
||||
if msg.token_cost is not None:
|
||||
return msg.token_cost
|
||||
|
||||
total_tokens = 4
|
||||
|
||||
if isinstance(msg, ToolMessage):
|
||||
total_tokens += 40
|
||||
|
||||
if isinstance(msg.content, str):
|
||||
total_tokens += cls._count_text(msg.content)
|
||||
elif isinstance(msg.content, list):
|
||||
for part in msg.content:
|
||||
if isinstance(part, TextPart) and part.text:
|
||||
total_tokens += cls._count_text(part.text)
|
||||
elif isinstance(part, ImagePart):
|
||||
total_tokens += cls._count_image(
|
||||
getattr(part, "media_resolution", None), model_name
|
||||
)
|
||||
elif isinstance(part, VideoPart | AudioPart | FilePart):
|
||||
total_tokens += 1032
|
||||
elif isinstance(part, ThoughtPart) and part.thought_text:
|
||||
total_tokens += cls._count_text(part.thought_text)
|
||||
elif isinstance(part, ToolCallPart) and part.args:
|
||||
total_tokens += cls._count_text(str(part.args))
|
||||
elif isinstance(part, ToolReturnPart) and part.output:
|
||||
total_tokens += cls._count_text(str(part.output))
|
||||
|
||||
msg.token_cost = total_tokens
|
||||
return total_tokens
|
||||
|
||||
@classmethod
|
||||
def count_context(
|
||||
cls, messages: Sequence[AgentMessage], model_name: str, base_overhead: int = 0
|
||||
) -> int:
|
||||
"""计算整个对话历史上下文的 Token 总和。"""
|
||||
if not messages:
|
||||
return base_overhead
|
||||
|
||||
total = base_overhead
|
||||
for msg in messages:
|
||||
if isinstance(msg, AgentEvent):
|
||||
try:
|
||||
res = msg.to_llm_message(None)
|
||||
if res is None:
|
||||
continue
|
||||
if isinstance(res, str):
|
||||
total += cls.count_message(LLMMessage.system(res), model_name)
|
||||
elif isinstance(res, list):
|
||||
total += sum(cls.count_message(m, model_name) for m in res)
|
||||
elif isinstance(res, LLMMessage):
|
||||
total += cls.count_message(res, model_name)
|
||||
except Exception:
|
||||
pass
|
||||
else:
|
||||
total += cls.count_message(msg, model_name)
|
||||
return total
|
||||
|
||||
|
||||
token_counter = TokenCounter
|
||||
|
||||
|
||||
def parse_usage_info(usage_info: dict | None) -> UsageInfo:
|
||||
"""
|
||||
全协议统一遥测解析器 (Universal Telemetry Parser)
|
||||
兼容 OpenAI Standard、OpenAI Responses (v1/responses) 以及 Gemini (usageMetadata)。
|
||||
"""
|
||||
if not usage_info or not isinstance(usage_info, dict):
|
||||
return UsageInfo()
|
||||
|
||||
prompt = 0
|
||||
completion = 0
|
||||
total = 0
|
||||
cache_hit = 0
|
||||
cache_miss = 0
|
||||
reasoning = 0
|
||||
|
||||
if "promptTokenCount" in usage_info or "candidatesTokenCount" in usage_info:
|
||||
prompt = usage_info.get("promptTokenCount", 0)
|
||||
completion = usage_info.get("candidatesTokenCount", 0)
|
||||
total = usage_info.get("totalTokenCount", 0)
|
||||
reasoning = usage_info.get("thoughtsTokenCount", 0)
|
||||
cache_hit = usage_info.get("cachedContentTokenCount", 0)
|
||||
|
||||
elif "input_tokens" in usage_info or "output_tokens" in usage_info:
|
||||
prompt = usage_info.get("input_tokens", 0)
|
||||
completion = usage_info.get("output_tokens", 0)
|
||||
total = usage_info.get("total_tokens", 0)
|
||||
cache_hit = (usage_info.get("input_tokens_details") or {}).get(
|
||||
"cached_tokens", 0
|
||||
)
|
||||
reasoning = (usage_info.get("output_tokens_details") or {}).get(
|
||||
"reasoning_tokens", 0
|
||||
)
|
||||
|
||||
else:
|
||||
prompt = usage_info.get("prompt_tokens", 0)
|
||||
completion = usage_info.get("completion_tokens", 0)
|
||||
total = usage_info.get("total_tokens", 0)
|
||||
|
||||
cache_hit = usage_info.get("prompt_cache_hit_tokens") or (
|
||||
usage_info.get("prompt_tokens_details") or {}
|
||||
).get("cached_tokens", 0)
|
||||
cache_miss = usage_info.get("prompt_cache_miss_tokens", 0)
|
||||
reasoning = (usage_info.get("completion_tokens_details") or {}).get(
|
||||
"reasoning_tokens", 0
|
||||
)
|
||||
|
||||
if cache_miss == 0 and prompt > 0:
|
||||
cache_miss = max(0, prompt - cache_hit)
|
||||
|
||||
return UsageInfo(
|
||||
prompt_tokens=prompt,
|
||||
completion_tokens=completion,
|
||||
total_tokens=total,
|
||||
prompt_cache_hit_tokens=cache_hit,
|
||||
prompt_cache_miss_tokens=cache_miss,
|
||||
reasoning_tokens=reasoning,
|
||||
)
|
||||
Reference in New Issue
Block a user