mirror of
https://github.com/jxxghp/MoviePilot.git
synced 2026-09-05 23:47:41 +08:00
Record streaming summaries for agent tool calls
This commit is contained in:
@@ -1023,6 +1023,7 @@ class MoviePilotAgent:
|
|||||||
skills_middleware = SkillsMiddleware(
|
skills_middleware = SkillsMiddleware(
|
||||||
sources=[str(agent_runtime_manager.skills_dir)],
|
sources=[str(agent_runtime_manager.skills_dir)],
|
||||||
bundled_skills_dir=str(settings.ROOT_PATH / "skills"),
|
bundled_skills_dir=str(settings.ROOT_PATH / "skills"),
|
||||||
|
stream_handler=self.stream_handler,
|
||||||
)
|
)
|
||||||
skill_tools = list(getattr(skills_middleware, "tools", []) or [])
|
skill_tools = list(getattr(skills_middleware, "tools", []) or [])
|
||||||
activity_log_middleware = None
|
activity_log_middleware = None
|
||||||
@@ -1030,6 +1031,7 @@ class MoviePilotAgent:
|
|||||||
if self.has_message_context:
|
if self.has_message_context:
|
||||||
activity_log_middleware = ActivityLogMiddleware(
|
activity_log_middleware = ActivityLogMiddleware(
|
||||||
activity_dir=str(agent_runtime_manager.activity_dir),
|
activity_dir=str(agent_runtime_manager.activity_dir),
|
||||||
|
stream_handler=self.stream_handler,
|
||||||
)
|
)
|
||||||
activity_log_tools = list(
|
activity_log_tools = list(
|
||||||
getattr(activity_log_middleware, "tools", []) or []
|
getattr(activity_log_middleware, "tools", []) or []
|
||||||
|
|||||||
@@ -22,6 +22,7 @@ from langchain.agents.middleware.types import (
|
|||||||
ModelResponse,
|
ModelResponse,
|
||||||
PrivateStateAttr, # noqa
|
PrivateStateAttr, # noqa
|
||||||
ResponseT,
|
ResponseT,
|
||||||
|
ToolCallRequest,
|
||||||
)
|
)
|
||||||
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
|
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
|
||||||
from langchain_core.tools import StructuredTool
|
from langchain_core.tools import StructuredTool
|
||||||
@@ -495,10 +496,13 @@ class ActivityLogMiddleware(AgentMiddleware[ActivityLogState, ContextT, Response
|
|||||||
activity_dir: str,
|
activity_dir: str,
|
||||||
retention_days: int = DEFAULT_RETENTION_DAYS,
|
retention_days: int = DEFAULT_RETENTION_DAYS,
|
||||||
prompt_load_days: int = PROMPT_LOAD_DAYS,
|
prompt_load_days: int = PROMPT_LOAD_DAYS,
|
||||||
|
stream_handler: Optional[Any] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
|
"""初始化活动日志中间件。"""
|
||||||
self.activity_dir = activity_dir
|
self.activity_dir = activity_dir
|
||||||
self.retention_days = retention_days
|
self.retention_days = retention_days
|
||||||
self.prompt_load_days = prompt_load_days
|
self.prompt_load_days = prompt_load_days
|
||||||
|
self.stream_handler = stream_handler
|
||||||
self._tool_provider = _ActivityLogToolProvider(activity_dir=activity_dir)
|
self._tool_provider = _ActivityLogToolProvider(activity_dir=activity_dir)
|
||||||
self.tools = [
|
self.tools = [
|
||||||
StructuredTool.from_function(
|
StructuredTool.from_function(
|
||||||
@@ -646,6 +650,39 @@ class ActivityLogMiddleware(AgentMiddleware[ActivityLogState, ContextT, Response
|
|||||||
modified_request = self.modify_request(request)
|
modified_request = self.modify_request(request)
|
||||||
return await handler(modified_request)
|
return await handler(modified_request)
|
||||||
|
|
||||||
|
async def awrap_tool_call(
|
||||||
|
self,
|
||||||
|
request: ToolCallRequest,
|
||||||
|
handler: Callable[[ToolCallRequest], Awaitable[Any]],
|
||||||
|
) -> Any:
|
||||||
|
"""在活动日志查询工具执行时记录聚合摘要。"""
|
||||||
|
tool = request.tool
|
||||||
|
tool_name = getattr(tool, "name", None)
|
||||||
|
if tool_name != QUERY_ACTIVITY_LOG_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"开始执行活动日志查询工具: keyword={tool_args.get('keyword') or '-'}, "
|
||||||
|
f"date={tool_args.get('date') or '-'}"
|
||||||
|
)
|
||||||
|
if self.stream_handler and getattr(self.stream_handler, "is_streaming", False):
|
||||||
|
self.stream_handler.record_tool_call(
|
||||||
|
tool_name=QUERY_ACTIVITY_LOG_TOOL_NAME,
|
||||||
|
tool_message=QUERY_ACTIVITY_LOG_TOOL_DESCRIPTION,
|
||||||
|
tool_kwargs=tool_args,
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
result = await handler(request)
|
||||||
|
except Exception as err:
|
||||||
|
logger.error(f"活动日志查询工具执行失败: error={err}")
|
||||||
|
raise
|
||||||
|
logger.info("活动日志查询工具执行完成")
|
||||||
|
return result
|
||||||
|
|
||||||
async def aafter_agent(
|
async def aafter_agent(
|
||||||
self, state: ActivityLogState, runtime: Runtime
|
self, state: ActivityLogState, runtime: Runtime
|
||||||
) -> Optional[dict[str, Any]]:
|
) -> Optional[dict[str, Any]]:
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ import re
|
|||||||
import shutil
|
import shutil
|
||||||
from collections.abc import Awaitable, Callable
|
from collections.abc import Awaitable, Callable
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Annotated, List, Optional
|
from typing import Annotated, Any, List, Optional
|
||||||
from typing import NotRequired, TypedDict
|
from typing import NotRequired, TypedDict
|
||||||
|
|
||||||
import yaml # noqa
|
import yaml # noqa
|
||||||
@@ -15,6 +15,7 @@ from langchain.agents.middleware.types import (
|
|||||||
ModelRequest,
|
ModelRequest,
|
||||||
ModelResponse,
|
ModelResponse,
|
||||||
ResponseT,
|
ResponseT,
|
||||||
|
ToolCallRequest,
|
||||||
)
|
)
|
||||||
from langchain.agents.middleware.types import PrivateStateAttr # noqa
|
from langchain.agents.middleware.types import PrivateStateAttr # noqa
|
||||||
from langchain_core.runnables import RunnableConfig
|
from langchain_core.runnables import RunnableConfig
|
||||||
@@ -525,6 +526,7 @@ class SkillsMiddleware(AgentMiddleware[SkillsState, ContextT, ResponseT]): # no
|
|||||||
*,
|
*,
|
||||||
sources: list[str],
|
sources: list[str],
|
||||||
bundled_skills_dir: str | None = None,
|
bundled_skills_dir: str | None = None,
|
||||||
|
stream_handler: Optional[Any] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""初始化 Skill 中间件。
|
"""初始化 Skill 中间件。
|
||||||
|
|
||||||
@@ -535,9 +537,12 @@ class SkillsMiddleware(AgentMiddleware[SkillsState, ContextT, ResponseT]): # no
|
|||||||
bundled_skills_dir : str | None
|
bundled_skills_dir : str | None
|
||||||
项目内置技能目录路径。若提供,在首次加载前会将其中不存在于
|
项目内置技能目录路径。若提供,在首次加载前会将其中不存在于
|
||||||
sources 首个目录的技能自动复制过去。
|
sources 首个目录的技能自动复制过去。
|
||||||
|
stream_handler : Optional[Any]
|
||||||
|
流式输出处理器,用于记录 skill 工具调用摘要。
|
||||||
"""
|
"""
|
||||||
self.sources = sources
|
self.sources = sources
|
||||||
self.bundled_skills_dir = bundled_skills_dir
|
self.bundled_skills_dir = bundled_skills_dir
|
||||||
|
self.stream_handler = stream_handler
|
||||||
self.system_prompt_template = SKILLS_SYSTEM_PROMPT
|
self.system_prompt_template = SKILLS_SYSTEM_PROMPT
|
||||||
self._skill_provider = _SkillToolProvider(sources=sources)
|
self._skill_provider = _SkillToolProvider(sources=sources)
|
||||||
self.tools = [
|
self.tools = [
|
||||||
@@ -584,7 +589,8 @@ class SkillsMiddleware(AgentMiddleware[SkillsState, ContextT, ResponseT]): # no
|
|||||||
skills_catalog=_format_skill_tool_catalog(skills)
|
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:
|
if not skills:
|
||||||
return "(No skills available yet.)"
|
return "(No skills available yet.)"
|
||||||
@@ -657,5 +663,38 @@ class SkillsMiddleware(AgentMiddleware[SkillsState, ContextT, ResponseT]): # no
|
|||||||
modified_request = self.modify_request(request)
|
modified_request = self.modify_request(request)
|
||||||
return await handler(modified_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"]
|
__all__ = ["SKILL_TOOL_NAME", "SkillMetadata", "SkillsMiddleware"]
|
||||||
|
|||||||
@@ -313,6 +313,31 @@ def _format_datetime(value: Optional[datetime]) -> Optional[str]:
|
|||||||
return value.strftime("%Y-%m-%d %H:%M:%S")
|
return value.strftime("%Y-%m-%d %H:%M:%S")
|
||||||
|
|
||||||
|
|
||||||
|
def _extract_tool_call_args(request: ToolCallRequest) -> dict[str, Any]:
|
||||||
|
"""提取工具调用参数,并规整为字典。"""
|
||||||
|
tool_call = request.tool_call or {}
|
||||||
|
tool_args = tool_call.get("args") or {}
|
||||||
|
if not isinstance(tool_args, dict):
|
||||||
|
return {}
|
||||||
|
return tool_args
|
||||||
|
|
||||||
|
|
||||||
|
def _record_subagent_tool_call(
|
||||||
|
*,
|
||||||
|
stream_handler: Any,
|
||||||
|
tool_name: str,
|
||||||
|
tool_args: dict[str, Any],
|
||||||
|
) -> None:
|
||||||
|
"""在流式处理器中记录子代理工具调用摘要。"""
|
||||||
|
if not stream_handler or not getattr(stream_handler, "is_streaming", False):
|
||||||
|
return
|
||||||
|
stream_handler.record_tool_call(
|
||||||
|
tool_name=tool_name,
|
||||||
|
tool_message="Subagent invoked",
|
||||||
|
tool_kwargs=tool_args,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class _SubAgentAgentProvider:
|
class _SubAgentAgentProvider:
|
||||||
"""子代理图懒加载与执行器。"""
|
"""子代理图懒加载与执行器。"""
|
||||||
|
|
||||||
@@ -409,8 +434,10 @@ class MoviePilotSubAgentMiddleware(AgentMiddleware):
|
|||||||
tools: list[BaseTool],
|
tools: list[BaseTool],
|
||||||
system_prompt: str = SUBAGENT_PARENT_PROMPT,
|
system_prompt: str = SUBAGENT_PARENT_PROMPT,
|
||||||
task_description: str = SUBAGENT_TASK_DESCRIPTION,
|
task_description: str = SUBAGENT_TASK_DESCRIPTION,
|
||||||
|
stream_handler: Any = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
self.system_prompt = system_prompt
|
self.system_prompt = system_prompt
|
||||||
|
self.stream_handler = stream_handler
|
||||||
self._provider = _SubAgentAgentProvider(
|
self._provider = _SubAgentAgentProvider(
|
||||||
model=model,
|
model=model,
|
||||||
profiles=profiles,
|
profiles=profiles,
|
||||||
@@ -453,6 +480,35 @@ class MoviePilotSubAgentMiddleware(AgentMiddleware):
|
|||||||
)
|
)
|
||||||
return await handler(request.override(system_message=new_system_message))
|
return await handler(request.override(system_message=new_system_message))
|
||||||
|
|
||||||
|
async def awrap_tool_call(
|
||||||
|
self,
|
||||||
|
request: ToolCallRequest,
|
||||||
|
handler: Callable[[ToolCallRequest], Awaitable[Any]],
|
||||||
|
) -> Any:
|
||||||
|
"""在 task 子代理工具执行时记录聚合摘要。"""
|
||||||
|
tool = request.tool
|
||||||
|
tool_name = getattr(tool, "name", None)
|
||||||
|
if tool_name != SUBAGENT_TASK_TOOL_NAME:
|
||||||
|
return await handler(request)
|
||||||
|
|
||||||
|
tool_args = _extract_tool_call_args(request)
|
||||||
|
logger.info(
|
||||||
|
f"开始执行子代理工具: tool_name={tool_name}, "
|
||||||
|
f"subagent_type={tool_args.get('subagent_type') or '-'}"
|
||||||
|
)
|
||||||
|
_record_subagent_tool_call(
|
||||||
|
stream_handler=self.stream_handler,
|
||||||
|
tool_name=SUBAGENT_TASK_TOOL_NAME,
|
||||||
|
tool_args=tool_args,
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
result = await handler(request)
|
||||||
|
except Exception as err:
|
||||||
|
logger.error(f"子代理工具执行失败: tool_name={tool_name}, error={err}")
|
||||||
|
raise
|
||||||
|
logger.info(f"子代理工具执行完成: tool_name={tool_name}")
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
class SubAgentTaskControlMiddleware(AgentMiddleware):
|
class SubAgentTaskControlMiddleware(AgentMiddleware):
|
||||||
"""提供异步子代理任务调度工具的中间件。"""
|
"""提供异步子代理任务调度工具的中间件。"""
|
||||||
@@ -464,8 +520,10 @@ class SubAgentTaskControlMiddleware(AgentMiddleware):
|
|||||||
profiles: tuple[_SubAgentProfile, ...],
|
profiles: tuple[_SubAgentProfile, ...],
|
||||||
tools: list[BaseTool],
|
tools: list[BaseTool],
|
||||||
task_description: str = SUBAGENT_CONTROL_DESCRIPTION,
|
task_description: str = SUBAGENT_CONTROL_DESCRIPTION,
|
||||||
|
stream_handler: Any = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""初始化异步子代理调度中间件。"""
|
"""初始化异步子代理调度中间件。"""
|
||||||
|
self.stream_handler = stream_handler
|
||||||
self._provider = _SubAgentAgentProvider(
|
self._provider = _SubAgentAgentProvider(
|
||||||
model=model,
|
model=model,
|
||||||
profiles=profiles,
|
profiles=profiles,
|
||||||
@@ -987,98 +1045,37 @@ class SubAgentTaskControlMiddleware(AgentMiddleware):
|
|||||||
logger.info(f"Agent 结束,取消未完成子代理任务: tasks={len(unfinished_records)}")
|
logger.info(f"Agent 结束,取消未完成子代理任务: tasks={len(unfinished_records)}")
|
||||||
await self._cancel_records(unfinished_records)
|
await self._cancel_records(unfinished_records)
|
||||||
|
|
||||||
|
|
||||||
class SubAgentCallSummaryMiddleware(AgentMiddleware):
|
|
||||||
"""记录子代理调用次数的中间件。"""
|
|
||||||
|
|
||||||
def __init__(self, *, stream_handler: Any = None) -> None:
|
|
||||||
self.stream_handler = stream_handler
|
|
||||||
self.tools = []
|
|
||||||
|
|
||||||
async def awrap_tool_call(
|
async def awrap_tool_call(
|
||||||
self,
|
self,
|
||||||
request: ToolCallRequest,
|
request: ToolCallRequest,
|
||||||
handler: Callable[[ToolCallRequest], Awaitable[Any]],
|
handler: Callable[[ToolCallRequest], Awaitable[Any]],
|
||||||
) -> Any:
|
) -> Any:
|
||||||
"""在子代理任务工具执行时记录聚合摘要。"""
|
"""在 subagent_task 子代理工具执行时记录聚合摘要。"""
|
||||||
tool = request.tool
|
tool = request.tool
|
||||||
tool_name = getattr(tool, "name", None)
|
tool_name = getattr(tool, "name", None)
|
||||||
is_subagent_tool = tool_name in {
|
if tool_name != SUBAGENT_CONTROL_TOOL_NAME:
|
||||||
SUBAGENT_TASK_TOOL_NAME,
|
return await handler(request)
|
||||||
SUBAGENT_CONTROL_TOOL_NAME,
|
|
||||||
}
|
tool_args = _extract_tool_call_args(request)
|
||||||
if is_subagent_tool:
|
logger.info(
|
||||||
tool_call = request.tool_call or {}
|
f"开始执行子代理工具: tool_name={tool_name}, "
|
||||||
tool_args = tool_call.get("args") or {}
|
f"action={tool_args.get('action') or '-'}, "
|
||||||
if not isinstance(tool_args, dict):
|
f"subagent_type={tool_args.get('subagent_type') or '-'}"
|
||||||
tool_args = {}
|
)
|
||||||
logger.info(
|
_record_subagent_tool_call(
|
||||||
f"开始执行子代理工具: tool_name={tool_name}, "
|
stream_handler=self.stream_handler,
|
||||||
f"action={tool_args.get('action') or '-'}, "
|
tool_name=SUBAGENT_CONTROL_TOOL_NAME,
|
||||||
f"subagent_type={tool_args.get('subagent_type') or '-'}"
|
tool_args=tool_args,
|
||||||
)
|
)
|
||||||
if (
|
|
||||||
self.stream_handler
|
|
||||||
and getattr(self.stream_handler, "is_streaming", False)
|
|
||||||
):
|
|
||||||
self.stream_handler.record_tool_call(
|
|
||||||
tool_name=tool_name or SUBAGENT_TASK_TOOL_NAME,
|
|
||||||
tool_message="Subagent invoked",
|
|
||||||
tool_kwargs=tool_args,
|
|
||||||
)
|
|
||||||
try:
|
try:
|
||||||
result = await handler(request)
|
result = await handler(request)
|
||||||
except Exception as err:
|
except Exception as err:
|
||||||
if is_subagent_tool:
|
logger.error(f"子代理工具执行失败: tool_name={tool_name}, error={err}")
|
||||||
logger.error(f"子代理工具执行失败: tool_name={tool_name}, error={err}")
|
|
||||||
raise
|
raise
|
||||||
if is_subagent_tool:
|
logger.info(f"子代理工具执行完成: tool_name={tool_name}")
|
||||||
logger.info(f"子代理工具执行完成: tool_name={tool_name}")
|
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
def _deepagents_spec(
|
|
||||||
profiles: tuple[_SubAgentProfile, ...], tools: list[BaseTool]
|
|
||||||
) -> list[dict[str, Any]]:
|
|
||||||
"""将内置定义转换为 Deep Agents 子代理配置。"""
|
|
||||||
specs = []
|
|
||||||
for profile in profiles:
|
|
||||||
specs.append(
|
|
||||||
{
|
|
||||||
"name": profile.name,
|
|
||||||
"description": profile.description,
|
|
||||||
"prompt": profile.prompt,
|
|
||||||
"tools": _select_tools(tools, profile),
|
|
||||||
}
|
|
||||||
)
|
|
||||||
return specs
|
|
||||||
|
|
||||||
|
|
||||||
def _try_create_deepagents_middleware(
|
|
||||||
*,
|
|
||||||
profiles: tuple[_SubAgentProfile, ...],
|
|
||||||
tools: list[BaseTool],
|
|
||||||
model: BaseChatModel,
|
|
||||||
) -> Optional[AgentMiddleware]:
|
|
||||||
"""优先创建 Deep Agents 官方子代理中间件。"""
|
|
||||||
try:
|
|
||||||
from deepagents.backends import StateBackend
|
|
||||||
from deepagents.middleware.subagents import SubAgentMiddleware
|
|
||||||
|
|
||||||
return SubAgentMiddleware(
|
|
||||||
backend=StateBackend(),
|
|
||||||
subagents=_deepagents_spec(profiles, tools),
|
|
||||||
default_model=model,
|
|
||||||
system_prompt=SUBAGENT_PARENT_PROMPT,
|
|
||||||
task_description=SUBAGENT_TASK_DESCRIPTION,
|
|
||||||
)
|
|
||||||
except ImportError:
|
|
||||||
return None
|
|
||||||
except Exception as err:
|
|
||||||
logger.debug(f"Deep Agents 子代理中间件不可用,使用本地实现: {err}")
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
def create_subagent_middlewares(
|
def create_subagent_middlewares(
|
||||||
*,
|
*,
|
||||||
model: BaseChatModel,
|
model: BaseChatModel,
|
||||||
@@ -1089,21 +1086,17 @@ def create_subagent_middlewares(
|
|||||||
_builtin_subagent_profiles.cache_clear()
|
_builtin_subagent_profiles.cache_clear()
|
||||||
builtin_subagent_names.cache_clear()
|
builtin_subagent_names.cache_clear()
|
||||||
profiles = _builtin_subagent_profiles()
|
profiles = _builtin_subagent_profiles()
|
||||||
subagent_middleware = _try_create_deepagents_middleware(
|
subagent_middleware = MoviePilotSubAgentMiddleware(
|
||||||
|
model=model,
|
||||||
profiles=profiles,
|
profiles=profiles,
|
||||||
tools=tools,
|
tools=tools,
|
||||||
model=model,
|
stream_handler=stream_handler,
|
||||||
)
|
)
|
||||||
if subagent_middleware is None:
|
|
||||||
subagent_middleware = MoviePilotSubAgentMiddleware(
|
|
||||||
model=model,
|
|
||||||
profiles=profiles,
|
|
||||||
tools=tools,
|
|
||||||
)
|
|
||||||
control_middleware = SubAgentTaskControlMiddleware(
|
control_middleware = SubAgentTaskControlMiddleware(
|
||||||
model=model,
|
model=model,
|
||||||
profiles=profiles,
|
profiles=profiles,
|
||||||
tools=tools,
|
tools=tools,
|
||||||
|
stream_handler=stream_handler,
|
||||||
)
|
)
|
||||||
|
|
||||||
task_tools = [
|
task_tools = [
|
||||||
@@ -1113,7 +1106,6 @@ def create_subagent_middlewares(
|
|||||||
return [
|
return [
|
||||||
subagent_middleware,
|
subagent_middleware,
|
||||||
control_middleware,
|
control_middleware,
|
||||||
SubAgentCallSummaryMiddleware(stream_handler=stream_handler),
|
|
||||||
], task_tools
|
], task_tools
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ from langchain_core.messages import AIMessage, HumanMessage, SystemMessage, Tool
|
|||||||
|
|
||||||
from app.agent.middleware.activity_log import (
|
from app.agent.middleware.activity_log import (
|
||||||
ActivityLogMiddleware,
|
ActivityLogMiddleware,
|
||||||
|
QUERY_ACTIVITY_LOG_TOOL_DESCRIPTION,
|
||||||
QUERY_ACTIVITY_LOG_TOOL_NAME,
|
QUERY_ACTIVITY_LOG_TOOL_NAME,
|
||||||
_summarize_with_llm,
|
_summarize_with_llm,
|
||||||
load_activity_log_index,
|
load_activity_log_index,
|
||||||
@@ -260,6 +261,51 @@ def test_activity_log_middleware_query_tool_returns_json_payload(tmp_path):
|
|||||||
assert payload["entries"][0]["summary"] == "帮用户整理了电影 A"
|
assert payload["entries"][0]["summary"] == "帮用户整理了电影 A"
|
||||||
|
|
||||||
|
|
||||||
|
def test_activity_log_tool_call_records_streaming_summary(tmp_path):
|
||||||
|
"""query_activity_log 工具执行时应记录流式聚合摘要。"""
|
||||||
|
|
||||||
|
async def _run_test():
|
||||||
|
calls = []
|
||||||
|
stream_handler = SimpleNamespace(
|
||||||
|
is_streaming=True,
|
||||||
|
record_tool_call=lambda **kwargs: calls.append(kwargs),
|
||||||
|
)
|
||||||
|
middleware = ActivityLogMiddleware(
|
||||||
|
activity_dir=str(tmp_path),
|
||||||
|
stream_handler=stream_handler,
|
||||||
|
)
|
||||||
|
request = SimpleNamespace(
|
||||||
|
tool=SimpleNamespace(name=QUERY_ACTIVITY_LOG_TOOL_NAME),
|
||||||
|
tool_call={
|
||||||
|
"args": {
|
||||||
|
"keyword": "整理",
|
||||||
|
"date": "2026-06-18",
|
||||||
|
}
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _fake_handler(_request):
|
||||||
|
"""返回模拟工具结果。"""
|
||||||
|
return "ok"
|
||||||
|
|
||||||
|
result = await middleware.awrap_tool_call(request, _fake_handler)
|
||||||
|
return result, calls
|
||||||
|
|
||||||
|
result, calls = asyncio.run(_run_test())
|
||||||
|
|
||||||
|
assert result == "ok"
|
||||||
|
assert calls == [
|
||||||
|
{
|
||||||
|
"tool_name": QUERY_ACTIVITY_LOG_TOOL_NAME,
|
||||||
|
"tool_message": QUERY_ACTIVITY_LOG_TOOL_DESCRIPTION,
|
||||||
|
"tool_kwargs": {
|
||||||
|
"keyword": "整理",
|
||||||
|
"date": "2026-06-18",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
def test_factory_does_not_register_activity_log_tool():
|
def test_factory_does_not_register_activity_log_tool():
|
||||||
"""活动日志查询工具应由中间件注册,不应进入全局工具工厂。"""
|
"""活动日志查询工具应由中间件注册,不应进入全局工具工厂。"""
|
||||||
with patch(
|
with patch(
|
||||||
|
|||||||
@@ -1,4 +1,5 @@
|
|||||||
import json
|
import json
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from anyio import Path as AsyncPath
|
from anyio import Path as AsyncPath
|
||||||
@@ -111,3 +112,45 @@ def test_modify_request_instructs_model_to_use_skill_tool_without_paths(tmp_path
|
|||||||
assert "moviepilot-cli" in system_content
|
assert "moviepilot-cli" in system_content
|
||||||
assert "Read `" not in system_content
|
assert "Read `" not in system_content
|
||||||
assert str(tmp_path) not in system_content
|
assert str(tmp_path) not in system_content
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_skill_tool_call_records_streaming_summary(tmp_path):
|
||||||
|
"""skill 工具执行时应记录流式聚合摘要。"""
|
||||||
|
_write_skill(tmp_path, "moviepilot-cli")
|
||||||
|
calls = []
|
||||||
|
stream_handler = SimpleNamespace(
|
||||||
|
is_streaming=True,
|
||||||
|
record_tool_call=lambda **kwargs: calls.append(kwargs),
|
||||||
|
)
|
||||||
|
middleware = SkillsMiddleware(
|
||||||
|
sources=[str(tmp_path)],
|
||||||
|
stream_handler=stream_handler,
|
||||||
|
)
|
||||||
|
request = SimpleNamespace(
|
||||||
|
tool=SimpleNamespace(name=SKILL_TOOL_NAME),
|
||||||
|
tool_call={
|
||||||
|
"args": {
|
||||||
|
"name": "moviepilot-cli",
|
||||||
|
"explanation": "测试加载技能",
|
||||||
|
}
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _fake_handler(_request):
|
||||||
|
"""返回模拟工具结果。"""
|
||||||
|
return "ok"
|
||||||
|
|
||||||
|
result = await middleware.awrap_tool_call(request, _fake_handler)
|
||||||
|
|
||||||
|
assert result == "ok"
|
||||||
|
assert calls == [
|
||||||
|
{
|
||||||
|
"tool_name": SKILL_TOOL_NAME,
|
||||||
|
"tool_message": "Skill loaded",
|
||||||
|
"tool_kwargs": {
|
||||||
|
"name": "moviepilot-cli",
|
||||||
|
"explanation": "测试加载技能",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
]
|
||||||
|
|||||||
@@ -11,7 +11,6 @@ from app.agent.middleware.subagents import (
|
|||||||
MoviePilotSubAgentMiddleware,
|
MoviePilotSubAgentMiddleware,
|
||||||
SUBAGENT_CONTROL_TOOL_NAME,
|
SUBAGENT_CONTROL_TOOL_NAME,
|
||||||
SUBAGENT_TASK_TOOL_NAME,
|
SUBAGENT_TASK_TOOL_NAME,
|
||||||
SubAgentCallSummaryMiddleware,
|
|
||||||
SubAgentTaskControlMiddleware,
|
SubAgentTaskControlMiddleware,
|
||||||
create_subagent_middlewares,
|
create_subagent_middlewares,
|
||||||
)
|
)
|
||||||
@@ -28,7 +27,9 @@ def test_create_subagent_middlewares_registers_task_tool():
|
|||||||
stream_handler=None,
|
stream_handler=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
assert len(middlewares) == 3
|
assert len(middlewares) == 2
|
||||||
|
assert isinstance(middlewares[0], MoviePilotSubAgentMiddleware)
|
||||||
|
assert isinstance(middlewares[1], SubAgentTaskControlMiddleware)
|
||||||
assert [tool.name for tool in task_tools] == [
|
assert [tool.name for tool in task_tools] == [
|
||||||
SUBAGENT_TASK_TOOL_NAME,
|
SUBAGENT_TASK_TOOL_NAME,
|
||||||
SUBAGENT_CONTROL_TOOL_NAME,
|
SUBAGENT_CONTROL_TOOL_NAME,
|
||||||
@@ -140,17 +141,27 @@ def test_builtin_tools_declare_tags_in_implementation():
|
|||||||
assert missing_tools == []
|
assert missing_tools == []
|
||||||
|
|
||||||
|
|
||||||
def test_call_summary_middleware_logs_subagent_tool_operations():
|
def test_task_tool_call_records_streaming_summary():
|
||||||
"""子代理工具包装层应记录工具执行开始和完成日志。"""
|
"""task 子代理工具执行时应记录流式聚合摘要。"""
|
||||||
|
|
||||||
async def _run_test():
|
async def _run_test():
|
||||||
middleware = SubAgentCallSummaryMiddleware()
|
calls = []
|
||||||
|
stream_handler = SimpleNamespace(
|
||||||
|
is_streaming=True,
|
||||||
|
record_tool_call=lambda **kwargs: calls.append(kwargs),
|
||||||
|
)
|
||||||
|
middleware = MoviePilotSubAgentMiddleware(
|
||||||
|
model=FakeListChatModel(responses=["ok"]),
|
||||||
|
profiles=subagent_module._builtin_subagent_profiles(),
|
||||||
|
tools=[],
|
||||||
|
stream_handler=stream_handler,
|
||||||
|
)
|
||||||
request = SimpleNamespace(
|
request = SimpleNamespace(
|
||||||
tool=SimpleNamespace(name=SUBAGENT_CONTROL_TOOL_NAME),
|
tool=SimpleNamespace(name=SUBAGENT_TASK_TOOL_NAME),
|
||||||
tool_call={
|
tool_call={
|
||||||
"args": {
|
"args": {
|
||||||
"action": "status",
|
"description": "检查媒体信息",
|
||||||
"subagent_type": "general-purpose",
|
"subagent_type": "media-researcher",
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
@@ -158,15 +169,74 @@ def test_call_summary_middleware_logs_subagent_tool_operations():
|
|||||||
async def _fake_handler(_request):
|
async def _fake_handler(_request):
|
||||||
return "ok"
|
return "ok"
|
||||||
|
|
||||||
with patch.object(subagent_module.logger, "info") as log_info:
|
result = await middleware.awrap_tool_call(request, _fake_handler)
|
||||||
result = await middleware.awrap_tool_call(request, _fake_handler)
|
return result, calls
|
||||||
|
|
||||||
messages = [call.args[0] for call in log_info.call_args_list]
|
result, calls = asyncio.run(_run_test())
|
||||||
assert result == "ok"
|
|
||||||
assert any("开始执行子代理工具" in message for message in messages)
|
|
||||||
assert any("子代理工具执行完成" in message for message in messages)
|
|
||||||
|
|
||||||
asyncio.run(_run_test())
|
assert result == "ok"
|
||||||
|
assert calls == [
|
||||||
|
{
|
||||||
|
"tool_name": SUBAGENT_TASK_TOOL_NAME,
|
||||||
|
"tool_message": "Subagent invoked",
|
||||||
|
"tool_kwargs": {
|
||||||
|
"description": "检查媒体信息",
|
||||||
|
"subagent_type": "media-researcher",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def test_control_tool_call_records_streaming_summary():
|
||||||
|
"""subagent_task 子代理工具执行时应记录流式聚合摘要。"""
|
||||||
|
|
||||||
|
async def _run_test():
|
||||||
|
calls = []
|
||||||
|
stream_handler = SimpleNamespace(
|
||||||
|
is_streaming=True,
|
||||||
|
record_tool_call=lambda **kwargs: calls.append(kwargs),
|
||||||
|
)
|
||||||
|
middleware = SubAgentTaskControlMiddleware(
|
||||||
|
model=FakeListChatModel(responses=["ok"]),
|
||||||
|
profiles=subagent_module._builtin_subagent_profiles(),
|
||||||
|
tools=[],
|
||||||
|
stream_handler=stream_handler,
|
||||||
|
)
|
||||||
|
request = SimpleNamespace(
|
||||||
|
tool=SimpleNamespace(name=SUBAGENT_CONTROL_TOOL_NAME),
|
||||||
|
tool_call={
|
||||||
|
"args": {
|
||||||
|
"action": "start",
|
||||||
|
"tasks": [
|
||||||
|
{"subagent_type": "media-researcher"},
|
||||||
|
{"subagent_type": "download-diagnostician"},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _fake_handler(_request):
|
||||||
|
return "ok"
|
||||||
|
|
||||||
|
result = await middleware.awrap_tool_call(request, _fake_handler)
|
||||||
|
return result, calls
|
||||||
|
|
||||||
|
result, calls = asyncio.run(_run_test())
|
||||||
|
|
||||||
|
assert result == "ok"
|
||||||
|
assert calls == [
|
||||||
|
{
|
||||||
|
"tool_name": SUBAGENT_CONTROL_TOOL_NAME,
|
||||||
|
"tool_message": "Subagent invoked",
|
||||||
|
"tool_kwargs": {
|
||||||
|
"action": "start",
|
||||||
|
"tasks": [
|
||||||
|
{"subagent_type": "media-researcher"},
|
||||||
|
{"subagent_type": "download-diagnostician"},
|
||||||
|
],
|
||||||
|
},
|
||||||
|
}
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
def test_control_tool_starts_tasks_concurrently_and_waits():
|
def test_control_tool_starts_tasks_concurrently_and_waits():
|
||||||
|
|||||||
Reference in New Issue
Block a user