mirror of
https://github.com/jxxghp/MoviePilot.git
synced 2026-09-06 16:07:01 +08:00
Merge remote-tracking branch 'origin/v3' into HEAD
This commit is contained in:
+48
-6
@@ -39,6 +39,7 @@ from app.agent.middleware.jobs import (
|
|||||||
)
|
)
|
||||||
from app.agent.middleware.memory import MemoryMiddleware
|
from app.agent.middleware.memory import MemoryMiddleware
|
||||||
from app.agent.middleware.patch_tool_calls import PatchToolCallsMiddleware
|
from app.agent.middleware.patch_tool_calls import PatchToolCallsMiddleware
|
||||||
|
from app.agent.middleware.policy import AgentPolicyMiddleware
|
||||||
from app.agent.middleware.runtime_config import RuntimeConfigMiddleware
|
from app.agent.middleware.runtime_config import RuntimeConfigMiddleware
|
||||||
from app.agent.middleware.skills import SKILL_TOOL_NAME, SkillsMiddleware
|
from app.agent.middleware.skills import SKILL_TOOL_NAME, SkillsMiddleware
|
||||||
from app.agent.middleware.subagents import (
|
from app.agent.middleware.subagents import (
|
||||||
@@ -50,6 +51,12 @@ from app.agent.middleware.subagents import (
|
|||||||
from app.agent.middleware.tool_selection import ToolSelectorMiddleware
|
from app.agent.middleware.tool_selection import ToolSelectorMiddleware
|
||||||
from app.agent.middleware.usage import UsageMiddleware
|
from app.agent.middleware.usage import UsageMiddleware
|
||||||
from app.agent.prompt import prompt_manager
|
from app.agent.prompt import prompt_manager
|
||||||
|
from app.agent.policy import (
|
||||||
|
AuthSource,
|
||||||
|
PrincipalType,
|
||||||
|
ToolOrigin,
|
||||||
|
ToolPolicyContext,
|
||||||
|
)
|
||||||
from app.agent.runtime import agent_runtime_manager
|
from app.agent.runtime import agent_runtime_manager
|
||||||
from app.agent.mcp import agent_mcp_manager
|
from app.agent.mcp import agent_mcp_manager
|
||||||
from app.agent.tools.factory import MoviePilotToolFactory
|
from app.agent.tools.factory import MoviePilotToolFactory
|
||||||
@@ -727,6 +734,40 @@ class MoviePilotAgent:
|
|||||||
"original_chat_id": None if self.is_background else self.original_chat_id,
|
"original_chat_id": None if self.is_background else self.original_chat_id,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
def _build_policy_context(self) -> ToolPolicyContext:
|
||||||
|
"""根据宿主入口建立模型参数无法伪造的策略上下文。"""
|
||||||
|
if not self.has_message_context:
|
||||||
|
origin = ToolOrigin.BACKGROUND
|
||||||
|
principal_type = PrincipalType.BACKGROUND
|
||||||
|
auth_source = AuthSource.INTERNAL
|
||||||
|
elif self.channel == MessageChannel.Web.value and self.source in {
|
||||||
|
"openai",
|
||||||
|
"openai.responses",
|
||||||
|
"anthropic",
|
||||||
|
}:
|
||||||
|
origin = ToolOrigin.AGENT_API
|
||||||
|
principal_type = PrincipalType.SYSTEM_ADMIN_INTEGRATION
|
||||||
|
auth_source = AuthSource.API_TOKEN
|
||||||
|
else:
|
||||||
|
origin = ToolOrigin.AGENT_INTERACTIVE
|
||||||
|
principal_type = PrincipalType.HUMAN
|
||||||
|
auth_source = (
|
||||||
|
AuthSource.WEB_SESSION
|
||||||
|
if self.channel
|
||||||
|
in {MessageChannel.Web.value, MessageChannel.WebAgent.value}
|
||||||
|
else AuthSource.CHANNEL
|
||||||
|
)
|
||||||
|
return ToolPolicyContext(
|
||||||
|
session_id=self.session_id,
|
||||||
|
user_id=str(self.user_id or self.username or principal_type.value),
|
||||||
|
origin=origin,
|
||||||
|
principal_type=principal_type,
|
||||||
|
auth_source=auth_source,
|
||||||
|
agent_context=self._tool_context,
|
||||||
|
channel=self.channel,
|
||||||
|
source=self.source,
|
||||||
|
)
|
||||||
|
|
||||||
def _should_stream(self) -> bool:
|
def _should_stream(self) -> bool:
|
||||||
"""
|
"""
|
||||||
判断是否应启用流式输出:
|
判断是否应启用流式输出:
|
||||||
@@ -1255,6 +1296,7 @@ class MoviePilotAgent:
|
|||||||
# LLM 模型(用于 agent 执行)
|
# LLM 模型(用于 agent 执行)
|
||||||
agent_model = await self._initialize_llm(streaming=streaming)
|
agent_model = await self._initialize_llm(streaming=streaming)
|
||||||
self._sync_model_profile(agent_model)
|
self._sync_model_profile(agent_model)
|
||||||
|
# 供应商原生工具不进入本地 ToolNode,宿主策略只覆盖 client-side tools。
|
||||||
server_tools = LLMHelper.get_server_tools(agent_model)
|
server_tools = LLMHelper.get_server_tools(agent_model)
|
||||||
use_local_web_search = LLMHelper.should_use_local_web_search(agent_model)
|
use_local_web_search = LLMHelper.should_use_local_web_search(agent_model)
|
||||||
|
|
||||||
@@ -1292,11 +1334,13 @@ class MoviePilotAgent:
|
|||||||
enabled=use_local_web_search,
|
enabled=use_local_web_search,
|
||||||
)
|
)
|
||||||
subagent_tools.extend(await self._initialize_subagent_mcp_tools())
|
subagent_tools.extend(await self._initialize_subagent_mcp_tools())
|
||||||
|
policy_context = self._build_policy_context()
|
||||||
subagent_middlewares, subagent_task_tools = create_subagent_middlewares(
|
subagent_middlewares, subagent_task_tools = create_subagent_middlewares(
|
||||||
model=non_streaming_model,
|
model=non_streaming_model,
|
||||||
tools=subagent_tools,
|
tools=subagent_tools,
|
||||||
server_tools=server_tools,
|
server_tools=server_tools,
|
||||||
stream_handler=self.stream_handler,
|
stream_handler=self.stream_handler,
|
||||||
|
policy_context=policy_context.for_subagent(),
|
||||||
)
|
)
|
||||||
max_tools = settings.LLM_MAX_TOOLS
|
max_tools = settings.LLM_MAX_TOOLS
|
||||||
always_include_tools = (
|
always_include_tools = (
|
||||||
@@ -1324,6 +1368,8 @@ class MoviePilotAgent:
|
|||||||
|
|
||||||
# 中间件
|
# 中间件
|
||||||
middlewares = [
|
middlewares = [
|
||||||
|
# 宿主策略必须位于最外层,确保插件覆盖工具基类也不能绕过。
|
||||||
|
AgentPolicyMiddleware(context=policy_context),
|
||||||
# Skills
|
# Skills
|
||||||
skills_middleware,
|
skills_middleware,
|
||||||
# Jobs 任务管理
|
# Jobs 任务管理
|
||||||
@@ -1334,6 +1380,8 @@ class MoviePilotAgent:
|
|||||||
RuntimeConfigMiddleware(),
|
RuntimeConfigMiddleware(),
|
||||||
# 记忆管理
|
# 记忆管理
|
||||||
MemoryMiddleware(memory_dir=str(agent_runtime_manager.memory_dir)),
|
MemoryMiddleware(memory_dir=str(agent_runtime_manager.memory_dir)),
|
||||||
|
# 活动日志依赖记忆上下文,并应在摘要压缩前完成读取与记录。
|
||||||
|
*([activity_log_middleware] if activity_log_middleware else []),
|
||||||
# 上下文压缩
|
# 上下文压缩
|
||||||
SummarizationMiddleware(
|
SummarizationMiddleware(
|
||||||
model=non_streaming_model, trigger=("fraction", 0.85)
|
model=non_streaming_model, trigger=("fraction", 0.85)
|
||||||
@@ -1346,12 +1394,6 @@ class MoviePilotAgent:
|
|||||||
UsageMiddleware(on_usage=self._record_usage),
|
UsageMiddleware(on_usage=self._record_usage),
|
||||||
]
|
]
|
||||||
|
|
||||||
if self.has_message_context:
|
|
||||||
middlewares.insert(
|
|
||||||
4,
|
|
||||||
activity_log_middleware,
|
|
||||||
)
|
|
||||||
|
|
||||||
# 工具选择
|
# 工具选择
|
||||||
if max_tools > 0:
|
if max_tools > 0:
|
||||||
middlewares.append(
|
middlewares.append(
|
||||||
|
|||||||
@@ -4,6 +4,7 @@ from typing import Any, Optional, Tuple
|
|||||||
|
|
||||||
from fastapi.concurrency import run_in_threadpool
|
from fastapi.concurrency import run_in_threadpool
|
||||||
|
|
||||||
|
from app.agent.policy import sanitize_for_host
|
||||||
from app.chain import ChainBase
|
from app.chain import ChainBase
|
||||||
from app.log import logger
|
from app.log import logger
|
||||||
from app.schemas import Notification
|
from app.schemas import Notification
|
||||||
@@ -256,10 +257,14 @@ class StreamingHandler:
|
|||||||
"""
|
"""
|
||||||
记录一次工具调用,供非啰嗦模式下延迟汇总输出。
|
记录一次工具调用,供非啰嗦模式下延迟汇总输出。
|
||||||
"""
|
"""
|
||||||
|
recorded_message = sanitize_for_host(tool_message) if tool_message else tool_message
|
||||||
|
recorded_args = sanitize_for_host(tool_kwargs or {})
|
||||||
|
if not isinstance(recorded_args, dict):
|
||||||
|
recorded_args = {}
|
||||||
category, target = self._classify_tool_call(
|
category, target = self._classify_tool_call(
|
||||||
tool_name=tool_name,
|
tool_name=tool_name,
|
||||||
tool_message=tool_message,
|
tool_message=recorded_message,
|
||||||
tool_kwargs=tool_kwargs or {},
|
tool_kwargs=recorded_args,
|
||||||
)
|
)
|
||||||
target_values = []
|
target_values = []
|
||||||
if isinstance(target, (list, tuple, set)):
|
if isinstance(target, (list, tuple, set)):
|
||||||
|
|||||||
@@ -33,6 +33,7 @@ from langgraph.runtime import Runtime
|
|||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
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 sanitize_for_host, summarize_error, summarize_result
|
||||||
from app.agent.tools.tags import ToolTag
|
from app.agent.tools.tags import ToolTag
|
||||||
from app.log import logger
|
from app.log import logger
|
||||||
|
|
||||||
@@ -181,7 +182,9 @@ def load_activity_log_index(activity_dir: str, days: int = PROMPT_LOAD_DAYS) ->
|
|||||||
try:
|
try:
|
||||||
content = log_path.read_text(encoding="utf-8", errors="replace")
|
content = log_path.read_text(encoding="utf-8", errors="replace")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning(f"读取活动日志索引失败 {log_path}: {e}")
|
logger.warning(
|
||||||
|
f"读取活动日志索引失败 {log_path}: {summarize_error(e)}"
|
||||||
|
)
|
||||||
continue
|
continue
|
||||||
entry_count = len(_parse_activity_entries(date_str, content))
|
entry_count = len(_parse_activity_entries(date_str, content))
|
||||||
if entry_count:
|
if entry_count:
|
||||||
@@ -245,7 +248,7 @@ def query_activity_logs(
|
|||||||
try:
|
try:
|
||||||
content = log_path.read_text(encoding="utf-8", errors="replace")
|
content = log_path.read_text(encoding="utf-8", errors="replace")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning(f"读取活动日志失败 {log_path}: {e}")
|
logger.warning(f"读取活动日志失败 {log_path}: {summarize_error(e)}")
|
||||||
continue
|
continue
|
||||||
for entry in _parse_activity_entries(date_str, content):
|
for entry in _parse_activity_entries(date_str, content):
|
||||||
if normalized_keyword and not _activity_summary_matches_keyword(
|
if normalized_keyword and not _activity_summary_matches_keyword(
|
||||||
@@ -287,14 +290,16 @@ class _ActivityLogToolProvider:
|
|||||||
limit: Optional[int] = DEFAULT_QUERY_LIMIT,
|
limit: Optional[int] = DEFAULT_QUERY_LIMIT,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""查询活动日志并返回 JSON 字符串。"""
|
"""查询活动日志并返回 JSON 字符串。"""
|
||||||
logger.info(
|
logged_args = sanitize_for_host(
|
||||||
"查询活动日志: keyword=%s, use_regex=%s, date=%s, days=%s, limit=%s",
|
{
|
||||||
keyword,
|
"keyword": keyword,
|
||||||
use_regex,
|
"use_regex": use_regex,
|
||||||
date,
|
"date": date,
|
||||||
days,
|
"days": days,
|
||||||
limit,
|
"limit": limit,
|
||||||
|
}
|
||||||
)
|
)
|
||||||
|
logger.info(f"查询活动日志: args={logged_args}")
|
||||||
try:
|
try:
|
||||||
payload = query_activity_logs(
|
payload = query_activity_logs(
|
||||||
self._activity_dir,
|
self._activity_dir,
|
||||||
@@ -306,11 +311,12 @@ class _ActivityLogToolProvider:
|
|||||||
)
|
)
|
||||||
return json.dumps(payload, ensure_ascii=False, indent=2)
|
return json.dumps(payload, ensure_ascii=False, indent=2)
|
||||||
except Exception as err:
|
except Exception as err:
|
||||||
logger.error(f"查询活动日志失败: {err}", exc_info=True)
|
error_summary = summarize_error(err)
|
||||||
|
logger.error(f"查询活动日志失败: {error_summary}")
|
||||||
return json.dumps(
|
return json.dumps(
|
||||||
{
|
{
|
||||||
"success": False,
|
"success": False,
|
||||||
"message": f"查询活动日志时发生错误: {str(err)}",
|
"message": f"查询活动日志时发生错误: {error_summary}",
|
||||||
},
|
},
|
||||||
ensure_ascii=False,
|
ensure_ascii=False,
|
||||||
)
|
)
|
||||||
@@ -454,7 +460,7 @@ async def _summarize_with_llm(conversation_text: str) -> Optional[str]:
|
|||||||
return None
|
return None
|
||||||
return summary if summary else None
|
return summary if summary else None
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.debug(f"LLM 活动摘要生成失败: {e}")
|
logger.debug(f"LLM 活动摘要生成失败: {summarize_error(e)}")
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
@@ -571,9 +577,9 @@ class ActivityLogMiddleware(AgentMiddleware[ActivityLogState, ContextT, Response
|
|||||||
else:
|
else:
|
||||||
with os.fdopen(fd, "w", encoding="utf-8") as stream:
|
with os.fdopen(fd, "w", encoding="utf-8") as stream:
|
||||||
stream.write(header + entry)
|
stream.write(header + entry)
|
||||||
logger.debug(f"Activity logged: {summary[:80]}")
|
logger.debug(f"Activity logged: {summarize_result(summary, max_chars=80)}")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning(f"Failed to append activity log: {e}")
|
logger.warning(f"Failed to append activity log: {summarize_error(e)}")
|
||||||
|
|
||||||
async def _cleanup_old_logs(self) -> None:
|
async def _cleanup_old_logs(self) -> None:
|
||||||
"""清理超过保留天数的旧日志文件。"""
|
"""清理超过保留天数的旧日志文件。"""
|
||||||
@@ -599,7 +605,9 @@ class ActivityLogMiddleware(AgentMiddleware[ActivityLogState, ContextT, Response
|
|||||||
except ValueError:
|
except ValueError:
|
||||||
continue
|
continue
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning(f"Failed to cleanup old activity logs: {e}")
|
logger.warning(
|
||||||
|
f"Failed to cleanup old activity logs: {summarize_error(e)}"
|
||||||
|
)
|
||||||
|
|
||||||
def _schedule_activity_recording(self, messages: list) -> None:
|
def _schedule_activity_recording(self, messages: list) -> None:
|
||||||
"""提交后台活动记录任务,不阻塞当前 Agent 会话结束。"""
|
"""提交后台活动记录任务,不阻塞当前 Agent 会话结束。"""
|
||||||
@@ -615,7 +623,7 @@ class ActivityLogMiddleware(AgentMiddleware[ActivityLogState, ContextT, Response
|
|||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
logger.debug("活动日志后台记录任务已取消")
|
logger.debug("活动日志后台记录任务已取消")
|
||||||
except Exception as err:
|
except Exception as err:
|
||||||
logger.warning(f"活动日志后台记录任务失败: {err}")
|
logger.warning(f"活动日志后台记录任务失败: {summarize_error(err)}")
|
||||||
|
|
||||||
async def _record_activity(self, messages: list) -> None:
|
async def _record_activity(self, messages: list) -> None:
|
||||||
"""在后台生成本轮活动摘要并写入活动日志。"""
|
"""在后台生成本轮活动摘要并写入活动日志。"""
|
||||||
@@ -637,7 +645,7 @@ class ActivityLogMiddleware(AgentMiddleware[ActivityLogState, ContextT, Response
|
|||||||
if summary:
|
if summary:
|
||||||
await self._append_activity(summary)
|
await self._append_activity(summary)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning(f"Failed to record activity: {e}")
|
logger.warning(f"Failed to record activity: {summarize_error(e)}")
|
||||||
|
|
||||||
async def abefore_agent(
|
async def abefore_agent(
|
||||||
self, state: ActivityLogState, runtime: Runtime
|
self, state: ActivityLogState, runtime: Runtime
|
||||||
@@ -686,9 +694,12 @@ class ActivityLogMiddleware(AgentMiddleware[ActivityLogState, ContextT, Response
|
|||||||
tool_args = tool_call.get("args") or {}
|
tool_args = tool_call.get("args") or {}
|
||||||
if not isinstance(tool_args, dict):
|
if not isinstance(tool_args, dict):
|
||||||
tool_args = {}
|
tool_args = {}
|
||||||
|
logged_args = sanitize_for_host(tool_args)
|
||||||
|
if not isinstance(logged_args, dict):
|
||||||
|
logged_args = {}
|
||||||
logger.info(
|
logger.info(
|
||||||
f"开始执行活动日志查询工具: keyword={tool_args.get('keyword') or '-'}, "
|
f"开始执行活动日志查询工具: keyword={logged_args.get('keyword') or '-'}, "
|
||||||
f"date={tool_args.get('date') or '-'}"
|
f"date={logged_args.get('date') or '-'}"
|
||||||
)
|
)
|
||||||
if self.stream_handler and getattr(self.stream_handler, "is_streaming", False):
|
if self.stream_handler and getattr(self.stream_handler, "is_streaming", False):
|
||||||
self.stream_handler.record_tool_call(
|
self.stream_handler.record_tool_call(
|
||||||
@@ -699,7 +710,9 @@ class ActivityLogMiddleware(AgentMiddleware[ActivityLogState, ContextT, Response
|
|||||||
try:
|
try:
|
||||||
result = await handler(request)
|
result = await handler(request)
|
||||||
except Exception as err:
|
except Exception as err:
|
||||||
logger.error(f"活动日志查询工具执行失败: error={err}")
|
logger.error(
|
||||||
|
f"活动日志查询工具执行失败: error={summarize_error(err)}"
|
||||||
|
)
|
||||||
raise
|
raise
|
||||||
logger.info("活动日志查询工具执行完成")
|
logger.info("活动日志查询工具执行完成")
|
||||||
return result
|
return result
|
||||||
@@ -714,7 +727,7 @@ class ActivityLogMiddleware(AgentMiddleware[ActivityLogState, ContextT, Response
|
|||||||
return None
|
return None
|
||||||
self._schedule_activity_recording(list(messages))
|
self._schedule_activity_recording(list(messages))
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning(f"Failed to record activity: {e}")
|
logger.warning(f"Failed to record activity: {summarize_error(e)}")
|
||||||
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,72 @@
|
|||||||
|
"""LangChain 工具调用的 MoviePilot 宿主策略中间件。"""
|
||||||
|
|
||||||
|
from collections.abc import Awaitable, Callable
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from langchain.agents.middleware import AgentMiddleware, ToolCallRequest
|
||||||
|
|
||||||
|
from app.agent.policy import (
|
||||||
|
DEFAULT_TOOL_POLICY_ORCHESTRATOR,
|
||||||
|
AgentToolPolicyOrchestrator,
|
||||||
|
ToolPolicyContext,
|
||||||
|
call_policy_hook,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class AgentPolicyMiddleware(AgentMiddleware):
|
||||||
|
"""观测进入本地 ToolNode 的 client-side 工具调用和结果。
|
||||||
|
|
||||||
|
模型供应商原生 server tools 在供应商侧执行,不经过本地 middleware,
|
||||||
|
因而不具备这里生成的 start/finish/fail 回执。
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
context: ToolPolicyContext,
|
||||||
|
orchestrator: AgentToolPolicyOrchestrator = DEFAULT_TOOL_POLICY_ORCHESTRATOR,
|
||||||
|
) -> None:
|
||||||
|
"""绑定宿主可信上下文和共享策略编排器。"""
|
||||||
|
self.context = context
|
||||||
|
self.orchestrator = orchestrator
|
||||||
|
|
||||||
|
async def awrap_tool_call(
|
||||||
|
self,
|
||||||
|
request: ToolCallRequest,
|
||||||
|
handler: Callable[[ToolCallRequest], Awaitable[Any]],
|
||||||
|
) -> Any:
|
||||||
|
"""在 handler 外层生成 shadow 决策和 secret-safe 回执摘要。"""
|
||||||
|
tool_call = request.tool_call or {}
|
||||||
|
arguments = tool_call.get("args") or {}
|
||||||
|
if not isinstance(arguments, dict):
|
||||||
|
arguments = {}
|
||||||
|
observation = call_policy_hook(
|
||||||
|
"start",
|
||||||
|
self.orchestrator.start,
|
||||||
|
context=self.context,
|
||||||
|
tool=request.tool,
|
||||||
|
arguments=arguments,
|
||||||
|
invocation_id=tool_call.get("id"),
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
result = await handler(request)
|
||||||
|
except Exception as error:
|
||||||
|
if observation is not None:
|
||||||
|
call_policy_hook(
|
||||||
|
"fail",
|
||||||
|
self.orchestrator.fail,
|
||||||
|
observation,
|
||||||
|
error,
|
||||||
|
)
|
||||||
|
raise
|
||||||
|
if observation is not None:
|
||||||
|
call_policy_hook(
|
||||||
|
"finish",
|
||||||
|
self.orchestrator.finish,
|
||||||
|
observation,
|
||||||
|
result,
|
||||||
|
)
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = ["AgentPolicyMiddleware"]
|
||||||
@@ -24,6 +24,7 @@ from langgraph.runtime import Runtime
|
|||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
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 sanitize_for_host, summarize_error
|
||||||
from app.agent.tools.tags import ToolTag
|
from app.agent.tools.tags import ToolTag
|
||||||
from app.log import logger
|
from app.log import logger
|
||||||
|
|
||||||
@@ -124,7 +125,7 @@ def _parse_skill_metadata( # noqa: C901
|
|||||||
try:
|
try:
|
||||||
frontmatter_data = yaml.safe_load(frontmatter_str)
|
frontmatter_data = yaml.safe_load(frontmatter_str)
|
||||||
except yaml.YAMLError as e:
|
except yaml.YAMLError as e:
|
||||||
logger.warning("Invalid YAML in %s: %s", skill_path, e)
|
logger.warning("Invalid YAML in %s: %s", skill_path, summarize_error(e))
|
||||||
return None
|
return None
|
||||||
|
|
||||||
if not isinstance(frontmatter_data, dict):
|
if not isinstance(frontmatter_data, dict):
|
||||||
@@ -339,7 +340,7 @@ def _extract_version(skill_md: Path) -> int:
|
|||||||
try:
|
try:
|
||||||
content = skill_md.read_text(encoding="utf-8", errors="replace")
|
content = skill_md.read_text(encoding="utf-8", errors="replace")
|
||||||
except Exception as err:
|
except Exception as err:
|
||||||
logger.debug(f"读取技能版本失败: {err}")
|
logger.debug(f"读取技能版本失败: {summarize_error(err)}")
|
||||||
return 0
|
return 0
|
||||||
match = re.match(r"^---\s*\n(.*?)\n---\s*\n", content, re.DOTALL)
|
match = re.match(r"^---\s*\n(.*?)\n---\s*\n", content, re.DOTALL)
|
||||||
if not match:
|
if not match:
|
||||||
@@ -397,7 +398,11 @@ def _sync_bundled_skills(bundled_dir: Path, target_dir: Path) -> None:
|
|||||||
"已自动复制内置技能 '%s' -> '%s'", skill_src.name, skill_dst
|
"已自动复制内置技能 '%s' -> '%s'", skill_src.name, skill_dst
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning("复制内置技能 '%s' 失败: %s", skill_src.name, e)
|
logger.warning(
|
||||||
|
"复制内置技能 '%s' 失败: %s",
|
||||||
|
sanitize_for_host(skill_src.name),
|
||||||
|
summarize_error(e),
|
||||||
|
)
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# 目标已存在,比较版本号
|
# 目标已存在,比较版本号
|
||||||
@@ -424,7 +429,11 @@ def _sync_bundled_skills(bundled_dir: Path, target_dir: Path) -> None:
|
|||||||
bundled_version,
|
bundled_version,
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning("更新内置技能 '%s' 失败: %s", skill_src.name, e)
|
logger.warning(
|
||||||
|
"更新内置技能 '%s' 失败: %s",
|
||||||
|
sanitize_for_host(skill_src.name),
|
||||||
|
summarize_error(e),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class _SkillToolProvider:
|
class _SkillToolProvider:
|
||||||
@@ -519,7 +528,7 @@ class _SkillToolProvider:
|
|||||||
|
|
||||||
async def load_skill(self, name: str) -> str:
|
async def load_skill(self, name: str) -> str:
|
||||||
"""加载指定 Skill 的完整说明并返回 JSON 字符串。"""
|
"""加载指定 Skill 的完整说明并返回 JSON 字符串。"""
|
||||||
logger.info(f"加载 Skill: name={name}")
|
logger.info(f"加载 Skill: name={sanitize_for_host(name)}")
|
||||||
try:
|
try:
|
||||||
skill = await self._find_skill(name)
|
skill = await self._find_skill(name)
|
||||||
if not skill:
|
if not skill:
|
||||||
@@ -547,11 +556,12 @@ class _SkillToolProvider:
|
|||||||
}
|
}
|
||||||
)
|
)
|
||||||
except Exception as err:
|
except Exception as err:
|
||||||
logger.error(f"加载 Skill 失败: {err}", exc_info=True)
|
error_summary = summarize_error(err)
|
||||||
|
logger.error(f"加载 Skill 失败: {error_summary}")
|
||||||
return json.dumps(
|
return json.dumps(
|
||||||
{
|
{
|
||||||
"success": False,
|
"success": False,
|
||||||
"message": f"加载 Skill 时发生错误: {str(err)}",
|
"message": f"加载 Skill 时发生错误: {error_summary}",
|
||||||
},
|
},
|
||||||
ensure_ascii=False,
|
ensure_ascii=False,
|
||||||
)
|
)
|
||||||
@@ -623,7 +633,7 @@ class SkillsMiddleware(AgentMiddleware[SkillsState, ContextT, ResponseT]): # no
|
|||||||
try:
|
try:
|
||||||
_sync_bundled_skills(bundled, target)
|
_sync_bundled_skills(bundled, target)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning("同步内置技能失败: %s", e)
|
logger.warning(f"同步内置技能失败: {summarize_error(e)}")
|
||||||
|
|
||||||
def _load_skills_metadata(self) -> list[SkillMetadata]:
|
def _load_skills_metadata(self) -> list[SkillMetadata]:
|
||||||
"""同步加载当前配置目录中的 Skill 元数据。"""
|
"""同步加载当前配置目录中的 Skill 元数据。"""
|
||||||
@@ -728,8 +738,11 @@ class SkillsMiddleware(AgentMiddleware[SkillsState, ContextT, ResponseT]): # no
|
|||||||
tool_args = tool_call.get("args") or {}
|
tool_args = tool_call.get("args") or {}
|
||||||
if not isinstance(tool_args, dict):
|
if not isinstance(tool_args, dict):
|
||||||
tool_args = {}
|
tool_args = {}
|
||||||
|
logged_args = sanitize_for_host(tool_args)
|
||||||
|
if not isinstance(logged_args, dict):
|
||||||
|
logged_args = {}
|
||||||
logger.info(
|
logger.info(
|
||||||
f"开始执行 Skill 工具: name={tool_args.get('name') or '-'}"
|
f"开始执行 Skill 工具: name={logged_args.get('name') or '-'}"
|
||||||
)
|
)
|
||||||
if self.stream_handler and getattr(self.stream_handler, "is_streaming", False):
|
if self.stream_handler and getattr(self.stream_handler, "is_streaming", False):
|
||||||
self.stream_handler.record_tool_call(
|
self.stream_handler.record_tool_call(
|
||||||
@@ -740,7 +753,7 @@ class SkillsMiddleware(AgentMiddleware[SkillsState, ContextT, ResponseT]): # no
|
|||||||
try:
|
try:
|
||||||
result = await handler(request)
|
result = await handler(request)
|
||||||
except Exception as err:
|
except Exception as err:
|
||||||
logger.error(f"Skill 工具执行失败: error={err}")
|
logger.error(f"Skill 工具执行失败: error={summarize_error(err)}")
|
||||||
raise
|
raise
|
||||||
logger.info("Skill 工具执行完成")
|
logger.info("Skill 工具执行完成")
|
||||||
return result
|
return result
|
||||||
|
|||||||
@@ -24,7 +24,16 @@ 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 import LLMHelper
|
||||||
|
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 (
|
||||||
|
AuthSource,
|
||||||
|
PrincipalType,
|
||||||
|
ToolOrigin,
|
||||||
|
ToolPolicyContext,
|
||||||
|
sanitize_for_host,
|
||||||
|
summarize_error,
|
||||||
|
)
|
||||||
from app.agent.runtime import SubAgentDefinition, agent_runtime_manager
|
from app.agent.runtime import SubAgentDefinition, agent_runtime_manager
|
||||||
from app.agent.tools.tags import ToolTag
|
from app.agent.tools.tags import ToolTag
|
||||||
from app.log import logger
|
from app.log import logger
|
||||||
@@ -93,6 +102,36 @@ Requirements:
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
def _default_subagent_policy_context(tools: list[BaseTool]) -> ToolPolicyContext:
|
||||||
|
"""从已注入工具继承会话归属,确保独立构造的子图也经过宿主策略。"""
|
||||||
|
for tool in tools:
|
||||||
|
session_id = getattr(tool, "_session_id", None)
|
||||||
|
user_id = getattr(tool, "_user_id", None)
|
||||||
|
if not session_id and not user_id:
|
||||||
|
continue
|
||||||
|
agent_context = getattr(tool, "_agent_context", None)
|
||||||
|
if not isinstance(agent_context, dict):
|
||||||
|
agent_context = {}
|
||||||
|
return ToolPolicyContext(
|
||||||
|
session_id=str(session_id or "subagent"),
|
||||||
|
user_id=str(user_id or "subagent"),
|
||||||
|
origin=ToolOrigin.SUBAGENT,
|
||||||
|
principal_type=PrincipalType.SUBAGENT,
|
||||||
|
auth_source=AuthSource.INTERNAL,
|
||||||
|
agent_context=agent_context,
|
||||||
|
channel=getattr(tool, "_channel", None),
|
||||||
|
source=getattr(tool, "_source", None),
|
||||||
|
)
|
||||||
|
return ToolPolicyContext(
|
||||||
|
session_id="subagent",
|
||||||
|
user_id="subagent",
|
||||||
|
origin=ToolOrigin.SUBAGENT,
|
||||||
|
principal_type=PrincipalType.SUBAGENT,
|
||||||
|
auth_source=AuthSource.INTERNAL,
|
||||||
|
agent_context={},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
class _SubAgentProfile:
|
class _SubAgentProfile:
|
||||||
"""子代理运行时定义。"""
|
"""子代理运行时定义。"""
|
||||||
@@ -378,12 +417,14 @@ class _SubAgentAgentProvider:
|
|||||||
profiles: tuple[_SubAgentProfile, ...],
|
profiles: tuple[_SubAgentProfile, ...],
|
||||||
tools: list[BaseTool],
|
tools: list[BaseTool],
|
||||||
server_tools: Optional[list[dict[str, Any]]] = None,
|
server_tools: Optional[list[dict[str, Any]]] = None,
|
||||||
|
policy_context: Optional[ToolPolicyContext] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""初始化子代理执行器。"""
|
"""初始化子代理执行器。"""
|
||||||
self._model = model
|
self._model = model
|
||||||
self._profiles = {profile.name: profile for profile in profiles}
|
self._profiles = {profile.name: profile for profile in profiles}
|
||||||
self._tools = tools
|
self._tools = tools
|
||||||
self._server_tools = server_tools or []
|
self._server_tools = server_tools or []
|
||||||
|
self._policy_context = policy_context or _default_subagent_policy_context(tools)
|
||||||
self._agents = {}
|
self._agents = {}
|
||||||
self._default_agent_name = "general-purpose"
|
self._default_agent_name = "general-purpose"
|
||||||
|
|
||||||
@@ -409,6 +450,7 @@ class _SubAgentAgentProvider:
|
|||||||
tools=[*subagent_tools, *self._server_tools],
|
tools=[*subagent_tools, *self._server_tools],
|
||||||
system_prompt=profile.prompt,
|
system_prompt=profile.prompt,
|
||||||
name=profile.name,
|
name=profile.name,
|
||||||
|
middleware=[AgentPolicyMiddleware(context=self._policy_context)],
|
||||||
)
|
)
|
||||||
self._agents[profile.name] = agent
|
self._agents[profile.name] = agent
|
||||||
return profile.name, agent
|
return profile.name, agent
|
||||||
@@ -444,7 +486,7 @@ class _SubAgentAgentProvider:
|
|||||||
except Exception as err:
|
except Exception as err:
|
||||||
logger.error(
|
logger.error(
|
||||||
f"子代理调用失败: subagent_type={agent_name}, "
|
f"子代理调用失败: subagent_type={agent_name}, "
|
||||||
f"task_id={log_task_id}, error={err}"
|
f"task_id={log_task_id}, error={summarize_error(err)}"
|
||||||
)
|
)
|
||||||
raise
|
raise
|
||||||
final_text = _extract_final_text(result)
|
final_text = _extract_final_text(result)
|
||||||
@@ -468,6 +510,7 @@ class MoviePilotSubAgentMiddleware(AgentMiddleware):
|
|||||||
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,
|
stream_handler: Any = None,
|
||||||
|
policy_context: Optional[ToolPolicyContext] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""初始化同步子代理中间件。"""
|
"""初始化同步子代理中间件。"""
|
||||||
self.system_prompt = system_prompt
|
self.system_prompt = system_prompt
|
||||||
@@ -477,6 +520,7 @@ class MoviePilotSubAgentMiddleware(AgentMiddleware):
|
|||||||
profiles=profiles,
|
profiles=profiles,
|
||||||
tools=tools,
|
tools=tools,
|
||||||
server_tools=server_tools,
|
server_tools=server_tools,
|
||||||
|
policy_context=policy_context,
|
||||||
)
|
)
|
||||||
self.tools = [
|
self.tools = [
|
||||||
StructuredTool.from_function(
|
StructuredTool.from_function(
|
||||||
@@ -527,9 +571,12 @@ class MoviePilotSubAgentMiddleware(AgentMiddleware):
|
|||||||
return await handler(request)
|
return await handler(request)
|
||||||
|
|
||||||
tool_args = _extract_tool_call_args(request)
|
tool_args = _extract_tool_call_args(request)
|
||||||
|
logged_args = sanitize_for_host(tool_args)
|
||||||
|
if not isinstance(logged_args, dict):
|
||||||
|
logged_args = {}
|
||||||
logger.info(
|
logger.info(
|
||||||
f"开始执行子代理工具: tool_name={tool_name}, "
|
f"开始执行子代理工具: tool_name={tool_name}, "
|
||||||
f"subagent_type={tool_args.get('subagent_type') or '-'}"
|
f"subagent_type={logged_args.get('subagent_type') or '-'}"
|
||||||
)
|
)
|
||||||
_record_subagent_tool_call(
|
_record_subagent_tool_call(
|
||||||
stream_handler=self.stream_handler,
|
stream_handler=self.stream_handler,
|
||||||
@@ -539,7 +586,10 @@ class MoviePilotSubAgentMiddleware(AgentMiddleware):
|
|||||||
try:
|
try:
|
||||||
result = await handler(request)
|
result = await handler(request)
|
||||||
except Exception as err:
|
except Exception as err:
|
||||||
logger.error(f"子代理工具执行失败: tool_name={tool_name}, error={err}")
|
logger.error(
|
||||||
|
f"子代理工具执行失败: tool_name={tool_name}, "
|
||||||
|
f"error={summarize_error(err)}"
|
||||||
|
)
|
||||||
raise
|
raise
|
||||||
logger.info(f"子代理工具执行完成: tool_name={tool_name}")
|
logger.info(f"子代理工具执行完成: tool_name={tool_name}")
|
||||||
return result
|
return result
|
||||||
@@ -557,6 +607,7 @@ class SubAgentTaskControlMiddleware(AgentMiddleware):
|
|||||||
server_tools: Optional[list[dict[str, Any]]] = None,
|
server_tools: Optional[list[dict[str, Any]]] = None,
|
||||||
task_description: str = SUBAGENT_CONTROL_DESCRIPTION,
|
task_description: str = SUBAGENT_CONTROL_DESCRIPTION,
|
||||||
stream_handler: Any = None,
|
stream_handler: Any = None,
|
||||||
|
policy_context: Optional[ToolPolicyContext] = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""初始化异步子代理调度中间件。"""
|
"""初始化异步子代理调度中间件。"""
|
||||||
self.stream_handler = stream_handler
|
self.stream_handler = stream_handler
|
||||||
@@ -565,6 +616,7 @@ class SubAgentTaskControlMiddleware(AgentMiddleware):
|
|||||||
profiles=profiles,
|
profiles=profiles,
|
||||||
tools=tools,
|
tools=tools,
|
||||||
server_tools=server_tools,
|
server_tools=server_tools,
|
||||||
|
policy_context=policy_context,
|
||||||
)
|
)
|
||||||
self._semaphore = asyncio.Semaphore(SUBAGENT_MAX_CONCURRENT_TASKS)
|
self._semaphore = asyncio.Semaphore(SUBAGENT_MAX_CONCURRENT_TASKS)
|
||||||
self._tasks: dict[str, _SubAgentRuntimeTask] = {}
|
self._tasks: dict[str, _SubAgentRuntimeTask] = {}
|
||||||
@@ -628,7 +680,7 @@ class SubAgentTaskControlMiddleware(AgentMiddleware):
|
|||||||
|
|
||||||
error = record.task.exception()
|
error = record.task.exception()
|
||||||
if error:
|
if error:
|
||||||
payload["error"] = str(error)
|
payload["error"] = summarize_error(error)
|
||||||
return payload
|
return payload
|
||||||
|
|
||||||
result, result_truncated = _clip_text(
|
result, result_truncated = _clip_text(
|
||||||
@@ -733,7 +785,10 @@ class SubAgentTaskControlMiddleware(AgentMiddleware):
|
|||||||
)
|
)
|
||||||
raise
|
raise
|
||||||
except Exception as err:
|
except Exception as err:
|
||||||
logger.error(f"子代理任务执行失败: task_id={record.task_id}, error={err}")
|
logger.error(
|
||||||
|
f"子代理任务执行失败: task_id={record.task_id}, "
|
||||||
|
f"error={summarize_error(err)}"
|
||||||
|
)
|
||||||
raise
|
raise
|
||||||
|
|
||||||
def _mark_task_finished(self, task_id: str, task: asyncio.Task) -> None:
|
def _mark_task_finished(self, task_id: str, task: asyncio.Task) -> None:
|
||||||
@@ -901,7 +956,10 @@ class SubAgentTaskControlMiddleware(AgentMiddleware):
|
|||||||
)
|
)
|
||||||
raise
|
raise
|
||||||
except Exception as err:
|
except Exception as err:
|
||||||
logger.error(f"管道子代理任务执行失败: task_id={record.task_id}, error={err}")
|
logger.error(
|
||||||
|
f"管道子代理任务执行失败: task_id={record.task_id}, "
|
||||||
|
f"error={summarize_error(err)}"
|
||||||
|
)
|
||||||
raise
|
raise
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
@@ -975,8 +1033,13 @@ class SubAgentTaskControlMiddleware(AgentMiddleware):
|
|||||||
)
|
)
|
||||||
return records, error
|
return records, error
|
||||||
except Exception as err:
|
except Exception as err:
|
||||||
error = f"第 {step_index} 个管道子代理任务执行失败: {err}"
|
error = (
|
||||||
logger.info(f"{error} task_id={record.task_id}")
|
f"第 {step_index} 个管道子代理任务执行失败: "
|
||||||
|
f"{summarize_error(err)}"
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
f"{error} task_id={record.task_id}"
|
||||||
|
)
|
||||||
return records, error
|
return records, error
|
||||||
|
|
||||||
previous_results.append((record, result))
|
previous_results.append((record, result))
|
||||||
@@ -1003,7 +1066,10 @@ class SubAgentTaskControlMiddleware(AgentMiddleware):
|
|||||||
tasks=tasks,
|
tasks=tasks,
|
||||||
)
|
)
|
||||||
if error:
|
if error:
|
||||||
logger.info(f"子代理管控操作未启动任务: action={action}, error={error}")
|
logger.info(
|
||||||
|
f"子代理管控操作未启动任务: action={action}, "
|
||||||
|
f"error={sanitize_for_host(error)}"
|
||||||
|
)
|
||||||
return self._json_response({"success": False, "error": error})
|
return self._json_response({"success": False, "error": error})
|
||||||
|
|
||||||
logger.info(f"准备启动子代理任务: action={action}, tasks={len(specs)}")
|
logger.info(f"准备启动子代理任务: action={action}, tasks={len(specs)}")
|
||||||
@@ -1095,10 +1161,13 @@ class SubAgentTaskControlMiddleware(AgentMiddleware):
|
|||||||
return await handler(request)
|
return await handler(request)
|
||||||
|
|
||||||
tool_args = _extract_tool_call_args(request)
|
tool_args = _extract_tool_call_args(request)
|
||||||
|
logged_args = sanitize_for_host(tool_args)
|
||||||
|
if not isinstance(logged_args, dict):
|
||||||
|
logged_args = {}
|
||||||
logger.info(
|
logger.info(
|
||||||
f"开始执行子代理工具: tool_name={tool_name}, "
|
f"开始执行子代理工具: tool_name={tool_name}, "
|
||||||
f"action={tool_args.get('action') or '-'}, "
|
f"action={logged_args.get('action') or '-'}, "
|
||||||
f"subagent_type={tool_args.get('subagent_type') or '-'}"
|
f"subagent_type={logged_args.get('subagent_type') or '-'}"
|
||||||
)
|
)
|
||||||
_record_subagent_tool_call(
|
_record_subagent_tool_call(
|
||||||
stream_handler=self.stream_handler,
|
stream_handler=self.stream_handler,
|
||||||
@@ -1108,7 +1177,10 @@ class SubAgentTaskControlMiddleware(AgentMiddleware):
|
|||||||
try:
|
try:
|
||||||
result = await handler(request)
|
result = await handler(request)
|
||||||
except Exception as err:
|
except Exception as err:
|
||||||
logger.error(f"子代理工具执行失败: tool_name={tool_name}, error={err}")
|
logger.error(
|
||||||
|
f"子代理工具执行失败: tool_name={tool_name}, "
|
||||||
|
f"error={summarize_error(err)}"
|
||||||
|
)
|
||||||
raise
|
raise
|
||||||
logger.info(f"子代理工具执行完成: tool_name={tool_name}")
|
logger.info(f"子代理工具执行完成: tool_name={tool_name}")
|
||||||
return result
|
return result
|
||||||
@@ -1120,6 +1192,7 @@ def create_subagent_middlewares(
|
|||||||
tools: list[BaseTool],
|
tools: list[BaseTool],
|
||||||
server_tools: Optional[list[dict[str, Any]]] = None,
|
server_tools: Optional[list[dict[str, Any]]] = None,
|
||||||
stream_handler: Any = None,
|
stream_handler: Any = None,
|
||||||
|
policy_context: Optional[ToolPolicyContext] = None,
|
||||||
) -> tuple[list[AgentMiddleware], list[BaseTool]]:
|
) -> tuple[list[AgentMiddleware], list[BaseTool]]:
|
||||||
"""创建子代理中间件列表和任务工具列表。"""
|
"""创建子代理中间件列表和任务工具列表。"""
|
||||||
runtime_signature = agent_runtime_manager.current_signature()
|
runtime_signature = agent_runtime_manager.current_signature()
|
||||||
@@ -1130,6 +1203,7 @@ def create_subagent_middlewares(
|
|||||||
tools=tools,
|
tools=tools,
|
||||||
server_tools=server_tools or [],
|
server_tools=server_tools or [],
|
||||||
stream_handler=stream_handler,
|
stream_handler=stream_handler,
|
||||||
|
policy_context=policy_context,
|
||||||
)
|
)
|
||||||
control_middleware = SubAgentTaskControlMiddleware(
|
control_middleware = SubAgentTaskControlMiddleware(
|
||||||
model=model,
|
model=model,
|
||||||
@@ -1137,6 +1211,7 @@ def create_subagent_middlewares(
|
|||||||
tools=tools,
|
tools=tools,
|
||||||
server_tools=server_tools or [],
|
server_tools=server_tools or [],
|
||||||
stream_handler=stream_handler,
|
stream_handler=stream_handler,
|
||||||
|
policy_context=policy_context,
|
||||||
)
|
)
|
||||||
|
|
||||||
task_tools = [
|
task_tools = [
|
||||||
|
|||||||
@@ -0,0 +1,66 @@
|
|||||||
|
"""MoviePilot Agent 宿主策略公共内部入口。"""
|
||||||
|
|
||||||
|
from app.agent.policy.contracts import (
|
||||||
|
ActionEffect,
|
||||||
|
ActionPolicy,
|
||||||
|
AuthSource,
|
||||||
|
ConfirmationMode,
|
||||||
|
ExecutionOutcome,
|
||||||
|
ExecutionReceipt,
|
||||||
|
MigrationState,
|
||||||
|
PolicyDecision,
|
||||||
|
PolicyObservation,
|
||||||
|
PolicyPrincipal,
|
||||||
|
PrincipalRole,
|
||||||
|
PrincipalType,
|
||||||
|
RecoveryMode,
|
||||||
|
ResultSensitivity,
|
||||||
|
ToolInvocation,
|
||||||
|
ToolOrigin,
|
||||||
|
ToolPolicyContext,
|
||||||
|
)
|
||||||
|
from app.agent.policy.orchestrator import (
|
||||||
|
DEFAULT_TOOL_POLICY_ORCHESTRATOR,
|
||||||
|
AgentToolPolicyOrchestrator,
|
||||||
|
call_policy_hook,
|
||||||
|
)
|
||||||
|
from app.agent.policy.registry import DEFAULT_TOOL_POLICY_REGISTRY, ToolPolicyRegistry
|
||||||
|
from app.agent.policy.sanitizer import (
|
||||||
|
REDACTED_VALUE,
|
||||||
|
sanitize_for_host,
|
||||||
|
stable_type_name,
|
||||||
|
summarize_error,
|
||||||
|
summarize_input,
|
||||||
|
summarize_result,
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"ActionEffect",
|
||||||
|
"ActionPolicy",
|
||||||
|
"AgentToolPolicyOrchestrator",
|
||||||
|
"AuthSource",
|
||||||
|
"ConfirmationMode",
|
||||||
|
"DEFAULT_TOOL_POLICY_ORCHESTRATOR",
|
||||||
|
"DEFAULT_TOOL_POLICY_REGISTRY",
|
||||||
|
"ExecutionOutcome",
|
||||||
|
"ExecutionReceipt",
|
||||||
|
"MigrationState",
|
||||||
|
"PolicyDecision",
|
||||||
|
"PolicyObservation",
|
||||||
|
"PolicyPrincipal",
|
||||||
|
"PrincipalRole",
|
||||||
|
"PrincipalType",
|
||||||
|
"REDACTED_VALUE",
|
||||||
|
"RecoveryMode",
|
||||||
|
"ResultSensitivity",
|
||||||
|
"ToolInvocation",
|
||||||
|
"ToolOrigin",
|
||||||
|
"ToolPolicyContext",
|
||||||
|
"ToolPolicyRegistry",
|
||||||
|
"call_policy_hook",
|
||||||
|
"sanitize_for_host",
|
||||||
|
"stable_type_name",
|
||||||
|
"summarize_error",
|
||||||
|
"summarize_input",
|
||||||
|
"summarize_result",
|
||||||
|
]
|
||||||
@@ -0,0 +1,247 @@
|
|||||||
|
"""MoviePilot Agent 宿主策略的内部契约。"""
|
||||||
|
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from enum import Enum
|
||||||
|
from typing import Any, Mapping, MutableMapping, Optional
|
||||||
|
|
||||||
|
|
||||||
|
class ToolOrigin(str, Enum):
|
||||||
|
"""工具调用的宿主可信入口。"""
|
||||||
|
|
||||||
|
AGENT_INTERACTIVE = "agent_interactive"
|
||||||
|
AGENT_API = "agent_api"
|
||||||
|
OPERATOR_DIRECT = "operator_direct"
|
||||||
|
BACKGROUND = "background"
|
||||||
|
SUBAGENT = "subagent"
|
||||||
|
|
||||||
|
|
||||||
|
class PrincipalType(str, Enum):
|
||||||
|
"""调用主体类型,用于区分人、管理员集成和内部运行时。"""
|
||||||
|
|
||||||
|
HUMAN = "human"
|
||||||
|
SYSTEM_ADMIN_INTEGRATION = "system_admin_integration"
|
||||||
|
SCOPED_AGENT = "scoped_agent"
|
||||||
|
BACKGROUND = "background"
|
||||||
|
SUBAGENT = "subagent"
|
||||||
|
|
||||||
|
|
||||||
|
class AuthSource(str, Enum):
|
||||||
|
"""主体身份的宿主认证来源。"""
|
||||||
|
|
||||||
|
CHANNEL = "channel"
|
||||||
|
WEB_SESSION = "web_session"
|
||||||
|
API_TOKEN = "api_token"
|
||||||
|
INTERNAL = "internal"
|
||||||
|
AGENT_TOKEN = "agent_token"
|
||||||
|
|
||||||
|
|
||||||
|
class PrincipalRole(str, Enum):
|
||||||
|
"""策略授权使用的角色层级。"""
|
||||||
|
|
||||||
|
USER = "user"
|
||||||
|
CHANNEL_ADMIN = "channel_admin"
|
||||||
|
SYSTEM_ADMIN = "system_admin"
|
||||||
|
SYSTEM_INTERNAL = "system_internal"
|
||||||
|
|
||||||
|
|
||||||
|
class ActionEffect(str, Enum):
|
||||||
|
"""工具调用的实际副作用类别。"""
|
||||||
|
|
||||||
|
SAFE_READ = "safe_read"
|
||||||
|
SENSITIVE_READ = "sensitive_read"
|
||||||
|
REVERSIBLE_WRITE = "reversible_write"
|
||||||
|
DESTRUCTIVE_WRITE = "destructive_write"
|
||||||
|
EXTERNAL_SIDE_EFFECT = "external_side_effect"
|
||||||
|
ARBITRARY_EXECUTION = "arbitrary_execution"
|
||||||
|
UNKNOWN = "unknown"
|
||||||
|
|
||||||
|
|
||||||
|
class ConfirmationMode(str, Enum):
|
||||||
|
"""动作在完成授权后所需的确认方式。"""
|
||||||
|
|
||||||
|
NONE = "none"
|
||||||
|
REQUIRED = "required"
|
||||||
|
UNSUPPORTED = "unsupported"
|
||||||
|
|
||||||
|
|
||||||
|
class RecoveryMode(str, Enum):
|
||||||
|
"""动作可提供的执行恢复保证。"""
|
||||||
|
|
||||||
|
NONE = "none"
|
||||||
|
TRANSACTION = "transaction"
|
||||||
|
BEFORE_STATE = "before_state"
|
||||||
|
RECOVERABLE_DELETE = "recoverable_delete"
|
||||||
|
IDEMPOTENT = "idempotent"
|
||||||
|
RECONCILE = "reconcile"
|
||||||
|
MANUAL_ONLY = "manual_only"
|
||||||
|
|
||||||
|
|
||||||
|
class ResultSensitivity(str, Enum):
|
||||||
|
"""工具结果进入模型、记忆和日志时的敏感等级。"""
|
||||||
|
|
||||||
|
NORMAL = "normal"
|
||||||
|
PRIVATE = "private"
|
||||||
|
SECRET = "secret"
|
||||||
|
UNKNOWN = "unknown"
|
||||||
|
|
||||||
|
|
||||||
|
class MigrationState(str, Enum):
|
||||||
|
"""工具策略从兼容观测迁移到宿主执行的状态。"""
|
||||||
|
|
||||||
|
ENFORCED = "enforced"
|
||||||
|
LEGACY_SHADOW = "legacy_shadow"
|
||||||
|
|
||||||
|
|
||||||
|
class ExecutionOutcome(str, Enum):
|
||||||
|
"""P1-G1 handler 生命周期终态;成功不代表业务授权或副作用已完成。"""
|
||||||
|
|
||||||
|
SUCCEEDED = "succeeded"
|
||||||
|
FAILED = "failed"
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class PolicyPrincipal:
|
||||||
|
"""由可信入口建立、不可由工具参数覆盖的调用主体。"""
|
||||||
|
|
||||||
|
principal_id: str
|
||||||
|
principal_type: PrincipalType
|
||||||
|
auth_source: AuthSource
|
||||||
|
role: PrincipalRole
|
||||||
|
scopes: tuple[str, ...] = ()
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class ToolInvocation:
|
||||||
|
"""一次进入宿主策略层的规范化工具调用。"""
|
||||||
|
|
||||||
|
invocation_id: str
|
||||||
|
tool_name: str
|
||||||
|
arguments: Mapping[str, Any]
|
||||||
|
principal: PolicyPrincipal
|
||||||
|
session_id: str
|
||||||
|
origin: ToolOrigin
|
||||||
|
channel: Optional[str] = None
|
||||||
|
source: Optional[str] = None
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class ActionPolicy:
|
||||||
|
"""参数级动作策略及其兼容迁移状态。"""
|
||||||
|
|
||||||
|
effect: ActionEffect
|
||||||
|
required_role: PrincipalRole
|
||||||
|
confirmation: ConfirmationMode
|
||||||
|
recovery: RecoveryMode
|
||||||
|
result_sensitivity: ResultSensitivity
|
||||||
|
migration_state: MigrationState
|
||||||
|
policy_version: str = "p1-g1-v1"
|
||||||
|
interactive_allowed: bool = True
|
||||||
|
machine_allowed: bool = True
|
||||||
|
background_allowed: bool = True
|
||||||
|
subagent_allowed: bool = True
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class PolicyDecision:
|
||||||
|
"""宿主策略层决定;shadow allow 仅表示新策略不拦截,旧门禁仍是授权事实源。"""
|
||||||
|
|
||||||
|
allowed: bool
|
||||||
|
confirmation_required: bool
|
||||||
|
shadow: bool
|
||||||
|
reason_code: str
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class PolicyObservation:
|
||||||
|
"""调用开始时生成、供完成或失败回执复用的观测对象。"""
|
||||||
|
|
||||||
|
invocation: ToolInvocation
|
||||||
|
policy: ActionPolicy
|
||||||
|
decision: PolicyDecision
|
||||||
|
input_summary: str
|
||||||
|
started_at: float
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class ExecutionReceipt:
|
||||||
|
"""P1-G1 的非持久化脱敏回执 envelope。"""
|
||||||
|
|
||||||
|
invocation_id: str
|
||||||
|
tool_name: str
|
||||||
|
origin: ToolOrigin
|
||||||
|
decision: PolicyDecision
|
||||||
|
outcome: ExecutionOutcome
|
||||||
|
input_summary: str
|
||||||
|
result_summary: Optional[str] = None
|
||||||
|
error_summary: Optional[str] = None
|
||||||
|
duration_ms: int = 0
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True)
|
||||||
|
class ToolPolicyContext:
|
||||||
|
"""宿主入口上下文;管理员状态引用会随缓存图的每轮执行刷新。"""
|
||||||
|
|
||||||
|
session_id: str
|
||||||
|
user_id: str
|
||||||
|
origin: ToolOrigin
|
||||||
|
principal_type: PrincipalType
|
||||||
|
auth_source: AuthSource
|
||||||
|
agent_context: MutableMapping[str, Any] = field(repr=False, compare=False)
|
||||||
|
channel: Optional[str] = None
|
||||||
|
source: Optional[str] = None
|
||||||
|
|
||||||
|
@property
|
||||||
|
def principal(self) -> PolicyPrincipal:
|
||||||
|
"""根据当前宿主上下文生成本次调用主体。"""
|
||||||
|
if self.principal_type in {PrincipalType.BACKGROUND, PrincipalType.SUBAGENT}:
|
||||||
|
default_role = PrincipalRole.SYSTEM_INTERNAL
|
||||||
|
else:
|
||||||
|
default_role = PrincipalRole.USER
|
||||||
|
role = (
|
||||||
|
PrincipalRole.SYSTEM_ADMIN
|
||||||
|
if bool(self.agent_context.get("is_admin"))
|
||||||
|
else default_role
|
||||||
|
)
|
||||||
|
raw_scopes = self.agent_context.get("policy_scopes") or ()
|
||||||
|
scopes = tuple(str(scope) for scope in raw_scopes if scope)
|
||||||
|
return PolicyPrincipal(
|
||||||
|
principal_id=str(self.user_id or self.principal_type.value),
|
||||||
|
principal_type=self.principal_type,
|
||||||
|
auth_source=self.auth_source,
|
||||||
|
role=role,
|
||||||
|
scopes=scopes,
|
||||||
|
)
|
||||||
|
|
||||||
|
def for_subagent(self) -> "ToolPolicyContext":
|
||||||
|
"""保留用户与会话归属,并切换为子代理可信来源。"""
|
||||||
|
return ToolPolicyContext(
|
||||||
|
session_id=self.session_id,
|
||||||
|
user_id=self.user_id,
|
||||||
|
origin=ToolOrigin.SUBAGENT,
|
||||||
|
principal_type=PrincipalType.SUBAGENT,
|
||||||
|
auth_source=AuthSource.INTERNAL,
|
||||||
|
agent_context=self.agent_context,
|
||||||
|
channel=self.channel,
|
||||||
|
source=self.source,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"ActionEffect",
|
||||||
|
"ActionPolicy",
|
||||||
|
"AuthSource",
|
||||||
|
"ConfirmationMode",
|
||||||
|
"ExecutionOutcome",
|
||||||
|
"ExecutionReceipt",
|
||||||
|
"MigrationState",
|
||||||
|
"PolicyDecision",
|
||||||
|
"PolicyObservation",
|
||||||
|
"PolicyPrincipal",
|
||||||
|
"PrincipalRole",
|
||||||
|
"PrincipalType",
|
||||||
|
"RecoveryMode",
|
||||||
|
"ResultSensitivity",
|
||||||
|
"ToolInvocation",
|
||||||
|
"ToolOrigin",
|
||||||
|
"ToolPolicyContext",
|
||||||
|
]
|
||||||
@@ -0,0 +1,191 @@
|
|||||||
|
"""Agent 工具策略观测、脱敏回执与共享执行边界。"""
|
||||||
|
|
||||||
|
import time
|
||||||
|
import uuid
|
||||||
|
from collections.abc import Callable
|
||||||
|
from typing import Any, Mapping, Optional, TypeVar
|
||||||
|
|
||||||
|
from langchain_core.messages import ToolMessage
|
||||||
|
from pydantic import ValidationError
|
||||||
|
|
||||||
|
from app.agent.policy.contracts import (
|
||||||
|
ExecutionOutcome,
|
||||||
|
ExecutionReceipt,
|
||||||
|
MigrationState,
|
||||||
|
PolicyDecision,
|
||||||
|
PolicyObservation,
|
||||||
|
ToolInvocation,
|
||||||
|
ToolPolicyContext,
|
||||||
|
)
|
||||||
|
from app.agent.policy.registry import DEFAULT_TOOL_POLICY_REGISTRY, ToolPolicyRegistry
|
||||||
|
from app.agent.policy.sanitizer import (
|
||||||
|
stable_type_name,
|
||||||
|
summarize_error,
|
||||||
|
summarize_input,
|
||||||
|
summarize_result,
|
||||||
|
)
|
||||||
|
from app.log import logger
|
||||||
|
|
||||||
|
|
||||||
|
_HookResult = TypeVar("_HookResult")
|
||||||
|
|
||||||
|
|
||||||
|
def call_policy_hook(
|
||||||
|
phase: str,
|
||||||
|
hook: Callable[..., _HookResult],
|
||||||
|
*args: Any,
|
||||||
|
**kwargs: Any,
|
||||||
|
) -> Optional[_HookResult]:
|
||||||
|
"""以 fail-open 方式调用 P1-G1 观测 hook,故障只记录稳定类型。"""
|
||||||
|
try:
|
||||||
|
return hook(*args, **kwargs)
|
||||||
|
except Exception as error:
|
||||||
|
try:
|
||||||
|
logger.warning(
|
||||||
|
f"Agent工具策略观测失败: phase={phase}, "
|
||||||
|
f"error_type={stable_type_name(error)}"
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_policy_arguments(tool: Any, arguments: Mapping[str, Any]) -> dict[str, Any]:
|
||||||
|
"""为策略生成 Pydantic 规范化副本,不改变真实执行参数。"""
|
||||||
|
raw_arguments = dict(arguments or {})
|
||||||
|
args_schema = getattr(tool, "args_schema", None)
|
||||||
|
if not args_schema:
|
||||||
|
return raw_arguments
|
||||||
|
try:
|
||||||
|
validated = args_schema.model_validate(raw_arguments)
|
||||||
|
return validated.model_dump(mode="json")
|
||||||
|
except (AttributeError, TypeError, ValueError, ValidationError):
|
||||||
|
# 实际 handler 仍负责既有参数错误语义;策略观测按原始值保守处理。
|
||||||
|
return raw_arguments
|
||||||
|
|
||||||
|
|
||||||
|
def _result_payload(result: Any) -> Any:
|
||||||
|
"""从 LangChain 工具消息中提取模型可见结果供脱敏摘要使用。"""
|
||||||
|
if isinstance(result, ToolMessage):
|
||||||
|
return result.content
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
class AgentToolPolicyOrchestrator:
|
||||||
|
"""让 Agent middleware 与 direct manager 复用同一策略生命周期。"""
|
||||||
|
|
||||||
|
def __init__(self, registry: ToolPolicyRegistry = DEFAULT_TOOL_POLICY_REGISTRY) -> None:
|
||||||
|
"""绑定固定工具迁移注册表。"""
|
||||||
|
self.registry = registry
|
||||||
|
|
||||||
|
def start(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
context: ToolPolicyContext,
|
||||||
|
tool: Any,
|
||||||
|
arguments: Mapping[str, Any],
|
||||||
|
invocation_id: Optional[str] = None,
|
||||||
|
) -> PolicyObservation:
|
||||||
|
"""解析调用策略,并创建不影响现有 allow 行为的观测对象。"""
|
||||||
|
tool_name = str(getattr(tool, "name", None) or "unknown_tool")
|
||||||
|
normalized_arguments = _normalize_policy_arguments(tool, arguments)
|
||||||
|
policy = self.registry.resolve(
|
||||||
|
tool_name=tool_name,
|
||||||
|
arguments=normalized_arguments,
|
||||||
|
requires_admin=bool(getattr(tool, "_require_admin", False)),
|
||||||
|
)
|
||||||
|
if policy.migration_state is MigrationState.LEGACY_SHADOW:
|
||||||
|
decision = PolicyDecision(
|
||||||
|
allowed=True,
|
||||||
|
confirmation_required=False,
|
||||||
|
shadow=True,
|
||||||
|
reason_code="legacy_shadow_allow",
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
decision = PolicyDecision(
|
||||||
|
allowed=True,
|
||||||
|
confirmation_required=False,
|
||||||
|
shadow=False,
|
||||||
|
reason_code="safe_read_allow",
|
||||||
|
)
|
||||||
|
invocation = ToolInvocation(
|
||||||
|
invocation_id=invocation_id or uuid.uuid4().hex,
|
||||||
|
tool_name=tool_name,
|
||||||
|
arguments=normalized_arguments,
|
||||||
|
principal=context.principal,
|
||||||
|
session_id=context.session_id,
|
||||||
|
origin=context.origin,
|
||||||
|
channel=context.channel,
|
||||||
|
source=context.source,
|
||||||
|
)
|
||||||
|
input_summary = summarize_input(normalized_arguments)
|
||||||
|
observation = PolicyObservation(
|
||||||
|
invocation=invocation,
|
||||||
|
policy=policy,
|
||||||
|
decision=decision,
|
||||||
|
input_summary=input_summary,
|
||||||
|
started_at=time.monotonic(),
|
||||||
|
)
|
||||||
|
logger.debug(
|
||||||
|
f"Agent工具策略: tool={tool_name}, origin={context.origin.value}, "
|
||||||
|
f"decision={decision.reason_code}, input={input_summary}"
|
||||||
|
)
|
||||||
|
return observation
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def finish(observation: PolicyObservation, result: Any) -> ExecutionReceipt:
|
||||||
|
"""生成成功回执 envelope,并只记录脱敏结果摘要。"""
|
||||||
|
result_summary = summarize_result(_result_payload(result))
|
||||||
|
receipt = ExecutionReceipt(
|
||||||
|
invocation_id=observation.invocation.invocation_id,
|
||||||
|
tool_name=observation.invocation.tool_name,
|
||||||
|
origin=observation.invocation.origin,
|
||||||
|
decision=observation.decision,
|
||||||
|
outcome=ExecutionOutcome.SUCCEEDED,
|
||||||
|
input_summary=observation.input_summary,
|
||||||
|
result_summary=result_summary,
|
||||||
|
duration_ms=max(
|
||||||
|
0,
|
||||||
|
int((time.monotonic() - observation.started_at) * 1000),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
f"Agent工具执行完成: tool={receipt.tool_name}, "
|
||||||
|
f"origin={receipt.origin.value}, shadow={receipt.decision.shadow}, "
|
||||||
|
f"duration_ms={receipt.duration_ms}, result={result_summary}"
|
||||||
|
)
|
||||||
|
return receipt
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def fail(observation: PolicyObservation, error: BaseException) -> ExecutionReceipt:
|
||||||
|
"""生成失败回执 envelope,不把异常中的凭据写入日志。"""
|
||||||
|
error_summary = summarize_error(error)
|
||||||
|
receipt = ExecutionReceipt(
|
||||||
|
invocation_id=observation.invocation.invocation_id,
|
||||||
|
tool_name=observation.invocation.tool_name,
|
||||||
|
origin=observation.invocation.origin,
|
||||||
|
decision=observation.decision,
|
||||||
|
outcome=ExecutionOutcome.FAILED,
|
||||||
|
input_summary=observation.input_summary,
|
||||||
|
error_summary=error_summary,
|
||||||
|
duration_ms=max(
|
||||||
|
0,
|
||||||
|
int((time.monotonic() - observation.started_at) * 1000),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
logger.error(
|
||||||
|
f"Agent工具执行失败: tool={receipt.tool_name}, "
|
||||||
|
f"origin={receipt.origin.value}, shadow={receipt.decision.shadow}, "
|
||||||
|
f"duration_ms={receipt.duration_ms}, error={error_summary}"
|
||||||
|
)
|
||||||
|
return receipt
|
||||||
|
|
||||||
|
|
||||||
|
DEFAULT_TOOL_POLICY_ORCHESTRATOR = AgentToolPolicyOrchestrator()
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"AgentToolPolicyOrchestrator",
|
||||||
|
"DEFAULT_TOOL_POLICY_ORCHESTRATOR",
|
||||||
|
"call_policy_hook",
|
||||||
|
]
|
||||||
@@ -0,0 +1,181 @@
|
|||||||
|
"""固定工具迁移注册表与参数级策略解析。"""
|
||||||
|
|
||||||
|
from typing import Any, Mapping
|
||||||
|
|
||||||
|
from app.agent.policy.contracts import (
|
||||||
|
ActionEffect,
|
||||||
|
ActionPolicy,
|
||||||
|
ConfirmationMode,
|
||||||
|
MigrationState,
|
||||||
|
PrincipalRole,
|
||||||
|
RecoveryMode,
|
||||||
|
ResultSensitivity,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# 这些读取已具备清晰的无副作用语义,用于证明新宿主边界不会改变正常结果。
|
||||||
|
SAFE_READ_TOOL_NAMES = frozenset(
|
||||||
|
{
|
||||||
|
"list_slash_commands",
|
||||||
|
"query_installed_plugins",
|
||||||
|
"query_personas",
|
||||||
|
"query_schedulers",
|
||||||
|
"query_workflows",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# 其余固定工具先显式处于兼容观测状态,待领域叶子 Goal 逐个迁移。
|
||||||
|
LEGACY_SHADOW_TOOL_NAMES = frozenset(
|
||||||
|
{
|
||||||
|
"add_custom_filter_rule",
|
||||||
|
"add_download_tasks",
|
||||||
|
"add_rule_group",
|
||||||
|
"add_subscribe",
|
||||||
|
"ask_user_choice",
|
||||||
|
"browse_webpage",
|
||||||
|
"create_agent_task",
|
||||||
|
"delete_agent_task",
|
||||||
|
"delete_custom_filter_rule",
|
||||||
|
"delete_download_history",
|
||||||
|
"delete_download_tasks",
|
||||||
|
"delete_rule_group",
|
||||||
|
"delete_subscribe",
|
||||||
|
"delete_transfer_history",
|
||||||
|
"edit_file",
|
||||||
|
"execute_command",
|
||||||
|
"get_recommendations",
|
||||||
|
"get_search_results",
|
||||||
|
"install_plugin",
|
||||||
|
"list_directory",
|
||||||
|
"query_agent_tasks",
|
||||||
|
"query_builtin_filter_rules",
|
||||||
|
"query_custom_filter_rules",
|
||||||
|
"query_custom_identifiers",
|
||||||
|
"query_directory_settings",
|
||||||
|
"query_doctor_report",
|
||||||
|
"query_download_tasks",
|
||||||
|
"query_downloaders",
|
||||||
|
"query_episode_schedule",
|
||||||
|
"query_library_exists",
|
||||||
|
"query_library_latest",
|
||||||
|
"query_market_plugins",
|
||||||
|
"query_media_detail",
|
||||||
|
"query_plugin_capabilities",
|
||||||
|
"query_plugin_config",
|
||||||
|
"query_plugin_data",
|
||||||
|
"query_popular_subscribes",
|
||||||
|
"query_rule_groups",
|
||||||
|
"query_site_userdata",
|
||||||
|
"query_sites",
|
||||||
|
"query_subscribe_history",
|
||||||
|
"query_subscribe_shares",
|
||||||
|
"query_subscribes",
|
||||||
|
"query_system_settings",
|
||||||
|
"query_transfer_history",
|
||||||
|
"read_file",
|
||||||
|
"recognize_captcha",
|
||||||
|
"recognize_media",
|
||||||
|
"reload_plugin",
|
||||||
|
"run_agent_task",
|
||||||
|
"run_scheduler",
|
||||||
|
"run_slash_command",
|
||||||
|
"run_workflow",
|
||||||
|
"scrape_metadata",
|
||||||
|
"search_media",
|
||||||
|
"search_person",
|
||||||
|
"search_person_credits",
|
||||||
|
"search_subscribe",
|
||||||
|
"search_torrents",
|
||||||
|
"search_web",
|
||||||
|
"send_local_file",
|
||||||
|
"send_message",
|
||||||
|
"send_voice_message",
|
||||||
|
"switch_persona",
|
||||||
|
"test_site",
|
||||||
|
"transfer_file",
|
||||||
|
"uninstall_plugin",
|
||||||
|
"update_agent_task",
|
||||||
|
"update_custom_filter_rule",
|
||||||
|
"update_custom_identifiers",
|
||||||
|
"update_download_tasks",
|
||||||
|
"update_persona_definition",
|
||||||
|
"update_plugin_config",
|
||||||
|
"update_rule_group",
|
||||||
|
"update_site",
|
||||||
|
"update_site_cookie",
|
||||||
|
"update_subscribe",
|
||||||
|
"update_system_settings",
|
||||||
|
"write_file",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class ToolPolicyRegistry:
|
||||||
|
"""解析固定和动态工具的 P1-G1 迁移策略。"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
safe_read_tool_names: frozenset[str] = SAFE_READ_TOOL_NAMES,
|
||||||
|
legacy_shadow_tool_names: frozenset[str] = LEGACY_SHADOW_TOOL_NAMES,
|
||||||
|
) -> None:
|
||||||
|
"""建立互斥的固定工具迁移表。"""
|
||||||
|
overlap = safe_read_tool_names & legacy_shadow_tool_names
|
||||||
|
if overlap:
|
||||||
|
raise ValueError(f"工具策略迁移表存在重复项: {sorted(overlap)}")
|
||||||
|
self._safe_read_tool_names = safe_read_tool_names
|
||||||
|
self._legacy_shadow_tool_names = legacy_shadow_tool_names
|
||||||
|
|
||||||
|
@property
|
||||||
|
def builtin_tool_names(self) -> set[str]:
|
||||||
|
"""返回注册表覆盖的全部固定工具名。"""
|
||||||
|
return set(self._safe_read_tool_names | self._legacy_shadow_tool_names)
|
||||||
|
|
||||||
|
def resolve(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
tool_name: str,
|
||||||
|
arguments: Mapping[str, Any],
|
||||||
|
requires_admin: bool,
|
||||||
|
) -> ActionPolicy:
|
||||||
|
"""根据工具名和宿主权限声明解析当前迁移策略。"""
|
||||||
|
del arguments # 参数级迁移由后续领域 Goal 逐项加入。
|
||||||
|
required_role = (
|
||||||
|
PrincipalRole.SYSTEM_ADMIN if requires_admin else PrincipalRole.USER
|
||||||
|
)
|
||||||
|
if tool_name in self._safe_read_tool_names:
|
||||||
|
return ActionPolicy(
|
||||||
|
effect=ActionEffect.SAFE_READ,
|
||||||
|
required_role=required_role,
|
||||||
|
confirmation=ConfirmationMode.NONE,
|
||||||
|
recovery=RecoveryMode.NONE,
|
||||||
|
result_sensitivity=ResultSensitivity.NORMAL,
|
||||||
|
# 角色门禁仍可能异步识别渠道管理员;G1 不复制旧授权事实源。
|
||||||
|
migration_state=(
|
||||||
|
MigrationState.LEGACY_SHADOW
|
||||||
|
if requires_admin
|
||||||
|
else MigrationState.ENFORCED
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
# 固定未迁移工具和动态工具都保持现有执行能力,但不得被视为安全读取。
|
||||||
|
return ActionPolicy(
|
||||||
|
effect=ActionEffect.UNKNOWN,
|
||||||
|
required_role=required_role,
|
||||||
|
confirmation=ConfirmationMode.REQUIRED,
|
||||||
|
recovery=RecoveryMode.MANUAL_ONLY,
|
||||||
|
result_sensitivity=ResultSensitivity.UNKNOWN,
|
||||||
|
migration_state=MigrationState.LEGACY_SHADOW,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
DEFAULT_TOOL_POLICY_REGISTRY = ToolPolicyRegistry()
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"DEFAULT_TOOL_POLICY_REGISTRY",
|
||||||
|
"LEGACY_SHADOW_TOOL_NAMES",
|
||||||
|
"SAFE_READ_TOOL_NAMES",
|
||||||
|
"ToolPolicyRegistry",
|
||||||
|
]
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,69 @@
|
|||||||
|
"""Agent 设置工具与宿主回执共用的敏感字段身份判定。"""
|
||||||
|
|
||||||
|
import re
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
|
||||||
|
_MAX_FIELD_NAME_CHARS = 256
|
||||||
|
_ACRONYM_BOUNDARY_PATTERN = re.compile(r"(?<=[A-Z])(?=[A-Z][a-z])")
|
||||||
|
_CAMEL_CASE_BOUNDARY_PATTERN = re.compile(r"(?<=[a-z0-9])(?=[A-Z])")
|
||||||
|
_SECRET_FIELD_NAMES = frozenset(
|
||||||
|
{
|
||||||
|
"access_token",
|
||||||
|
"api_key",
|
||||||
|
"apikey",
|
||||||
|
"api_token",
|
||||||
|
"auth_header",
|
||||||
|
"authorization",
|
||||||
|
"client_secret",
|
||||||
|
"cookie",
|
||||||
|
"passkey",
|
||||||
|
"passwd",
|
||||||
|
"password",
|
||||||
|
"private_key",
|
||||||
|
"pwd",
|
||||||
|
"refresh_token",
|
||||||
|
"secret",
|
||||||
|
"secret_access_key",
|
||||||
|
"secret_key",
|
||||||
|
"token",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
_SECRET_FIELD_ENDINGS = tuple(f"_{name}" for name in _SECRET_FIELD_NAMES)
|
||||||
|
_SECRET_SETTING_NAMES = frozenset(
|
||||||
|
{
|
||||||
|
# CookieCloud 的用户 key 没有类型后缀,但与密码共同构成端到端加密凭据。
|
||||||
|
"cookiecloud_key",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
_SECRET_SETTING_ENDINGS = (
|
||||||
|
"_encrypt_key",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_field_name(value: Any) -> str:
|
||||||
|
"""将短字段名规范化为 snake_case,非字符串不参与身份推导。"""
|
||||||
|
if type(value) is not str:
|
||||||
|
return ""
|
||||||
|
text = value.strip()
|
||||||
|
if len(text) > _MAX_FIELD_NAME_CHARS:
|
||||||
|
text = text[-_MAX_FIELD_NAME_CHARS:]
|
||||||
|
text = _ACRONYM_BOUNDARY_PATTERN.sub("_", text)
|
||||||
|
text = _CAMEL_CASE_BOUNDARY_PATTERN.sub("_", text)
|
||||||
|
return re.sub(r"[^a-z0-9]+", "_", text.lower()).strip("_")
|
||||||
|
|
||||||
|
|
||||||
|
def is_secret_setting_key(key: Any) -> bool:
|
||||||
|
"""按完整字段或类型后缀识别凭据,避免误伤 token 统计与过期配置。"""
|
||||||
|
normalized = _normalize_field_name(key)
|
||||||
|
if not normalized:
|
||||||
|
return False
|
||||||
|
return (
|
||||||
|
normalized in _SECRET_FIELD_NAMES
|
||||||
|
or normalized in _SECRET_SETTING_NAMES
|
||||||
|
or normalized.endswith(_SECRET_FIELD_ENDINGS)
|
||||||
|
or normalized.endswith(_SECRET_SETTING_ENDINGS)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = ["is_secret_setting_key"]
|
||||||
+19
-13
@@ -11,6 +11,11 @@ from langchain_core.tools import BaseTool
|
|||||||
from pydantic import PrivateAttr
|
from pydantic import PrivateAttr
|
||||||
|
|
||||||
from app.agent import StreamingHandler
|
from app.agent import StreamingHandler
|
||||||
|
from app.agent.policy.sanitizer import (
|
||||||
|
summarize_error,
|
||||||
|
summarize_input,
|
||||||
|
summarize_result,
|
||||||
|
)
|
||||||
from app.agent.tools.tags import ToolTag
|
from app.agent.tools.tags import ToolTag
|
||||||
from app.chain import ChainBase
|
from app.chain import ChainBase
|
||||||
from app.core.config import settings
|
from app.core.config import settings
|
||||||
@@ -39,7 +44,9 @@ def serialize_tool_result_for_agent(result: Any) -> str:
|
|||||||
try:
|
try:
|
||||||
return json.dumps(result, ensure_ascii=False, indent=2, default=str)
|
return json.dumps(result, ensure_ascii=False, indent=2, default=str)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning(f"工具结果转换为JSON失败: {e}, 使用字符串表示")
|
logger.warning(
|
||||||
|
f"工具结果转换为JSON失败: {summarize_error(e)}, 使用字符串表示"
|
||||||
|
)
|
||||||
return str(result)
|
return str(result)
|
||||||
|
|
||||||
|
|
||||||
@@ -303,27 +310,26 @@ class MoviePilotTool(BaseTool, metaclass=ABCMeta):
|
|||||||
# 未启用流式传输,不发送任何工具消息内容
|
# 未启用流式传输,不发送任何工具消息内容
|
||||||
pass
|
pass
|
||||||
|
|
||||||
logger.debug(f"Executing tool {self.name} with args: {kwargs}")
|
logger.debug(
|
||||||
|
f"Executing tool {self.name} with input summary: {summarize_input(kwargs)}"
|
||||||
|
)
|
||||||
|
|
||||||
# 执行具体工具逻辑
|
# 执行具体工具逻辑
|
||||||
try:
|
try:
|
||||||
result = await self.run_with_timeout(**kwargs)
|
result = await self.run_with_timeout(**kwargs)
|
||||||
|
|
||||||
# 记录工具执行结果摘要日志
|
logger.info(
|
||||||
str_result = serialize_tool_result_for_agent(result)
|
f"Agent工具 {self.name} 执行完成,"
|
||||||
if len(str_result) > 500:
|
f"结果摘要: {summarize_result(result)}"
|
||||||
summary = str_result[:500] + f"...(已截断,总长度: {len(str_result)})"
|
)
|
||||||
else:
|
|
||||||
summary = str_result
|
|
||||||
logger.info(f"Agent工具 {self.name} 执行完成,结果摘要: {summary}")
|
|
||||||
|
|
||||||
except ToolExecutionTimeoutError as e:
|
except ToolExecutionTimeoutError as e:
|
||||||
error_message = str(e)
|
error_message = summarize_error(e)
|
||||||
logger.warning(error_message)
|
logger.warning(error_message)
|
||||||
result = error_message
|
result = error_message
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
error_message = f"工具执行异常 ({type(e).__name__}): {str(e)}"
|
error_message = f"工具执行异常: {summarize_error(e)}"
|
||||||
logger.error(f"Tool {self.name} execution failed: {e}", exc_info=True)
|
logger.error(f"Tool {self.name} execution failed: {summarize_error(e)}")
|
||||||
result = error_message
|
result = error_message
|
||||||
|
|
||||||
return format_tool_result_for_agent(
|
return format_tool_result_for_agent(
|
||||||
@@ -625,7 +631,7 @@ class MoviePilotTool(BaseTool, metaclass=ABCMeta):
|
|||||||
|
|
||||||
return False
|
return False
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"检查权限失败: {e}")
|
logger.error(f"检查权限失败: {summarize_error(e)}")
|
||||||
|
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|||||||
@@ -3,6 +3,7 @@
|
|||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import Any, Optional
|
from typing import Any, Optional
|
||||||
|
|
||||||
|
from app.agent.policy.secret_fields import is_secret_setting_key
|
||||||
from app.core.config import Settings
|
from app.core.config import Settings
|
||||||
from app.schemas.types import SystemConfigKey
|
from app.schemas.types import SystemConfigKey
|
||||||
|
|
||||||
@@ -368,26 +369,6 @@ def get_default_list_match_field(setting_key: str) -> Optional[str]:
|
|||||||
return LIST_ITEM_MATCH_FIELD_DEFAULTS.get(setting_key)
|
return LIST_ITEM_MATCH_FIELD_DEFAULTS.get(setting_key)
|
||||||
|
|
||||||
|
|
||||||
SECRET_KEYWORDS = (
|
|
||||||
"api_key",
|
|
||||||
"apikey",
|
|
||||||
"token",
|
|
||||||
"secret",
|
|
||||||
"password",
|
|
||||||
"passwd",
|
|
||||||
"cookie",
|
|
||||||
"authorization",
|
|
||||||
"refresh_token",
|
|
||||||
"access_token",
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def is_secret_setting_key(key: str) -> bool:
|
|
||||||
"""判断设置键名是否疑似敏感字段。"""
|
|
||||||
normalized = _normalize_token(key)
|
|
||||||
return any(keyword in normalized for keyword in SECRET_KEYWORDS)
|
|
||||||
|
|
||||||
|
|
||||||
def redact_secret_value(value: Any, *, redact_scalar: bool = False) -> Any:
|
def redact_secret_value(value: Any, *, redact_scalar: bool = False) -> Any:
|
||||||
"""递归脱敏配置值中的密钥、Cookie、Token 等敏感字段。"""
|
"""递归脱敏配置值中的密钥、Cookie、Token 等敏感字段。"""
|
||||||
if isinstance(value, dict):
|
if isinstance(value, dict):
|
||||||
|
|||||||
+52
-15
@@ -3,6 +3,16 @@ import threading
|
|||||||
import uuid
|
import uuid
|
||||||
from typing import Any, Dict, List, Optional
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
|
from app.agent.policy import (
|
||||||
|
DEFAULT_TOOL_POLICY_ORCHESTRATOR,
|
||||||
|
AgentToolPolicyOrchestrator,
|
||||||
|
AuthSource,
|
||||||
|
PrincipalType,
|
||||||
|
ToolOrigin,
|
||||||
|
ToolPolicyContext,
|
||||||
|
call_policy_hook,
|
||||||
|
summarize_error,
|
||||||
|
)
|
||||||
from app.agent.tools.base import ToolExecutionTimeoutError, format_tool_result_for_agent
|
from app.agent.tools.base import ToolExecutionTimeoutError, format_tool_result_for_agent
|
||||||
from app.agent.tools.factory import MoviePilotToolFactory
|
from app.agent.tools.factory import MoviePilotToolFactory
|
||||||
from app.core.plugin import PluginManager
|
from app.core.plugin import PluginManager
|
||||||
@@ -30,6 +40,7 @@ class MoviePilotToolsManager:
|
|||||||
user_id: str = "api_user",
|
user_id: str = "api_user",
|
||||||
session_id: str = uuid.uuid4(),
|
session_id: str = uuid.uuid4(),
|
||||||
is_admin: bool = True,
|
is_admin: bool = True,
|
||||||
|
policy_orchestrator: Optional[AgentToolPolicyOrchestrator] = None,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
初始化工具管理器
|
初始化工具管理器
|
||||||
@@ -41,6 +52,19 @@ class MoviePilotToolsManager:
|
|||||||
self.user_id = user_id
|
self.user_id = user_id
|
||||||
self.session_id = session_id
|
self.session_id = session_id
|
||||||
self.is_admin = is_admin
|
self.is_admin = is_admin
|
||||||
|
self.policy_orchestrator = (
|
||||||
|
policy_orchestrator or DEFAULT_TOOL_POLICY_ORCHESTRATOR
|
||||||
|
)
|
||||||
|
self._policy_context = ToolPolicyContext(
|
||||||
|
session_id=session_id,
|
||||||
|
user_id=user_id,
|
||||||
|
origin=ToolOrigin.OPERATOR_DIRECT,
|
||||||
|
principal_type=PrincipalType.SYSTEM_ADMIN_INTEGRATION,
|
||||||
|
auth_source=AuthSource.API_TOKEN,
|
||||||
|
channel=None,
|
||||||
|
source="api",
|
||||||
|
agent_context={"is_admin": is_admin},
|
||||||
|
)
|
||||||
self.tools: List[Any] = []
|
self.tools: List[Any] = []
|
||||||
self._tools_lock = threading.Lock()
|
self._tools_lock = threading.Lock()
|
||||||
self._plugin_agent_tools_revision = -1
|
self._plugin_agent_tools_revision = -1
|
||||||
@@ -74,7 +98,7 @@ class MoviePilotToolsManager:
|
|||||||
self._plugin_agent_tools_revision = plugin_tools_revision
|
self._plugin_agent_tools_revision = plugin_tools_revision
|
||||||
logger.info(f"成功加载 {len(self.tools)} 个工具")
|
logger.info(f"成功加载 {len(self.tools)} 个工具")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"加载工具失败: {e}", exc_info=True)
|
logger.error(f"加载工具失败: {summarize_error(e)}")
|
||||||
self.tools = []
|
self.tools = []
|
||||||
self._plugin_agent_tools_revision = -1
|
self._plugin_agent_tools_revision = -1
|
||||||
|
|
||||||
@@ -231,7 +255,7 @@ class MoviePilotToolsManager:
|
|||||||
schema = args_schema.model_json_schema()
|
schema = args_schema.model_json_schema()
|
||||||
properties = schema.get("properties", {})
|
properties = schema.get("properties", {})
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning(f"获取工具schema失败: {e}")
|
logger.warning(f"获取工具schema失败: {summarize_error(e)}")
|
||||||
return arguments
|
return arguments
|
||||||
|
|
||||||
# 规范化参数
|
# 规范化参数
|
||||||
@@ -286,6 +310,7 @@ class MoviePilotToolsManager:
|
|||||||
)
|
)
|
||||||
return error_msg
|
return error_msg
|
||||||
|
|
||||||
|
observation = None
|
||||||
try:
|
try:
|
||||||
permission_error = self._check_tool_permission(tool_instance)
|
permission_error = self._check_tool_permission(tool_instance)
|
||||||
if permission_error:
|
if permission_error:
|
||||||
@@ -293,38 +318,50 @@ class MoviePilotToolsManager:
|
|||||||
|
|
||||||
# 规范化参数类型
|
# 规范化参数类型
|
||||||
normalized_arguments = self._normalize_arguments(tool_instance, arguments)
|
normalized_arguments = self._normalize_arguments(tool_instance, arguments)
|
||||||
|
self._policy_context.agent_context["is_admin"] = self.is_admin
|
||||||
|
observation = call_policy_hook(
|
||||||
|
"start",
|
||||||
|
self.policy_orchestrator.start,
|
||||||
|
context=self._policy_context,
|
||||||
|
tool=tool_instance,
|
||||||
|
arguments=normalized_arguments,
|
||||||
|
)
|
||||||
|
|
||||||
# 调用工具的run方法。HTTP/MCP 工具调用不会经过 BaseTool._arun,
|
# 调用工具的run方法。HTTP/MCP 工具调用不会经过 BaseTool._arun,
|
||||||
# 因此这里也必须复用同一套返回值格式化和兜底截断逻辑。
|
# 因此这里也必须复用同一套返回值格式化和兜底截断逻辑。
|
||||||
result = await tool_instance.run_with_timeout(**normalized_arguments)
|
result = await tool_instance.run_with_timeout(**normalized_arguments)
|
||||||
|
|
||||||
# 记录工具执行结果摘要日志
|
|
||||||
str_result = format_tool_result_for_agent(
|
str_result = format_tool_result_for_agent(
|
||||||
result,
|
result,
|
||||||
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),
|
||||||
)
|
)
|
||||||
if len(str_result) > 500:
|
|
||||||
summary = str_result[:500] + f"...(已截断,总长度: {len(str_result)})"
|
|
||||||
else:
|
|
||||||
summary = str_result
|
|
||||||
logger.info(f"Agent工具 {tool_name} 执行完成,结果摘要: {summary}")
|
|
||||||
|
|
||||||
return str_result
|
|
||||||
except ToolExecutionTimeoutError as e:
|
except ToolExecutionTimeoutError as e:
|
||||||
logger.warning(str(e))
|
if observation:
|
||||||
|
call_policy_hook("fail", self.policy_orchestrator.fail, observation, e)
|
||||||
|
logger.warning(summarize_error(e))
|
||||||
return format_tool_result_for_agent(
|
return format_tool_result_for_agent(
|
||||||
str(e),
|
summarize_error(e),
|
||||||
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 Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"调用工具 {tool_name} 时发生错误: {e}", exc_info=True)
|
if observation:
|
||||||
|
call_policy_hook("fail", self.policy_orchestrator.fail, observation, e)
|
||||||
|
error_summary = summarize_error(e)
|
||||||
|
logger.error(f"调用工具 {tool_name} 时发生错误: {error_summary}")
|
||||||
error_msg = json.dumps(
|
error_msg = json.dumps(
|
||||||
{"error": f"调用工具 '{tool_name}' 时发生错误: {str(e)}"},
|
{"error": f"调用工具 '{tool_name}' 时发生错误: {error_summary}"},
|
||||||
ensure_ascii=False,
|
ensure_ascii=False,
|
||||||
)
|
)
|
||||||
return error_msg
|
return error_msg
|
||||||
|
if observation:
|
||||||
|
call_policy_hook(
|
||||||
|
"finish",
|
||||||
|
self.policy_orchestrator.finish,
|
||||||
|
observation,
|
||||||
|
str_result,
|
||||||
|
)
|
||||||
|
return str_result
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _convert_to_json_schema(args_schema: Any) -> Dict[str, Any]:
|
def _convert_to_json_schema(args_schema: Any) -> Dict[str, Any]:
|
||||||
|
|||||||
@@ -2,7 +2,7 @@ import asyncio
|
|||||||
import json
|
import json
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import AsyncMock, patch
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
from langchain_core.messages import AIMessage, HumanMessage, SystemMessage, ToolMessage
|
from langchain_core.messages import AIMessage, HumanMessage, SystemMessage, ToolMessage
|
||||||
|
|
||||||
@@ -418,6 +418,67 @@ def test_activity_log_tool_call_records_streaming_summary(tmp_path):
|
|||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def test_activity_log_middleware_sanitizes_its_own_logs(tmp_path):
|
||||||
|
"""活动日志中间件读取参数和异常写日志时必须脱敏。"""
|
||||||
|
|
||||||
|
async def _run_test():
|
||||||
|
secret_marker = "activity-secret-marker-6825"
|
||||||
|
stream_handler = SimpleNamespace(
|
||||||
|
is_streaming=True,
|
||||||
|
record_tool_call=MagicMock(),
|
||||||
|
)
|
||||||
|
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": f"token={secret_marker}"}},
|
||||||
|
)
|
||||||
|
mock_logger = MagicMock()
|
||||||
|
|
||||||
|
async def _failing_handler(_request):
|
||||||
|
raise RuntimeError(f"Authorization: Bearer {secret_marker}")
|
||||||
|
|
||||||
|
with patch("app.agent.middleware.activity_log.logger", mock_logger):
|
||||||
|
try:
|
||||||
|
await middleware.awrap_tool_call(request, _failing_handler)
|
||||||
|
except RuntimeError:
|
||||||
|
pass
|
||||||
|
else:
|
||||||
|
raise AssertionError("middleware should re-raise handler errors")
|
||||||
|
|
||||||
|
return secret_marker, mock_logger
|
||||||
|
|
||||||
|
secret_marker, mock_logger = asyncio.run(_run_test())
|
||||||
|
|
||||||
|
assert secret_marker not in str(mock_logger.method_calls)
|
||||||
|
assert "***" in str(mock_logger.method_calls)
|
||||||
|
|
||||||
|
|
||||||
|
def test_activity_log_provider_error_does_not_echo_secret(tmp_path):
|
||||||
|
"""活动日志 provider 内部异常不能进入日志或模型错误结果。"""
|
||||||
|
secret_marker = "activity-provider-secret-3584"
|
||||||
|
middleware = ActivityLogMiddleware(activity_dir=str(tmp_path))
|
||||||
|
mock_logger = MagicMock()
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"app.agent.middleware.activity_log.query_activity_logs",
|
||||||
|
side_effect=RuntimeError(f"OPENAI_API_KEY={secret_marker}"),
|
||||||
|
),
|
||||||
|
patch("app.agent.middleware.activity_log.logger", mock_logger),
|
||||||
|
):
|
||||||
|
result = asyncio.run(
|
||||||
|
middleware._tool_provider.query_activity_log(keyword="visible")
|
||||||
|
)
|
||||||
|
|
||||||
|
assert secret_marker not in result
|
||||||
|
assert secret_marker not in str(mock_logger.method_calls)
|
||||||
|
assert "***" in result
|
||||||
|
assert "***" in str(mock_logger.method_calls)
|
||||||
|
|
||||||
|
|
||||||
def test_factory_does_not_register_activity_log_tool():
|
def test_factory_does_not_register_activity_log_tool():
|
||||||
"""活动日志查询工具应由中间件注册,不应进入全局工具工厂。"""
|
"""活动日志查询工具应由中间件注册,不应进入全局工具工厂。"""
|
||||||
with patch(
|
with patch(
|
||||||
|
|||||||
@@ -431,6 +431,7 @@ class AgentBackgroundOutputTest(unittest.IsolatedAsyncioTestCase):
|
|||||||
|
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
[
|
[
|
||||||
|
"AgentPolicyMiddleware",
|
||||||
"skills",
|
"skills",
|
||||||
"jobs",
|
"jobs",
|
||||||
"runtime",
|
"runtime",
|
||||||
@@ -546,6 +547,7 @@ class AgentBackgroundOutputTest(unittest.IsolatedAsyncioTestCase):
|
|||||||
|
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
[
|
[
|
||||||
|
"AgentPolicyMiddleware",
|
||||||
"skills",
|
"skills",
|
||||||
"jobs",
|
"jobs",
|
||||||
"runtime",
|
"runtime",
|
||||||
@@ -749,6 +751,7 @@ class AgentBackgroundOutputTest(unittest.IsolatedAsyncioTestCase):
|
|||||||
|
|
||||||
self.assertEqual(
|
self.assertEqual(
|
||||||
[
|
[
|
||||||
|
"AgentPolicyMiddleware",
|
||||||
"skills",
|
"skills",
|
||||||
"jobs",
|
"jobs",
|
||||||
"runtime",
|
"runtime",
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
import json
|
import json
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
from anyio import Path as AsyncPath
|
from anyio import Path as AsyncPath
|
||||||
@@ -171,3 +172,59 @@ async def test_skill_tool_call_records_streaming_summary(tmp_path):
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_skill_middleware_sanitizes_its_own_logs(tmp_path):
|
||||||
|
"""Skill 中间件读取参数和异常写日志时必须脱敏。"""
|
||||||
|
secret_marker = "skill-secret-marker-2471"
|
||||||
|
stream_handler = SimpleNamespace(
|
||||||
|
is_streaming=True,
|
||||||
|
record_tool_call=MagicMock(),
|
||||||
|
)
|
||||||
|
middleware = SkillsMiddleware(
|
||||||
|
sources=[str(tmp_path)],
|
||||||
|
stream_handler=stream_handler,
|
||||||
|
)
|
||||||
|
request = SimpleNamespace(
|
||||||
|
tool=SimpleNamespace(name=SKILL_TOOL_NAME),
|
||||||
|
tool_call={"args": {"name": f"api_key={secret_marker}"}},
|
||||||
|
)
|
||||||
|
mock_logger = MagicMock()
|
||||||
|
|
||||||
|
async def _failing_handler(_request):
|
||||||
|
raise RuntimeError(f"Authorization: Bearer {secret_marker}")
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch("app.agent.middleware.skills.logger", mock_logger),
|
||||||
|
pytest.raises(RuntimeError),
|
||||||
|
):
|
||||||
|
await middleware.awrap_tool_call(request, _failing_handler)
|
||||||
|
|
||||||
|
assert secret_marker not in str(mock_logger.method_calls)
|
||||||
|
assert "***" in str(mock_logger.method_calls)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.anyio
|
||||||
|
async def test_skill_provider_error_does_not_echo_secret(tmp_path):
|
||||||
|
"""Skill provider 内部捕获的异常不能进入日志或模型错误结果。"""
|
||||||
|
secret_marker = "skill-provider-secret-6518"
|
||||||
|
middleware = SkillsMiddleware(sources=[str(tmp_path)])
|
||||||
|
mock_logger = MagicMock()
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch.object(
|
||||||
|
middleware._skill_provider,
|
||||||
|
"_find_skill",
|
||||||
|
new=AsyncMock(
|
||||||
|
side_effect=RuntimeError(f"DATABASE_PASSWORD={secret_marker}")
|
||||||
|
),
|
||||||
|
),
|
||||||
|
patch("app.agent.middleware.skills.logger", mock_logger),
|
||||||
|
):
|
||||||
|
result = await middleware._skill_provider.load_skill("visible-skill")
|
||||||
|
|
||||||
|
assert secret_marker not in result
|
||||||
|
assert secret_marker not in str(mock_logger.method_calls)
|
||||||
|
assert "***" in result
|
||||||
|
assert "***" in str(mock_logger.method_calls)
|
||||||
|
|||||||
@@ -2,11 +2,12 @@ import asyncio
|
|||||||
import json
|
import json
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import patch
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
from langchain_core.language_models.fake_chat_models import FakeListChatModel
|
from langchain_core.language_models.fake_chat_models import FakeListChatModel
|
||||||
|
|
||||||
import app.agent.middleware.subagents as subagent_module
|
import app.agent.middleware.subagents as subagent_module
|
||||||
|
from app.agent.middleware.policy import AgentPolicyMiddleware
|
||||||
from app.agent.middleware.subagents import (
|
from app.agent.middleware.subagents import (
|
||||||
MoviePilotSubAgentMiddleware,
|
MoviePilotSubAgentMiddleware,
|
||||||
SUBAGENT_CONTROL_TOOL_NAME,
|
SUBAGENT_CONTROL_TOOL_NAME,
|
||||||
@@ -14,6 +15,7 @@ from app.agent.middleware.subagents import (
|
|||||||
SubAgentTaskControlMiddleware,
|
SubAgentTaskControlMiddleware,
|
||||||
create_subagent_middlewares,
|
create_subagent_middlewares,
|
||||||
)
|
)
|
||||||
|
from app.agent.policy import AuthSource, PrincipalType, ToolOrigin, ToolPolicyContext
|
||||||
from app.agent.tools.tags import ToolTag
|
from app.agent.tools.tags import ToolTag
|
||||||
|
|
||||||
|
|
||||||
@@ -77,6 +79,40 @@ def test_subagent_tools_are_selected_by_tags():
|
|||||||
assert [tool.name for tool in captured["tools"]] == ["custom_media_lookup"]
|
assert [tool.name for tool in captured["tools"]] == ["custom_media_lookup"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_subagent_graph_registers_policy_middleware_as_outermost():
|
||||||
|
"""懒加载的子代理图必须继承宿主上下文并先经过 policy middleware。"""
|
||||||
|
model = FakeListChatModel(responses=["ok"])
|
||||||
|
context = ToolPolicyContext(
|
||||||
|
session_id="subagent-session",
|
||||||
|
user_id="user-1",
|
||||||
|
origin=ToolOrigin.SUBAGENT,
|
||||||
|
principal_type=PrincipalType.SUBAGENT,
|
||||||
|
auth_source=AuthSource.INTERNAL,
|
||||||
|
agent_context={"is_admin": True},
|
||||||
|
channel="Telegram",
|
||||||
|
source="telegram",
|
||||||
|
)
|
||||||
|
captured = {}
|
||||||
|
|
||||||
|
def _fake_create_agent(**kwargs):
|
||||||
|
captured.update(kwargs)
|
||||||
|
return kwargs
|
||||||
|
|
||||||
|
middleware = MoviePilotSubAgentMiddleware(
|
||||||
|
model=model,
|
||||||
|
profiles=subagent_module._builtin_subagent_profiles(),
|
||||||
|
tools=[],
|
||||||
|
policy_context=context,
|
||||||
|
)
|
||||||
|
|
||||||
|
with patch.object(subagent_module, "create_agent", side_effect=_fake_create_agent):
|
||||||
|
middleware._get_agent("general-purpose")
|
||||||
|
|
||||||
|
assert isinstance(captured["middleware"][0], AgentPolicyMiddleware)
|
||||||
|
assert captured["middleware"][0].context is context
|
||||||
|
assert captured["middleware"][0].context.origin is ToolOrigin.SUBAGENT
|
||||||
|
|
||||||
|
|
||||||
def test_moviepilot_explorer_selects_code_and_settings_tools():
|
def test_moviepilot_explorer_selects_code_and_settings_tools():
|
||||||
"""MoviePilot 探索子代理应能读取代码、目录、设置和命令诊断工具。"""
|
"""MoviePilot 探索子代理应能读取代码、目录、设置和命令诊断工具。"""
|
||||||
model = FakeListChatModel(responses=["ok"])
|
model = FakeListChatModel(responses=["ok"])
|
||||||
@@ -187,6 +223,51 @@ def test_task_tool_call_records_streaming_summary():
|
|||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def test_task_middleware_sanitizes_its_own_logs():
|
||||||
|
"""子代理中间件读取任务参数和异常写日志时必须脱敏。"""
|
||||||
|
|
||||||
|
async def _run_test():
|
||||||
|
secret_marker = "subagent-secret-marker-7316"
|
||||||
|
stream_handler = SimpleNamespace(
|
||||||
|
is_streaming=True,
|
||||||
|
record_tool_call=MagicMock(),
|
||||||
|
)
|
||||||
|
middleware = MoviePilotSubAgentMiddleware(
|
||||||
|
model=FakeListChatModel(responses=["ok"]),
|
||||||
|
profiles=subagent_module._builtin_subagent_profiles(),
|
||||||
|
tools=[],
|
||||||
|
stream_handler=stream_handler,
|
||||||
|
)
|
||||||
|
request = SimpleNamespace(
|
||||||
|
tool=SimpleNamespace(name=SUBAGENT_TASK_TOOL_NAME),
|
||||||
|
tool_call={
|
||||||
|
"args": {
|
||||||
|
"description": f"password={secret_marker}",
|
||||||
|
"subagent_type": "media-researcher",
|
||||||
|
}
|
||||||
|
},
|
||||||
|
)
|
||||||
|
mock_logger = MagicMock()
|
||||||
|
|
||||||
|
async def _failing_handler(_request):
|
||||||
|
raise RuntimeError(f"Authorization: Bearer {secret_marker}")
|
||||||
|
|
||||||
|
with patch.object(subagent_module, "logger", mock_logger):
|
||||||
|
try:
|
||||||
|
await middleware.awrap_tool_call(request, _failing_handler)
|
||||||
|
except RuntimeError:
|
||||||
|
pass
|
||||||
|
else:
|
||||||
|
raise AssertionError("middleware should re-raise handler errors")
|
||||||
|
|
||||||
|
return secret_marker, mock_logger
|
||||||
|
|
||||||
|
secret_marker, mock_logger = asyncio.run(_run_test())
|
||||||
|
|
||||||
|
assert secret_marker not in str(mock_logger.method_calls)
|
||||||
|
assert "***" in str(mock_logger.method_calls)
|
||||||
|
|
||||||
|
|
||||||
def test_control_tool_call_records_streaming_summary():
|
def test_control_tool_call_records_streaming_summary():
|
||||||
"""subagent_task 子代理工具执行时应记录流式聚合摘要。"""
|
"""subagent_task 子代理工具执行时应记录流式聚合摘要。"""
|
||||||
|
|
||||||
@@ -392,11 +473,14 @@ def test_control_tool_pipeline_stops_after_failed_step():
|
|||||||
tools=[],
|
tools=[],
|
||||||
)
|
)
|
||||||
calls = []
|
calls = []
|
||||||
|
secret_marker = "subagent-runtime-secret-9042"
|
||||||
|
|
||||||
async def _fake_run_task(self, *, description, subagent_type, task_id=None):
|
async def _fake_run_task(self, *, description, subagent_type, task_id=None):
|
||||||
calls.append(subagent_type)
|
calls.append(subagent_type)
|
||||||
if subagent_type == "download-diagnostician":
|
if subagent_type == "download-diagnostician":
|
||||||
raise RuntimeError("下载器不可用")
|
raise RuntimeError(
|
||||||
|
f"下载器不可用 DATABASE_PASSWORD={secret_marker}"
|
||||||
|
)
|
||||||
return f"{subagent_type}:ok"
|
return f"{subagent_type}:ok"
|
||||||
|
|
||||||
with patch.object(
|
with patch.object(
|
||||||
@@ -433,6 +517,10 @@ def test_control_tool_pipeline_stops_after_failed_step():
|
|||||||
"failed",
|
"failed",
|
||||||
]
|
]
|
||||||
assert "下载器不可用" in payload["tasks"][1]["error"]
|
assert "下载器不可用" in payload["tasks"][1]["error"]
|
||||||
|
assert secret_marker not in payload["error"]
|
||||||
|
assert secret_marker not in payload["tasks"][1]["error"]
|
||||||
|
assert "***" in payload["error"]
|
||||||
|
assert "***" in payload["tasks"][1]["error"]
|
||||||
|
|
||||||
asyncio.run(_run_test())
|
asyncio.run(_run_test())
|
||||||
|
|
||||||
|
|||||||
@@ -3,6 +3,8 @@ import json
|
|||||||
import unittest
|
import unittest
|
||||||
from unittest.mock import AsyncMock, MagicMock, patch
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
from app.agent.tools.impl._system_setting_utils import list_setting_specs
|
from app.agent.tools.impl._system_setting_utils import list_setting_specs
|
||||||
from app.agent.tools.impl.query_system_settings import QuerySystemSettingsTool
|
from app.agent.tools.impl.query_system_settings import QuerySystemSettingsTool
|
||||||
from app.agent.tools.impl.update_system_settings import UpdateSystemSettingsTool
|
from app.agent.tools.impl.update_system_settings import UpdateSystemSettingsTool
|
||||||
@@ -295,3 +297,44 @@ class TestAgentSystemSettingsTools(unittest.TestCase):
|
|||||||
payload = json.loads(result)
|
payload = json.loads(result)
|
||||||
self.assertIn("error", payload)
|
self.assertIn("error", payload)
|
||||||
self.assertIn("系统管理员", payload["error"])
|
self.assertIn("系统管理员", payload["error"])
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("setting_key", "should_redact"),
|
||||||
|
[
|
||||||
|
("API_TOKEN", True),
|
||||||
|
("LLM_API_KEY", True),
|
||||||
|
("COOKIECLOUD_KEY", True),
|
||||||
|
("COOKIECLOUD_AUTH_HEADER", True),
|
||||||
|
("SUPERUSER_PASSWORD", True),
|
||||||
|
("DB_POSTGRESQL_PASSWORD", True),
|
||||||
|
("GITHUB_TOKEN", True),
|
||||||
|
("FEISHU_VERIFICATION_TOKEN", True),
|
||||||
|
("SECRET_KEY", True),
|
||||||
|
("RESOURCE_SECRET_KEY", True),
|
||||||
|
("PROJECT_NAME", False),
|
||||||
|
("ACCESS_TOKEN_EXPIRE_MINUTES", False),
|
||||||
|
("LLM_MAX_CONTEXT_TOKENS", False),
|
||||||
|
("COOKIECLOUD_INTERVAL", False),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_query_system_settings_uses_precise_secret_identity_matrix(
|
||||||
|
setting_key: str,
|
||||||
|
should_redact: bool,
|
||||||
|
) -> None:
|
||||||
|
"""设置查询应隐藏真实凭据,同时保留仅名称相似的普通设置。"""
|
||||||
|
marker = "credential-marker" if should_redact else "visible-marker"
|
||||||
|
tool = QuerySystemSettingsTool(session_id="session-1", user_id="10001")
|
||||||
|
|
||||||
|
with patch.object(
|
||||||
|
QuerySystemSettingsTool,
|
||||||
|
"_load_setting_value",
|
||||||
|
return_value=marker,
|
||||||
|
):
|
||||||
|
payload = json.loads(asyncio.run(tool.run(setting_key=setting_key)))
|
||||||
|
|
||||||
|
item = payload["settings"][0]
|
||||||
|
assert item["redacted"] is should_redact
|
||||||
|
assert item["value"] == ("***" if should_redact else marker)
|
||||||
|
if should_redact:
|
||||||
|
assert marker not in json.dumps(payload)
|
||||||
|
|||||||
@@ -0,0 +1,744 @@
|
|||||||
|
import asyncio
|
||||||
|
import json
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from langchain.agents.middleware import SummarizationMiddleware
|
||||||
|
from langchain_core.messages import ToolMessage
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
|
import app.agent as agent_module
|
||||||
|
from app.agent.middleware.activity_log import ActivityLogMiddleware
|
||||||
|
from app.agent.middleware.memory import MemoryMiddleware
|
||||||
|
from app.agent.middleware.policy import AgentPolicyMiddleware
|
||||||
|
from app.agent.policy import (
|
||||||
|
DEFAULT_TOOL_POLICY_ORCHESTRATOR,
|
||||||
|
DEFAULT_TOOL_POLICY_REGISTRY,
|
||||||
|
ActionEffect,
|
||||||
|
AuthSource,
|
||||||
|
MigrationState,
|
||||||
|
PolicyPrincipal,
|
||||||
|
PrincipalRole,
|
||||||
|
PrincipalType,
|
||||||
|
ResultSensitivity,
|
||||||
|
ToolOrigin,
|
||||||
|
ToolPolicyContext,
|
||||||
|
call_policy_hook,
|
||||||
|
)
|
||||||
|
from app.agent.tools.base import MoviePilotTool
|
||||||
|
from app.agent.tools.factory import MoviePilotToolFactory
|
||||||
|
from app.agent.tools.impl.ask_user_choice import AskUserChoiceTool
|
||||||
|
from app.agent.tools.impl.query_system_settings import QuerySystemSettingsTool
|
||||||
|
from app.agent.tools.impl.send_local_file import SendLocalFileTool
|
||||||
|
from app.agent.tools.impl.send_voice_message import SendVoiceMessageTool
|
||||||
|
from app.agent.tools.manager import MoviePilotToolsManager
|
||||||
|
from app.schemas.types import MessageChannel
|
||||||
|
|
||||||
|
|
||||||
|
class _EchoInput(BaseModel):
|
||||||
|
"""策略测试工具的输入契约。"""
|
||||||
|
|
||||||
|
query: str = Field(description="待回显文本")
|
||||||
|
|
||||||
|
|
||||||
|
class _EchoTool(MoviePilotTool):
|
||||||
|
"""返回输入文本的测试工具。"""
|
||||||
|
|
||||||
|
name: str = "policy_echo"
|
||||||
|
description: str = "Echo policy test input."
|
||||||
|
args_schema: type[BaseModel] = _EchoInput
|
||||||
|
|
||||||
|
async def run(self, query: str) -> str:
|
||||||
|
"""返回输入文本。"""
|
||||||
|
return query
|
||||||
|
|
||||||
|
|
||||||
|
class _OverrideArunTool(_EchoTool):
|
||||||
|
"""模拟插件覆盖宿主工具基类 `_arun` 的实现。"""
|
||||||
|
|
||||||
|
name: str = "plugin_override_arun"
|
||||||
|
|
||||||
|
def __init__(self, events: list[str], **kwargs):
|
||||||
|
super().__init__(**kwargs)
|
||||||
|
self._events = events
|
||||||
|
|
||||||
|
async def _arun(self, *args, **kwargs) -> str:
|
||||||
|
"""绕过基类实现并记录真实执行顺序。"""
|
||||||
|
self._events.append("tool")
|
||||||
|
return str(kwargs.get("query") or "ok")
|
||||||
|
|
||||||
|
|
||||||
|
class _FailingTool(_EchoTool):
|
||||||
|
"""抛出固定异常以验证 direct manager 的 observation fail-open。"""
|
||||||
|
|
||||||
|
name: str = "policy_failure"
|
||||||
|
|
||||||
|
async def run(self, query: str) -> str:
|
||||||
|
"""模拟真实工具失败。"""
|
||||||
|
raise ValueError(f"tool-error:{query}")
|
||||||
|
|
||||||
|
|
||||||
|
class _AdminSafeReadTool(_EchoTool):
|
||||||
|
"""复用 safe-read 名称并保留旧 require_admin 门禁的测试工具。"""
|
||||||
|
|
||||||
|
name: str = "list_slash_commands"
|
||||||
|
require_admin: bool = True
|
||||||
|
|
||||||
|
def __init__(self, events: list[str], **kwargs):
|
||||||
|
super().__init__(**kwargs)
|
||||||
|
self._events = events
|
||||||
|
|
||||||
|
async def run(self, query: str) -> str:
|
||||||
|
"""记录实际执行,区分旧门禁放行与拒绝。"""
|
||||||
|
self._events.append("run")
|
||||||
|
return query
|
||||||
|
|
||||||
|
|
||||||
|
def _tool_class_name(tool_class: type[MoviePilotTool]) -> str:
|
||||||
|
"""从 Pydantic 字段默认值读取工具类的稳定名称。"""
|
||||||
|
return str(tool_class.model_fields["name"].default)
|
||||||
|
|
||||||
|
|
||||||
|
def _interactive_context(*, is_admin: bool = True) -> ToolPolicyContext:
|
||||||
|
"""构造会随本轮管理员上下文刷新的交互式策略上下文。"""
|
||||||
|
return ToolPolicyContext(
|
||||||
|
session_id="session-1",
|
||||||
|
user_id="user-1",
|
||||||
|
origin=ToolOrigin.AGENT_INTERACTIVE,
|
||||||
|
principal_type=PrincipalType.HUMAN,
|
||||||
|
auth_source=AuthSource.CHANNEL,
|
||||||
|
channel="telegram",
|
||||||
|
source="user",
|
||||||
|
agent_context={"is_admin": is_admin},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_builtin_policy_registry_covers_every_fixed_tool() -> None:
|
||||||
|
"""固定内置工具必须全部具有显式 migration registry 条目。"""
|
||||||
|
fixed_tool_names = {
|
||||||
|
_tool_class_name(tool_class)
|
||||||
|
for tool_class in MoviePilotToolFactory.BUILTIN_TOOL_CLASSES
|
||||||
|
}
|
||||||
|
fixed_tool_names.update(
|
||||||
|
{
|
||||||
|
_tool_class_name(AskUserChoiceTool),
|
||||||
|
_tool_class_name(SendLocalFileTool),
|
||||||
|
_tool_class_name(SendVoiceMessageTool),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
assert DEFAULT_TOOL_POLICY_REGISTRY.builtin_tool_names == fixed_tool_names
|
||||||
|
|
||||||
|
|
||||||
|
def test_registry_separates_safe_read_from_legacy_shadow() -> None:
|
||||||
|
"""少量安全读取可直接迁移,其余高影响工具保持 shadow allow。"""
|
||||||
|
safe_policy = DEFAULT_TOOL_POLICY_REGISTRY.resolve(
|
||||||
|
tool_name="query_personas",
|
||||||
|
arguments={},
|
||||||
|
requires_admin=False,
|
||||||
|
)
|
||||||
|
admin_safe_policy = DEFAULT_TOOL_POLICY_REGISTRY.resolve(
|
||||||
|
tool_name="list_slash_commands",
|
||||||
|
arguments={},
|
||||||
|
requires_admin=True,
|
||||||
|
)
|
||||||
|
shadow_policy = DEFAULT_TOOL_POLICY_REGISTRY.resolve(
|
||||||
|
tool_name="update_system_settings",
|
||||||
|
arguments={"updates": []},
|
||||||
|
requires_admin=True,
|
||||||
|
)
|
||||||
|
dynamic_policy = DEFAULT_TOOL_POLICY_REGISTRY.resolve(
|
||||||
|
tool_name="plugin_unknown_tool",
|
||||||
|
arguments={"action": "custom"},
|
||||||
|
requires_admin=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert safe_policy.effect is ActionEffect.SAFE_READ
|
||||||
|
assert safe_policy.result_sensitivity is ResultSensitivity.NORMAL
|
||||||
|
assert safe_policy.migration_state is MigrationState.ENFORCED
|
||||||
|
|
||||||
|
assert admin_safe_policy.effect is ActionEffect.SAFE_READ
|
||||||
|
assert admin_safe_policy.required_role is PrincipalRole.SYSTEM_ADMIN
|
||||||
|
assert admin_safe_policy.migration_state is MigrationState.LEGACY_SHADOW
|
||||||
|
|
||||||
|
assert shadow_policy.migration_state is MigrationState.LEGACY_SHADOW
|
||||||
|
assert shadow_policy.effect is ActionEffect.UNKNOWN
|
||||||
|
assert shadow_policy.required_role is PrincipalRole.SYSTEM_ADMIN
|
||||||
|
|
||||||
|
assert dynamic_policy.migration_state is MigrationState.LEGACY_SHADOW
|
||||||
|
assert dynamic_policy.effect is ActionEffect.UNKNOWN
|
||||||
|
assert dynamic_policy.result_sensitivity is ResultSensitivity.UNKNOWN
|
||||||
|
|
||||||
|
|
||||||
|
def test_legacy_shadow_decision_allows_without_claiming_enforcement() -> None:
|
||||||
|
"""G1 的 shadow 决策只能观测,不能拒绝或要求确认。"""
|
||||||
|
context = _interactive_context()
|
||||||
|
tool = _EchoTool(session_id="session-1", user_id="user-1")
|
||||||
|
|
||||||
|
observation = DEFAULT_TOOL_POLICY_ORCHESTRATOR.start(
|
||||||
|
context=context,
|
||||||
|
tool=tool,
|
||||||
|
arguments={"query": "hello"},
|
||||||
|
)
|
||||||
|
|
||||||
|
assert observation.decision.allowed is True
|
||||||
|
assert observation.decision.shadow is True
|
||||||
|
assert observation.decision.reason_code == "legacy_shadow_allow"
|
||||||
|
|
||||||
|
|
||||||
|
def test_policy_context_reads_mutable_admin_state_without_model_fields() -> None:
|
||||||
|
"""缓存图复用时权限取当前宿主状态,模型参数不能伪造 principal。"""
|
||||||
|
context = _interactive_context(is_admin=False)
|
||||||
|
forged_arguments = {
|
||||||
|
"query": "hello",
|
||||||
|
"origin": "operator_direct",
|
||||||
|
"principal_type": "system_admin_integration",
|
||||||
|
"is_admin": True,
|
||||||
|
}
|
||||||
|
|
||||||
|
assert context.principal.role is PrincipalRole.USER
|
||||||
|
context.agent_context["is_admin"] = True
|
||||||
|
assert context.principal.role is PrincipalRole.SYSTEM_ADMIN
|
||||||
|
assert "origin" not in PolicyPrincipal.__dataclass_fields__
|
||||||
|
assert forged_arguments["origin"] != context.origin.value
|
||||||
|
|
||||||
|
|
||||||
|
def test_policy_context_maps_trusted_host_origins() -> None:
|
||||||
|
"""各入口必须由宿主稳定映射 origin、主体类型和认证来源。"""
|
||||||
|
cases = [
|
||||||
|
(
|
||||||
|
{"channel": MessageChannel.Web.value, "source": "openai"},
|
||||||
|
ToolOrigin.AGENT_API,
|
||||||
|
PrincipalType.SYSTEM_ADMIN_INTEGRATION,
|
||||||
|
AuthSource.API_TOKEN,
|
||||||
|
),
|
||||||
|
(
|
||||||
|
{"channel": MessageChannel.Web.value, "source": "openai.responses"},
|
||||||
|
ToolOrigin.AGENT_API,
|
||||||
|
PrincipalType.SYSTEM_ADMIN_INTEGRATION,
|
||||||
|
AuthSource.API_TOKEN,
|
||||||
|
),
|
||||||
|
(
|
||||||
|
{"channel": MessageChannel.Web.value, "source": "anthropic"},
|
||||||
|
ToolOrigin.AGENT_API,
|
||||||
|
PrincipalType.SYSTEM_ADMIN_INTEGRATION,
|
||||||
|
AuthSource.API_TOKEN,
|
||||||
|
),
|
||||||
|
(
|
||||||
|
{"channel": MessageChannel.Web.value, "source": "browser"},
|
||||||
|
ToolOrigin.AGENT_INTERACTIVE,
|
||||||
|
PrincipalType.HUMAN,
|
||||||
|
AuthSource.WEB_SESSION,
|
||||||
|
),
|
||||||
|
(
|
||||||
|
{"channel": MessageChannel.WebAgent.value, "source": "web-agent"},
|
||||||
|
ToolOrigin.AGENT_INTERACTIVE,
|
||||||
|
PrincipalType.HUMAN,
|
||||||
|
AuthSource.WEB_SESSION,
|
||||||
|
),
|
||||||
|
(
|
||||||
|
{"channel": MessageChannel.Telegram.value, "source": "telegram"},
|
||||||
|
ToolOrigin.AGENT_INTERACTIVE,
|
||||||
|
PrincipalType.HUMAN,
|
||||||
|
AuthSource.CHANNEL,
|
||||||
|
),
|
||||||
|
(
|
||||||
|
{"channel": MessageChannel.Feishu.value, "source": "feishu"},
|
||||||
|
ToolOrigin.AGENT_INTERACTIVE,
|
||||||
|
PrincipalType.HUMAN,
|
||||||
|
AuthSource.CHANNEL,
|
||||||
|
),
|
||||||
|
(
|
||||||
|
{"channel": None, "source": None, "output_callback": lambda _text: None},
|
||||||
|
ToolOrigin.BACKGROUND,
|
||||||
|
PrincipalType.BACKGROUND,
|
||||||
|
AuthSource.INTERNAL,
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
for kwargs, expected_origin, expected_principal, expected_auth_source in cases:
|
||||||
|
context = agent_module.MoviePilotAgent(
|
||||||
|
session_id="origin-session",
|
||||||
|
user_id="user-1",
|
||||||
|
**kwargs,
|
||||||
|
)._build_policy_context()
|
||||||
|
|
||||||
|
assert context.origin is expected_origin
|
||||||
|
assert context.principal_type is expected_principal
|
||||||
|
assert context.auth_source is expected_auth_source
|
||||||
|
|
||||||
|
subagent_context = agent_module.MoviePilotAgent(
|
||||||
|
session_id="subagent-session",
|
||||||
|
user_id="user-1",
|
||||||
|
channel=MessageChannel.Telegram.value,
|
||||||
|
source="telegram",
|
||||||
|
)._build_policy_context().for_subagent()
|
||||||
|
assert subagent_context.origin is ToolOrigin.SUBAGENT
|
||||||
|
assert subagent_context.principal_type is PrincipalType.SUBAGENT
|
||||||
|
assert subagent_context.auth_source is AuthSource.INTERNAL
|
||||||
|
|
||||||
|
|
||||||
|
def test_host_middleware_observes_plugin_before_overridden_arun() -> None:
|
||||||
|
"""插件覆盖 `_arun` 时,宿主 middleware 仍必须先执行策略。"""
|
||||||
|
events: list[str] = []
|
||||||
|
tool = _OverrideArunTool(
|
||||||
|
events,
|
||||||
|
session_id="session-1",
|
||||||
|
user_id="user-1",
|
||||||
|
)
|
||||||
|
middleware = AgentPolicyMiddleware(context=_interactive_context())
|
||||||
|
request = SimpleNamespace(
|
||||||
|
tool=tool,
|
||||||
|
tool_call={
|
||||||
|
"id": "call-1",
|
||||||
|
"name": tool.name,
|
||||||
|
"args": {"query": "ok"},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
original_start = DEFAULT_TOOL_POLICY_ORCHESTRATOR.start
|
||||||
|
|
||||||
|
def _record_start(**kwargs):
|
||||||
|
events.append("policy")
|
||||||
|
return original_start(**kwargs)
|
||||||
|
|
||||||
|
async def _handler(_request):
|
||||||
|
result = await tool._arun(query="ok")
|
||||||
|
return ToolMessage(content=result, tool_call_id="call-1")
|
||||||
|
|
||||||
|
with patch.object(
|
||||||
|
DEFAULT_TOOL_POLICY_ORCHESTRATOR,
|
||||||
|
"start",
|
||||||
|
side_effect=_record_start,
|
||||||
|
):
|
||||||
|
result = asyncio.run(middleware.awrap_tool_call(request, _handler))
|
||||||
|
|
||||||
|
assert result.content == "ok"
|
||||||
|
assert events == ["policy", "tool"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("failed_phase", ["start", "finish"])
|
||||||
|
def test_middleware_observation_failure_does_not_replace_success(
|
||||||
|
failed_phase: str,
|
||||||
|
) -> None:
|
||||||
|
"""shadow start/finish 故障不能阻止 handler 或替换成功结果。"""
|
||||||
|
orchestrator = MagicMock()
|
||||||
|
orchestrator.start.return_value = object()
|
||||||
|
getattr(orchestrator, failed_phase).side_effect = RuntimeError(
|
||||||
|
f"policy-{failed_phase}-failure"
|
||||||
|
)
|
||||||
|
middleware = AgentPolicyMiddleware(
|
||||||
|
context=_interactive_context(),
|
||||||
|
orchestrator=orchestrator,
|
||||||
|
)
|
||||||
|
request = SimpleNamespace(
|
||||||
|
tool=_EchoTool(session_id="session-1", user_id="user-1"),
|
||||||
|
tool_call={"id": "call-1", "args": {"query": "same"}},
|
||||||
|
)
|
||||||
|
expected = ToolMessage(content="same", tool_call_id="call-1")
|
||||||
|
handler_called = False
|
||||||
|
|
||||||
|
async def _handler(_request):
|
||||||
|
nonlocal handler_called
|
||||||
|
handler_called = True
|
||||||
|
return expected
|
||||||
|
|
||||||
|
result = asyncio.run(middleware.awrap_tool_call(request, _handler))
|
||||||
|
|
||||||
|
assert handler_called is True
|
||||||
|
assert result is expected
|
||||||
|
|
||||||
|
|
||||||
|
def test_middleware_fail_observation_does_not_mask_tool_error() -> None:
|
||||||
|
"""shadow fail hook 故障后仍必须抛出原始工具异常。"""
|
||||||
|
orchestrator = MagicMock()
|
||||||
|
orchestrator.start.return_value = object()
|
||||||
|
orchestrator.fail.side_effect = RuntimeError("policy-fail-hook-failure")
|
||||||
|
middleware = AgentPolicyMiddleware(
|
||||||
|
context=_interactive_context(),
|
||||||
|
orchestrator=orchestrator,
|
||||||
|
)
|
||||||
|
request = SimpleNamespace(
|
||||||
|
tool=_EchoTool(session_id="session-1", user_id="user-1"),
|
||||||
|
tool_call={"id": "call-1", "args": {"query": "same"}},
|
||||||
|
)
|
||||||
|
tool_error = ValueError("original-tool-failure")
|
||||||
|
|
||||||
|
async def _handler(_request):
|
||||||
|
raise tool_error
|
||||||
|
|
||||||
|
with pytest.raises(ValueError) as error_info:
|
||||||
|
asyncio.run(middleware.awrap_tool_call(request, _handler))
|
||||||
|
|
||||||
|
assert error_info.value is tool_error
|
||||||
|
|
||||||
|
|
||||||
|
def test_policy_hook_failure_logs_only_stable_type_information() -> None:
|
||||||
|
"""fail-open 诊断只记录阶段和异常类型,不读取可能含凭据的异常文本。"""
|
||||||
|
mock_logger = MagicMock()
|
||||||
|
|
||||||
|
def _fail() -> None:
|
||||||
|
raise RuntimeError("DATABASE_PASSWORD=policy-secret-marker")
|
||||||
|
|
||||||
|
with patch("app.agent.policy.orchestrator.logger", mock_logger):
|
||||||
|
result = call_policy_hook("start", _fail)
|
||||||
|
|
||||||
|
logged = "\n".join(str(call) for call in mock_logger.warning.call_args_list)
|
||||||
|
assert result is None
|
||||||
|
assert "phase=start" in logged
|
||||||
|
assert "RuntimeError" in logged
|
||||||
|
assert "policy-secret-marker" not in logged
|
||||||
|
|
||||||
|
|
||||||
|
def test_policy_hook_hostile_error_type_is_fail_open() -> None:
|
||||||
|
"""异常类型名协议不可信时,观测故障仍不得逃出 fail-open 边界。"""
|
||||||
|
secret_marker = "hostile-policy-type-secret-6274"
|
||||||
|
|
||||||
|
class _HostileMeta(type):
|
||||||
|
def __getattribute__(cls, name):
|
||||||
|
if name == "__name__":
|
||||||
|
raise RuntimeError(f"DATABASE_PASSWORD={secret_marker}")
|
||||||
|
return super().__getattribute__(name)
|
||||||
|
|
||||||
|
class _HostilePolicyError(RuntimeError, metaclass=_HostileMeta):
|
||||||
|
pass
|
||||||
|
|
||||||
|
mock_logger = MagicMock()
|
||||||
|
|
||||||
|
def _fail() -> None:
|
||||||
|
raise _HostilePolicyError("visible policy failure")
|
||||||
|
|
||||||
|
escaped = False
|
||||||
|
with patch("app.agent.policy.orchestrator.logger", mock_logger):
|
||||||
|
try:
|
||||||
|
result = call_policy_hook("start", _fail)
|
||||||
|
except BaseException:
|
||||||
|
escaped = True
|
||||||
|
result = None
|
||||||
|
|
||||||
|
logged = "\n".join(str(call) for call in mock_logger.warning.call_args_list)
|
||||||
|
assert escaped is False
|
||||||
|
assert result is None
|
||||||
|
assert "phase=start" in logged
|
||||||
|
assert secret_marker not in logged
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("legacy_admin", "expected_result", "expected_run_count"),
|
||||||
|
[
|
||||||
|
(True, "same", 1),
|
||||||
|
(
|
||||||
|
False,
|
||||||
|
(
|
||||||
|
"抱歉,您没有执行此工具的权限。"
|
||||||
|
"只有渠道管理员或系统管理员才能执行工具操作。"
|
||||||
|
"如需执行工具,请联系管理员将您的用户ID添加到渠道管理员列表中(设定 -> 通知 -> 对应渠道配置 -> 管理员名单),"
|
||||||
|
"或联系系统管理员为您设置管理员权限。"
|
||||||
|
),
|
||||||
|
0,
|
||||||
|
),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_agent_admin_safe_read_keeps_legacy_authorization_authority(
|
||||||
|
legacy_admin: bool,
|
||||||
|
expected_result: str,
|
||||||
|
expected_run_count: int,
|
||||||
|
) -> None:
|
||||||
|
"""渠道管理员与普通用户结果保持旧门禁语义,policy 只做 shadow passthrough。"""
|
||||||
|
events: list[str] = []
|
||||||
|
observations = []
|
||||||
|
original_start = DEFAULT_TOOL_POLICY_ORCHESTRATOR.start
|
||||||
|
tool = _AdminSafeReadTool(
|
||||||
|
events,
|
||||||
|
session_id="session-1",
|
||||||
|
user_id="user-1",
|
||||||
|
)
|
||||||
|
tool.set_message_attr(
|
||||||
|
channel=MessageChannel.Telegram.value,
|
||||||
|
source="user",
|
||||||
|
username="member",
|
||||||
|
)
|
||||||
|
tool.set_agent_context({"is_admin": False})
|
||||||
|
middleware = AgentPolicyMiddleware(context=_interactive_context(is_admin=False))
|
||||||
|
request = SimpleNamespace(
|
||||||
|
tool=tool,
|
||||||
|
tool_call={"id": "call-1", "args": {"query": "same"}},
|
||||||
|
)
|
||||||
|
|
||||||
|
def _capture_start(**kwargs):
|
||||||
|
observation = original_start(**kwargs)
|
||||||
|
observations.append(observation)
|
||||||
|
return observation
|
||||||
|
|
||||||
|
async def _handler(_request):
|
||||||
|
content = await tool._arun(query="same")
|
||||||
|
return ToolMessage(content=content, tool_call_id="call-1")
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch.object(
|
||||||
|
_AdminSafeReadTool,
|
||||||
|
"is_admin_user",
|
||||||
|
new=AsyncMock(return_value=legacy_admin),
|
||||||
|
),
|
||||||
|
patch.object(
|
||||||
|
DEFAULT_TOOL_POLICY_ORCHESTRATOR,
|
||||||
|
"start",
|
||||||
|
side_effect=_capture_start,
|
||||||
|
),
|
||||||
|
):
|
||||||
|
result = asyncio.run(middleware.awrap_tool_call(request, _handler))
|
||||||
|
|
||||||
|
assert result.content == expected_result
|
||||||
|
assert events.count("run") == expected_run_count
|
||||||
|
assert len(observations) == 1
|
||||||
|
assert observations[0].policy.effect is ActionEffect.SAFE_READ
|
||||||
|
assert observations[0].policy.migration_state is MigrationState.LEGACY_SHADOW
|
||||||
|
assert observations[0].decision.allowed is True
|
||||||
|
assert observations[0].decision.shadow is True
|
||||||
|
assert observations[0].decision.reason_code == "legacy_shadow_allow"
|
||||||
|
|
||||||
|
|
||||||
|
def test_direct_non_admin_safe_read_rejects_before_policy_or_schema() -> None:
|
||||||
|
"""direct 未授权请求保持原 JSON 拒绝格式且不提前触发 policy/schema。"""
|
||||||
|
events: list[str] = []
|
||||||
|
orchestrator = MagicMock()
|
||||||
|
tool = _AdminSafeReadTool(
|
||||||
|
events,
|
||||||
|
session_id="session-1",
|
||||||
|
user_id="user-1",
|
||||||
|
)
|
||||||
|
manager = MoviePilotToolsManager(
|
||||||
|
is_admin=False,
|
||||||
|
policy_orchestrator=orchestrator,
|
||||||
|
)
|
||||||
|
manager.tools = [tool]
|
||||||
|
|
||||||
|
result = json.loads(
|
||||||
|
asyncio.run(manager.call_tool(tool.name, {"query": "same"}))
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result == {
|
||||||
|
"error": "抱歉,您没有执行此工具的权限。只有系统管理员才能执行工具操作。"
|
||||||
|
}
|
||||||
|
assert events == []
|
||||||
|
orchestrator.start.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("failed_phase", ["start", "finish"])
|
||||||
|
def test_direct_manager_observation_failure_does_not_replace_success(
|
||||||
|
failed_phase: str,
|
||||||
|
) -> None:
|
||||||
|
"""direct manager 的 shadow start/finish 故障不能改写真实返回值。"""
|
||||||
|
orchestrator = MagicMock()
|
||||||
|
orchestrator.start.return_value = object()
|
||||||
|
getattr(orchestrator, failed_phase).side_effect = RuntimeError(
|
||||||
|
f"policy-{failed_phase}-failure"
|
||||||
|
)
|
||||||
|
tool = _EchoTool(session_id="session-1", user_id="user-1")
|
||||||
|
manager = MoviePilotToolsManager(
|
||||||
|
is_admin=True,
|
||||||
|
policy_orchestrator=orchestrator,
|
||||||
|
)
|
||||||
|
manager.tools = [tool]
|
||||||
|
|
||||||
|
result = asyncio.run(manager.call_tool(tool.name, {"query": "same"}))
|
||||||
|
|
||||||
|
assert result == "same"
|
||||||
|
|
||||||
|
|
||||||
|
def test_direct_manager_fail_observation_does_not_mask_tool_error() -> None:
|
||||||
|
"""direct manager 的 fail hook 故障不能替换既有工具错误格式。"""
|
||||||
|
orchestrator = MagicMock()
|
||||||
|
orchestrator.start.return_value = object()
|
||||||
|
orchestrator.fail.side_effect = RuntimeError("policy-fail-hook-failure")
|
||||||
|
tool = _FailingTool(session_id="session-1", user_id="user-1")
|
||||||
|
manager = MoviePilotToolsManager(
|
||||||
|
is_admin=True,
|
||||||
|
policy_orchestrator=orchestrator,
|
||||||
|
)
|
||||||
|
manager.tools = [tool]
|
||||||
|
|
||||||
|
result = asyncio.run(manager.call_tool(tool.name, {"query": "same"}))
|
||||||
|
|
||||||
|
assert "ValueError" in result
|
||||||
|
assert "tool-error:same" in result
|
||||||
|
assert "policy-fail-hook-failure" not in result
|
||||||
|
|
||||||
|
|
||||||
|
def test_direct_manager_and_agent_middleware_share_policy_resolution() -> None:
|
||||||
|
"""相同工具参数在 Agent 与 direct 入口应获得相同动作策略。"""
|
||||||
|
observations = []
|
||||||
|
original_start = DEFAULT_TOOL_POLICY_ORCHESTRATOR.start
|
||||||
|
|
||||||
|
def _capture_start(**kwargs):
|
||||||
|
observation = original_start(**kwargs)
|
||||||
|
observations.append(observation)
|
||||||
|
return observation
|
||||||
|
|
||||||
|
tool = _EchoTool(session_id="session-1", user_id="user-1")
|
||||||
|
middleware = AgentPolicyMiddleware(context=_interactive_context())
|
||||||
|
request = SimpleNamespace(
|
||||||
|
tool=tool,
|
||||||
|
tool_call={
|
||||||
|
"id": "call-agent",
|
||||||
|
"name": tool.name,
|
||||||
|
"args": {"query": "same"},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _handler(_request):
|
||||||
|
return ToolMessage(content="same", tool_call_id="call-agent")
|
||||||
|
|
||||||
|
manager = MoviePilotToolsManager(
|
||||||
|
user_id="api-user",
|
||||||
|
session_id="api-session",
|
||||||
|
is_admin=True,
|
||||||
|
policy_orchestrator=DEFAULT_TOOL_POLICY_ORCHESTRATOR,
|
||||||
|
)
|
||||||
|
manager.tools = [tool]
|
||||||
|
|
||||||
|
async def _run_both():
|
||||||
|
await middleware.awrap_tool_call(request, _handler)
|
||||||
|
return await manager.call_tool(tool.name, {"query": "same"})
|
||||||
|
|
||||||
|
with patch.object(
|
||||||
|
DEFAULT_TOOL_POLICY_ORCHESTRATOR,
|
||||||
|
"start",
|
||||||
|
side_effect=_capture_start,
|
||||||
|
):
|
||||||
|
direct_result = asyncio.run(_run_both())
|
||||||
|
|
||||||
|
assert direct_result == "same"
|
||||||
|
assert [item.invocation.origin for item in observations] == [
|
||||||
|
ToolOrigin.AGENT_INTERACTIVE,
|
||||||
|
ToolOrigin.OPERATOR_DIRECT,
|
||||||
|
]
|
||||||
|
assert observations[0].policy == observations[1].policy
|
||||||
|
|
||||||
|
|
||||||
|
def test_agent_middleware_secret_setting_result_stays_out_of_receipt_logs() -> None:
|
||||||
|
"""Agent ToolNode 可接收管理员请求的原值,但策略回执不得记录该值。"""
|
||||||
|
secret_marker = "middleware-secret-setting-marker"
|
||||||
|
tool = QuerySystemSettingsTool(session_id="session-1", user_id="admin")
|
||||||
|
tool.set_agent_context({"is_admin": True})
|
||||||
|
middleware = AgentPolicyMiddleware(context=_interactive_context())
|
||||||
|
request = SimpleNamespace(
|
||||||
|
tool=tool,
|
||||||
|
tool_call={
|
||||||
|
"id": "call-secret-setting",
|
||||||
|
"name": tool.name,
|
||||||
|
"args": {"setting_key": "COOKIECLOUD_KEY", "show_secrets": True},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
mock_logger = MagicMock()
|
||||||
|
|
||||||
|
async def _handler(_request):
|
||||||
|
result = await tool._arun(
|
||||||
|
setting_key="COOKIECLOUD_KEY",
|
||||||
|
show_secrets=True,
|
||||||
|
)
|
||||||
|
return ToolMessage(content=result, tool_call_id="call-secret-setting")
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch.object(
|
||||||
|
QuerySystemSettingsTool,
|
||||||
|
"_load_setting_value",
|
||||||
|
return_value=secret_marker,
|
||||||
|
),
|
||||||
|
patch("app.agent.policy.orchestrator.logger", mock_logger),
|
||||||
|
):
|
||||||
|
result = asyncio.run(middleware.awrap_tool_call(request, _handler))
|
||||||
|
|
||||||
|
assert secret_marker in result.content
|
||||||
|
logged = "\n".join(
|
||||||
|
str(call)
|
||||||
|
for call in (
|
||||||
|
mock_logger.debug.call_args_list + mock_logger.info.call_args_list
|
||||||
|
)
|
||||||
|
)
|
||||||
|
assert secret_marker not in logged
|
||||||
|
assert '"value": "***"' in logged
|
||||||
|
assert '"value_preview": "***"' in logged
|
||||||
|
|
||||||
|
|
||||||
|
def test_main_agent_registers_policy_middleware_as_outermost() -> None:
|
||||||
|
"""主 Agent 必须把宿主策略中间件放在 middleware 链最外层。"""
|
||||||
|
agent = agent_module.MoviePilotAgent(session_id="session-1", user_id="user-1")
|
||||||
|
fake_llm = SimpleNamespace(
|
||||||
|
_llm_type="openai-chat",
|
||||||
|
model="fake",
|
||||||
|
profile={"max_input_tokens": 64000},
|
||||||
|
)
|
||||||
|
captured = {}
|
||||||
|
|
||||||
|
def _fake_create_agent(**kwargs):
|
||||||
|
captured.update(kwargs)
|
||||||
|
return object()
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch.object(agent, "_initialize_llm", return_value=fake_llm),
|
||||||
|
patch.object(agent, "_initialize_tools", return_value=[]),
|
||||||
|
patch.object(agent_module.prompt_manager, "get_agent_prompt", return_value="prompt"),
|
||||||
|
patch.object(agent_module, "create_subagent_middlewares", return_value=([], [])),
|
||||||
|
patch.object(agent_module, "create_agent", side_effect=_fake_create_agent),
|
||||||
|
patch.object(agent_module.settings, "LLM_MAX_TOOLS", 0),
|
||||||
|
):
|
||||||
|
asyncio.run(agent._create_agent(streaming=False))
|
||||||
|
|
||||||
|
assert isinstance(captured["middleware"][0], AgentPolicyMiddleware)
|
||||||
|
|
||||||
|
|
||||||
|
def test_main_agent_preserves_activity_log_middleware_order() -> None:
|
||||||
|
"""策略层加入后,ActivityLog 仍应位于 Memory 后、摘要前。"""
|
||||||
|
agent = agent_module.MoviePilotAgent(
|
||||||
|
session_id="session-1",
|
||||||
|
user_id="user-1",
|
||||||
|
channel=MessageChannel.WebAgent.value,
|
||||||
|
source="web-agent",
|
||||||
|
)
|
||||||
|
fake_llm = SimpleNamespace(
|
||||||
|
_llm_type="openai-chat",
|
||||||
|
model="fake",
|
||||||
|
profile={"max_input_tokens": 64000},
|
||||||
|
)
|
||||||
|
captured = {}
|
||||||
|
|
||||||
|
def _fake_create_agent(**kwargs):
|
||||||
|
captured.update(kwargs)
|
||||||
|
return object()
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch.object(agent, "_initialize_llm", return_value=fake_llm),
|
||||||
|
patch.object(agent, "_initialize_tools", return_value=[]),
|
||||||
|
patch.object(agent_module.prompt_manager, "get_agent_prompt", return_value="prompt"),
|
||||||
|
patch.object(agent_module, "create_subagent_middlewares", return_value=([], [])),
|
||||||
|
patch.object(agent_module, "create_agent", side_effect=_fake_create_agent),
|
||||||
|
patch.object(agent_module.settings, "LLM_MAX_TOOLS", 0),
|
||||||
|
):
|
||||||
|
asyncio.run(agent._create_agent(streaming=False))
|
||||||
|
|
||||||
|
middlewares = captured["middleware"]
|
||||||
|
policy_index = next(
|
||||||
|
index
|
||||||
|
for index, middleware in enumerate(middlewares)
|
||||||
|
if isinstance(middleware, AgentPolicyMiddleware)
|
||||||
|
)
|
||||||
|
memory_index = next(
|
||||||
|
index
|
||||||
|
for index, middleware in enumerate(middlewares)
|
||||||
|
if isinstance(middleware, MemoryMiddleware)
|
||||||
|
)
|
||||||
|
activity_index = next(
|
||||||
|
index
|
||||||
|
for index, middleware in enumerate(middlewares)
|
||||||
|
if isinstance(middleware, ActivityLogMiddleware)
|
||||||
|
)
|
||||||
|
summary_index = next(
|
||||||
|
index
|
||||||
|
for index, middleware in enumerate(middlewares)
|
||||||
|
if isinstance(middleware, SummarizationMiddleware)
|
||||||
|
)
|
||||||
|
|
||||||
|
assert policy_index == 0
|
||||||
|
assert activity_index == memory_index + 1
|
||||||
|
assert summary_index == activity_index + 1
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -149,6 +149,24 @@ class TestAgentToolStreaming:
|
|||||||
|
|
||||||
assert buffered_message == "处理中:\n\n(执行了 2 次搜索,读取了 2 个文件)\n\n继续分析"
|
assert buffered_message == "处理中:\n\n(执行了 2 次搜索,读取了 2 个文件)\n\n继续分析"
|
||||||
|
|
||||||
|
def test_tool_summary_does_not_retain_secret_arguments(self):
|
||||||
|
"""流式聚合状态只能保留递归脱敏后的工具参数。"""
|
||||||
|
secret_marker = "stream-secret-marker-1947"
|
||||||
|
handler = StreamingHandler()
|
||||||
|
|
||||||
|
handler.record_tool_call(
|
||||||
|
tool_name="execute_command",
|
||||||
|
tool_message=f"Authorization: Bearer {secret_marker}",
|
||||||
|
tool_kwargs={
|
||||||
|
"command": f"curl -H 'Authorization: Bearer {secret_marker}'",
|
||||||
|
"headers": {"X-API-Key": secret_marker},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
pending_state = str(handler._pending_tool_stats)
|
||||||
|
assert secret_marker not in pending_state
|
||||||
|
assert "***" in pending_state
|
||||||
|
|
||||||
def test_non_verbose_tool_summary_counts_subagents(self):
|
def test_non_verbose_tool_summary_counts_subagents(self):
|
||||||
"""校验非详细模式统计子代理调用次数。"""
|
"""校验非详细模式统计子代理调用次数。"""
|
||||||
async def _run():
|
async def _run():
|
||||||
|
|||||||
Reference in New Issue
Block a user