fix(agent): 收敛超时与取消后的副作用生命周期 (#6388)

This commit is contained in:
InfinityPacer
2026-08-22 12:11:15 +08:00
committed by GitHub
parent 7b6353337d
commit ed3c0541c7
16 changed files with 952 additions and 72 deletions
+35 -10
View File
@@ -1,16 +1,19 @@
"""LangChain 工具调用的 MoviePilot 宿主策略中间件。""" """LangChain 工具调用的 MoviePilot 宿主策略中间件。"""
import asyncio
from collections.abc import Awaitable, Callable from collections.abc import Awaitable, Callable
from typing import Any from typing import Any
from langchain.agents.middleware import AgentMiddleware, ToolCallRequest, hook_config from langchain.agents.middleware import AgentMiddleware, ToolCallRequest, hook_config
from langchain_core.messages import AIMessage, ToolMessage from langchain_core.messages import AIMessage, ToolMessage
from app.agent.policy import ( from app.agent.policy.contracts import (
ToolOrigin,
ToolPolicyContext,
)
from app.agent.policy.orchestrator import (
DEFAULT_TOOL_POLICY_ORCHESTRATOR, DEFAULT_TOOL_POLICY_ORCHESTRATOR,
AgentToolPolicyOrchestrator, AgentToolPolicyOrchestrator,
ToolPolicyContext,
ToolOrigin,
call_policy_hook, call_policy_hook,
) )
from app.agent.tools.catalog import ToolCatalogSnapshot from app.agent.tools.catalog import ToolCatalogSnapshot
@@ -19,6 +22,10 @@ from app.agent.tools.impl.query_system_settings import QuerySystemSettingsTool
POLICY_DENIED_MESSAGE = "当前宿主策略不允许执行该工具。" POLICY_DENIED_MESSAGE = "当前宿主策略不允许执行该工具。"
POLICY_UNAVAILABLE_MESSAGE = "宿主策略暂时不可用,未执行该工具。" POLICY_UNAVAILABLE_MESSAGE = "宿主策略暂时不可用,未执行该工具。"
TOOL_TIMEOUT_MESSAGE = (
"工具执行超时,已停止等待结果;"
"若工具包含外部写操作,操作可能仍在继续,请先确认实际状态再重试。"
)
class AgentPolicyMiddleware(AgentMiddleware): class AgentPolicyMiddleware(AgentMiddleware):
@@ -107,13 +114,22 @@ class AgentPolicyMiddleware(AgentMiddleware):
arguments = tool_call.get("args") or {} arguments = tool_call.get("args") or {}
if not isinstance(arguments, dict): if not isinstance(arguments, dict):
arguments = {} arguments = {}
_, result = await self.execute_tool_call( try:
tool=request.tool, _, result = await self.execute_tool_call(
arguments=arguments, tool=request.tool,
invocation_id=tool_call.get("id"), arguments=arguments,
handler=lambda: handler(request), invocation_id=tool_call.get("id"),
enforce_decision=False, handler=lambda: handler(request),
) enforce_decision=False,
)
except TimeoutError:
tool_name = str(getattr(request.tool, "name", None) or "unknown")
return ToolMessage(
content=TOOL_TIMEOUT_MESSAGE,
tool_call_id=str(tool_call.get("id") or ""),
name=tool_name,
status="error",
)
# 普通 ToolNode 保持 shadow 观测;已确认调用使用默认的强制决策语义。 # 普通 ToolNode 保持 shadow 观测;已确认调用使用默认的强制决策语义。
return result return result
@@ -144,6 +160,15 @@ class AgentPolicyMiddleware(AgentMiddleware):
return False, POLICY_DENIED_MESSAGE return False, POLICY_DENIED_MESSAGE
try: try:
result = await handler() result = await handler()
except asyncio.CancelledError as error:
if observation is not None:
call_policy_hook(
"cancel",
self.orchestrator.fail,
observation,
error,
)
raise
except Exception as error: except Exception as error:
if observation is not None: if observation is not None:
call_policy_hook( call_policy_hook(
+57 -26
View File
@@ -23,14 +23,16 @@ from langchain_core.messages import AIMessage, HumanMessage
from langchain_core.tools import BaseTool, StructuredTool from langchain_core.tools import BaseTool, StructuredTool
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from app.agent.llm import LLMHelper from app.agent.llm.helper import LLMHelper
from app.agent.middleware.policy import AgentPolicyMiddleware from app.agent.middleware.policy import AgentPolicyMiddleware
from app.agent.middleware.utils import append_to_system_message from app.agent.middleware.utils import append_to_system_message
from app.agent.policy import ( from app.agent.policy.contracts import (
AuthSource, AuthSource,
PrincipalType, PrincipalType,
ToolOrigin, ToolOrigin,
ToolPolicyContext, ToolPolicyContext,
)
from app.agent.policy.sanitizer import (
sanitize_for_host, sanitize_for_host,
summarize_error, summarize_error,
) )
@@ -46,6 +48,7 @@ SUBAGENT_STREAM_MARKER_KEY = "ls_agent_type"
SUBAGENT_STREAM_MARKER_VALUE = "subagent" SUBAGENT_STREAM_MARKER_VALUE = "subagent"
SUBAGENT_DEFAULT_WAIT_TIMEOUT_MS = 60000 SUBAGENT_DEFAULT_WAIT_TIMEOUT_MS = 60000
SUBAGENT_MAX_WAIT_TIMEOUT_MS = 300000 SUBAGENT_MAX_WAIT_TIMEOUT_MS = 300000
SUBAGENT_CANCEL_GRACE_SECONDS = 5.0
SUBAGENT_MAX_ACTIVE_TASKS = 8 SUBAGENT_MAX_ACTIVE_TASKS = 8
SUBAGENT_MAX_CONCURRENT_TASKS = 4 SUBAGENT_MAX_CONCURRENT_TASKS = 4
SUBAGENT_RESULT_MAX_CHARS = 12000 SUBAGENT_RESULT_MAX_CHARS = 12000
@@ -890,8 +893,10 @@ class SubAgentTaskControlMiddleware(AgentMiddleware):
) )
@staticmethod @staticmethod
async def _cancel_records(records: list[_SubAgentRuntimeTask]) -> None: async def _cancel_records(
"""取消一组尚未完成的任务。""" records: list[_SubAgentRuntimeTask],
) -> list[_SubAgentRuntimeTask]:
"""取消一组任务,并返回等待上限内仍未收敛的记录。"""
cancellable_tasks = [ cancellable_tasks = [
record.task for record in records if not record.task.done() record.task for record in records if not record.task.done()
] ]
@@ -899,9 +904,35 @@ class SubAgentTaskControlMiddleware(AgentMiddleware):
logger.info(f"开始取消子代理任务: tasks={len(cancellable_tasks)}") logger.info(f"开始取消子代理任务: tasks={len(cancellable_tasks)}")
for task in cancellable_tasks: for task in cancellable_tasks:
task.cancel() task.cancel()
if cancellable_tasks: if not cancellable_tasks:
await asyncio.gather(*cancellable_tasks, return_exceptions=True) return []
done, pending = await asyncio.wait(
cancellable_tasks,
timeout=SUBAGENT_CANCEL_GRACE_SECONDS,
)
if done:
await asyncio.gather(*done, return_exceptions=True)
if pending:
logger.warning(
f"子代理任务取消等待超时: pending={len(pending)}, "
f"timeout={SUBAGENT_CANCEL_GRACE_SECONDS}s"
)
else:
logger.info(f"子代理任务取消完成: tasks={len(cancellable_tasks)}") logger.info(f"子代理任务取消完成: tasks={len(cancellable_tasks)}")
return [record for record in records if record.task in pending]
async def close(self) -> None:
"""取消脱离当前 Agent 回合的子代理任务。"""
unfinished_records = [
record for record in self._tasks.values() if not record.task.done()
]
if unfinished_records:
logger.info(
f"关闭子代理任务控制器,取消未完成任务: tasks={len(unfinished_records)}"
)
await self._cancel_records(unfinished_records)
self._tasks.clear()
@staticmethod @staticmethod
def _pipeline_description( def _pipeline_description(
@@ -1039,14 +1070,16 @@ class SubAgentTaskControlMiddleware(AgentMiddleware):
f"subagent_type={record.subagent_type}" f"subagent_type={record.subagent_type}"
) )
try: done, pending = await asyncio.wait({task}, timeout=timeout)
result = await asyncio.wait_for(task, timeout=timeout) if pending:
except asyncio.TimeoutError: task.cancel()
error = f"{step_index} 个管道子代理任务等待超时。" error = f"{step_index} 个管道子代理任务等待超时。"
logger.info( logger.info(
f"{error} task_id={record.task_id}, timeout_ms={normalized_timeout_ms}" f"{error} task_id={record.task_id}, timeout_ms={normalized_timeout_ms}"
) )
return records, error return records, error
try:
result = next(iter(done)).result()
except Exception as err: except Exception as err:
error = ( error = (
f"{step_index} 个管道子代理任务执行失败: " f"{step_index} 个管道子代理任务执行失败: "
@@ -1125,6 +1158,7 @@ class SubAgentTaskControlMiddleware(AgentMiddleware):
active_only=action in {"wait", "cancel"} and not task_ids and not task_id, active_only=action in {"wait", "cancel"} and not task_ids and not task_id,
) )
cancellation_pending: list[_SubAgentRuntimeTask] = []
if action == "wait": if action == "wait":
logger.info( logger.info(
f"准备等待子代理任务: selected={len(records)}, missing={len(missing_ids)}" f"准备等待子代理任务: selected={len(records)}, missing={len(missing_ids)}"
@@ -1138,31 +1172,28 @@ class SubAgentTaskControlMiddleware(AgentMiddleware):
logger.info( logger.info(
f"准备取消子代理任务: selected={len(records)}, missing={len(missing_ids)}" f"准备取消子代理任务: selected={len(records)}, missing={len(missing_ids)}"
) )
await self._cancel_records(records) cancellation_pending = await self._cancel_records(records)
elif action == "status": elif action == "status":
logger.info( logger.info(
f"查询子代理任务状态: selected={len(records)}, missing={len(missing_ids)}" f"查询子代理任务状态: selected={len(records)}, missing={len(missing_ids)}"
) )
return self._json_response( response = {
{ "success": not cancellation_pending,
"success": True, "action": action,
"action": action, "wait_mode": wait_mode if action == "wait" else None,
"wait_mode": wait_mode if action == "wait" else None, "missing_task_ids": missing_ids,
"missing_task_ids": missing_ids, "tasks": [self._task_output(record) for record in records],
"tasks": [self._task_output(record) for record in records], }
} if action == "cancel":
) response["cancel_pending_task_ids"] = [
record.task_id for record in cancellation_pending
]
return self._json_response(response)
async def aafter_agent(self, state: Any, runtime: Any) -> None: async def aafter_agent(self, state: Any, runtime: Any) -> None:
"""Agent 结束时取消未完成的子代理任务,避免后台泄漏。""" """Agent 结束时取消未完成的子代理任务,避免后台泄漏。"""
unfinished_records = [ await self.close()
record for record in self._tasks.values() if not record.task.done()
]
if unfinished_records:
logger.info(f"Agent 结束,取消未完成子代理任务: tasks={len(unfinished_records)}")
await self._cancel_records(unfinished_records)
self._tasks.clear()
async def awrap_tool_call( async def awrap_tool_call(
self, self,
+59 -6
View File
@@ -1,5 +1,6 @@
import asyncio import asyncio
import hashlib import hashlib
import inspect
import json import json
import re import re
import traceback import traceback
@@ -395,6 +396,7 @@ class MoviePilotAgent:
self._llm_provider_selection: Dict[str, Any] = {} self._llm_provider_selection: Dict[str, Any] = {}
self._agent_started_at: Optional[datetime] = None self._agent_started_at: Optional[datetime] = None
self._compiled_agent_bundle: Optional[_CompiledAgentBundle] = None self._compiled_agent_bundle: Optional[_CompiledAgentBundle] = None
self._subagent_middlewares: tuple[Any, ...] = ()
self._last_agent_cache_hit = False self._last_agent_cache_hit = False
# 流式token管理 # 流式token管理
@@ -1693,7 +1695,23 @@ class MoviePilotAgent:
return bundle.agent return bundle.agent
return None return None
def _cache_agent( @staticmethod
async def _close_subagent_middleware_instances(
middlewares: tuple[Any, ...],
) -> None:
"""释放不再由 Agent 图持有的子代理控制器。"""
for middleware in middlewares:
close = getattr(middleware, "close", None)
if not callable(close):
continue
try:
result = close()
if inspect.isawaitable(result):
await result
except Exception as error:
logger.debug(f"关闭子代理中间件失败: {error}")
async def _cache_agent(
self, self,
*, *,
signature: tuple[Any, ...], signature: tuple[Any, ...],
@@ -1702,8 +1720,18 @@ class MoviePilotAgent:
tool_catalog: ToolCatalogSnapshot, tool_catalog: ToolCatalogSnapshot,
subagent_catalog: ToolCatalogSnapshot, subagent_catalog: ToolCatalogSnapshot,
mcp_config_signature: str, mcp_config_signature: str,
subagent_middlewares: tuple[Any, ...] = (),
) -> Any: ) -> Any:
"""保存当前会话可复用的 Agent 图。""" """保存当前会话可复用的 Agent 图。"""
previous_middlewares = tuple(
middleware
for middleware in self._subagent_middlewares
if not any(
middleware is replacement
for replacement in subagent_middlewares
)
)
await self._close_subagent_middleware_instances(previous_middlewares)
self._compiled_agent_bundle = _CompiledAgentBundle( self._compiled_agent_bundle = _CompiledAgentBundle(
signature=signature, signature=signature,
agent=agent, agent=agent,
@@ -1715,8 +1743,16 @@ class MoviePilotAgent:
mcp_config_signature=mcp_config_signature, mcp_config_signature=mcp_config_signature,
catalog_checked_at=datetime.now(), catalog_checked_at=datetime.now(),
) )
self._subagent_middlewares = subagent_middlewares
return agent return agent
async def _invalidate_cached_agent(self) -> None:
"""使当前图失效,并释放只属于该图的子代理控制器。"""
subagent_middlewares = self._subagent_middlewares
self._subagent_middlewares = ()
self._compiled_agent_bundle = None
await self._close_subagent_middleware_instances(subagent_middlewares)
@staticmethod @staticmethod
def _latest_turn_messages(messages: List[BaseMessage]) -> List[BaseMessage]: def _latest_turn_messages(messages: List[BaseMessage]) -> List[BaseMessage]:
"""从完整历史中提取本轮新增用户消息。""" """从完整历史中提取本轮新增用户消息。"""
@@ -1785,6 +1821,7 @@ class MoviePilotAgent:
创建 LangGraph Agent(使用 create_agent + SummarizationMiddleware 创建 LangGraph Agent(使用 create_agent + SummarizationMiddleware
:param streaming: 是否启用流式输出 :param streaming: 是否启用流式输出
""" """
temporary_subagent_middlewares: tuple[Any, ...] = ()
try: try:
runtime_config = await self._resolve_llm_runtime_config() runtime_config = await self._resolve_llm_runtime_config()
plugin_revision = _get_plugin_tools_revision() plugin_revision = _get_plugin_tools_revision()
@@ -1888,6 +1925,7 @@ class MoviePilotAgent:
policy_context=policy_context.for_subagent(), policy_context=policy_context.for_subagent(),
catalog=subagent_catalog, catalog=subagent_catalog,
) )
temporary_subagent_middlewares = tuple(subagent_middlewares)
# 严格目录必须覆盖 LangGraph ToolNode 可执行的全部 client-side 工具。 # 严格目录必须覆盖 LangGraph ToolNode 可执行的全部 client-side 工具。
tool_catalog = ToolCatalogSnapshot.from_tools( tool_catalog = ToolCatalogSnapshot.from_tools(
[ [
@@ -1910,6 +1948,10 @@ class MoviePilotAgent:
if cached_agent: if cached_agent:
# 签名相同表示已编译图中的精确工具实例仍有效;新建快照仅用于复核。 # 签名相同表示已编译图中的精确工具实例仍有效;新建快照仅用于复核。
cached_bundle.catalog_checked_at = datetime.now() cached_bundle.catalog_checked_at = datetime.now()
await self._close_subagent_middleware_instances(
temporary_subagent_middlewares
)
temporary_subagent_middlewares = ()
logger.debug(f"复用会话内 Agent 图: session_id={self.session_id}") logger.debug(f"复用会话内 Agent 图: session_id={self.session_id}")
return cached_agent return cached_agent
max_tools = settings.LLM_MAX_TOOLS max_tools = settings.LLM_MAX_TOOLS
@@ -2014,17 +2056,28 @@ class MoviePilotAgent:
middleware=middlewares, middleware=middlewares,
checkpointer=InMemorySaver(), checkpointer=InMemorySaver(),
) )
return self._cache_agent( cached_agent = await self._cache_agent(
signature=bundle_signature, signature=bundle_signature,
agent=agent, agent=agent,
streaming=streaming, streaming=streaming,
tool_catalog=tool_catalog, tool_catalog=tool_catalog,
subagent_catalog=subagent_catalog, subagent_catalog=subagent_catalog,
mcp_config_signature=mcp_config_signature, mcp_config_signature=mcp_config_signature,
subagent_middlewares=tuple(subagent_middlewares),
) )
temporary_subagent_middlewares = ()
return cached_agent
except asyncio.CancelledError:
await self._close_subagent_middleware_instances(
temporary_subagent_middlewares
)
raise
except Exception as e: except Exception as e:
await self._close_subagent_middleware_instances(
temporary_subagent_middlewares
)
logger.error(f"创建 Agent 失败: {e}") logger.error(f"创建 Agent 失败: {e}")
raise e raise
async def process( async def process(
self, self,
@@ -2365,11 +2418,11 @@ class MoviePilotAgent:
except asyncio.CancelledError: except asyncio.CancelledError:
logger.info(f"Agent执行被取消: session_id={self.session_id}") logger.info(f"Agent执行被取消: session_id={self.session_id}")
self._compiled_agent_bundle = None await self._invalidate_cached_agent()
execution_error = "任务已取消" execution_error = "任务已取消"
raise raise
except Exception as e: except Exception as e:
self._compiled_agent_bundle = None await self._invalidate_cached_agent()
execution_error = str(e) execution_error = str(e)
if self._messages_have_image_input(messages) and self._is_unsupported_image_input_error(e): if self._messages_have_image_input(messages) and self._is_unsupported_image_input_error(e):
logger.warning( logger.warning(
@@ -2422,9 +2475,9 @@ class MoviePilotAgent:
""" """
清理智能体资源 清理智能体资源
""" """
await self._invalidate_cached_agent()
self._pending_secret_confirmation = None self._pending_secret_confirmation = None
self.protected_output_callback = None self.protected_output_callback = None
self._compiled_agent_bundle = None
logger.info(f"MoviePilot智能体已清理: session_id={self.session_id}") logger.info(f"MoviePilot智能体已清理: session_id={self.session_id}")
+2
View File
@@ -184,6 +184,8 @@ class ExecutionReceipt:
result_summary: Optional[str] = None result_summary: Optional[str] = None
error_summary: Optional[str] = None error_summary: Optional[str] = None
duration_ms: int = 0 duration_ms: int = 0
external_may_continue: bool = False # 中断后外部操作仍可能继续,不能视为已停止。
needs_reconcile: bool = False # 调用方需查询外部实际状态后再决定补偿或重试。
@dataclass(frozen=True) @dataclass(frozen=True)
+30 -1
View File
@@ -1,5 +1,6 @@
"""Agent 工具策略观测、脱敏回执与共享执行边界。""" """Agent 工具策略观测、脱敏回执与共享执行边界。"""
import asyncio
import time import time
import uuid import uuid
from collections.abc import Callable from collections.abc import Callable
@@ -9,12 +10,14 @@ from langchain_core.messages import ToolMessage
from pydantic import ValidationError from pydantic import ValidationError
from app.agent.policy.contracts import ( from app.agent.policy.contracts import (
ActionEffect,
ConfirmationMode, ConfirmationMode,
ExecutionOutcome, ExecutionOutcome,
ExecutionReceipt, ExecutionReceipt,
MigrationState, MigrationState,
PolicyDecision, PolicyDecision,
PolicyObservation, PolicyObservation,
RecoveryMode,
ToolInvocation, ToolInvocation,
ToolPolicyContext, ToolPolicyContext,
) )
@@ -165,10 +168,32 @@ class AgentToolPolicyOrchestrator:
) )
return receipt return receipt
@staticmethod
def _uncertain_external_state(
observation: PolicyObservation,
error: BaseException,
) -> tuple[bool, bool]:
"""标记取消或超时后无法确认的写操作终态。"""
interrupted = isinstance(error, (asyncio.CancelledError, TimeoutError))
read_only = observation.policy.effect in {
ActionEffect.SAFE_READ,
ActionEffect.SENSITIVE_READ,
}
if not interrupted or read_only:
return False, False
needs_reconcile = observation.policy.recovery not in {
RecoveryMode.TRANSACTION,
RecoveryMode.IDEMPOTENT,
}
return True, needs_reconcile
@staticmethod @staticmethod
def fail(observation: PolicyObservation, error: BaseException) -> ExecutionReceipt: def fail(observation: PolicyObservation, error: BaseException) -> ExecutionReceipt:
"""生成失败回执 envelope,不把异常中的凭据写入日志。""" """生成失败回执 envelope,不把异常中的凭据写入日志。"""
error_summary = summarize_error(error) error_summary = summarize_error(error)
external_may_continue, needs_reconcile = (
AgentToolPolicyOrchestrator._uncertain_external_state(observation, error)
)
receipt = ExecutionReceipt( receipt = ExecutionReceipt(
invocation_id=observation.invocation.invocation_id, invocation_id=observation.invocation.invocation_id,
tool_name=observation.invocation.tool_name, tool_name=observation.invocation.tool_name,
@@ -181,11 +206,15 @@ class AgentToolPolicyOrchestrator:
0, 0,
int((time.monotonic() - observation.started_at) * 1000), int((time.monotonic() - observation.started_at) * 1000),
), ),
external_may_continue=external_may_continue,
needs_reconcile=needs_reconcile,
) )
logger.error( logger.error(
f"Agent工具执行失败: tool={receipt.tool_name}, " f"Agent工具执行失败: tool={receipt.tool_name}, "
f"origin={receipt.origin.value}, shadow={receipt.decision.shadow}, " f"origin={receipt.origin.value}, shadow={receipt.decision.shadow}, "
f"duration_ms={receipt.duration_ms}, error={error_summary}" f"duration_ms={receipt.duration_ms}, error={error_summary}, "
f"external_may_continue={receipt.external_may_continue}, "
f"needs_reconcile={receipt.needs_reconcile}"
) )
return receipt return receipt
+14 -1
View File
@@ -2,6 +2,7 @@
from __future__ import annotations from __future__ import annotations
import sys
import threading import threading
from typing import Any from typing import Any
@@ -132,6 +133,18 @@ def is_tool_factory_materialized() -> bool:
) )
async def close_materialized_terminal_sessions() -> None:
"""关闭已物化的终端会话管理器,不触发新的 Agent 工具导入。"""
module = sys.modules.get("app.agent.tools.impl._terminal_session")
manager = getattr(module, "terminal_session_manager", None) if module else None
close = getattr(manager, "close", None)
if callable(close):
await close()
async def begin_agent_shutdown() -> None: async def begin_agent_shutdown() -> None:
"""不可逆关闭首用闸门,并等待全部同步及异步能力释放。""" """不可逆关闭首用闸门,并等待全部同步及异步能力释放。"""
await _ensure_runtime().shutdown_async(reason="application_shutdown") try:
await _ensure_runtime().shutdown_async(reason="application_shutdown")
finally:
await close_materialized_terminal_sessions()
+2 -1
View File
@@ -377,7 +377,7 @@ class MoviePilotTool(BaseTool, metaclass=ABCMeta):
except ToolExecutionTimeoutError as e: except ToolExecutionTimeoutError as e:
error_message = summarize_error(e) error_message = summarize_error(e)
logger.warning(error_message) logger.warning(error_message)
result = error_message raise
except Exception as e: except Exception as e:
error_message = f"工具执行异常: {summarize_error(e)}" error_message = f"工具执行异常: {summarize_error(e)}"
logger.error(f"Tool {self.name} execution failed: {summarize_error(e)}") logger.error(f"Tool {self.name} execution failed: {summarize_error(e)}")
@@ -415,6 +415,7 @@ class MoviePilotTool(BaseTool, metaclass=ABCMeta):
except asyncio.TimeoutError as err: except asyncio.TimeoutError as err:
raise ToolExecutionTimeoutError( raise ToolExecutionTimeoutError(
f"工具 {self.name} 执行超时(超过 {timeout:g} 秒),已停止等待结果。" f"工具 {self.name} 执行超时(超过 {timeout:g} 秒),已停止等待结果。"
"若工具包含外部写操作,操作可能仍在继续,请先确认实际状态再重试。"
) from err ) from err
@staticmethod @staticmethod
+104 -8
View File
@@ -128,6 +128,11 @@ class _TerminalSessionManager:
"""初始化会话表和并发保护锁。""" """初始化会话表和并发保护锁。"""
self._sessions: dict[str, _TerminalSession] = {} self._sessions: dict[str, _TerminalSession] = {}
self._lock = asyncio.Lock() self._lock = asyncio.Lock()
self._close_lock = asyncio.Lock()
self._closed = False
self._starting = 0
self._starts_idle = asyncio.Event()
self._starts_idle.set()
@staticmethod @staticmethod
def _normalize_bool(value: Any, default: bool = True) -> bool: def _normalize_bool(value: Any, default: bool = True) -> bool:
@@ -211,20 +216,55 @@ class _TerminalSessionManager:
should_use_pty = self._normalize_bool(use_pty, default=True) and os.name == "posix" should_use_pty = self._normalize_bool(use_pty, default=True) and os.name == "posix"
async with self._lock: async with self._lock:
if self._closed:
raise RuntimeError("终端会话管理器已关闭")
self._cleanup_finished_sessions_locked() self._cleanup_finished_sessions_locked()
if self._active_session_count_locked() >= TERMINAL_CONCURRENCY_LIMIT: if (
self._active_session_count_locked() + self._starting
>= TERMINAL_CONCURRENCY_LIMIT
):
raise RuntimeError( raise RuntimeError(
f"后台终端会话数已达到上限 {TERMINAL_CONCURRENCY_LIMIT}" f"后台终端会话数已达到上限 {TERMINAL_CONCURRENCY_LIMIT}"
) )
self._starting += 1
self._starts_idle.clear()
session = ( session: Optional[_TerminalSession] = None
await self._start_pty_session(command, normalized_cwd, normalized_env) reject_session = False
if should_use_pty session_registered = False
else await self._start_pipe_session(command, normalized_cwd, normalized_env) session_released = False
) try:
session = (
await self._start_pty_session(command, normalized_cwd, normalized_env)
if should_use_pty
else await self._start_pipe_session(
command, normalized_cwd, normalized_env
)
)
async with self._lock: async with self._lock:
self._sessions[session.session_id] = session reject_session = self._closed
if not reject_session:
self._sessions[session.session_id] = session
session_registered = True
if reject_session:
await self._terminate_session(session)
session_released = True
raise RuntimeError("终端会话管理器已关闭")
except BaseException:
if session is not None and not session_registered and not session_released:
cleanup_task = asyncio.create_task(self._terminate_session(session))
try:
await asyncio.shield(cleanup_task)
except asyncio.CancelledError:
await cleanup_task
raise
finally:
async with self._lock:
self._starting -= 1
if self._starting == 0:
self._starts_idle.set()
logger.info( logger.info(
"启动后台终端会话: session_id=%s, pid=%s, use_pty=%s, command=%s", "启动后台终端会话: session_id=%s, pid=%s, use_pty=%s, command=%s",
@@ -473,6 +513,62 @@ class _TerminalSessionManager:
return self._session_payload(session, output="", output_truncated=False) return self._session_payload(session, output="", output_truncated=False)
async def close(self) -> None:
"""停止所有后台终端会话并释放 PTY、读取任务和会话记录。"""
async with self._close_lock:
async with self._lock:
self._closed = True
await self._starts_idle.wait()
async with self._lock:
sessions = list(self._sessions.values())
await asyncio.gather(
*(self._terminate_session(session) for session in sessions),
return_exceptions=True,
)
async with self._lock:
for session in sessions:
session.close_pty()
self._sessions.clear()
async def _terminate_session(self, session: _TerminalSession) -> None:
"""以有限等待停止进程,并在必要时升级为 SIGKILL。"""
if session.status == "running":
session.kill_requested = True
self._send_signal(session, signal.SIGTERM)
wait_task = session.wait_task
if wait_task and not wait_task.done():
try:
await asyncio.wait_for(
asyncio.shield(wait_task),
timeout=TERMINAL_KILL_GRACE_SECONDS,
)
except asyncio.TimeoutError:
force_signal = getattr(signal, "SIGKILL", signal.SIGTERM)
self._send_signal(session, force_signal)
try:
await asyncio.wait_for(
asyncio.shield(wait_task),
timeout=TERMINAL_KILL_GRACE_SECONDS,
)
except asyncio.TimeoutError:
logger.error(
"终端会话关闭超时: session_id=%s, pid=%s",
session.session_id,
session.pid,
)
for task in session.reader_tasks:
if not task.done():
task.cancel()
if session.reader_tasks:
await asyncio.gather(*session.reader_tasks, return_exceptions=True)
session.close_pty()
def get_session(self, session_id: str) -> _TerminalSession: def get_session(self, session_id: str) -> _TerminalSession:
"""按 ID 获取会话,不存在时抛出清晰错误。""" """按 ID 获取会话,不存在时抛出清晰错误。"""
session = self._sessions.get(session_id) session = self._sessions.get(session_id)
+11 -5
View File
@@ -1,5 +1,6 @@
from __future__ import annotations from __future__ import annotations
import asyncio
import json import json
import threading import threading
import uuid import uuid
@@ -8,7 +9,8 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional
from app.runtime.log import logger from app.runtime.log import logger
if TYPE_CHECKING: if TYPE_CHECKING:
from app.agent.policy import AgentToolPolicyOrchestrator, ToolPolicyContext from app.agent.policy.contracts import ToolPolicyContext
from app.agent.policy.orchestrator import AgentToolPolicyOrchestrator
from app.agent.tools.catalog import ToolCatalogSnapshot from app.agent.tools.catalog import ToolCatalogSnapshot
@@ -57,7 +59,7 @@ class MoviePilotToolsManager:
@staticmethod @staticmethod
def _summarize_error(error: Exception) -> str: def _summarize_error(error: Exception) -> str:
"""仅在错误路径加载策略脱敏器,保持默认导入轻量。""" """仅在错误路径加载策略脱敏器,保持默认导入轻量。"""
from app.agent.policy import summarize_error from app.agent.policy.sanitizer import summarize_error
return summarize_error(error) return summarize_error(error)
@@ -143,13 +145,13 @@ class MoviePilotToolsManager:
if policy_orchestrator is not None and policy_context is not None: if policy_orchestrator is not None and policy_context is not None:
return policy_orchestrator, policy_context return policy_orchestrator, policy_context
from app.agent.policy import ( from app.agent.policy.contracts import (
DEFAULT_TOOL_POLICY_ORCHESTRATOR,
AuthSource, AuthSource,
PrincipalType, PrincipalType,
ToolOrigin, ToolOrigin,
ToolPolicyContext, ToolPolicyContext,
) )
from app.agent.policy.orchestrator import DEFAULT_TOOL_POLICY_ORCHESTRATOR
if policy_orchestrator is None: if policy_orchestrator is None:
policy_orchestrator = DEFAULT_TOOL_POLICY_ORCHESTRATOR policy_orchestrator = DEFAULT_TOOL_POLICY_ORCHESTRATOR
@@ -381,7 +383,7 @@ class MoviePilotToolsManager:
) )
return error_msg return error_msg
from app.agent.policy import call_policy_hook from app.agent.policy.orchestrator import call_policy_hook
from app.agent.tools.base import ( from app.agent.tools.base import (
ToolExecutionTimeoutError, ToolExecutionTimeoutError,
format_tool_result_for_agent, format_tool_result_for_agent,
@@ -414,6 +416,10 @@ class MoviePilotToolsManager:
tool_name=tool_name, tool_name=tool_name,
max_chars=getattr(tool_instance, "result_max_chars", None), max_chars=getattr(tool_instance, "result_max_chars", None),
) )
except asyncio.CancelledError as e:
if observation is not None and policy_orchestrator is not None:
call_policy_hook("cancel", policy_orchestrator.fail, observation, e)
raise
except ToolExecutionTimeoutError as e: except ToolExecutionTimeoutError as e:
if observation is not None and policy_orchestrator is not None: if observation is not None and policy_orchestrator is not None:
call_policy_hook("fail", policy_orchestrator.fail, observation, e) call_policy_hook("fail", policy_orchestrator.fail, observation, e)
+6 -3
View File
@@ -3,6 +3,7 @@ from typing import Any
from app.agent.runtime_loader import ( from app.agent.runtime_loader import (
activate_agent_service, activate_agent_service,
begin_agent_shutdown, begin_agent_shutdown,
close_materialized_terminal_sessions,
get_agent_manager as get_runtime_agent_manager, get_agent_manager as get_runtime_agent_manager,
get_running_agent_manager as get_runtime_running_agent_manager, get_running_agent_manager as get_runtime_running_agent_manager,
is_tool_factory_materialized, is_tool_factory_materialized,
@@ -70,14 +71,14 @@ def _get_prompt_manager() -> Any:
def _get_capability_manager() -> Any: def _get_capability_manager() -> Any:
"""首个多模态调用才导入 Agent 能力管理器。""" """首个多模态调用才导入 Agent 能力管理器。"""
from app.agent.llm import AgentCapabilityManager from app.agent.llm.capability import AgentCapabilityManager
return AgentCapabilityManager return AgentCapabilityManager
def _get_llm_helper() -> Any: def _get_llm_helper() -> Any:
"""首个模型能力查询才导入 LLM helper。""" """首个模型能力查询才导入 LLM helper。"""
from app.agent.llm import LLMHelper from app.agent.llm.helper import LLMHelper
return LLMHelper return LLMHelper
@@ -216,6 +217,8 @@ async def stop_agent():
if is_tool_factory_materialized(): if is_tool_factory_materialized():
from app.agent.tools.base import shutdown_blocking_executors from app.agent.tools.base import shutdown_blocking_executors
shutdown_blocking_executors(cancel_futures=True) shutdown_blocking_executors(wait=False, cancel_futures=True)
except Exception as e: except Exception as e:
logger.error(f"停止AI智能体时发生错误: {e}") logger.error(f"停止AI智能体时发生错误: {e}")
finally:
await close_materialized_terminal_sessions()
+12 -2
View File
@@ -13,8 +13,8 @@
"runtime_to_db": [], "runtime_to_db": [],
"workflow_to_db": [] "workflow_to_db": []
}, },
"edge_count": 6354, "edge_count": 6364,
"edge_sha256": "5c7bd53e742c806fdd9fa129e35ab979008d54e568013df61a4f60bffa347ddf", "edge_sha256": "0f241a46ed27309c2071d4c99f00ba067ad02b3974581cc206ddec4c522a982c",
"edges": [ "edges": [
"app -> app.runtime", "app -> app.runtime",
"app -> app.runtime.compat", "app -> app.runtime.compat",
@@ -258,6 +258,8 @@
"app.agent.middleware.memory -> app.runtime.log", "app.agent.middleware.memory -> app.runtime.log",
"app.agent.middleware.policy -> app.agent", "app.agent.middleware.policy -> app.agent",
"app.agent.middleware.policy -> app.agent.policy", "app.agent.middleware.policy -> app.agent.policy",
"app.agent.middleware.policy -> app.agent.policy.contracts",
"app.agent.middleware.policy -> app.agent.policy.orchestrator",
"app.agent.middleware.policy -> app.agent.tools", "app.agent.middleware.policy -> app.agent.tools",
"app.agent.middleware.policy -> app.agent.tools.catalog", "app.agent.middleware.policy -> app.agent.tools.catalog",
"app.agent.middleware.policy -> app.agent.tools.impl", "app.agent.middleware.policy -> app.agent.tools.impl",
@@ -278,10 +280,13 @@
"app.agent.middleware.skills -> app.runtime.log", "app.agent.middleware.skills -> app.runtime.log",
"app.agent.middleware.subagents -> app.agent", "app.agent.middleware.subagents -> app.agent",
"app.agent.middleware.subagents -> app.agent.llm", "app.agent.middleware.subagents -> app.agent.llm",
"app.agent.middleware.subagents -> app.agent.llm.helper",
"app.agent.middleware.subagents -> app.agent.middleware", "app.agent.middleware.subagents -> app.agent.middleware",
"app.agent.middleware.subagents -> app.agent.middleware.policy", "app.agent.middleware.subagents -> app.agent.middleware.policy",
"app.agent.middleware.subagents -> app.agent.middleware.utils", "app.agent.middleware.subagents -> app.agent.middleware.utils",
"app.agent.middleware.subagents -> app.agent.policy", "app.agent.middleware.subagents -> app.agent.policy",
"app.agent.middleware.subagents -> app.agent.policy.contracts",
"app.agent.middleware.subagents -> app.agent.policy.sanitizer",
"app.agent.middleware.subagents -> app.agent.runtime", "app.agent.middleware.subagents -> app.agent.runtime",
"app.agent.middleware.subagents -> app.agent.tools", "app.agent.middleware.subagents -> app.agent.tools",
"app.agent.middleware.subagents -> app.agent.tools.catalog", "app.agent.middleware.subagents -> app.agent.tools.catalog",
@@ -1471,6 +1476,9 @@
"app.agent.tools.impl.write_file -> app.runtime.log", "app.agent.tools.impl.write_file -> app.runtime.log",
"app.agent.tools.manager -> app.agent", "app.agent.tools.manager -> app.agent",
"app.agent.tools.manager -> app.agent.policy", "app.agent.tools.manager -> app.agent.policy",
"app.agent.tools.manager -> app.agent.policy.contracts",
"app.agent.tools.manager -> app.agent.policy.orchestrator",
"app.agent.tools.manager -> app.agent.policy.sanitizer",
"app.agent.tools.manager -> app.agent.runtime_loader", "app.agent.tools.manager -> app.agent.runtime_loader",
"app.agent.tools.manager -> app.agent.tools", "app.agent.tools.manager -> app.agent.tools",
"app.agent.tools.manager -> app.agent.tools.base", "app.agent.tools.manager -> app.agent.tools.base",
@@ -5883,7 +5891,9 @@
"app.sdk.utilities -> app.sdk.string", "app.sdk.utilities -> app.sdk.string",
"app.startup.agent_initializer -> app.agent", "app.startup.agent_initializer -> app.agent",
"app.startup.agent_initializer -> app.agent.llm", "app.startup.agent_initializer -> app.agent.llm",
"app.startup.agent_initializer -> app.agent.llm.capability",
"app.startup.agent_initializer -> app.agent.llm.gateway", "app.startup.agent_initializer -> app.agent.llm.gateway",
"app.startup.agent_initializer -> app.agent.llm.helper",
"app.startup.agent_initializer -> app.agent.llm.provider", "app.startup.agent_initializer -> app.agent.llm.provider",
"app.startup.agent_initializer -> app.agent.prompt", "app.startup.agent_initializer -> app.agent.prompt",
"app.startup.agent_initializer -> app.agent.prompt.transfer_redo", "app.startup.agent_initializer -> app.agent.prompt.transfer_redo",
+85 -1
View File
@@ -157,6 +157,7 @@ async def test_expired_unchanged_catalog_renews_freshness() -> None:
model="fake", model="fake",
profile={"max_input_tokens": 64000}, profile={"max_input_tokens": 64000},
) )
temporary_middleware = SimpleNamespace(close=AsyncMock())
with patch.object( with patch.object(
agent, agent,
@@ -195,7 +196,7 @@ async def test_expired_unchanged_catalog_renews_freshness() -> None:
return_value=SimpleNamespace(name="skills", tools=[]), return_value=SimpleNamespace(name="skills", tools=[]),
), patch( ), patch(
"app.agent.orchestrator.create_subagent_middlewares", "app.agent.orchestrator.create_subagent_middlewares",
return_value=([], []), return_value=([temporary_middleware], []),
), patch( ), patch(
"app.agent.orchestrator._get_plugin_tools_revision", "app.agent.orchestrator._get_plugin_tools_revision",
return_value=0, return_value=0,
@@ -210,6 +211,89 @@ async def test_expired_unchanged_catalog_renews_freshness() -> None:
assert graph is cached_graph assert graph is cached_graph
assert agent._compiled_agent_bundle.catalog_checked_at > expired_at assert agent._compiled_agent_bundle.catalog_checked_at > expired_at
temporary_middleware.close.assert_awaited_once()
@pytest.mark.anyio
async def test_create_agent_cancellation_closes_temporary_subagent_middleware() -> None:
"""构图取消时必须释放尚未被缓存接管的子代理控制器。"""
catalog = ToolCatalogSnapshot.from_tools(
[], plugin_revision=0, factory_revision="factory-v1"
)
fake_llm = SimpleNamespace(
_llm_type="openai-chat",
model="fake",
profile={"max_input_tokens": 64000},
)
temporary_middleware = SimpleNamespace(close=AsyncMock())
signature_started = __import__("asyncio").Event()
agent = MoviePilotAgent(session_id="cancel-create", user_id="user-1")
async def _wait_for_signature(*_args, **_kwargs):
signature_started.set()
await __import__("asyncio").Future()
with patch.object(
agent,
"_resolve_llm_runtime_config",
new=AsyncMock(return_value={"provider": "openai", "model": "fake"}),
), patch.object(
agent,
"_initialize_local_tool_catalogs",
return_value=(catalog, catalog),
), patch.object(
agent,
"_initialize_mcp_tools",
new=AsyncMock(return_value=[]),
), patch.object(
agent,
"_initialize_subagent_mcp_tools",
new=AsyncMock(return_value=[]),
), patch.object(
agent,
"_initialize_llm",
new=AsyncMock(return_value=fake_llm),
), patch.object(
agent,
"_sync_model_profile",
), patch.object(
agent,
"_agent_bundle_signature",
new=_wait_for_signature,
), patch(
"app.agent.orchestrator._get_plugin_tools_revision",
return_value=0,
), patch(
"app.agent.orchestrator.agent_mcp_manager.config_signature",
return_value="mcp-config",
), patch(
"app.agent.orchestrator.agent_mcp_manager.list_enabled_tool_specs",
new=AsyncMock(return_value=[]),
), patch(
"app.agent.orchestrator.ServerToolRegistry.resolve_web_search",
return_value=SimpleNamespace(use_local_web_search=True),
), patch(
"app.agent.orchestrator.LLMHelper.get_server_tools",
return_value=[],
), patch(
"app.agent.orchestrator.prompt_manager.get_agent_prompt",
return_value="prompt",
), patch(
"app.agent.orchestrator.SkillsMiddleware",
return_value=SimpleNamespace(name="skills", tools=[]),
), patch(
"app.agent.orchestrator.create_subagent_middlewares",
return_value=([temporary_middleware], []),
):
create_task = __import__("asyncio").create_task(
agent._create_agent(streaming=False)
)
await signature_started.wait()
create_task.cancel()
with pytest.raises(__import__("asyncio").CancelledError):
await create_task
temporary_middleware.close.assert_awaited_once()
@pytest.mark.anyio @pytest.mark.anyio
+40 -1
View File
@@ -1,6 +1,8 @@
from __future__ import annotations from __future__ import annotations
import asyncio
import sys import sys
import threading
import types import types
from unittest.mock import AsyncMock, MagicMock from unittest.mock import AsyncMock, MagicMock
@@ -242,4 +244,41 @@ async def test_stop_closes_tool_executor_after_factory_materialization(
await agent_initializer.stop_agent() await agent_initializer.stop_agent()
cleanup.assert_called_once_with(cancel_futures=True) cleanup.assert_called_once_with(wait=False, cancel_futures=True)
@pytest.mark.anyio
async def test_stop_does_not_wait_for_running_blocking_tool(monkeypatch) -> None:
"""应用关闭不得等待已经进入线程池且尚未返回的工具调用。"""
from app.agent.tools.base import MoviePilotTool
started = threading.Event()
release = threading.Event()
def _blocking_call() -> str:
started.set()
release.wait()
return "done"
worker = asyncio.create_task(
MoviePilotTool.run_blocking("web", _blocking_call)
)
assert await asyncio.wait_for(asyncio.to_thread(started.wait), timeout=1)
monkeypatch.setattr(agent_initializer, "begin_agent_shutdown", AsyncMock())
monkeypatch.setattr(
agent_initializer,
"is_tool_factory_materialized",
lambda: True,
)
monkeypatch.setattr(
agent_initializer,
"agent_initializer",
agent_initializer.AgentInitializer(),
)
try:
await asyncio.wait_for(agent_initializer.stop_agent(), timeout=0.2)
assert worker.done() is False
finally:
release.set()
assert await asyncio.wait_for(worker, timeout=1) == "done"
+436
View File
@@ -0,0 +1,436 @@
import asyncio
import json
import os
import shlex
import subprocess
import sys
from datetime import datetime
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from langchain_core.messages import ToolMessage
from app.agent.middleware.policy import AgentPolicyMiddleware
from app.agent.middleware.subagents import SubAgentTaskControlMiddleware
from app.agent.orchestrator import MoviePilotAgent
from app.agent.policy import (
AuthSource,
PrincipalType,
ToolOrigin,
ToolPolicyContext,
)
from app.agent.policy.orchestrator import DEFAULT_TOOL_POLICY_ORCHESTRATOR
from app.agent.tools.base import MoviePilotTool
from app.agent.tools.catalog import ToolCatalogSnapshot
from app.agent.tools.impl._terminal_session import (
_TerminalSession,
_TerminalSessionManager,
)
class _SlowWriteTool(MoviePilotTool):
"""模拟超时后外部写操作仍可能继续的工具。"""
name: str = "plugin_write"
description: str = "Test a slow write tool."
async def run(self, **kwargs) -> str:
"""等待足够久以触发测试超时。"""
await asyncio.sleep(1)
return "finished"
def _policy_context() -> ToolPolicyContext:
"""构造策略观测所需的最小宿主上下文。"""
return ToolPolicyContext(
session_id="session-1",
user_id="user-1",
origin=ToolOrigin.OPERATOR_DIRECT,
principal_type=PrincipalType.HUMAN,
auth_source=AuthSource.INTERNAL,
agent_context={"is_admin": True},
)
def _shell_command(code: str) -> str:
"""构造跨平台的短生命周期 Python 子进程命令。"""
args = [sys.executable, "-c", code]
if os.name == "nt":
return subprocess.list2cmdline(args)
return " ".join(shlex.quote(arg) for arg in args)
def test_timeout_marks_unknown_external_state_for_write_tools() -> None:
"""写类工具超时后必须明确提示外部状态可能仍在继续。"""
tool = type("DynamicWriteTool", (), {"name": "plugin_write", "args_schema": None})()
observation = DEFAULT_TOOL_POLICY_ORCHESTRATOR.start(
context=_policy_context(),
tool=tool,
arguments={},
)
receipt = DEFAULT_TOOL_POLICY_ORCHESTRATOR.fail(
observation,
TimeoutError("tool timeout"),
)
assert receipt.external_may_continue is True
assert receipt.needs_reconcile is True
def test_timeout_does_not_mark_safe_reads_for_reconciliation() -> None:
"""只读工具超时不应伪造外部副作用终态。"""
tool = type("SafeReadTool", (), {"name": "query_personas", "args_schema": None})()
observation = DEFAULT_TOOL_POLICY_ORCHESTRATOR.start(
context=_policy_context(),
tool=tool,
arguments={},
)
receipt = DEFAULT_TOOL_POLICY_ORCHESTRATOR.fail(
observation,
TimeoutError("tool timeout"),
)
assert receipt.external_may_continue is False
assert receipt.needs_reconcile is False
@pytest.mark.anyio
async def test_terminal_manager_close_terminates_running_pipe_session() -> None:
"""应用关闭时终端管理器必须终止仍在运行的管道进程。"""
manager = _TerminalSessionManager()
payload = await manager.start(
command=_shell_command("import time; time.sleep(30)"),
use_pty=False,
)
session = manager.get_session(payload["session_id"])
await manager.close()
assert session.process is not None
assert session.process.returncode is not None
assert session.status == "killed"
assert manager._sessions == {}
@pytest.mark.anyio
async def test_terminal_manager_close_waits_for_starting_session() -> None:
"""关闭必须接管已经获准但尚未登记的终端启动。"""
manager = _TerminalSessionManager()
start_entered = asyncio.Event()
allow_start = asyncio.Event()
session = _TerminalSession(
session_id="term-starting",
command="sleep",
cwd=".",
pid=12345,
use_pty=False,
)
async def _start_session(*_args) -> _TerminalSession:
start_entered.set()
await allow_start.wait()
return session
manager._start_pipe_session = _start_session
manager._terminate_session = AsyncMock()
start_task = asyncio.create_task(
manager.start(command="sleep", use_pty=False)
)
await start_entered.wait()
close_task = asyncio.create_task(manager.close())
await asyncio.sleep(0)
assert close_task.done() is False
allow_start.set()
with pytest.raises(RuntimeError, match="已关闭"):
await start_task
await close_task
manager._terminate_session.assert_awaited_once_with(session)
assert manager._sessions == {}
@pytest.mark.anyio
async def test_terminal_manager_rejects_start_after_close() -> None:
"""应用关闭后的终端管理器不得重新创建外部进程。"""
manager = _TerminalSessionManager()
await manager.close()
with pytest.raises(RuntimeError, match="已关闭"):
await manager.start(command="echo closed", use_pty=False)
@pytest.mark.anyio
async def test_terminal_manager_cancellation_terminates_unregistered_session() -> None:
"""调用方取消时必须回收已经创建但尚未登记的终端进程。"""
manager = _TerminalSessionManager()
session_created = asyncio.Event()
registration_locked = asyncio.Event()
release_registration = asyncio.Event()
termination_started = asyncio.Event()
session = _TerminalSession(
session_id="term-cancelled",
command="sleep",
cwd=".",
pid=12345,
use_pty=False,
)
async def _hold_registration_lock() -> None:
await session_created.wait()
async with manager._lock:
registration_locked.set()
await release_registration.wait()
async def _start_session(*_args) -> _TerminalSession:
session_created.set()
await registration_locked.wait()
return session
async def _terminate_session(_session: _TerminalSession) -> None:
termination_started.set()
lock_holder = asyncio.create_task(_hold_registration_lock())
manager._start_pipe_session = _start_session
manager._terminate_session = AsyncMock(side_effect=_terminate_session)
start_task = asyncio.create_task(
manager.start(command="sleep", use_pty=False)
)
await registration_locked.wait()
await asyncio.sleep(0)
start_task.cancel()
await termination_started.wait()
release_registration.set()
with pytest.raises(asyncio.CancelledError):
await start_task
await lock_holder
manager._terminate_session.assert_awaited_once_with(session)
assert manager._sessions == {}
assert manager._starting == 0
@pytest.mark.anyio
async def test_langchain_timeout_records_policy_failure() -> None:
"""LangChain 工具超时必须形成不泄露异常凭据的失败消息。"""
tool = _SlowWriteTool(session_id="session-1", user_id="user-1")
orchestrator = MagicMock()
orchestrator.start.side_effect = DEFAULT_TOOL_POLICY_ORCHESTRATOR.start
orchestrator.fail.side_effect = DEFAULT_TOOL_POLICY_ORCHESTRATOR.fail
orchestrator.finish.side_effect = DEFAULT_TOOL_POLICY_ORCHESTRATOR.finish
middleware = AgentPolicyMiddleware(
context=_policy_context(),
orchestrator=orchestrator,
tools=[tool],
)
request = SimpleNamespace(
tool=tool,
tool_call={"id": "call-timeout", "name": tool.name, "args": {}},
)
async def _handler(_request):
result = await tool._arun()
return ToolMessage(content=result, tool_call_id="call-timeout")
with patch("app.agent.tools.base.settings.LLM_TOOL_TIMEOUT", 0.01):
result = await middleware.awrap_tool_call(request, _handler)
assert "工具执行超时" in result.content
assert "若工具包含外部写操作" in result.content
assert "请先确认实际状态再重试" in result.content
assert result.status == "error"
orchestrator.fail.assert_called_once()
orchestrator.finish.assert_not_called()
@pytest.mark.anyio
async def test_langchain_timeout_message_sanitizes_dynamic_error() -> None:
"""非宿主工具抛出的超时异常不得把凭据带入模型上下文。"""
tool = SimpleNamespace(name="dynamic_tool")
middleware = AgentPolicyMiddleware(
context=_policy_context(),
tools=[tool],
)
request = SimpleNamespace(
tool=tool,
tool_call={"id": "call-timeout", "name": tool.name, "args": {}},
)
async def _handler(_request):
raise TimeoutError("Authorization: Bearer secret-value")
result = await middleware.awrap_tool_call(request, _handler)
assert result.status == "error"
assert "secret-value" not in result.content
@pytest.mark.anyio
async def test_agent_cleanup_closes_subagent_middlewares() -> None:
"""会话资源清理必须覆盖脱离当前回合的 subagent 控制器。"""
agent = MoviePilotAgent(session_id="session-1", user_id="user-1")
closed = []
class _Middleware:
async def close(self) -> None:
closed.append(True)
agent._subagent_middlewares = (_Middleware(),)
await agent.cleanup()
assert closed == [True]
assert agent._subagent_middlewares == ()
@pytest.mark.anyio
async def test_agent_cache_replacement_closes_previous_subagent_middleware() -> None:
"""Agent 图被替换时必须释放旧图持有的子代理控制器。"""
agent = MoviePilotAgent(session_id="session-1", user_id="user-1")
old_middleware = SimpleNamespace(close=AsyncMock())
new_middleware = SimpleNamespace(close=AsyncMock())
catalog = ToolCatalogSnapshot.from_tools(
[], plugin_revision=0, factory_revision="factory-v1"
)
agent._subagent_middlewares = (old_middleware,)
await agent._cache_agent(
signature=("new",),
agent=object(),
streaming=False,
tool_catalog=catalog,
subagent_catalog=catalog,
mcp_config_signature="mcp-config",
subagent_middlewares=(new_middleware,),
)
old_middleware.close.assert_awaited_once()
new_middleware.close.assert_not_awaited()
assert agent._subagent_middlewares == (new_middleware,)
@pytest.mark.anyio
async def test_agent_execution_failure_closes_cached_subagent_middleware() -> None:
"""图执行失败失效缓存时必须同步释放该图持有的子代理控制器。"""
agent = MoviePilotAgent(session_id="session-1", user_id="user-1")
middleware = SimpleNamespace(close=AsyncMock())
graph = SimpleNamespace(ainvoke=AsyncMock(side_effect=RuntimeError("failed")))
agent._compiled_agent_bundle = SimpleNamespace(agent=graph)
agent._subagent_middlewares = (middleware,)
agent._should_stream = lambda: False
agent._create_agent = AsyncMock(return_value=graph)
agent._dispatch_execution_notice = AsyncMock()
agent.stream_handler = SimpleNamespace(
stop_streaming=AsyncMock(return_value=(False, ""))
)
result, _ = await agent._execute_agent([])
assert "failed" in result
middleware.close.assert_awaited_once()
assert agent._compiled_agent_bundle is None
assert agent._subagent_middlewares == ()
def test_subagent_control_middleware_close_is_idempotent() -> None:
"""子代理控制器的全局关闭路径应可重复调用。"""
middleware = object.__new__(SubAgentTaskControlMiddleware)
middleware._tasks = {}
async def _close() -> None:
await middleware.close()
await middleware.close()
asyncio.run(_close())
assert middleware._tasks == {}
@pytest.mark.anyio
async def test_subagent_close_has_bounded_cancel_wait() -> None:
"""子代理忽略首次取消时,控制器关闭仍必须在上限内返回。"""
middleware = object.__new__(SubAgentTaskControlMiddleware)
release = asyncio.Event()
cancelled = asyncio.Event()
async def _ignore_first_cancel() -> None:
try:
await asyncio.Future()
except asyncio.CancelledError:
cancelled.set()
await release.wait()
task = asyncio.create_task(_ignore_first_cancel())
record = SimpleNamespace(
task_id="subagent-stubborn",
description="stubborn",
subagent_type="general-purpose",
task=task,
created_at=datetime.now(),
started_at=datetime.now(),
finished_at=None,
)
middleware._tasks = {record.task_id: record}
await asyncio.sleep(0)
with patch(
"app.agent.middleware.subagents.SUBAGENT_CANCEL_GRACE_SECONDS",
0.01,
):
await asyncio.wait_for(middleware.close(), timeout=0.2)
assert cancelled.is_set()
assert task.done() is False
assert middleware._tasks == {}
release.set()
await asyncio.wait_for(task, timeout=0.2)
@pytest.mark.anyio
async def test_subagent_cancel_reports_tasks_still_stopping() -> None:
"""取消上限到达后不得把仍运行的子代理报告为取消成功。"""
middleware = object.__new__(SubAgentTaskControlMiddleware)
release = asyncio.Event()
async def _ignore_first_cancel() -> None:
try:
await asyncio.Future()
except asyncio.CancelledError:
await release.wait()
task = asyncio.create_task(_ignore_first_cancel())
record = SimpleNamespace(
task_id="subagent-stopping",
description="stopping",
subagent_type="general-purpose",
task=task,
created_at=datetime.now(),
started_at=datetime.now(),
finished_at=None,
)
middleware._tasks = {record.task_id: record}
await asyncio.sleep(0)
with patch(
"app.agent.middleware.subagents.SUBAGENT_CANCEL_GRACE_SECONDS",
0.01,
):
payload = await middleware._control_task(
action="cancel",
task_id=record.task_id,
)
result = json.loads(payload)
assert result["success"] is False
assert result["cancel_pending_task_ids"] == [record.task_id]
assert result["tasks"][0]["status"] == "running"
release.set()
await asyncio.wait_for(task, timeout=0.2)
+49
View File
@@ -525,6 +525,55 @@ def test_control_tool_pipeline_stops_after_failed_step():
asyncio.run(_run_test()) asyncio.run(_run_test())
def test_control_tool_pipeline_timeout_is_bounded_when_task_ignores_cancel():
"""管道步骤忽略取消时,等待上限仍必须按时返回失败。"""
async def _run_test():
model = FakeListChatModel(responses=["ok"])
middleware = SubAgentTaskControlMiddleware(
model=model,
profiles=subagent_module._builtin_subagent_profiles(),
tools=[],
)
release = asyncio.Event()
cancelled = asyncio.Event()
async def _ignore_cancel(self, *, description, subagent_type, task_id=None):
try:
await asyncio.Future()
except asyncio.CancelledError:
cancelled.set()
await release.wait()
return "late-result"
with patch.object(
subagent_module._SubAgentAgentProvider,
"run_task",
new=_ignore_cancel,
):
pipeline = asyncio.create_task(
middleware._control_task(
action="pipeline",
description="慢任务",
timeout_ms=10,
)
)
payload = json.loads(await asyncio.wait_for(pipeline, timeout=0.2))
assert payload["success"] is False
assert "等待超时" in payload["error"]
assert payload["tasks"][0]["status"] == "running"
assert cancelled.is_set()
release.set()
await asyncio.wait_for(
middleware._tasks[payload["tasks"][0]["task_id"]].task,
timeout=0.2,
)
asyncio.run(_run_test())
def test_after_agent_cancels_unfinished_tasks(): def test_after_agent_cancels_unfinished_tasks():
"""Agent 结束时应取消仍在运行的异步子代理任务。""" """Agent 结束时应取消仍在运行的异步子代理任务。"""
+10 -7
View File
@@ -4,7 +4,12 @@ from unittest.mock import patch
import pytest import pytest
from app.agent.tools.base import MoviePilotTool, _blocking_executors, shutdown_blocking_executors from app.agent.tools.base import (
MoviePilotTool,
ToolExecutionTimeoutError,
_blocking_executors,
shutdown_blocking_executors,
)
from app.agent.tools.manager import MoviePilotToolsManager from app.agent.tools.manager import MoviePilotToolsManager
@@ -31,18 +36,16 @@ class BlockingAgentTool(MoviePilotTool):
return "unused" return "unused"
def test_arun_returns_timeout_message_when_tool_exceeds_limit(): def test_arun_raises_timeout_when_tool_exceeds_limit():
"""LangChain 工具入口应按 LLM_TOOL_TIMEOUT 停止等待慢工具""" """底层工具入口应把超时交给宿主策略记录失败终态"""
tool = SlowAgentTool(session_id="session-1", user_id="10001") tool = SlowAgentTool(session_id="session-1", user_id="10001")
async def _run_tool(): async def _run_tool():
with patch("app.agent.tools.base.settings.LLM_TOOL_TIMEOUT", 0.05): with patch("app.agent.tools.base.settings.LLM_TOOL_TIMEOUT", 0.05):
return await tool._arun() return await tool._arun()
result = asyncio.run(_run_tool()) with pytest.raises(ToolExecutionTimeoutError, match="超过 0.05 秒"):
asyncio.run(_run_tool())
assert "工具 slow_agent_tool 执行超时" in result
assert "超过 0.05 秒" in result
def test_http_tool_manager_uses_same_timeout_guard(): def test_http_tool_manager_uses_same_timeout_guard():