Record streaming summaries for agent tool calls

This commit is contained in:
jxxghp
2026-06-21 21:43:44 +08:00
parent 3c74f1bf58
commit 7358b4df14
7 changed files with 333 additions and 104 deletions
+41 -2
View File
@@ -3,7 +3,7 @@ import re
import shutil
from collections.abc import Awaitable, Callable
from pathlib import Path
from typing import Annotated, List, Optional
from typing import Annotated, Any, List, Optional
from typing import NotRequired, TypedDict
import yaml # noqa
@@ -15,6 +15,7 @@ from langchain.agents.middleware.types import (
ModelRequest,
ModelResponse,
ResponseT,
ToolCallRequest,
)
from langchain.agents.middleware.types import PrivateStateAttr # noqa
from langchain_core.runnables import RunnableConfig
@@ -525,6 +526,7 @@ class SkillsMiddleware(AgentMiddleware[SkillsState, ContextT, ResponseT]): # no
*,
sources: list[str],
bundled_skills_dir: str | None = None,
stream_handler: Optional[Any] = None,
) -> None:
"""初始化 Skill 中间件。
@@ -535,9 +537,12 @@ class SkillsMiddleware(AgentMiddleware[SkillsState, ContextT, ResponseT]): # no
bundled_skills_dir : str | None
项目内置技能目录路径。若提供,在首次加载前会将其中不存在于
sources 首个目录的技能自动复制过去。
stream_handler : Optional[Any]
流式输出处理器,用于记录 skill 工具调用摘要。
"""
self.sources = sources
self.bundled_skills_dir = bundled_skills_dir
self.stream_handler = stream_handler
self.system_prompt_template = SKILLS_SYSTEM_PROMPT
self._skill_provider = _SkillToolProvider(sources=sources)
self.tools = [
@@ -584,7 +589,8 @@ class SkillsMiddleware(AgentMiddleware[SkillsState, ContextT, ResponseT]): # no
skills_catalog=_format_skill_tool_catalog(skills)
)
def _format_skills_list(self, skills: list[SkillMetadata]) -> str:
@staticmethod
def _format_skills_list(skills: list[SkillMetadata]) -> str:
"""格式化技能元数据列表用于系统提示词。"""
if not skills:
return "(No skills available yet.)"
@@ -657,5 +663,38 @@ class SkillsMiddleware(AgentMiddleware[SkillsState, ContextT, ResponseT]): # no
modified_request = self.modify_request(request)
return await handler(modified_request)
async def awrap_tool_call(
self,
request: ToolCallRequest,
handler: Callable[[ToolCallRequest], Awaitable[Any]],
) -> Any:
"""在 skill 工具执行时记录聚合摘要。"""
tool = request.tool
tool_name = getattr(tool, "name", None)
if tool_name != SKILL_TOOL_NAME:
return await handler(request)
tool_call = request.tool_call or {}
tool_args = tool_call.get("args") or {}
if not isinstance(tool_args, dict):
tool_args = {}
logger.info(
f"开始执行 Skill 工具: name={tool_args.get('name') or '-'}, "
f"explanation={tool_args.get('explanation') or '-'}"
)
if self.stream_handler and getattr(self.stream_handler, "is_streaming", False):
self.stream_handler.record_tool_call(
tool_name=SKILL_TOOL_NAME,
tool_message="Skill loaded",
tool_kwargs=tool_args,
)
try:
result = await handler(request)
except Exception as err:
logger.error(f"Skill 工具执行失败: error={err}")
raise
logger.info("Skill 工具执行完成")
return result
__all__ = ["SKILL_TOOL_NAME", "SkillMetadata", "SkillsMiddleware"]