diff --git a/zhenxun/services/ai/capabilities/builtin.py b/zhenxun/services/ai/capabilities/builtin.py index b00ce96a..92a4276d 100644 --- a/zhenxun/services/ai/capabilities/builtin.py +++ b/zhenxun/services/ai/capabilities/builtin.py @@ -1,7 +1,8 @@ from __future__ import annotations +import hashlib import json -from typing import Any +from typing import Any, cast from zhenxun.models.user_console import UserConsole from zhenxun.services.ai.capabilities import ( @@ -9,18 +10,24 @@ from zhenxun.services.ai.capabilities import ( WrapModelRequestHandler, WrapToolExecuteHandler, ) +from zhenxun.services.ai.config import get_llm_config from zhenxun.services.ai.core.exceptions import ( + AbortException, + ControlFlowExit, GuardrailViolationError, LLMException, ModelRetry, ResponseParseException, SchemaParseError, ToolFatalError, + ToolFinishException, + ToolRetryError, UpstreamServerException, ) from zhenxun.services.ai.core.messages import ( ChatRequest, ChatResponse, + LLMMessage, ToolCallPart, ) from zhenxun.services.ai.core.models import LLMContext @@ -49,8 +56,6 @@ class StuckDetectionCapability(AbstractCapability): llm_context: LLMContext[ChatRequest, ChatResponse], handler: WrapModelRequestHandler, ) -> ChatResponse: - import hashlib - max_repeated_errors = 3 action_hashes = [] messages = list(llm_context.request.messages) @@ -131,13 +136,9 @@ class GlobalCycleLimitCapability(AbstractCapability): global_max = llm_context.request.extra.get("__global_max_cycles__") if global_max is None: - from zhenxun.services.ai.config import get_llm_config - global_max = get_llm_config().agent_settings.global_max_cycles if global_max is not None and global_cycles > global_max: - from zhenxun.services.ai.core.exceptions import AbortException - logger.error( "🚨 触发全局防护:整个流水线执行步数已达到全局上限 " f"({global_max}),强制熔断!" @@ -175,8 +176,6 @@ class PermissionCapability(AbstractCapability): f"🛡️ [Capability] 权限拦截: 用户 {user_id} 尝试调用 " f"{getattr(tool, 'name', 'unknown')}" ) - from zhenxun.services.ai.core.exceptions import ToolFatalError - raise ToolFatalError( msg, display_content=f"❌ 权限不足: 需要等级 {admin_level}" ) @@ -216,8 +215,6 @@ class BillingCapability(AbstractCapability): f"💰 [Capability] 金币拦截: 用户 {user_id} 尝试调用 " f"{getattr(tool, 'name', 'unknown')}" ) - from zhenxun.services.ai.core.exceptions import ToolFatalError - raise ToolFatalError( msg, display_content=f"❌ 余额不足: 需要 {cost_gold} 金币" ) @@ -241,12 +238,6 @@ class ToolRetryAndReflectionCapability(AbstractCapability): try: return await handler(arguments) except Exception as e: - from zhenxun.services.ai.core.exceptions import ( - AbortException, - ControlFlowExit, - ToolFatalError, - ToolFinishException, - ) from zhenxun.services.ai.tools.engine.executor import ToolExecutionPolicy from zhenxun.services.ai.tools.models import ToolResult @@ -257,8 +248,6 @@ class ToolRetryAndReflectionCapability(AbstractCapability): retries += 1 context.run.tool_retries[tool_name] = retries - from typing import cast - from zhenxun.services.ai.tools.core.tool import BaseTool tool = cast(BaseTool, context.call.current_tool) @@ -280,7 +269,7 @@ class ToolRetryAndReflectionCapability(AbstractCapability): class ReflexionCapability(AbstractCapability): """自愈反思与验证引擎 (Reflexion Engine)。 - 统一处理结构化解析失败 and 语义护栏拦截。""" + 统一处理结构化解析失败和语义护栏拦截。""" async def wrap_tool_execute(self, context, tool_name, arguments, handler): try: @@ -289,7 +278,6 @@ class ReflexionCapability(AbstractCapability): from zhenxun.services.ai.core.engine.structured_parser import ( DEFAULT_IVR_TEMPLATE, ) - from zhenxun.services.ai.core.exceptions import ModelRetry, ToolRetryError from zhenxun.services.ai.tools.models import ToolResult if isinstance(error, ToolRetryError | ModelRetry): @@ -301,6 +289,82 @@ class ReflexionCapability(AbstractCapability): ).as_error() raise error.with_traceback(None) from None + def _extract_error_info( + self, e: Exception, current_response_text: str + ) -> tuple[str, str, bool]: + """ + 统一提取异常报错详情与可恢复标记 + 返回 (error_msg, raw_response, is_recoverable) + """ + is_model_retry = isinstance(e, ModelRetry) + is_llm_error = isinstance(e, LLMException) + llm_error = cast(LLMException, e) if is_llm_error else None + + if ( + not is_model_retry + and llm_error + and not isinstance( + llm_error, ResponseParseException | UpstreamServerException + ) + ): + raise e + + if is_model_retry: + error_msg = getattr(e, "message", str(e)) + raw_response = current_response_text + else: + error_msg = ( + llm_error.details.get("validation_error", str(e)) + if llm_error + else str(e) + ) + raw_response = current_response_text or ( + llm_error.details.get("raw_response", "") if llm_error else "" + ) + + is_recoverable = getattr(llm_error, "recoverable", True) if llm_error else True + return error_msg, raw_response, is_recoverable + + def _generate_feedback_prompt( + self, e: Exception, error_msg: str, error_template: str | None + ) -> str: + """ + 根据不同的异常类型生成针对性的自愈反思提示词 (Feedback Prompt) + """ + if isinstance(e, SchemaParseError): + if "数据内容未通过规则校验" in error_msg: + return f"""### ⚠️ [数据内容校验失败] +你输出的 JSON 格式完全正确,但部分字段的内容未能通过业务规则约束。 + +**失败详情:** +> {error_msg} + +**修正要求:** 请仔细阅读上述失败详情,你必须打破先前的部分指令限制以满足上述规则,调整报错字段的值并重新输出。""" # noqa: E501 + else: + return f"""### ❌ [格式解析失败] +你输出的结构化数据(JSON)格式损坏或字段类型不匹配,未能通过校验。 + +**解析错误报告:** +> {error_msg} + +**修正要求:** 请仔细检查缺失的必填字段、错误的数据类型或未闭合的括号,严格参考你可用的 Schema 定义,重新输出正确格式的数据。""" # noqa: E501 + elif isinstance(e, GuardrailViolationError): + return f"""### 🛡️ [业务护栏违规] +你输出的数据格式完全正确,但在业务逻辑层触发了合规/风控护栏。 + +**拦截原因报告:** +> {error_msg} + +**修正要求:** 请结合上述反馈报告,反思你的决策逻辑或内容生成,在保持数据格式正确的前提下,重新生成符合护栏规范的内容。""" # noqa: E501 + else: + if error_template: + return error_template.format(error_msg=error_msg) + from zhenxun.services.ai.core.engine.structured_parser import ( + DEFAULT_IVR_TEMPLATE, + ) + + return DEFAULT_IVR_TEMPLATE.format(error_msg=error_msg) + async def wrap_model_request( self, context: RunContext, @@ -335,8 +399,6 @@ class ReflexionCapability(AbstractCapability): llm_context.request.messages, context ) - from typing import cast - response = await handler(llm_context) current_response_text = response.text @@ -354,8 +416,6 @@ class ReflexionCapability(AbstractCapability): resp_out, final_obj_out = await pipeline.run_output_pipeline( response, final_obj, context ) - from typing import cast - response = cast("ChatResponse", resp_out) final_obj = final_obj_out current_response_text = response.text @@ -364,42 +424,15 @@ class ReflexionCapability(AbstractCapability): return response except Exception as e: - from typing import cast - - from zhenxun.services.ai.core.messages import LLMMessage - - is_model_retry = isinstance(e, ModelRetry) - is_llm_error = isinstance(e, LLMException) - llm_error: LLMException | None = ( - cast(LLMException, e) if is_llm_error else None - ) last_exception = e - - if ( - not is_model_retry - and llm_error - and not isinstance( - llm_error, ResponseParseException | UpstreamServerException + try: + error_msg, raw_response, is_recoverable = self._extract_error_info( + e, current_response_text ) - ): - raise e + except Exception as fatal_e: + raise fatal_e.with_traceback(None) from None if attempt < max_retries: - if is_model_retry: - error_msg = getattr(e, "message", str(e)) - raw_response = current_response_text - else: - error_msg = ( - llm_error.details.get("validation_error", str(e)) - if llm_error - else str(e) - ) - raw_response = current_response_text or ( - llm_error.details.get("raw_response", "") - if llm_error - else "" - ) - logger.warning( "输出校验未通过 " f"(尝试 {attempt + 1}/{max_retries + 1})。" @@ -414,45 +447,16 @@ class ReflexionCapability(AbstractCapability): ) ) - if isinstance(e, SchemaParseError): - feedback_prompt = ( - "### ❌ [格式解析失败]\n" - "你输出的结构化数据(JSON)格式损坏或字段不匹配," - "未能通过 Schema 校验。\n\n" - "**解析错误报告:**\n" - f"> {error_msg}\n\n" - "**修正要求:** 请仔细检查缺失的必填字段、错误的数据类型或" - "未闭合的括号,严格参考你可用的工具 Schema 定义," - "重新输出正确格式的数据。" - ) - elif isinstance(e, GuardrailViolationError): - feedback_prompt = ( - "### 🛡️ [业务护栏违规]\n" - "你输出的数据格式完全正确,但在业务逻辑层触发了合规/风控护栏。\n\n" - "**拦截原因报告:**\n" - f"> {error_msg}\n\n" - "**修正要求:** 请结合上述反馈报告," - "反思你的决策逻辑或内容生成," - "在保持数据格式正确的前提下,重新生成符合护栏规范的内容。" - ) - else: - if output_processor and error_template: - feedback_prompt = error_template.format(error_msg=error_msg) - else: - from zhenxun.services.ai.core.engine import ( - structured_parser as sp, - ) - - feedback_prompt = sp.DEFAULT_IVR_TEMPLATE.format( - error_msg=error_msg - ) + feedback_prompt = self._generate_feedback_prompt( + e, error_msg, error_template + ) ivr_messages.append( cast(LLMMessage, LLMMessage.user(feedback_prompt)) ) continue - if llm_error and not getattr(llm_error, "recoverable", True): - raise llm_error.with_traceback(None) from None + if not is_recoverable: + raise last_exception.with_traceback(None) from None if last_exception: raise last_exception.with_traceback(None) from None diff --git a/zhenxun/services/ai/capabilities/wrappers.py b/zhenxun/services/ai/capabilities/wrappers.py index 7905def6..e8d27926 100644 --- a/zhenxun/services/ai/capabilities/wrappers.py +++ b/zhenxun/services/ai/capabilities/wrappers.py @@ -1,9 +1,12 @@ from __future__ import annotations from collections.abc import Callable +import copy import graphlib from typing import TYPE_CHECKING, Any +from nonebot.utils import is_coroutine_callable + from zhenxun.services.ai.core.messages import ChatRequest, ChatResponse from zhenxun.services.ai.core.options import GenerationConfig @@ -91,10 +94,11 @@ class CombinedCapability(AbstractCapability): """ 组合能力容器。 将多个 Capability 按顺序融合成一个复合的洋葱模型, - 处理生命周期的正序/倒序和链式调用。 + 处理生命周期的正序/倒序 and 链式调用。 """ def __init__(self, capabilities: list[AbstractCapability]): + """初始化组合能力,对传入能力集进行展平去重和拓扑排序""" flat = [] for c in capabilities: if isinstance(c, CombinedCapability): @@ -111,6 +115,7 @@ class CombinedCapability(AbstractCapability): self.capabilities = sort_capabilities(deduped) async def for_run(self, context: RunContext) -> "AbstractCapability": + """为当前运行实例解析并更新所包裹的能力列表""" new_caps = [] changed = False for cap in self.capabilities: @@ -126,6 +131,7 @@ class CombinedCapability(AbstractCapability): async def get_generation_config( self, context: RunContext ) -> GenerationConfig | None: + """获取组合中所有能力合并后的生成配置""" final_config = None for cap in self.capabilities: cap_config = await cap.get_generation_config(context) @@ -137,12 +143,14 @@ class CombinedCapability(AbstractCapability): return final_config async def get_system_prompts(self, context: RunContext) -> list[str]: + """获取组合中所有能力提供的系统提示词列表""" prompts = [] for cap in self.capabilities: prompts.extend(await cap.get_system_prompts(context)) return prompts async def get_tools(self, context: RunContext) -> list[Any]: + """获取组合中所有能力附带注册的工具列表""" tools = [] for cap in self.capabilities: tools.extend(await cap.get_tools(context)) @@ -151,6 +159,7 @@ class CombinedCapability(AbstractCapability): async def prepare_tools( self, context: RunContext, tool_defs: list[Any] ) -> list[Any]: + """依次调用组合中所有能力的准备工具钩子处理工具定义""" current_defs = list(tool_defs) for cap in self.capabilities: res = await cap.prepare_tools(context, current_defs) @@ -161,6 +170,7 @@ class CombinedCapability(AbstractCapability): async def wrap_run( self, context: RunContext, handler: WrapRunHandler ) -> "AgentRunResult[Any]": + """串联能力组合的 wrap_run 洋葱模型拦截器链""" chain = handler for cap in reversed(self.capabilities): chain = _make_wrap_link(cap, "wrap_run", context, {}, chain, None) @@ -172,6 +182,7 @@ class CombinedCapability(AbstractCapability): llm_context: LLMContext[ChatRequest, ChatResponse], handler: WrapModelRequestHandler, ) -> ChatResponse: + """串联能力组合的 wrap_model_request 洋葱模型拦截器链""" chain = handler for cap in reversed(self.capabilities): chain = _make_wrap_link( @@ -186,6 +197,7 @@ class CombinedCapability(AbstractCapability): args: str | dict[str, Any], handler: WrapToolValidateHandler, ) -> dict[str, Any]: + """串联能力组合的 wrap_tool_validate 洋葱模型拦截器链""" chain = handler for cap in reversed(self.capabilities): chain = _make_wrap_link( @@ -205,6 +217,7 @@ class CombinedCapability(AbstractCapability): arguments: dict[str, Any], handler: WrapToolExecuteHandler, ) -> Any: + """串联能力组合的 wrap_tool_execute 洋葱模型拦截器链""" chain = handler for cap in reversed(self.capabilities): chain = _make_wrap_link( @@ -250,15 +263,16 @@ class DynamicCapability(AbstractCapability): """动态能力注入:允许在运行时基于上下文生成真正的 Capability""" def __init__(self, capability_func: Callable): + """初始化动态能力""" self.capability_func = capability_func @classmethod def get_serialization_name(cls) -> str | None: + """获取反序列化标识""" return None async def for_run(self, context: RunContext) -> "AbstractCapability": - from nonebot.utils import is_coroutine_callable - + """在运行时基于当前上下文动态实例化并执行真正的 Capability""" if is_coroutine_callable(self.capability_func): cap = await self.capability_func(context) else: @@ -275,17 +289,19 @@ class WrapperCapability(AbstractCapability): """ def __init__(self, wrapped: AbstractCapability): + """初始化代理包装器""" self.wrapped = wrapped @classmethod def get_serialization_name(cls) -> str | None: + """获取反序列化标识""" return None async def for_run(self, context: RunContext) -> "AbstractCapability": + """对内部包裹的实例执行运行时解析并深度克隆""" new_wrapped = await self.wrapped.for_run(context) if new_wrapped is self.wrapped: return self - import copy new_self = copy.copy(self) new_self.wrapped = new_wrapped @@ -294,22 +310,27 @@ class WrapperCapability(AbstractCapability): async def get_generation_config( self, context: RunContext ) -> GenerationConfig | None: + """透传获取内部包裹实例的生成配置""" return await self.wrapped.get_generation_config(context) async def get_system_prompts(self, context: RunContext) -> list[str]: + """透传获取内部包裹实例的系统提示词""" return await self.wrapped.get_system_prompts(context) async def get_tools(self, context: RunContext) -> list[Any]: + """透传获取内部包裹实例的工具列表""" return await self.wrapped.get_tools(context) async def prepare_tools( self, context: RunContext, tool_defs: list[Any] ) -> list[Any]: + """透传执行内部包裹实例的准备工具钩子""" return await self.wrapped.prepare_tools(context, tool_defs) async def wrap_run( self, context: RunContext, handler: WrapRunHandler ) -> "AgentRunResult[Any]": + """透传执行内部包裹实例的 wrap_run 拦截器""" return await self.wrapped.wrap_run(context, handler) async def wrap_model_request( @@ -318,6 +339,7 @@ class WrapperCapability(AbstractCapability): llm_context: LLMContext[ChatRequest, ChatResponse], handler: WrapModelRequestHandler, ) -> ChatResponse: + """透传执行内部包裹实例的 wrap_model_request 拦截器""" return await self.wrapped.wrap_model_request(context, llm_context, handler) async def wrap_tool_validate( @@ -327,6 +349,7 @@ class WrapperCapability(AbstractCapability): args: str | dict[str, Any], handler: WrapToolValidateHandler, ) -> dict[str, Any]: + """透传执行内部包裹实例的 wrap_tool_validate 拦截器""" return await self.wrapped.wrap_tool_validate(context, tool_name, args, handler) async def wrap_tool_execute( @@ -336,6 +359,7 @@ class WrapperCapability(AbstractCapability): arguments: dict[str, Any], handler: WrapToolExecuteHandler, ) -> Any: + """透传执行内部包裹实例的 wrap_tool_execute 拦截器""" return await self.wrapped.wrap_tool_execute( context, tool_name, arguments, handler ) diff --git a/zhenxun/services/ai/context/memory/storage/backends.py b/zhenxun/services/ai/context/memory/storage/backends.py index 8211c5b5..ce81a79a 100644 --- a/zhenxun/services/ai/context/memory/storage/backends.py +++ b/zhenxun/services/ai/context/memory/storage/backends.py @@ -2,6 +2,7 @@ import asyncio import base64 from collections.abc import Callable import datetime +from pathlib import Path import time from typing import TYPE_CHECKING, Any, cast @@ -26,6 +27,7 @@ from zhenxun.services.ai.core.messages import ( LLMContentPart, LLMMessage, SystemMessage, + TextPart, ToolMessage, UserMessage, ) @@ -34,13 +36,47 @@ from zhenxun.services.db_context import Model from zhenxun.utils.pydantic_compat import TypeAdapter, model_dump +class AbstractMemoryRecord(Model): + """Tortoise ORM 短期记忆持久化基类 (Mixin)。""" + + id = fields.UUIDField(pk=True, description="主键") + session_id = fields.CharField(max_length=255, index=True) + role = fields.CharField(max_length=32) + content = fields.JSONField() + api_context = fields.JSONField(null=True) + created_at = fields.DatetimeField(auto_now_add=True) + metadata = fields.JSONField(null=True) + + class Meta: # type: ignore + abstract = True + + +class AbstractSlotRecord(Model): + """Tortoise ORM 记忆槽持久化基类 (Mixin)。""" + + id = fields.CharField( + pk=True, max_length=128, description="复合主键: session_id + label" + ) + session_id = fields.CharField(max_length=255, index=True) + label = fields.CharField(max_length=64, index=True) + content = fields.TextField() + size_limit = fields.IntField(default=2000) + pinned = fields.BooleanField(default=True) + scope = fields.CharField(max_length=255) + description = fields.CharField(max_length=255, default="") + created_at = fields.FloatField() + updated_at = fields.FloatField() + + class Meta: # type: ignore + abstract = True + + class DBMessageSerializer: """将 LLMMessage 与数据库 JSON 格式进行序列化/反序列化的帮助类""" @staticmethod def deserialize_content(content_raw: Any) -> list[LLMContentPart]: - from zhenxun.services.ai.core.messages import TextPart - + """反序列化数据库中的 JSON 数据为 LLMMessage 消息内容部件列表""" content_parts: list[LLMContentPart] = [] if isinstance(content_raw, list): adapter = TypeAdapter(LLMContentPart) @@ -59,8 +95,7 @@ class DBMessageSerializer: @staticmethod def serialize_content(content_payload: Any) -> list[dict[str, Any]]: - from pathlib import Path - + """将 LLMMessage 消息内容序列化为可存储于数据库的 JSON 格式""" if isinstance(content_payload, str): return [{"type": "text", "text": content_payload}] elif isinstance(content_payload, list): @@ -96,6 +131,7 @@ class MemoryScope: self, rag_client: "ScopedRAGClient", ): + """初始化长期记忆作用域与 RAG 客户端""" self.rag_client = rag_client self._background_tasks: set[Any] = set() @@ -164,13 +200,13 @@ class MemoryScope: async def forget( self, session: SessionMetadata, record_ids: list[str] | None = None ) -> int: + """从 RAG 向量数据库中删除指定的记忆记录""" return await self.rag_client.delete( record_ids=record_ids, ) async def _reinforce_memories(self, records: list[BaseRecord]): - import time - + """惰性强化记忆:更新被检索记忆的访问次数和最后访问时间""" now = time.time() for r in records: r.metadata["access_count"] = r.metadata.get("access_count", 0) + 1 @@ -179,15 +215,20 @@ class MemoryScope: class InMemoryChatContext(BaseChatContext): + """基于内存的聊天上下文存储后端""" + def __init__(self): + """初始化内存聊天上下文""" self._messages: dict[str, list[LLMMessage]] = {} async def get_messages(self, session: SessionMetadata) -> list[LLMMessage]: + """获取指定会话的所有短期历史消息""" return list(self._messages.get(session.session_id, [])) async def search( self, query: str, session: SessionMetadata, limit: int = 10 ) -> list[LLMMessage]: + """在内存中简单检索包含查询词的历史消息""" results = [] for msg in self._messages.get(session.session_id, []): if query in msg.extract_text: @@ -199,6 +240,7 @@ class InMemoryChatContext(BaseChatContext): async def add_messages( self, session: SessionMetadata, messages: list[LLMMessage] ) -> None: + """向指定会话中追加历史消息""" if session.session_id not in self._messages: self._messages[session.session_id] = [] self._messages[session.session_id].extend(messages) @@ -206,9 +248,11 @@ class InMemoryChatContext(BaseChatContext): async def set_messages( self, session: SessionMetadata, messages: list[LLMMessage] ) -> None: + """覆盖设置指定会话的历史消息""" self._messages[session.session_id] = list(messages) async def clear(self, session: SessionMetadata) -> None: + """清空指定会话的全部历史消息""" self._messages.pop(session.session_id, None) async def clear_by_query(self, query: ScopeSelector) -> None: @@ -221,22 +265,9 @@ class InMemoryChatContext(BaseChatContext): self._messages.pop(sid, None) -class AbstractMemoryRecord(Model): - """Tortoise ORM 短期记忆持久化基类 (Mixin)。""" - - id = fields.UUIDField(pk=True, description="主键") - session_id = fields.CharField(max_length=255, index=True) - role = fields.CharField(max_length=32) - content = fields.JSONField() - api_context = fields.JSONField(null=True) - created_at = fields.DatetimeField(auto_now_add=True) - metadata = fields.JSONField(null=True) - - class Meta: # type: ignore - abstract = True - - class TortoiseChatContext(BaseChatContext): + """基于 Tortoise ORM 的聊天上下文存储后端""" + def __init__( self, model_class: type[AbstractMemoryRecord], @@ -245,10 +276,12 @@ class TortoiseChatContext(BaseChatContext): ] | None = None, ): + """初始化 Tortoise ORM 聊天上下文存储后端""" self.model_class = model_class self.custom_save_hook = custom_save_hook def _row_to_message(self, row: AbstractMemoryRecord) -> LLMMessage: + """将数据库记录转换为 LLMMessage 实例""" content_parts = DBMessageSerializer.deserialize_content(row.content) metadata: dict[str, Any] | None = ( row.metadata if isinstance(row.metadata, dict) else None @@ -270,6 +303,7 @@ class TortoiseChatContext(BaseChatContext): return cast(LLMMessage, LLMMessage(role=role, **kwargs)) async def get_messages(self, session: SessionMetadata) -> list[LLMMessage]: + """从数据库中查询并获取指定会话的短期历史消息""" rows = ( await self.model_class.filter(session_id=session.session_id) .order_by("created_at") @@ -280,6 +314,7 @@ class TortoiseChatContext(BaseChatContext): async def search( self, query: str, session: SessionMetadata, limit: int = 10 ) -> list[LLMMessage]: + """在数据库中检索包含查询词的历史消息""" rows = ( await self.model_class.filter( session_id=session.session_id, content__icontains=query @@ -293,6 +328,7 @@ class TortoiseChatContext(BaseChatContext): async def add_messages( self, session: SessionMetadata, messages: list[LLMMessage] ) -> None: + """向数据库中批量追加指定会话的历史消息""" if not messages: return @@ -331,10 +367,12 @@ class TortoiseChatContext(BaseChatContext): async def set_messages( self, session: SessionMetadata, messages: list[LLMMessage] ) -> None: + """覆盖设置指定会话的数据库历史消息""" await self.clear(session) await self.add_messages(session, messages) async def clear(self, session: SessionMetadata) -> None: + """删除指定会话在数据库中的全部历史消息""" await self.model_class.filter(session_id=session.session_id).delete() async def clear_by_query(self, query: ScopeSelector) -> None: @@ -357,31 +395,15 @@ def get_orm_chat_context( ) -class AbstractSlotRecord(Model): - """Tortoise ORM 记忆槽持久化基类 (Mixin)。""" - - id = fields.CharField( - pk=True, max_length=128, description="复合主键: session_id + label" - ) - session_id = fields.CharField(max_length=255, index=True) - label = fields.CharField(max_length=64, index=True) - content = fields.TextField() - size_limit = fields.IntField(default=2000) - pinned = fields.BooleanField(default=True) - scope = fields.CharField(max_length=255) - description = fields.CharField(max_length=255, default="") - created_at = fields.FloatField() - updated_at = fields.FloatField() - - class Meta: # type: ignore - abstract = True - - class TortoiseSlotContext(BaseSlotContext): + """基于 Tortoise ORM 的记忆槽存储后端""" + def __init__(self, model_class: type[AbstractSlotRecord]): + """初始化 Tortoise ORM 记忆槽存储后端""" self.model_class = model_class def _row_to_slot(self, row: AbstractSlotRecord) -> MemorySlot: + """将数据库记忆槽记录转换为 MemorySlot 实例""" return MemorySlot( label=row.label, content=row.content, @@ -394,6 +416,7 @@ class TortoiseSlotContext(BaseSlotContext): ) async def get_slot(self, session: SessionMetadata, label: str) -> MemorySlot | None: + """查询并获取指定会话及作用域下 label 对应的记忆槽""" rows = await self.model_class.filter( session_id__in=session.accessible_scopes, label=label ).all() @@ -405,6 +428,7 @@ class TortoiseSlotContext(BaseSlotContext): return None async def set_slot(self, session: SessionMetadata, slot: MemorySlot) -> None: + """保存或更新指定会话的记忆槽到数据库""" composite_id = f"{slot.scope}_{slot.label}" await self.model_class.update_or_create( @@ -425,10 +449,12 @@ class TortoiseSlotContext(BaseSlotContext): async def delete_slot( self, session: SessionMetadata, label: str, scope: str ) -> None: + """从数据库中删除指定作用域和 label 的记忆槽""" composite_id = f"{scope}_{label}" await self.model_class.filter(id=composite_id).delete() async def list_pinned_slots(self, session: SessionMetadata) -> list[MemorySlot]: + """获取指定会话所有可访问的、置顶且非空的记忆槽""" rows = await self.model_class.filter( session_id__in=session.accessible_scopes, pinned=True ).all() @@ -442,6 +468,7 @@ class TortoiseSlotContext(BaseSlotContext): return [s for s in merged.values() if s.content.strip()] async def list_all_slots(self, session: SessionMetadata) -> list[MemorySlot]: + """获取指定会话所有可访问的记忆槽列表""" rows = await self.model_class.filter( session_id__in=session.accessible_scopes ).all() @@ -454,6 +481,7 @@ class TortoiseSlotContext(BaseSlotContext): return list(merged.values()) async def clear_by_query(self, query: ScopeSelector) -> None: + """批量清理匹配指定前缀的所有记忆槽""" scope_prefix = query.scope_prefix await self.model_class.filter(session_id__startswith=scope_prefix).delete() diff --git a/zhenxun/services/ai/context/rag/ingestion.py b/zhenxun/services/ai/context/rag/ingestion.py index f0e36a1a..5f4e96d6 100644 --- a/zhenxun/services/ai/context/rag/ingestion.py +++ b/zhenxun/services/ai/context/rag/ingestion.py @@ -8,11 +8,15 @@ from zhenxun.services.log import logger class ChunkingStrategy(ABC): + """分块策略抽象基类,用于将长文本记录切分为多个短的 BaseRecord""" + @abstractmethod def chunk(self, record: BaseRecord) -> list[BaseRecord]: + """将输入的记录切分为子记录列表。由子类具体实现""" raise NotImplementedError def clean_text(self, text: str) -> str: + """清洗和规范化文本,去除多余的空行和空白字符""" cleaned_text = re.sub(r"\n+", "\n", text) cleaned_text = re.sub(r"[ \t]+", " ", cleaned_text) return cleaned_text.strip() @@ -20,6 +24,7 @@ class ChunkingStrategy(ABC): def _create_chunk_record( self, original_record: BaseRecord, chunk_number: int, content: str ) -> BaseRecord: + """根据原始记录创建分块后的 BaseRecord,并自动附带切片索引和父级ID等元数据""" meta_data = original_record.metadata.copy() meta_data["chunk_index"] = chunk_number meta_data["chunk_size"] = len(content) @@ -44,6 +49,7 @@ class DocumentChunking(ChunkingStrategy): self.chunk_size = chunk_size def chunk(self, record: BaseRecord) -> list[BaseRecord]: + """按双换行将内容切分为段落,并将相邻段落合并为符合最大字符长度限制的分块""" if len(record.content) <= self.chunk_size: return [ self._create_chunk_record(record, 0, self.clean_text(record.content)) @@ -113,7 +119,7 @@ class RecursiveCharacterChunking(ChunkingStrategy): ] def _split_text(self, text: str, separators: list[str]) -> list[str]: - """核心递归切分逻辑""" + """核心递归切分逻辑。尝试使用给定的分隔符列表按优先级切分文本""" final_chunks = [] separator = separators[-1] new_separators = [] @@ -187,6 +193,7 @@ class RecursiveCharacterChunking(ChunkingStrategy): return chunks def chunk(self, record: BaseRecord) -> list[BaseRecord]: + """使用递归字符切分方式,将文本切分为带有重叠部分的分块记录""" content = record.content.strip() if len(content) <= self.chunk_size: @@ -219,6 +226,7 @@ class RowChunking(ChunkingStrategy): self.rows_per_chunk = rows_per_chunk def chunk(self, record: BaseRecord) -> list[BaseRecord]: + """将表格行数据按指定行数切分为块,每个块都带有相同的表头首部""" lines = record.content.splitlines() lines = [line for line in lines if line.strip()] @@ -260,6 +268,7 @@ class DeduplicationProcessor: self.threshold = threshold async def process(self, records: list[BaseRecord]) -> list[BaseRecord]: + """对输入记录列表进行批处理内去重,过滤相似度达到或超过阈值的重复记录""" if not records or len(records) <= 1: return records @@ -297,7 +306,9 @@ class BaseBatchNode(ABC): """批处理节点基类:一次性接收并处理全部记录""" @abstractmethod - async def process_batch(self, records: list[BaseRecord]) -> list[BaseRecord]: ... + async def process_batch(self, records: list[BaseRecord]) -> list[BaseRecord]: + """批量处理 BaseRecord 记录列表""" + ... class BaseMapNode(ABC): @@ -306,7 +317,9 @@ class BaseMapNode(ABC): @abstractmethod async def process_one( self, record: BaseRecord - ) -> BaseRecord | list[BaseRecord] | None: ... + ) -> BaseRecord | list[BaseRecord] | None: + """处理单条 BaseRecord 记录,可返回修改后的记录、拆分后的多条记录,或 None(表示过滤该记录)""" # noqa: E501 + ... class DynamicChunkingNode(BaseBatchNode): @@ -329,6 +342,7 @@ class DynamicChunkingNode(BaseBatchNode): self.strategies = custom_strategies or {".csv": RowChunking(rows_per_chunk=30)} async def process_batch(self, records: list[BaseRecord]) -> list[BaseRecord]: + """根据记录的元数据扩展名动态匹配并执行切块策略""" chunks = [] for record in records: ext = record.metadata.get("extension", "") @@ -338,7 +352,7 @@ class DynamicChunkingNode(BaseBatchNode): class BaseEmbeddingBatchNode(BaseBatchNode): - """批量向量化抽象基类:提取文本、分批请求 API 并将结果写回的公共逻辑""" + """批量向量化抽象基类:提取文本、分批请求 API 并将结果 write 回的公共逻辑""" def __init__(self, embedder, batch_size: int = 80): """ @@ -357,6 +371,7 @@ class BaseEmbeddingBatchNode(BaseBatchNode): pass async def process_batch(self, records: list[BaseRecord]) -> list[BaseRecord]: + """批量将文本记录提取并请求向量化接口,写回向量数据""" if not self.embedder or not records: return records @@ -382,14 +397,15 @@ class BaseEmbeddingBatchNode(BaseBatchNode): class EmbeddingNode(BaseEmbeddingBatchNode): - """并发向量化初次构建节点""" + """并发向量化初次构建节点,只对有实际内容的记录进行向量化。""" def _filter_target_records(self, records: list[BaseRecord]) -> list[BaseRecord]: + """筛选出非空内容的记录进行向量化""" return [r for r in records if r.content.strip()] class DedupNode(BaseBatchNode): - """批次内查重节点""" + """批次内查重节点。在流水线中作为去重节点使用。""" def __init__(self, threshold: float): """ @@ -401,6 +417,7 @@ class DedupNode(BaseBatchNode): self.processor = DeduplicationProcessor(threshold=threshold) async def process_batch(self, records: list[BaseRecord]) -> list[BaseRecord]: + """调用去重处理器过滤本批次中的高度重复记录""" return await self.processor.process(records) @@ -417,6 +434,7 @@ class StorageCommitNode(BaseBatchNode): self.storage = storage async def process_batch(self, records: list[BaseRecord]) -> list[BaseRecord]: + """根据记录的操作类型(保存、更新或删除)分类,并批量提交至存储后端""" if not records: return records @@ -468,9 +486,11 @@ class IndexPipeline: self.max_workers = max_workers def add_node(self, node: BaseBatchNode | BaseMapNode): + """向处理流水线中追加一个节点""" self.nodes.append(node) async def run(self, records: list[BaseRecord]) -> list[BaseRecord]: + """并发调度处理引擎,运行并执行流水线中的所有处理节点,返回处理后的记录""" if not records: return [] diff --git a/zhenxun/services/ai/core/engine/context_renderer.py b/zhenxun/services/ai/core/engine/context_renderer.py index aa17e2ec..87304f0a 100644 --- a/zhenxun/services/ai/core/engine/context_renderer.py +++ b/zhenxun/services/ai/core/engine/context_renderer.py @@ -16,6 +16,16 @@ class ContextConverter: def flatten_to_llm_messages( messages: Sequence[AgentMessage], context: Any | None = None ) -> list[LLMMessage]: + """ + 将包含 AgentEvent 和 LLMMessage 的混合消息序列,扁平化转换为大模型 API 专用的 LLMMessage 列表。 + + 参数: + messages: 混合了 LLMMessage 和 AgentEvent 对象的原始业务消息序列。 + context: 渲染 AgentEvent 所需的运行时上下文对象(如 RunContext)。 + + 返回: + list[LLMMessage]: 扁平化转换后生成的纯底层大模型原生消息列表。 + """ # noqa: E501 flattened: list[LLMMessage] = [] for msg in messages: diff --git a/zhenxun/services/ai/core/engine/structured_parser.py b/zhenxun/services/ai/core/engine/structured_parser.py index d33434f0..21997fce 100644 --- a/zhenxun/services/ai/core/engine/structured_parser.py +++ b/zhenxun/services/ai/core/engine/structured_parser.py @@ -115,10 +115,21 @@ class BaseOutputProcessor(Generic[OutputDataT]): repaired_obj = json_repair.loads(text, skip_json_loads=True) return model_validate(self.target_model, repaired_obj) except Exception as repair_error: - logger.error( - f"LLM结构化输出校验最终失败: {repair_error}", - e=repair_error, + logger.debug( + "JSON修复或模型校验失败,将交由大模型进行反思自愈: " + f"{type(repair_error).__name__}" ) + if isinstance(repair_error, ValidationError): + error_msgs = [] + for err in repair_error.errors(): + loc = ".".join(str(x) for x in err["loc"]) or "root" + msg = err.get("msg", "") + error_msgs.append(f"字段 `{loc}`: {msg}") + clean_error_str = "\n".join(error_msgs) + raise SchemaParseError( + f"数据内容未通过规则校验:\n{clean_error_str}" + ) + raise SchemaParseError( f"JSON格式损坏或字段不匹配,未能通过Schema验证: {repair_error}" ) diff --git a/zhenxun/services/ai/core/exceptions.py b/zhenxun/services/ai/core/exceptions.py index fdc326eb..982e03e1 100644 --- a/zhenxun/services/ai/core/exceptions.py +++ b/zhenxun/services/ai/core/exceptions.py @@ -9,6 +9,12 @@ class ModelRetry(Exception): """用于通知大模型修正并重试的异常""" def __init__(self, message: str): + """ + 初始化用于通知大模型重试的异常。 + + 参数: + message: 用于提示大模型的具体重试和自我纠错信息。 + """ self.message = message super().__init__(message) @@ -17,6 +23,12 @@ class SchemaParseError(ModelRetry): """格式解析异常。当大模型返回的 JSON 损坏或不符合 Schema 时抛出。""" def __init__(self, message: str): + """ + 初始化 Schema 格式解析错误异常。 + + 参数: + message: 详细的 JSON 解析失败或 Schema 校验报错信息。 + """ super().__init__(message) @@ -24,6 +36,12 @@ class GuardrailViolationError(ModelRetry): """护栏违规异常。当大模型返回的数据格式正确,但违反业务规则时抛出。""" def __init__(self, message: str): + """ + 初始化安全护栏校验未通过的异常。 + + 参数: + message: 触发业务护栏违规拦截的详细原因说明。 + """ super().__init__(message) @@ -41,6 +59,13 @@ class ToolFatalError(ControlFlowExit): """ def __init__(self, message: str, display_content: str | None = None): + """ + 初始化不可恢复的致命工具执行异常。 + + 参数: + message: 供大模型及系统调试日志记录的底层致命错误详情。 + display_content: 直接向终端用户呈现的友好拦截文案。 + """ self.message = message self.display_content = display_content or f"❌ 工具遇到致命错误: {message}" super().__init__(self.message) @@ -50,6 +75,14 @@ class GuardrailFatalException(ControlFlowExit): """护栏致命拦截异常 (触发 ABORT/REJECT 时抛出)""" def __init__(self, guard_name: str, reason: str, display: str | None = None): + """ + 初始化护栏强制拦截中断异常。 + + 参数: + guard_name: 拦截本次执行的安全护栏规则名称。 + reason: 拦截或拒绝的底层业务决策详情。 + display: 直接反馈给用户的风控友好提示消息。 + """ self.guard_name = guard_name self.reason = reason self.display = display or f"🛡️ 安全拦截: {reason}" @@ -64,6 +97,12 @@ class ToolRetryError(Exception): """ def __init__(self, message: str): + """ + 初始化触发大模型反思与自愈的可恢复工具错误。 + + 参数: + message: 会被传递给大模型用于进行 Reflexion 的错误反馈 Prompt。 + """ self.message = message super().__init__(self.message) @@ -76,6 +115,13 @@ class ToolFinishException(ToolFatalError): """ def __init__(self, message: str, display_content: str | None = None): + """ + 初始化用于中断大模型思考循环并返回结果的结束异常。 + + 参数: + message: 内部记录的中断异常信息. + display_content: 中断执行流后,向用户展现的最终文本。 + """ super().__init__(message, display_content) @@ -83,6 +129,13 @@ class AbortException(ControlFlowExit): """异常中止当前 Agent 思考流。""" def __init__(self, reason: str, display: Any = None): + """ + 初始化用于强制中止 Agent 推理执行流的异常。 + + 参数: + reason: 触发强制中断的技术或业务原因。 + display: 中断后向用户展示的显示结果。 + """ self.reason = reason self.display = display super().__init__(f"Aborted: {reason}") @@ -96,6 +149,13 @@ class InterventionHandledException(ControlFlowExit): """ def __init__(self, message: str, display_content: str | None = None): + """ + 初始化干预处理成功以安全熔断生命周期的异常。 + + 参数: + message: 内部调试与审计的干预详情描述。 + display_content: 向发起干预的用户端展现的进度提醒提示。 + """ self.message = message self.display_content = display_content super().__init__(self.message) @@ -105,6 +165,13 @@ class ConcurrencyRejectException(ControlFlowExit): """并发拒绝异常。当 Agent 设置为 REJECT 且正在忙碌时抛出。""" def __init__(self, message: str, display: Any = None): + """ + 初始化并发调度拒绝接收新任务的异常。 + + 参数: + message: 系统内部拦截的并发冲突详细说明。 + display: 提示给并发用户的友好限流排队通知。 + """ self.message = message self.display = display or "⏳ 智能体正在处理您的上一个请求,请稍后再试~" super().__init__(message) @@ -114,6 +181,12 @@ class ConcurrencyInterruptException(ControlFlowExit): """并发打断异常。当 Agent 设置为 INTERRUPT 且被新请求打断时抛出。""" def __init__(self, message: str): + """ + 初始化并发抢占执行被打断的异常。 + + 参数: + message: 系统内部调度器生成的抢占与接管日志描述。 + """ self.message = message super().__init__(message) @@ -127,6 +200,14 @@ class NeedsInputException(Exception): def __init__( self, missing_field: str, missing_description: str, original_kwargs: dict ): + """ + 初始化 HITL 人机交互表单输入暂停请求的异常。 + + 参数: + missing_field: 缺失的必填参数字段名。 + missing_description: 字段的提示描述(通常由 Field 描述提取)。 + original_kwargs: 抛出异常前工具已成功收集的其它参数字典。 + """ self.missing_field = missing_field self.missing_description = missing_description self.original_kwargs = original_kwargs @@ -142,6 +223,13 @@ class NeedsAuthException(Exception): """ def __init__(self, provider: str, message: str): + """ + 初始化因凭证失效需重新发起用户鉴权挂起的异常。 + + 参数: + provider: 需要发起授权验证的外部 OAuth/API 服务商标识。 + message: 授权校验失败的诊断描述。 + """ self.provider = provider self.message = message super().__init__(f"Needs auth for: {provider} - {message}") @@ -151,6 +239,14 @@ class SandboxPathEscapeError(Exception): """当沙箱内的路径解析结果试图逃逸出允许的工作区根目录时抛出""" def __init__(self, path: str, resolved_path: str | None = None, reason: str = ""): + """ + 初始化路径安全越界逃逸拦截异常。 + + 参数: + path: 引起逃逸嫌疑的原始路径参数。 + resolved_path: 物理求值后的解析路径(如果有)。 + reason: 触发路径校验失败的底层判决依据。 + """ self.path = path self.resolved_path = resolved_path self.reason = reason @@ -166,6 +262,14 @@ class WorkspaceIOError(Exception): """沙箱文件系统读写操作失败""" def __init__(self, path: str, message: str, cause: Exception | None = None): + """ + 初始化沙箱文件系统底层读写操作失败的 IO 异常。 + + 参数: + path: 读写发生故障的物理或沙箱逻辑路径。 + message: 底层 IO 操作报错原因详细说明。 + cause: 触发该 IO 错误的根源 Python 底层 Exception 实例。 + """ self.path = path self.cause = cause super().__init__(f"沙箱 IO 异常 [{path}]: {message}") @@ -175,6 +279,13 @@ class SandboxFatalError(ToolFatalError): """沙箱底层容器发生致命崩溃(如 OOM, 被宿主机强杀等)""" def __init__(self, message: str, display_content: str | None = None): + """ + 初始化沙箱执行容器严重失联或崩溃的致命异常。 + + 参数: + message: 容器底层抛出的系统异常详情或心跳超时诊断。 + display_content: 向终端用户反馈的系统故障提醒。 + """ display = display_content or f"❌ 沙箱不可用: {message}" super().__init__(message, display_content=display) @@ -188,6 +299,14 @@ class LLMException(Exception): details: dict[str, Any] | None = None, cause: Exception | None = None, ): + """ + 初始化底层大模型 API 调用及服务异常。 + + 参数: + message: 通用的调用错误或失败总结说明。 + details: 包含接口名、重试指示、服务端回传原始信息的字典。 + cause: 触发此错误的根源协议请求异常实例。 + """ self.message = message self.details = details or {} self.cause = cause diff --git a/zhenxun/services/ai/core/options.py b/zhenxun/services/ai/core/options.py index b7f9ae85..e465b4db 100644 --- a/zhenxun/services/ai/core/options.py +++ b/zhenxun/services/ai/core/options.py @@ -298,9 +298,13 @@ class TTSConfig(BaseModel): """语速 (通用映射)""" openai_options: OpenAITTSOptions = Field(default_factory=OpenAITTSOptions) + """OpenAI 厂商专属请求参数集""" gemini_options: GeminiTTSOptions = Field(default_factory=GeminiTTSOptions) + """Gemini 厂商专属请求参数集""" minimax_options: MiniMaxTTSOptions = Field(default_factory=MiniMaxTTSOptions) + """MiniMax 厂商专属请求参数集""" mimo_options: MiMoTTSOptions = Field(default_factory=MiMoTTSOptions) + """MiMo 厂商专属请求参数集""" custom_kwargs: dict[str, Any] = Field(default_factory=dict) """兜底逃生舱,包含的键值对将直接透传至顶层请求体中 (可用于缓存 TTL)""" @@ -315,13 +319,20 @@ class GenerationConfig(BaseModel): """ common: CommonLLMConfig = Field(default_factory=CommonLLMConfig) + """大模型生成核心通用配置项""" output: OutputFormatConfig = Field(default_factory=OutputFormatConfig) + """输出格式与约束控制配置项""" tools: ToolCallConfig = Field(default_factory=ToolCallConfig) + """工具调用策略与函数声明配置项""" media: MediaGenerationConfig = Field(default_factory=MediaGenerationConfig) + """多媒体生成相关配置项""" openai_options: OpenAIOptions = Field(default_factory=OpenAIOptions) + """OpenAI 厂商专属请求参数集""" gemini_options: GeminiOptions = Field(default_factory=GeminiOptions) + """Gemini 厂商专属请求参数集""" deepseek_options: DeepSeekOptions = Field(default_factory=DeepSeekOptions) + """DeepSeek 厂商专属请求参数集""" enable_caching: bool | None = Field(default=None) """是否在此次生成中开启上下文缓存 (Context Caching)""" diff --git a/zhenxun/services/ai/flow/agent/agent.py b/zhenxun/services/ai/flow/agent/agent.py index 08ccf574..fcabfc94 100644 --- a/zhenxun/services/ai/flow/agent/agent.py +++ b/zhenxun/services/ai/flow/agent/agent.py @@ -6,10 +6,16 @@ from typing import Any, Generic, cast from zhenxun.services.ai.capabilities import ( AbstractCapability, + CombinedCapability, DynamicCapability, ) +from zhenxun.services.ai.config import get_llm_config from zhenxun.services.ai.context.knowledge.base import BaseKnowledge from zhenxun.services.ai.context.memory.builder import MemoryBuilder +from zhenxun.services.ai.context.memory.capabilities import ( + AgenticMemoryCapability, + SlotMemoryCapability, +) from zhenxun.services.ai.context.memory.models import MemoryConfig from zhenxun.services.ai.core.exceptions import ( ConcurrencyInterruptException, @@ -28,20 +34,15 @@ from zhenxun.services.ai.core.options import ( from zhenxun.services.ai.core.protocols.tool import ToolExecutable, ToolResolvable from zhenxun.services.ai.core.stream_events import AgentStreamEvent, EventBus from zhenxun.services.ai.core.templates import PromptTemplate -from zhenxun.services.ai.flow.agent.engine.builders import ToolBuilder -from zhenxun.services.ai.flow.agent.models import ( - AgentConfig, - AgentRunResources, - AgentState, - Persona, -) -from zhenxun.services.ai.flow.base import BaseRunnable -from zhenxun.services.ai.guardrails import GuardrailSource +from zhenxun.services.ai.flow.base import BaseRunnable, ConcurrencyPolicy +from zhenxun.services.ai.flow.concurrency import apply_concurrency_policy +from zhenxun.services.ai.guardrails import GuardrailSource, parse_guardrails from zhenxun.services.ai.llm.builder import IntentBuilder +from zhenxun.services.ai.message_builder import MessageBuilder from zhenxun.services.ai.run import ( AgentRunResult, + AgentTask, RunContext, - Task, ) from zhenxun.services.ai.run.context import AgentDepsT from zhenxun.services.ai.run.di import DependencyInjector @@ -52,10 +53,20 @@ from zhenxun.services.ai.run.models import ( OutputDataT, StreamedRunResult, ) -from zhenxun.services.ai.tools.core.tool import BaseTool +from zhenxun.services.ai.run.subscribers import ( + DefaultUISubscriber, + TelemetrySubscriber, +) +from zhenxun.services.ai.tools.bridges.delegate import DelegateTool +from zhenxun.services.ai.tools.core.tool import BaseTool, FunctionTool from zhenxun.services.ai.tools.core.toolkit import BaseToolkit -from zhenxun.services.ai.tools.models import Query +from zhenxun.services.ai.tools.models import Query, ToolOptions +from zhenxun.services.ai.tools.providers.builtin.hitl import HITLToolkit +from zhenxun.services.ai.tools.providers.skills.capabilities import ( + SkillCapability, +) from zhenxun.services.ai.tools.providers.skills.models import Skill, SkillSource +from zhenxun.services.ai.utils import ContextUtils from zhenxun.services.log import logger from zhenxun.utils.pydantic_compat import ( model_construct, @@ -65,7 +76,18 @@ from zhenxun.utils.pydantic_compat import ( ) from zhenxun.utils.utils import infer_plugin_namespace -from .engine.executor import BaseAgentExecutor +from .engine.builders import ( + AgentProfileResolver, + ContextBuilder, + ToolBuilder, +) +from .engine.executor import BaseAgentExecutor, StandardAgentExecutor +from .models import ( + AgentConfig, + AgentRunResources, + AgentState, + Persona, +) ToolSource = ( Callable | BaseTool | dict[str, Any] | str | BaseToolkit | ToolResolvable | Query @@ -413,8 +435,6 @@ class Agent( self.tool_filters = [] self.toolset_funcs = [] self._event_listeners: dict[type[AgentStreamEvent], list[Callable]] = {} - from zhenxun.services.ai.guardrails import parse_guardrails - self._guardrails = parse_guardrails(guardrails) self.memory_config = MemoryBuilder.resolve(memory) @@ -428,8 +448,6 @@ class Agent( self.engine_config = self.config if self.config.enable_hitl is None: - from zhenxun.services.ai.config import get_llm_config - self.config.enable_hitl = get_llm_config().agent_settings.enable_hitl self.config.stateless = not self.memory_config.short_term.enable @@ -450,19 +468,11 @@ class Agent( self.capabilities: list[AbstractCapability] = [] if self.memory_config.long_term.enable and self.memory_config.long_term.agentic: - from zhenxun.services.ai.context.memory.capabilities import ( - AgenticMemoryCapability, - ) - self.capabilities.append( AgenticMemoryCapability(self.memory_config, self.namespace) ) if self.memory_config.slots.enable: - from zhenxun.services.ai.context.memory.capabilities import ( - SlotMemoryCapability, - ) - self.capabilities.append( SlotMemoryCapability(self.memory_config, self.namespace) ) @@ -475,15 +485,9 @@ class Agent( self.capabilities.append(DynamicCapability(cap)) if self.config.enable_hitl: - from zhenxun.services.ai.tools.providers.builtin.hitl import HITLToolkit - self.tool_definitions.append(HITLToolkit()) if skills: - from zhenxun.services.ai.tools.providers.skills.capabilities import ( - SkillCapability, - ) - self.capabilities.append( SkillCapability(skills=skills, namespace=self.namespace) ) @@ -502,9 +506,6 @@ class Agent( """ def decorator(f: Callable): - from zhenxun.services.ai.tools.core.tool import FunctionTool - from zhenxun.services.ai.tools.models import ToolOptions - tool_name = name or f.__name__ tool_desc = description or f.__doc__ or "未提供描述" base_settings = settings or getattr(f, "__tool_settings__", ToolOptions()) @@ -566,15 +567,11 @@ class Agent( if func is None: def decorator(f: Callable): - from zhenxun.services.ai.guardrails import parse_guardrails - self._guardrails.extend(parse_guardrails([f])) return f return decorator else: - from zhenxun.services.ai.guardrails import parse_guardrails - self._guardrails.extend(parse_guardrails([func])) return func @@ -594,13 +591,11 @@ class Agent( async def __resolve_to_tools__(self) -> list[ToolExecutable]: """协议支持:将自身 Agent 转化为可被上级调用的工具""" - from zhenxun.services.ai.tools.bridges.delegate import DelegateTool - return [DelegateTool(self)] async def run( self, - prompt: PromptInput | Task | None = None, + prompt: PromptInput | AgentTask | None = None, *, config: AgentConfig | dict | None = None, deps: AgentDepsT | None = None, @@ -611,7 +606,7 @@ class Agent( 智能体单次运行阻塞核心入口,内部使用上下文管理器静默消费事件流直至执行结束。 参数: - prompt: 用户输入的消息内容或标准数据契约任务对象 (Task)。 + prompt: 用户输入的消息内容或标准数据契约任务对象 (AgentTask)。 deps: 强类型的外部依赖注入对象 (例如 NoneBot 的 Bot, Event)。 context: 显式传入的运行时与会话上下文 (RunContext)。 config: 单次运行时的动态配置覆盖字典或对象。 @@ -628,7 +623,7 @@ class Agent( @contextlib.asynccontextmanager async def run_stream( self, - prompt: PromptInput | Task | None = None, + prompt: PromptInput | AgentTask | None = None, *, config: AgentConfig | dict | None = None, deps: AgentDepsT | None = None, @@ -648,10 +643,6 @@ class Agent( effective_config = self.config.merge_with(override_conf) if effective_config.skills: - from zhenxun.services.ai.tools.providers.skills.capabilities import ( - SkillCapability, - ) - if effective_config.capabilities is None: effective_config.capabilities = [] effective_config.capabilities.append( @@ -661,11 +652,6 @@ class Agent( ) bus = event_bus or EventBus() - from zhenxun.services.ai.run.subscribers import ( - DefaultUISubscriber, - TelemetrySubscriber, - ) - TelemetrySubscriber().attach(bus) if self._event_listeners: @@ -674,8 +660,6 @@ class Agent( def _make_handler(callback_func: Callable) -> Callable: async def _di_handler(event: AgentStreamEvent): - from zhenxun.services.ai.run.di import DependencyInjector - await DependencyInjector.invoke( callback_func, {"stream_event": event}, safe_context ) @@ -700,8 +684,6 @@ class Agent( policy = getattr(self.config, "concurrency_policy", None) if policy is None: - from zhenxun.services.ai.flow.base import ConcurrencyPolicy - policy = ( ConcurrencyPolicy.ALLOW if getattr(self.config, "stateless", True) @@ -710,8 +692,6 @@ class Agent( intervention_policy = getattr(self.config, "intervention_policy", None) - from zhenxun.services.ai.utils import ContextUtils - lock_id = ContextUtils.extract_concurrency_lock_id( safe_context, getattr(self.config, "concurrency_scope", None), @@ -719,8 +699,6 @@ class Agent( ) async def _execution_task(): - from zhenxun.services.ai.flow.concurrency import apply_concurrency_policy - cancel_token = safe_context.run.cancellation_token or CancellationToken() safe_context.run.cancellation_token = cancel_token @@ -767,16 +745,16 @@ class Agent( task.cancel() def _parse_task_prompt( - self, prompt: PromptInput | Task | None - ) -> tuple[Task | None, Any | None, list[Any], Any, list[Any]]: - """解析输入意图,提取数据契约 (Task)""" + self, prompt: PromptInput | AgentTask | None + ) -> tuple[AgentTask | None, Any | None, list[Any], Any, list[Any]]: + """解析输入意图,提取数据契约 (AgentTask)""" task_obj = None final_prompt_payload = None extra_tools = [] run_output_type = self.response_model task_guardrails = [] - if isinstance(prompt, Task): + if isinstance(prompt, AgentTask): task_obj = prompt if task_obj.response_model: run_output_type = task_obj.response_model @@ -803,7 +781,7 @@ class Agent( async def on_state_init( self, - prompt: PromptInput | Task | None = None, + prompt: PromptInput | AgentTask | None = None, context: RunContext[AgentDepsT] | None = None, config: AgentConfig | None = None, cancellation_token: Any = None, @@ -886,11 +864,6 @@ class Agent( self, state: AgentState, resources: AgentRunResources ) -> None: """装配记忆与提示词上下文、解析可用工具集""" - from zhenxun.services.ai.capabilities import CombinedCapability - from zhenxun.services.ai.flow.agent.engine.builders import ( - AgentProfileResolver, - ContextBuilder, - ) context = resources.run_context reader = resources.memory_reader @@ -961,8 +934,6 @@ class Agent( messages_for_run.extend(context.run.messages) if final_prompt_payload is not None: - from zhenxun.services.ai.message_builder import MessageBuilder - if msgs := await MessageBuilder.normalize_to_llm_messages( final_prompt_payload, bot=context.get_bot(), event=context.get_event() ): @@ -986,8 +957,6 @@ class Agent( self, state: AgentState, resources: AgentRunResources ) -> AgentRunResult[OutputDataT]: """真正调度大模型执行器并执行记忆落盘""" - from zhenxun.services.ai.flow.agent.engine.executor import StandardAgentExecutor - context = resources.run_context for tk in resources.toolkits: @@ -1027,7 +996,7 @@ class Agent( async def _run_step( self, - prompt: PromptInput | Task | None = None, + prompt: PromptInput | AgentTask | None = None, *, context: RunContext[AgentDepsT], config: AgentConfig, @@ -1046,8 +1015,6 @@ class Agent( ) await self.on_context_build(state, resources) - from zhenxun.services.ai.capabilities import CombinedCapability - run_scoped_cap = ( resources.run_scoped_cap if isinstance(resources.run_scoped_cap, CombinedCapability) diff --git a/zhenxun/services/ai/flow/agent/capabilities.py b/zhenxun/services/ai/flow/agent/capabilities.py index f7374cb2..0b5e0f22 100644 --- a/zhenxun/services/ai/flow/agent/capabilities.py +++ b/zhenxun/services/ai/flow/agent/capabilities.py @@ -10,7 +10,8 @@ from zhenxun.services.ai.core.engine.structured_parser import ( from zhenxun.services.ai.core.exceptions import UpstreamServerException from zhenxun.services.ai.core.messages import TaskLifecycleEvent from zhenxun.services.ai.core.options import BaseOutputDefinition, ToolOutput -from zhenxun.services.ai.run import AgentRunResult, RunContext, Task +from zhenxun.services.ai.guardrails import parse_guardrails +from zhenxun.services.ai.run import AgentRunResult, AgentTask, RunContext from zhenxun.services.log import logger @@ -33,8 +34,6 @@ class OutputValidationCapability(AbstractCapability): self.output_type = output_type self.raw_schema = raw_schema - from zhenxun.services.ai.guardrails import parse_guardrails - self.guardrails = parse_guardrails(guardrails) self.processor = None self.submit_tool = None @@ -114,7 +113,7 @@ class OutputValidationCapability(AbstractCapability): class TaskTrackingCapability(AbstractCapability): """数据契约任务状态追踪与事件遥测组件""" - def __init__(self, task: Task, agent_name: str): + def __init__(self, task: AgentTask, agent_name: str): self.task = task self.agent_name = agent_name diff --git a/zhenxun/services/ai/flow/agent/engine/builders.py b/zhenxun/services/ai/flow/agent/engine/builders.py index 14cd476c..db56d883 100644 --- a/zhenxun/services/ai/flow/agent/engine/builders.py +++ b/zhenxun/services/ai/flow/agent/engine/builders.py @@ -5,19 +5,31 @@ from typing import Any, cast from nonebot.utils import is_coroutine_callable -from zhenxun.services.ai.capabilities import CombinedCapability +from zhenxun.services.ai.capabilities import ( + AbstractCapability, + CombinedCapability, + DynamicCapability, +) +from zhenxun.services.ai.context.memory.builder import MemoryBuilder +from zhenxun.services.ai.context.memory.engine import MemoryReader, MemoryWriter from zhenxun.services.ai.context.memory.models import MemoryConfig -from zhenxun.services.ai.core.messages import LLMMessage +from zhenxun.services.ai.context.memory.types import SessionMetadata +from zhenxun.services.ai.core.messages import LLMMessage, TextPart from zhenxun.services.ai.core.options import GenerationConfig from zhenxun.services.ai.core.templates import PromptTemplate +from zhenxun.services.ai.flow.agent.capabilities import ( + OutputValidationCapability, + TaskTrackingCapability, +) from zhenxun.services.ai.flow.agent.models import Persona -from zhenxun.services.ai.run import RunContext +from zhenxun.services.ai.run import GLOBAL_CAPABILITIES, RunContext from zhenxun.services.ai.run.di import DependencyInjector from zhenxun.services.ai.tools.engine.registry import ( ToolCollection, tool_provider_manager, ) from zhenxun.services.ai.tools.models import GlobalToolFilter, ResolvedToolPayload +from zhenxun.services.ai.utils.scope import ScopeSelector from zhenxun.utils.pydantic_compat import model_copy @@ -28,8 +40,17 @@ class AgentProfileResolver: def resolve_memory( agent_memory_config: MemoryConfig, override_memory: Any | None ) -> MemoryConfig: - from zhenxun.services.ai.context.memory.builder import MemoryBuilder + """ + 解析并合并 Memory 记忆域的配置。 + 支持从外部覆盖配置并重新构建。 + 参数: + agent_memory_config: 预置的 Agent 默认记忆域配置对象。 + override_memory: 运行时覆盖的记忆域配置,可为 dict, MemoryConfig 或其他合法结构。 + + 返回: + MemoryConfig: 合并并生成的运行时记忆域配置实例。 + """ # noqa: E501 if override_memory is not None: return MemoryBuilder.resolve(override_memory) return model_copy(agent_memory_config, deep=True) @@ -40,6 +61,18 @@ class AgentProfileResolver: cap_config: GenerationConfig | None, profile_config: GenerationConfig | None, ) -> GenerationConfig: + """ + 解析并合并多层 GenerationConfig 模型生成配置。 + 优先级顺序由低到高为:基础配置 -> 拦截器能力配置 -> 运行时 Profile 覆盖配置。 + + 参数: + base_config: 基础的模型生成配置对象。 + cap_config: 拦截器能力中提取出的模型生成参数配置。 + profile_config: 运行时传入的 Profile 覆盖参数配置。 + + 返回: + GenerationConfig: 合并多层配置后生成的最终运行时生成配置实例。 + """ final_gen_config = model_copy(base_config, deep=True) if cap_config: final_gen_config = final_gen_config.merge_with(cap_config) @@ -49,7 +82,7 @@ class AgentProfileResolver: class CapabilityBuilder: - """拦截器能力组装器:负责合并 Agent, Task, Profile 和全局的中间件""" + """拦截器能力组装器:负责合并 Agent, AgentTask, Profile 和全局的中间件""" @staticmethod async def build_for_run( @@ -64,15 +97,25 @@ class CapabilityBuilder: profile_capabilities: list | None, context: RunContext, ) -> CombinedCapability: - from zhenxun.services.ai.capabilities import ( - AbstractCapability, - DynamicCapability, - ) - from zhenxun.services.ai.flow.agent.capabilities import ( - OutputValidationCapability, - TaskTrackingCapability, - ) + """ + 为当前的 Agent 运行实例组装并实例化所有能力拦截器中间件。 + 整合全局能力、任务追踪、格式校验以及动态注入的能力。 + 参数: + agent_name: 执行推理的 Agent 标识名。 + namespace: 会话所归属的命名空间。 + output_type: 期待大模型返回的结构化 Pydantic 模型类型(支持 None 或 str)。 + raw_schema: 原始结构化 Schema 定义字典。 + agent_guardrails: Agent 自身定义的业务语义安全护栏列表。 + task_guardrails: 本次任务定义的业务语义安全护栏列表。 + task_obj: 被追踪的任务上下文对象实例。 + agent_capabilities: Agent 定义的静态拦截器列表。 + profile_capabilities: 运行时动态传入的能力或中间件列表。 + context: 运行上下文对象实例。 + + 返回: + CombinedCapability: 已经过运行初始化完毕的合并能力拦截器实例。 + """ dynamic_caps = [] combined_guardrails = agent_guardrails + task_guardrails @@ -100,10 +143,6 @@ class CapabilityBuilder: elif callable(cap): run_level_caps.append(DynamicCapability(cap)) - from zhenxun.services.ai.run import ( - GLOBAL_CAPABILITIES, - ) - base_caps = GLOBAL_CAPABILITIES.get("global", []).copy() if namespace != "global" and namespace in GLOBAL_CAPABILITIES: base_caps.extend(GLOBAL_CAPABILITIES[namespace]) @@ -129,8 +168,20 @@ class ContextBuilder: run_scoped_cap: CombinedCapability, persona: Persona | None = None, ) -> tuple[str, list[Any]]: - """解析提示词,返回 (静态系统提示词文本, 动态独立消息列表) 元组""" + """ + 解析、合并并渲染 Agent 的系统提示词和上下文记忆。 + 包含对依赖参数的动态注入和 Jinja 模板渲染。 + 参数: + instruction: 任务级别的初始指令或提示词模板。 + system_prompts: 系统提示词生成函数(支持依赖注入)列表。 + run_context: 运行上下文对象实例。 + run_scoped_cap: 运行域下的合并能力中间件。 + persona: 设定的 Agent 人设配置实例。 + + 返回: + tuple[str, list[Any]]: 包含 (渲染后的静态系统提示词文本, 渲染后的动态消息列表) 的元组。 + """ # noqa: E501 static_instructions = [] dynamic_messages = [] @@ -174,7 +225,7 @@ class ContextBuilder: static_instructions.append("\n\n".join(persona_parts)) if instruction: - static_instructions.append("## 本次任务指令 (Task)") + static_instructions.append("## 本次任务指令 (AgentTask)") if instruction: if isinstance(instruction, PromptTemplate): @@ -205,7 +256,6 @@ class ContextBuilder: render_context.update(run_context.state) rendered_dynamic_messages = [] - from zhenxun.services.ai.core.messages import TextPart for msg in dynamic_messages: if msg.role == "system": @@ -252,7 +302,22 @@ class ToolBuilder: run_context: RunContext, run_scoped_cap: CombinedCapability, ) -> ResolvedToolPayload: - """解析、合并并过滤工具集""" + """ + 解析并合并来自静态定义、动态函数依赖以及能力的工具列表。 + 通过工具提供者管理器完成工具的具体实例化及参数绑定。 + + 参数: + tool_definitions: 静态工具或工具集合的定义列表。 + toolset_funcs: 待依赖注入解析的工具集生成函数列表。 + system_tools: 系统默认强制集成的工具定义列表。 + namespace: 会话命名空间。 + tool_filter: 全局的工具过滤条件,此处为占位。 + run_context: 运行上下文对象实例. + run_scoped_cap: 运行域下的合并能力中间件,用于提供特定的能力工具。 + + 返回: + ResolvedToolPayload: 解析完毕并附带依赖绑定关系的工具负载载体。 + """ defs_to_resolve = list(tool_definitions) for ts_func in toolset_funcs: @@ -304,7 +369,18 @@ class ToolBuilder: tool_filters: list[Callable], run_scoped_cap: CombinedCapability, ) -> ToolCollection: - """处理生命周期:在工具发往执行器前,进行最终的 Schema 拦截和清洗""" + """ + 在将工具发往模型执行器之前,触发最终的过滤器与能力拦截,进行 schema 的清洗。 + + 参数: + effective_tools: 备选的工具执行实例列表。 + context: 运行上下文对象实例。 + tool_filters: 运行时自定义工具过滤与清洗函数列表。 + run_scoped_cap: 运行域下的合并能力中间件,提供拦截入口。 + + 返回: + ToolCollection: 准备就绪的、可直接发往模型的最终有效工具执行集。 + """ current_tool_defs = [] for t_exec in effective_tools: if hasattr(t_exec, "get_definition"): @@ -355,10 +431,18 @@ class SessionBuilder: agent_name: str, effective_memory: MemoryConfig, ) -> tuple[Any, Any, Any]: - from zhenxun.services.ai.context.memory.engine import MemoryReader, MemoryWriter - from zhenxun.services.ai.context.memory.types import SessionMetadata - from zhenxun.services.ai.utils.scope import ScopeSelector + """ + 根据当前用户、群组和平台标识,动态隔离前缀,计算并构建会话与记忆存储的读写门面。 + 参数: + context: 运行上下文对象实例。 + namespace: 当前会话的命名空间。 + agent_name: 执行推理的 Agent 标识名。 + effective_memory: 运行时最终生效的 MemoryConfig 配置对象。 + + 返回: + tuple[Any, Any, Any]: 包含 (SessionMetadata 会话元数据, MemoryReader 记忆读取器, MemoryWriter 记忆写入器) 的元组。 + """ # noqa: E501 bot_id = None bot_inst = context.get_bot() if bot_inst and hasattr(bot_inst, "self_id"): diff --git a/zhenxun/services/ai/flow/agent/engine/executor.py b/zhenxun/services/ai/flow/agent/engine/executor.py index 12265a6f..f6f8af40 100644 --- a/zhenxun/services/ai/flow/agent/engine/executor.py +++ b/zhenxun/services/ai/flow/agent/engine/executor.py @@ -3,6 +3,7 @@ import asyncio import json from typing import Any, cast +from zhenxun.services.ai.capabilities import CombinedCapability from zhenxun.services.ai.core.engine.context_renderer import ContextConverter from zhenxun.services.ai.core.engine.token_counter import ( parse_usage_info, @@ -36,9 +37,12 @@ from zhenxun.services.ai.core.stream_events import ( ) from zhenxun.services.ai.flow.agent.engine.directive import ( DirectiveHandlerFunc, + directive_manager, ) from zhenxun.services.ai.flow.agent.models import AgentRunResources, AgentState +from zhenxun.services.ai.llm.engine.router import LLMOrchestrator from zhenxun.services.ai.run import AgentRunResult, RunContext +from zhenxun.services.ai.run.session import session_manager from zhenxun.services.ai.tools.engine.executor import ToolExecutor from zhenxun.services.ai.tools.models import ToolResult from zhenxun.services.log import logger @@ -279,9 +283,6 @@ class StandardAgentExecutor(BaseAgentExecutor): extra: dict[str, Any] | None = None, cancellation_token: Any = None, ) -> ChatResponse: - from zhenxun.services.ai.capabilities import CombinedCapability - from zhenxun.services.ai.llm.engine.router import LLMOrchestrator - request = ChatRequest( messages=messages, config=config, @@ -315,8 +316,6 @@ class StandardAgentExecutor(BaseAgentExecutor): ) -> AgentRunResult[Any]: """覆盖基类的模板方法,实现灵活的 while 控制流和 FOLLOW_UP 合并""" await self.on_start(state, resources) - from zhenxun.services.ai.run.session import session_manager - session_info = await session_manager.get_or_create( resources.run_context.session_id or "default_session" ) @@ -388,9 +387,6 @@ class StandardAgentExecutor(BaseAgentExecutor): self, state: AgentState, resources: AgentRunResources ) -> None: run_context = resources.run_context - - from zhenxun.services.ai.run.session import session_manager - session_info = await session_manager.get_or_create( run_context.session_id or "default_session" ) @@ -499,8 +495,6 @@ class StandardAgentExecutor(BaseAgentExecutor): None, ) if return_part: - from zhenxun.services.ai.tools.models import ToolResult - box["result"] = ToolResult(output=return_part.output) else: client_tool_calls.append(call) @@ -606,8 +600,6 @@ class StandardAgentExecutor(BaseAgentExecutor): if not tool_calls or not tool_results: return - from zhenxun.services.ai.flow.agent.engine.directive import directive_manager - for i, res_or_exc in enumerate(tool_results): original_call = tool_calls[i] tool_res = None diff --git a/zhenxun/services/ai/flow/agent/models.py b/zhenxun/services/ai/flow/agent/models.py index 39c47325..7cf0cc3d 100644 --- a/zhenxun/services/ai/flow/agent/models.py +++ b/zhenxun/services/ai/flow/agent/models.py @@ -165,7 +165,7 @@ class AgentRunResources(BaseModel): memory_writer: Any | None = None """用于将对话历史安全落盘的写入器""" run_scoped_cap: Any | None = None - """聚合了 Agent/Task/全局 的复合能力拦截器 (CombinedCapability)""" + """聚合了 Agent/AgentTask/全局 的复合能力拦截器 (CombinedCapability)""" task_obj: Any | None = None """(如有) 解析后的结构化数据任务契约""" toolkits: list[Any] = Field(default_factory=list) diff --git a/zhenxun/services/ai/flow/concurrency.py b/zhenxun/services/ai/flow/concurrency.py index be0b90d1..bfdffded 100644 --- a/zhenxun/services/ai/flow/concurrency.py +++ b/zhenxun/services/ai/flow/concurrency.py @@ -2,9 +2,16 @@ import asyncio from contextlib import asynccontextmanager from typing import Any -from zhenxun.services.ai.core.exceptions import ConcurrencyRejectException +from zhenxun.services.ai.core.exceptions import ( + ConcurrencyRejectException, + InterventionHandledException, +) from zhenxun.services.ai.core.models import CancellationToken -from zhenxun.services.ai.flow.base import ConcurrencyPolicy +from zhenxun.services.ai.run.models import AgentTask +from zhenxun.services.ai.run.session import LockContext, session_manager +from zhenxun.services.log import logger + +from .base import ConcurrencyPolicy, InterventionPolicy @asynccontextmanager @@ -16,8 +23,21 @@ async def apply_concurrency_policy( intervention_policy: Any = None, message: Any = None, ): - """应用并发策略的中央调度上下文管理器""" - from zhenxun.services.ai.run.session import LockContext, session_manager + """ + 异步上下文管理器:对大模型执行流应用特定的并发及消息干预调度策略。 + 负责请求互斥锁竞争、任务中断/拒绝处理,以及运行时用户实时指令的插队控制。 + + 参数: + session_id: 当前会话的唯一标识,用于在会话管理器中隔离上下文。 + lock_id: 当前锁域标识,决定了哪些 Agent 或任务使用同一套互斥锁竞争机制。 + policy: 当发生并发锁占用时执行的策略(允许、拒绝、中断、排队)。 + cancel_token: 运行时用于监听取消请求的取消令牌实例。 + intervention_policy: 用户消息干预策略(转向、追加)。 + message: 并发竞争发生时新入站的用户请求消息或 AgentTask 载荷。 + + 返回: + AsyncGenerator: 返回异步生成器,供 async with 消费,包裹大模型的整个执行环节。 + """ current_task = asyncio.current_task() task_tuple = (cancel_token, current_task) @@ -30,22 +50,15 @@ async def apply_concurrency_policy( lock_ctx = session_manager.lock_contexts.setdefault(lock_id, LockContext()) if exec_lock.locked(): - from zhenxun.services.ai.flow.base import InterventionPolicy - if intervention_policy in ( InterventionPolicy.STEER, InterventionPolicy.FOLLOW_UP, ): - from zhenxun.services.ai.core.exceptions import ( - InterventionHandledException, - ) - session = await session_manager.get_or_create(session_id) actual_msg = message - from zhenxun.services.ai.run.models import Task - if isinstance(message, Task): + if isinstance(message, AgentTask): actual_msg = message.description elif hasattr(message, "extract_plain_text"): actual_msg = message.extract_plain_text() @@ -82,8 +95,6 @@ async def apply_concurrency_policy( elif policy == ConcurrencyPolicy.QUEUE: if exec_lock.locked(): - from zhenxun.services.log import logger - logger.info( f"⏳ [并发控制] 锁域 {lock_id} 被占用," "新请求已进入后台等待队列 (QUEUE)..." diff --git a/zhenxun/services/ai/flow/team/models.py b/zhenxun/services/ai/flow/team/models.py index 9eb4b18e..3ed1ac35 100644 --- a/zhenxun/services/ai/flow/team/models.py +++ b/zhenxun/services/ai/flow/team/models.py @@ -78,6 +78,7 @@ class ConcurrentCallAction(TeamAction): """ actions: list[CallAction] + """并发执行的呼叫动作列表""" class FinishAction(TeamAction): @@ -108,13 +109,21 @@ class SubTaskRecord(BaseModel): """单条子任务(工单)数据契约""" id: str = Field(default_factory=lambda: str(uuid.uuid4())[:8]) + """子任务的唯一标识 ID""" title: str = "" + """子任务标题""" description: str = "" + """子任务的具体描述""" assignee: str | None = None + """执行该任务的指派成员 Agent 名称""" dependencies: list[str] = Field(default_factory=list) + """该任务依赖的前置任务 ID 列表""" status: TaskNodeStatus = TaskNodeStatus.pending + """任务的当前执行状态""" result: str | None = None + """任务执行结果或输出信息""" notes: list[str] = Field(default_factory=list) + """执行过程中的追加备注或错误记录""" metadata: dict[str, Any] = Field(default_factory=dict) """附加元数据,供系统底层或第三方插件挂载隐式上下文,对大模型不可见""" @@ -126,8 +135,11 @@ class TaskBoardState(BaseModel): """ tasks: list[SubTaskRecord] = Field(default_factory=list) + """看板中存储的所有子任务记录列表""" is_goal_complete: bool = False + """团队的终极目标是否已宣告完成""" final_summary: str | None = None + """目标完成后的终结总结报告""" def create_task( self, @@ -151,6 +163,7 @@ class TaskBoardState(BaseModel): return task def get_task(self, task_id: str) -> SubTaskRecord | None: + """根据 ID 或标题获取对应的子任务记录""" return next( (t for t in self.tasks if t.id == task_id or t.title == task_id), None ) @@ -158,6 +171,7 @@ class TaskBoardState(BaseModel): def update_task_status( self, task_id: str, status: TaskNodeStatus, result: str | None = None ) -> SubTaskRecord | None: + """更新指定任务的状态以及执行结果""" task = self.get_task(task_id) if not task: return None @@ -221,7 +235,7 @@ class TaskBoardState(BaseModel): return available def all_terminal(self) -> bool: - """判断是否所有的任务都已经进入了终结状态(完成或失败)""" + """检查是否所有的任务均已进入终结状态(已完成或已失败)""" if not self.tasks: return False return all( diff --git a/zhenxun/services/ai/flow/team/router.py b/zhenxun/services/ai/flow/team/router.py index ac227ef8..b2cb3e9b 100644 --- a/zhenxun/services/ai/flow/team/router.py +++ b/zhenxun/services/ai/flow/team/router.py @@ -9,7 +9,7 @@ from nonebot.utils import is_coroutine_callable from zhenxun.services.ai.core.messages import AgentMessage from zhenxun.services.ai.core.templates import PromptTemplate from zhenxun.services.ai.flow.team.models import RouteDecision, Transition -from zhenxun.services.ai.run import RunContext, Task +from zhenxun.services.ai.run import AgentTask, RunContext from zhenxun.services.ai.run.di import DependencyInjector from zhenxun.services.log import logger @@ -22,7 +22,7 @@ class BaseRouter(ABC): self, context: RunContext, history: Sequence[AgentMessage], - prompt: str | Task | None = None, + prompt: str | AgentTask | None = None, ) -> RouteDecision | None: """核心路由方法""" pass @@ -46,11 +46,12 @@ class FunctionRouter(BaseRouter): self, context: RunContext, history: Sequence[AgentMessage], - prompt: str | Task | None = None, + prompt: str | AgentTask | None = None, ) -> RouteDecision | None: sig = inspect.signature(self.selector_func) call_kwargs = {"prompt": prompt, "context": context, "history": history} - if isinstance(prompt, Task): + if isinstance(prompt, AgentTask): + call_kwargs["agent_task"] = prompt call_kwargs["task"] = prompt kwargs_resolved = await DependencyInjector.resolve_all( @@ -95,11 +96,11 @@ class RegexRouter(BaseRouter): self, context: RunContext, history: Sequence[AgentMessage], - prompt: str | Task | None = None, + prompt: str | AgentTask | None = None, ) -> RouteDecision | None: text_to_match = ( prompt.description - if isinstance(prompt, Task) + if isinstance(prompt, AgentTask) else (prompt or context.run.user_input or "") ) @@ -125,7 +126,7 @@ class ChainRouter(BaseRouter): self, context: RunContext, history: Sequence[AgentMessage], - prompt: str | Task | None = None, + prompt: str | AgentTask | None = None, ) -> RouteDecision | None: for router in self.routers: decision = await router.route(context, history, prompt) @@ -174,7 +175,7 @@ class LLMRouter(BaseRouter): self, context: RunContext, history: Sequence[AgentMessage], - prompt: str | Task | None = None, + prompt: str | AgentTask | None = None, ) -> RouteDecision | None: from zhenxun.services.ai.flow.agent.agent import Agent from zhenxun.services.ai.flow.agent.models import AgentConfig diff --git a/zhenxun/services/ai/flow/team/runner.py b/zhenxun/services/ai/flow/team/runner.py index d59f0b4e..fadb6c0a 100644 --- a/zhenxun/services/ai/flow/team/runner.py +++ b/zhenxun/services/ai/flow/team/runner.py @@ -18,6 +18,7 @@ from zhenxun.services.ai.flow.team.strategy import BaseTeamStrategy from zhenxun.services.ai.run import AgentRunResult, RunContext from zhenxun.services.ai.run.models import AgentRunEnd from zhenxun.services.log import logger +from zhenxun.utils.pydantic_compat import model_construct class TeamRunner: @@ -228,7 +229,9 @@ class TeamRunner: logger.info(f"🏁 **团队 [{self.team.name}]** 协作圆满结束!") if not isinstance(final_result, AgentRunResult): - final_result = AgentRunResult(output=final_result, usage=cumulative_usage) + final_result = model_construct( + AgentRunResult, output=final_result, usage=cumulative_usage + ) else: final_result.usage += cumulative_usage diff --git a/zhenxun/services/ai/flow/team/strategy.py b/zhenxun/services/ai/flow/team/strategy.py index 58460988..be099734 100644 --- a/zhenxun/services/ai/flow/team/strategy.py +++ b/zhenxun/services/ai/flow/team/strategy.py @@ -17,7 +17,7 @@ from zhenxun.services.ai.flow.team.models import ( TeamAction, ) from zhenxun.services.ai.flow.team.router import BaseRouter -from zhenxun.services.ai.run import RunContext, Task +from zhenxun.services.ai.run import AgentTask, RunContext from zhenxun.services.ai.tools.bridges.delegate import DelegateTool from zhenxun.services.log import logger @@ -44,7 +44,11 @@ class BaseTeamStrategy(ABC): return PromptTemplate(template).render(**kwargs) async def generate_plan( - self, team: "Team", prompt: str | Task | None, context: RunContext, **kwargs + self, + team: "Team", + prompt: str | AgentTask | None, + context: RunContext, + **kwargs, ) -> AsyncGenerator[TeamAction, Any]: """ 核心决策生成器 (Action Yielding Pattern)。 @@ -137,7 +141,11 @@ class RouteStrategy(BaseTeamStrategy): self.state_flow = state_flow async def generate_plan( - self, team: "Team", prompt: str | Task | None, context: RunContext, **kwargs + self, + team: "Team", + prompt: str | AgentTask | None, + context: RunContext, + **kwargs, ) -> AsyncGenerator[TeamAction, Any]: router = self.router if not router: @@ -266,7 +274,7 @@ class RouteStrategy(BaseTeamStrategy): ) continue - yield FinishAction(result=run_result.output) + yield FinishAction(result=run_result) break @@ -297,7 +305,11 @@ class CoordinateStrategy(BaseTeamStrategy): self.leader_tools = leader_tools or [] async def generate_plan( - self, team: "Team", prompt: str | Task | None, context: RunContext, **kwargs + self, + team: "Team", + prompt: str | AgentTask | None, + context: RunContext, + **kwargs, ) -> AsyncGenerator[TeamAction, Any]: delegation_tools = [] for m in team.members: @@ -338,7 +350,7 @@ class CoordinateStrategy(BaseTeamStrategy): leader_res = yield CallAction(agent=leader_agent, task=prompt) - yield FinishAction(result=leader_res.output) + yield FinishAction(result=leader_res) class BroadcastStrategy(BaseTeamStrategy): @@ -367,10 +379,14 @@ class BroadcastStrategy(BaseTeamStrategy): self.leader_tools = leader_tools or [] async def generate_plan( - self, team: "Team", prompt: str | Task | None, context: RunContext, **kwargs + self, + team: "Team", + prompt: str | AgentTask | None, + context: RunContext, + **kwargs, ) -> AsyncGenerator[TeamAction, Any]: task_desc_str = ( - prompt.description if isinstance(prompt, Task) else (prompt or "") + prompt.description if isinstance(prompt, AgentTask) else (prompt or "") ) if context.run.event_bus: @@ -413,7 +429,7 @@ class BroadcastStrategy(BaseTeamStrategy): leader_res = yield CallAction(agent=leader_agent, task=synthesize_prompt) - yield FinishAction(result=leader_res.output) + yield FinishAction(result=leader_res) class TaskStrategy(BaseTeamStrategy): @@ -471,7 +487,11 @@ class TaskStrategy(BaseTeamStrategy): self.leader_tools.append(self.bb_toolkit) async def generate_plan( - self, team: "Team", prompt: str | Task | None, context: RunContext, **kwargs + self, + team: "Team", + prompt: str | AgentTask | None, + context: RunContext, + **kwargs, ) -> AsyncGenerator[TeamAction, Any]: from zhenxun.services.ai.flow.team.models import TaskBoardState, TaskNodeStatus from zhenxun.services.ai.flow.team.task_tools import TaskPlanningToolkit @@ -542,6 +562,10 @@ class TaskStrategy(BaseTeamStrategy): goal_str = getattr(prompt, "description", None) or ( str(prompt) if prompt else "" ) + expected_out = getattr(prompt, "expected_output", None) + if expected_out: + goal_str += f"\n\n### 🎯 [预期产出要求]\n{expected_out}" + planner_prompt = f"""### 🎯 用户的终极目标 (Original Goal) {goal_str} diff --git a/zhenxun/services/ai/flow/team/team.py b/zhenxun/services/ai/flow/team/team.py index 0996a35d..c96d51c4 100644 --- a/zhenxun/services/ai/flow/team/team.py +++ b/zhenxun/services/ai/flow/team/team.py @@ -19,7 +19,7 @@ from zhenxun.services.ai.flow.team.router import BaseRouter from zhenxun.services.ai.flow.team.strategy import ( BaseTeamStrategy, ) -from zhenxun.services.ai.run import AgentRunResult, RunContext, Task +from zhenxun.services.ai.run import AgentRunResult, AgentTask, RunContext from zhenxun.services.ai.tools.providers.skills.models import Skill, SkillSource from zhenxun.utils.utils import infer_plugin_namespace @@ -239,7 +239,7 @@ class Team(BaseRunnable[AgentRunResult[Any]]): async def run( self, - prompt: PromptInput | Task | None = None, + prompt: PromptInput | AgentTask | None = None, *, context: "RunContext | None" = None, capabilities: list[CapabilitySource] | None = None, @@ -250,7 +250,7 @@ class Team(BaseRunnable[AgentRunResult[Any]]): 团队级运行阻塞核心入口,内部静默分配任务给成员直至汇总结束。 参数: - prompt: 派发给多智能体团队的任务描述 or 契约对象 (Task)。 + prompt: 派发给多智能体团队的任务描述 or 契约对象 (AgentTask)。 context: 显式传入的会话与运行上下文。 capabilities: 仅针对本次团队执行动态注入的临时拦截器列表。 kwargs: 透传的其他附加参数。 @@ -278,7 +278,7 @@ class Team(BaseRunnable[AgentRunResult[Any]]): @contextlib.asynccontextmanager async def run_stream( self, - prompt: PromptInput | Task | None = None, + prompt: PromptInput | AgentTask | None = None, *, context: "RunContext | None" = None, capabilities: list[CapabilitySource] | None = None, diff --git a/zhenxun/services/ai/flow/workflow/__init__.py b/zhenxun/services/ai/flow/workflow/__init__.py index 504061c7..bfa8d934 100644 --- a/zhenxun/services/ai/flow/workflow/__init__.py +++ b/zhenxun/services/ai/flow/workflow/__init__.py @@ -1,5 +1,3 @@ -from .auto import AutoWorkflow -from .decorators import AND, OR, entry, listen, router from .engine import Workflow from .nodes import ( Condition, @@ -11,9 +9,6 @@ from .nodes import ( ) __all__ = [ - "AND", - "OR", - "AutoWorkflow", "Condition", "Loop", "Parallel", @@ -21,7 +16,4 @@ __all__ = [ "Step", "Steps", "Workflow", - "entry", - "listen", - "router", ] diff --git a/zhenxun/services/ai/flow/workflow/auto.py b/zhenxun/services/ai/flow/workflow/auto.py deleted file mode 100644 index 2a78944c..00000000 --- a/zhenxun/services/ai/flow/workflow/auto.py +++ /dev/null @@ -1,107 +0,0 @@ -import graphlib -from typing import Any - -from zhenxun.services.ai.flow.workflow.engine import Workflow -from zhenxun.services.ai.flow.workflow.nodes import NodeFactory, Parallel, Router -from zhenxun.services.log import logger - - -class AutoWorkflow(Workflow): - """ - 自动化声明式工作流 (Facade)。 - 允许开发者通过 @entry, @listen 装饰器定义类方法, - 在初始化时,底层编译器会自动分析依赖并推导为原生的 Steps 和 Parallel 节点图。 - """ - - def __init__(self, name: str | None = None, description: str = "", **kwargs: Any): - workflow_name = name or self.__class__.__name__ - compiled_steps = self._compile_graph() - - super().__init__( - name=workflow_name, steps=compiled_steps, description=description - ) - - self._auto_kwargs = kwargs - - def _compile_graph(self) -> list[Any]: - """核心图推导编译器:支持 Router 嵌套与 AND/OR 拓扑排序""" - methods_meta = {} - router_paths = set() - - for attr_name in dir(self): - if attr_name.startswith("_"): - continue - attr = getattr(self, attr_name) - if hasattr(attr, "__workflow_meta__"): - methods_meta[attr_name] = attr.__workflow_meta__ - if attr.__workflow_meta__.get("paths"): - router_paths.update(attr.__workflow_meta__["paths"]) - - if not methods_meta: - logger.warning( - f"AutoWorkflow '{self.__class__.__name__}' 没有检测到任何被装饰的方法!" - ) - return [] - - top_level_methods = { - k: v - for k, v in methods_meta.items() - if not any(t in router_paths for t in v.get("triggers", [])) - } - branch_methods = { - k: v - for k, v in methods_meta.items() - if any(t in router_paths for t in v.get("triggers", [])) - } - - ts = graphlib.TopologicalSorter() - for name, meta in top_level_methods.items(): - valid_triggers = [ - t for t in meta.get("triggers", []) if t in top_level_methods - ] - ts.add(name, *valid_triggers) - - try: - ts.prepare() - except graphlib.CycleError as e: - raise ValueError(f"AutoWorkflow 编译失败:检测到循环依赖!{e}") - - workflow_steps = [] - while ts.is_active(): - ready_nodes = ts.get_ready() - step_nodes = [] - for n in ready_nodes: - meta = top_level_methods[n] - if meta["type"] in ("router", "entry_router"): - choices = [] - for path in meta.get("paths", []): - branch_name = next( - ( - bk - for bk, bv in branch_methods.items() - if path in bv.get("triggers", []) - ), - None, - ) - if branch_name: - choices.append( - NodeFactory.build(getattr(self, branch_name), name=path) - ) - - step_nodes.append( - Router(name=n, selector=getattr(self, n), choices=choices) - ) - else: - step_nodes.append(NodeFactory.build(getattr(self, n), name=n)) - - if len(step_nodes) == 1: - workflow_steps.append(step_nodes[0]) - elif len(step_nodes) > 1: - workflow_steps.append( - Parallel(*step_nodes, name=f"Parallel_{'_'.join(ready_nodes)[:30]}") - ) - - for node in ready_nodes: - ts.done(node) - - return workflow_steps diff --git a/zhenxun/services/ai/flow/workflow/base.py b/zhenxun/services/ai/flow/workflow/base.py index 6f33e46d..8651387d 100644 --- a/zhenxun/services/ai/flow/workflow/base.py +++ b/zhenxun/services/ai/flow/workflow/base.py @@ -8,10 +8,12 @@ from zhenxun.services.ai.core.exceptions import ( ControlFlowExit, ToolFatalError, ) -from zhenxun.services.ai.flow.workflow.types import ( +from zhenxun.services.ai.flow.workflow.policies import ( AbortPolicy, BaseFailurePolicy, PolicyAction, +) +from zhenxun.services.ai.flow.workflow.types import ( StepInput, StepOutput, StepType, @@ -26,8 +28,6 @@ class BaseNode(ABC): def __init__( self, name: str, - requires_confirmation: bool = False, - confirmation_message: str | None = None, failure_policy: BaseFailurePolicy | None = None, ): """ @@ -35,14 +35,10 @@ class BaseNode(ABC): 参数: name: 节点的唯一名称标识。 - requires_confirmation: 标记该节点在执行前是否需要人工介入授权 (HITL),默认 False。 - confirmation_message: 挂起等待授权时,向前端/群聊展示的提示文案,默认 None。 failure_policy: 该节点执行失败时的错误恢复与自愈策略, 默认使用中断策略 (AbortPolicy)。 - """ # noqa: E501 + """ self.name = name - self.requires_confirmation = requires_confirmation - self.confirmation_message = confirmation_message self.failure_policy = failure_policy or AbortPolicy() @property @@ -79,7 +75,7 @@ class BaseNode(ABC): content=content, success=False, stop=True, - error=str(e) + error=f"{type(e).__name__}: {e}" if isinstance(e, AbortException | ToolFatalError) else None, ) @@ -110,10 +106,10 @@ class BaseNode(ABC): output = StepOutput( step_name=self.name, step_type=self.node_type, - content=f"节点执行失败,已通过策略自动跳过: {e}", + content=f"节点执行失败,已通过策略自动跳过: {type(e).__name__} - {e}", success=False, stop=False, - error=str(e), + error=f"{type(e).__name__}: {e}", ) return "break", output, None, None else: @@ -121,10 +117,10 @@ class BaseNode(ABC): output = StepOutput( step_name=self.name, step_type=self.node_type, - content=f"执行崩溃: {e}", + content=f"执行崩溃: {type(e).__name__} - {e}", success=False, stop=True, - error=str(e), + error=f"{type(e).__name__}: {e}", ) return "break", output, None, None @@ -168,37 +164,12 @@ class BaseNode(ABC): logger.debug(f" ⚙️ [节点] `{self.name}` 开始执行...") cached_out = context.state.get("__completed_steps__", {}).get(self.name) - if ( - cached_out - and cached_out.success - and not getattr(cached_out, "is_paused", False) - ): + if cached_out and cached_out.success: logger.debug(f"⏭️ 快进跳过已完成节点: {self.name}") yield cached_out return - if self.requires_confirmation: - if not context.state.get(f"__hitl_confirmed_{self.name}"): - msg = ( - self.confirmation_message - or f"⚠️ 工作流即将执行高危步骤:[{self.name}],等待授权..." - ) - logger.debug(f" ⏸️ **[节点挂起]** `{self.name}`: {msg}") - - output = StepOutput( - step_name=self.name, - step_type=self.node_type, - content="[任务已挂起,等待人工授权/输入]", - success=True, - stop=True, - is_paused=True, - pause_reason=msg, - ) - - yield output - return - current_input = step_input attempt = 1 diff --git a/zhenxun/services/ai/flow/workflow/decorators.py b/zhenxun/services/ai/flow/workflow/decorators.py deleted file mode 100644 index 9f1e212c..00000000 --- a/zhenxun/services/ai/flow/workflow/decorators.py +++ /dev/null @@ -1,91 +0,0 @@ -from collections.abc import Callable -from typing import Any - - -def AND(*triggers: str) -> dict[str, Any]: - """逻辑与:所有前置方法都完成才执行""" - return {"logic": "AND", "triggers": list(triggers)} - - -def OR(*triggers: str) -> dict[str, Any]: - """逻辑或:任意前置方法完成即执行""" - return {"logic": "OR", "triggers": list(triggers)} - - -def entry() -> Callable: - """ - 标记为工作流的入口节点。 - 执行工作流时会自动作为第一批任务执行。 - """ - - def decorator(func: Callable) -> Callable: - setattr( - func, - "__workflow_meta__", - { - "type": "entry", - "triggers": [], - "logic": "OR", - "paths": [], - }, - ) - return func - - return decorator - - -def listen(condition: str | dict[str, Any]) -> Callable: - """ - 监听其他节点的完成状态。 - - 用法: - @listen("step_a") - @listen(AND("step_a", "step_b")) - """ - - def decorator(func: Callable) -> Callable: - if isinstance(condition, str): - meta = { - "type": "listen", - "triggers": [condition], - "logic": "OR", - "paths": [], - } - elif isinstance(condition, dict): - meta = { - "type": "listen", - "triggers": condition["triggers"], - "logic": condition.get("logic", "OR"), - "paths": [], - } - else: - raise TypeError("listen condition 必须是字符串或 AND/OR 函数的返回值") - - setattr(func, "__workflow_meta__", meta) - return func - - return decorator - - -def router( - condition: str | dict[str, Any] | None = None, paths: list[str] | None = None -) -> Callable: - """ - 标记为路由节点。执行此方法后,会根据返回值走向对应的 paths。 - """ - - def decorator(func: Callable) -> Callable: - meta = {"type": "router", "triggers": [], "logic": "OR", "paths": paths or []} - if isinstance(condition, str): - meta["triggers"] = [condition] - elif isinstance(condition, dict): - meta["triggers"] = condition["triggers"] - meta["logic"] = condition.get("logic", "OR") - - if not condition: - meta["type"] = "entry_router" - - setattr(func, "__workflow_meta__", meta) - return func - - return decorator diff --git a/zhenxun/services/ai/flow/workflow/engine.py b/zhenxun/services/ai/flow/workflow/engine.py index f18dfa44..664b904b 100644 --- a/zhenxun/services/ai/flow/workflow/engine.py +++ b/zhenxun/services/ai/flow/workflow/engine.py @@ -1,12 +1,13 @@ import asyncio from collections.abc import AsyncIterator +import contextlib from typing import TYPE_CHECKING, Any import uuid +from nonebot.params import Depends + if TYPE_CHECKING: from zhenxun.services.ai.flow.workflow.nodes import NodeSource - from zhenxun.services.ai.run import StreamedRunResult - from zhenxun.services.ai.core.exceptions import ControlFlowExit, ToolRetryError from zhenxun.services.ai.core.messages import PromptInput, UsageInfo @@ -18,9 +19,16 @@ from zhenxun.services.ai.flow.workflow.types import ( StepOutput, WorkflowRunResult, ) -from zhenxun.services.ai.run import RunContext +from zhenxun.services.ai.run.context import RunContext +from zhenxun.services.ai.run.models import ( + AgentRunEnd, + AgentRunError, + AgentRunResult, + StreamedRunResult, +) from zhenxun.services.ai.tools.core.tool import FunctionTool from zhenxun.services.log import logger +from zhenxun.utils.message import MessageUtils class Workflow(BaseRunnable[WorkflowRunResult]): @@ -52,6 +60,7 @@ class Workflow(BaseRunnable[WorkflowRunResult]): safe_context: RunContext, final_output: StepOutput, ) -> WorkflowRunResult: + """根据执行链上的全量输出构建最终的工作流执行结果对象""" flat_outputs = {} def _extract(out: StepOutput): @@ -63,19 +72,7 @@ class Workflow(BaseRunnable[WorkflowRunResult]): if final_output: _extract(final_output) - paused_step = next( - ( - v.step_name - for v in reversed(list(flat_outputs.values())) - if getattr(v, "is_paused", False) and v.step_name - ), - None, - ) - status = ( - "paused" - if paused_step - else ("completed" if final_output and final_output.success else "error") - ) + status = "completed" if final_output and final_output.success else "error" return WorkflowRunResult( workflow_id=self.id, @@ -86,12 +83,10 @@ class Workflow(BaseRunnable[WorkflowRunResult]): step_outputs=flat_outputs, last_step_content=final_output.content if final_output else None, final_output=final_output, - paused_step_name=paused_step, ) def bind(self, **kwargs: Any) -> Any: """DI 注入语法糖""" - from nonebot.params import Depends async def _dependency() -> "Workflow": return self @@ -112,8 +107,6 @@ class Workflow(BaseRunnable[WorkflowRunResult]): 返回: WorkflowRunResult: 包含执行状态、断点快照、各节点产出的全量工作流结果对象。 """ - from zhenxun.utils.message import MessageUtils - ctx = RunContext() bot = ctx.get_bot() event = ctx.get_event() @@ -128,12 +121,6 @@ class Workflow(BaseRunnable[WorkflowRunResult]): else "执行完毕" ) await MessageUtils.build_message(msg).send(reply_to=reply_to) - elif res.status == "paused": - pause_msg = ( - f"⏸️ 工作流执行已被挂起,停在步骤: {res.paused_step_name}。" - "请提供授权或人工输入后继续。" - ) - await MessageUtils.build_message(pause_msg).send(reply_to=reply_to) elif res.status == "error": err_msg = res.final_output.error if res.final_output else "未知异常" await MessageUtils.build_message( @@ -186,8 +173,6 @@ class Workflow(BaseRunnable[WorkflowRunResult]): raise e - import contextlib - @contextlib.asynccontextmanager async def run_stream( self, @@ -197,9 +182,6 @@ class Workflow(BaseRunnable[WorkflowRunResult]): **kwargs: Any, ) -> AsyncIterator["StreamedRunResult[Any]"]: """对齐 BaseRunnable 接口的流式上下文管理器""" - from zhenxun.services.ai.run import StreamedRunResult - from zhenxun.services.ai.run.models import AgentRunError - event_bus = EventBus() if context: context.run.event_bus = event_bus @@ -226,7 +208,7 @@ class Workflow(BaseRunnable[WorkflowRunResult]): context: RunContext | None = None, **kwargs: Any, ) -> AsyncIterator[Any]: - """原 arun_stream 逻辑改名,供内部 _execution_task 调用""" + """流式执行工作流节点树的内部实现""" session_id = ( context.session_id if context and context.session_id else f"wf_{self.id}" ) @@ -251,9 +233,6 @@ class Workflow(BaseRunnable[WorkflowRunResult]): if final_output: logger.debug(f"🏭 **工作流 [{self.name}] 运行结束**") - from zhenxun.services.ai.run import AgentRunResult - from zhenxun.services.ai.run.models import AgentRunEnd - wf_result = self._build_result( initial_input, safe_context, final_output ) @@ -266,41 +245,9 @@ class Workflow(BaseRunnable[WorkflowRunResult]): except Exception: pass - async def acontinue_run( - self, - run_result: WorkflowRunResult, - user_auth_data: dict[str, Any] | None = None, - context: RunContext | None = None, - ) -> WorkflowRunResult: - safe_context = context or RunContext(session_id=f"wf_{self.id}") - safe_context.state.update(run_result.state) - - safe_context.state["__completed_steps__"] = run_result.step_outputs.copy() - for step_name, out in run_result.step_outputs.items(): - safe_context.upstream_results[step_name] = out.content - - if run_result.paused_step_name: - safe_context.state[f"__hitl_confirmed_{run_result.paused_step_name}"] = True - if user_auth_data: - safe_context.state[f"__hitl_input_{run_result.paused_step_name}"] = ( - user_auth_data - ) - - resume_input = StepInput( - input=run_result.original_input, - previous_step_content=run_result.last_step_content, - ) - - logger.debug( - f"🚀 工作流 [{self.name}] 状态已恢复," - f"正在快进到步骤: {run_result.paused_step_name}..." - ) - - final_output = await self.root_steps.aexecute(resume_input, safe_context) - - return self._build_result(resume_input, safe_context, final_output) - def as_tool(self, tool_name: str | None = None) -> FunctionTool: + """将工作流封装并导出为可供 Agent 直接调用的 FunctionTool 实例""" + async def _execute_workflow_tool(prompt: str, context: RunContext) -> str: run_result = await self.run(prompt=prompt, context=context) output = run_result.final_output diff --git a/zhenxun/services/ai/flow/workflow/nodes.py b/zhenxun/services/ai/flow/workflow/nodes.py index 58e36da1..b81d56d4 100644 --- a/zhenxun/services/ai/flow/workflow/nodes.py +++ b/zhenxun/services/ai/flow/workflow/nodes.py @@ -1,20 +1,20 @@ import asyncio from collections.abc import AsyncIterator, Callable, Sequence +import copy from typing import Any, cast -from pydantic import BaseModel, Field - from zhenxun.services.ai.core.messages import PromptInput from zhenxun.services.ai.flow.base import BaseRunnable from zhenxun.services.ai.flow.workflow.base import BaseNode +from zhenxun.services.ai.flow.workflow.policies import BaseFailurePolicy from zhenxun.services.ai.flow.workflow.types import ( - BaseFailurePolicy, StepInput, StepOutput, StepType, ) -from zhenxun.services.ai.run import RunContext +from zhenxun.services.ai.run import AgentTask, RunContext from zhenxun.services.ai.run.di import DependencyInjector +from zhenxun.services.ai.run.models import AgentRunEnd from zhenxun.services.log import logger NodeSource = BaseNode | BaseRunnable | Callable @@ -47,8 +47,6 @@ class Step(BaseNode): name: str | None = None, executor: NodeSource | None = None, prompt: PromptInput | None = None, - requires_confirmation: bool = False, - confirmation_message: str | None = None, failure_policy: BaseFailurePolicy | None = None, ): """ @@ -58,8 +56,6 @@ class Step(BaseNode): name: 步骤的名称,为空则自动取执行器的名称,默认 None。 executor: 该步骤要运行的核心执行器(支持 RunnableNode 或 Callable 依赖注入)。 prompt: 该步骤的初始输入或提示词定义,默认 None。 - requires_confirmation: 标记该节点在执行前是否需要人工介入授权,默认 False。 - confirmation_message: 挂起等待授权时展示的提示文案,默认 None。 failure_policy: 该节点执行失败时的错误处理策略,默认使用中断策略。 """ # noqa: E501 actual_name = name or getattr( @@ -67,8 +63,6 @@ class Step(BaseNode): ) super().__init__( name=actual_name, - requires_confirmation=requires_confirmation, - confirmation_message=confirmation_message, failure_policy=failure_policy, ) self.executor = executor @@ -83,9 +77,7 @@ class Step(BaseNode): ) -> AsyncIterator[Any]: if False: yield None - raise NotImplementedError( - "This is a facade. Real execution happens in subclasses." - ) + raise NotImplementedError("这是一个外观门面,实际的执行发生在子类中。") class RunnableNode(Step): @@ -94,27 +86,28 @@ class RunnableNode(Step): async def run_stream( self, step_input: StepInput, context: RunContext ) -> AsyncIterator[Any]: - import copy - - from zhenxun.services.ai.flow.base import BaseRunnable - from zhenxun.services.ai.run import Task - executor = cast(BaseRunnable, self.executor) - prompt_data = self.prompt if self.prompt is not None else step_input.input - if isinstance(prompt_data, Task): + if isinstance(step_input.previous_step_content, AgentTask): + prompt_data = step_input.previous_step_content + prev_content_to_append = None + else: + prompt_data = self.prompt if self.prompt is not None else step_input.input + prev_content_to_append = step_input.previous_step_content + + if isinstance(prompt_data, AgentTask): prompt_data = copy.copy(prompt_data) - if step_input.previous_step_content: - prev_content = str(step_input.previous_step_content) + if prev_content_to_append: + prev_content = str(prev_content_to_append) prompt_data.description = ( f"### 🔙 [上游节点执行输出]\n{prev_content}\n\n" f"### 🎯 [当前需执行的任务]\n{prompt_data.description}" ) context.run.user_input = prompt_data.description else: - if step_input.previous_step_content: + if prev_content_to_append: prompt_data = ( - f"[上游节点执行输出]:\n{step_input.previous_step_content}\n\n" + f"[上游节点执行输出]:\n{prev_content_to_append}\n\n" f"[当前需执行的任务]:\n{prompt_data or ''}" ) context.run.user_input = str(prompt_data) if prompt_data else "" @@ -126,8 +119,6 @@ class RunnableNode(Step): prompt=prompt_data, context=sandbox_context ) as stream_result: async for event in stream_result.stream_events(): - from zhenxun.services.ai.run.models import AgentRunEnd - if isinstance(event, AgentRunEnd): final_result = event.result yield event @@ -158,26 +149,6 @@ class FunctionNode(Step): yield StepOutput(content=res, success=True) -class StepMeta(BaseModel): - """承载工作流节点装饰器元数据的内部模型""" - - name: str | None = None - requires_confirmation: bool = False - confirmation_message: str | None = None - failure_policy: Any = None - - -class ConditionMeta(BaseModel): - name: str | None = None - if_true: list[Any] = Field(default_factory=list) - if_false: list[Any] = Field(default_factory=list) - - -class RouterMeta(BaseModel): - name: str | None = None - choices: list[Any] = Field(default_factory=list) - - class Steps(BaseNode): """串行执行的工作流容器。按照列表顺序依次执行。""" @@ -223,7 +194,6 @@ class Steps(BaseNode): yield StepOutput( content=all_outputs[-1].content if all_outputs else "No steps executed", success=all(o.success for o in all_outputs), - is_paused=any(getattr(o, "is_paused", False) for o in all_outputs), steps=all_outputs, ) @@ -432,7 +402,6 @@ class Loop(BaseNode): yield StepOutput( content=all_results[-1].content if all_results else "No iterations run", success=all(o.success for o in all_results), - is_paused=any(getattr(o, "is_paused", False) for o in all_results), steps=all_results, ) @@ -539,7 +508,6 @@ class Parallel(BaseNode): yield StepOutput( content="\n\n".join(aggregated_content_parts), success=not has_any_failure, - is_paused=any(getattr(o, "is_paused", False) for o in all_outputs), steps=all_outputs, stop=any(getattr(o, "stop", False) for o in all_outputs), ) @@ -555,18 +523,12 @@ class NodeFactory: cls, executor: NodeSource, name: str | None = None, - requires_confirmation: bool = False, - confirmation_message: str | None = None, failure_policy: Any = None, ) -> BaseNode: """底层物理实例化分发""" - from zhenxun.services.ai.flow.base import BaseRunnable - kwargs = { "name": name, "executor": executor, - "requires_confirmation": requires_confirmation, - "confirmation_message": confirmation_message, "failure_policy": failure_policy, } if isinstance(executor, BaseRunnable): @@ -590,36 +552,6 @@ class NodeFactory: return item if isinstance(item, BaseRunnable) or callable(item): - cond_meta = getattr(item, "__workflow_condition_meta__", None) - if cond_meta: - final_name = name or cond_meta.name or "ConditionGroup" - return Condition( - evaluator=item, - steps=cond_meta.if_true, - else_steps=cond_meta.if_false, - name=final_name, - ) - - router_meta = getattr(item, "__workflow_router_meta__", None) - if router_meta: - final_name = name or router_meta.name or "RouterGroup" - return Router( - selector=item, - choices=router_meta.choices, - name=final_name, - ) - - step_meta = getattr(item, "__workflow_step_meta__", None) - if step_meta: - final_name = name or step_meta.name - return NodeFactory._create_step( - executor=item, - name=final_name, - requires_confirmation=step_meta.requires_confirmation, - confirmation_message=step_meta.confirmation_message, - failure_policy=step_meta.failure_policy, - ) - return NodeFactory._create_step(executor=item, name=name) raise ValueError( diff --git a/zhenxun/services/ai/flow/workflow/policies.py b/zhenxun/services/ai/flow/workflow/policies.py new file mode 100644 index 00000000..ea48902d --- /dev/null +++ b/zhenxun/services/ai/flow/workflow/policies.py @@ -0,0 +1,184 @@ +import copy +from enum import Enum +from typing import Any + +from pydantic import BaseModel, Field + +from zhenxun.services.ai.flow.workflow.types import StepInput +from zhenxun.services.ai.llm.api import generate_structured +from zhenxun.services.log import logger + + +class PolicyAction(str, Enum): + RETRY = "retry" + CONTINUE = "continue" + ABORT = "abort" + FALLBACK = "fallback" + + +class PolicyResult(BaseModel): + """错误策略执行结果""" + + action: PolicyAction + """采取的具体恢复策略动作""" + delay: float = 0.0 + """执行延迟或重试前需要等待的缓冲秒数""" + new_input: StepInput | None = None + """用于动态纠错自愈时替换传入的新参数结构""" + fallback_node: Any | None = None + """策略裁定降级时所指定的备用工作流节点""" + healer_agent_name: str | None = None + """执行了高级自愈的大模型或修复者名称""" + + +class BaseFailurePolicy: + """错误处理策略抽象基类""" + + async def handle_failure( + self, node: Any, exception: BaseException, step_input: StepInput, context: Any + ) -> PolicyResult: + """ + 处理节点执行失败的策略入口方法。 + + 参数: + node: 发生异常的目标工作流节点。 + exception: 捕获到的具体异常实例。 + step_input: 节点执行时的原始输入数据。 + context: 当前工作流运行上下文。 + + 返回: + PolicyResult: 包含错误恢复动作、延迟时间以及备用参数等决策信息的策略结果对象。 + """ # noqa: E501 + raise NotImplementedError + + +class AbortPolicy(BaseFailurePolicy): + """直接中断策略""" + + async def handle_failure( + self, node: Any, exception: BaseException, step_input: StepInput, context: Any + ) -> PolicyResult: + return PolicyResult(action=PolicyAction.ABORT) + + +class SkipPolicy(BaseFailurePolicy): + """跳过并继续策略""" + + async def handle_failure( + self, node: Any, exception: BaseException, step_input: StepInput, context: Any + ) -> PolicyResult: + return PolicyResult(action=PolicyAction.CONTINUE) + + +class RetryPolicy(BaseFailurePolicy): + """退避重试策略""" + + def __init__(self, max_retries: int = 3, delay: float = 1.0): + """ + 初始化退避重试策略。 + + 参数: + max_retries: 最大允许重试的次数限制,默认 3。 + delay: 每次重试前需要等待和睡眠的秒数,默认 1.0。 + """ + self.max_retries = max_retries + self.delay = delay + + async def handle_failure( + self, node: Any, exception: BaseException, step_input: StepInput, context: Any + ) -> PolicyResult: + counts = context.state.setdefault("__retry_counts__", {}) + key = f"{node.name}_{id(self)}" + counts[key] = counts.get(key, 0) + 1 + + if counts[key] <= self.max_retries: + return PolicyResult(action=PolicyAction.RETRY, delay=self.delay) + return PolicyResult(action=PolicyAction.ABORT) + + +class FallbackPolicy(BaseFailurePolicy): + """降级路由策略""" + + def __init__(self, fallback_node: Any): + """ + 初始化降级路由策略。 + + 参数: + fallback_node: 当主节点发生致命故障时,直接转入执行的备用降级节点。 + """ + self.fallback_node = fallback_node + + async def handle_failure( + self, node: Any, exception: BaseException, step_input: StepInput, context: Any + ) -> PolicyResult: + return PolicyResult( + action=PolicyAction.FALLBACK, fallback_node=self.fallback_node + ) + + +class SelfHealingPolicy(BaseFailurePolicy): + """大模型高级自愈策略""" + + def __init__(self, healer_model: str, max_retries: int = 2): + """ + 初始化大模型高级自愈策略。 + + 参数: + healer_model: 用于分析错误原因并智能修复入参的大模型名称。 + max_retries: 最大尝试自愈修复的次数,默认 2。 + """ + self.healer_model = healer_model + self.max_retries = max_retries + + async def handle_failure( + self, node: Any, exception: BaseException, step_input: StepInput, context: Any + ) -> PolicyResult: + counts = context.state.setdefault("__heal_counts__", {}) + key = f"{node.name}_{id(self)}" + counts[key] = counts.get(key, 0) + 1 + + if counts[key] > self.max_retries: + logger.warning(f"节点 '{node.name}' 自愈次数达上限,宣告失败。") + return PolicyResult(action=PolicyAction.ABORT) + + class HealedInput(BaseModel): + """自愈后输入结构""" + + fixed_input: str = Field( + description="""修复后的输入参数 必须是完全合法的数据结构""" + ) + + prompt = f"""# Self-Healing Task + +请修复节点 `{node.name}` 的参数错误。 + +## Original Input +{step_input.input} + +## Exception +{exception} + +## Requirements +- 分析错误原因 +- 将输入修复为可被程序正确解析的格式 +- 只输出修复后的结果,不要输出额外解释 +""" + + try: + logger.info(f"🩹 触发 AI 自愈分析 (节点: {node.name})...") + res = await generate_structured( + prompt, response_model=HealedInput, model=self.healer_model + ) + + new_input = copy.copy(step_input) + new_input.input = res.fixed_input + + return PolicyResult( + action=PolicyAction.RETRY, + new_input=new_input, + healer_agent_name=self.healer_model, + ) + + except Exception as e: + logger.error(f"自愈过程发生大模型调用异常: {e}") + return PolicyResult(action=PolicyAction.ABORT) diff --git a/zhenxun/services/ai/flow/workflow/types.py b/zhenxun/services/ai/flow/workflow/types.py index ce5127f7..e891e445 100644 --- a/zhenxun/services/ai/flow/workflow/types.py +++ b/zhenxun/services/ai/flow/workflow/types.py @@ -1,4 +1,3 @@ -from abc import ABC, abstractmethod from enum import Enum from typing import Any @@ -50,10 +49,6 @@ class StepOutput(BaseModel): """执行失败时的异常详情""" stop: bool = False """标记是否触发了终止信号,以阻断后续流程的执行""" - is_paused: bool = False - """标记该步骤是否因等待外力交互 (HITL) 而处于挂起状态""" - pause_reason: str | None = None - """导致步骤挂起的原因描述""" steps: list["StepOutput"] | None = None """嵌套步骤的输出结果集合(如复合节点 Loop、Parallel 的内部产出)""" @@ -67,7 +62,7 @@ class WorkflowRunResult(BaseModel): workflow_name: str """工作流的名称""" status: str - """流水线的最终运行状态 (completed, paused, error 等)""" + """流水线的最终运行状态 (completed, error 等)""" original_input: Any """最初传入根节点的原始输入""" state: dict[str, Any] @@ -78,176 +73,3 @@ class WorkflowRunResult(BaseModel): """最后一个成功执行的步骤所产出的内容""" final_output: StepOutput | None = None """工作流根节点最终包装的完整产出对象""" - paused_step_name: str | None = None - """若流水线处于挂起态,记录是哪个步骤引发了挂起""" - - -class PolicyAction(str, Enum): - RETRY = "retry" - CONTINUE = "continue" - ABORT = "abort" - FALLBACK = "fallback" - - -class PolicyResult(BaseModel): - """错误策略执行结果""" - - action: PolicyAction - """采取的具体恢复策略动作""" - delay: float = 0.0 - """执行延迟或重试前需要等待的缓冲秒数""" - new_input: StepInput | None = None - """用于动态纠错自愈时替换传入的新参数结构""" - fallback_node: Any | None = None - """策略裁定降级时所指定的备用工作流节点""" - healer_agent_name: str | None = None - """执行了高级自愈的大模型或修复者名称""" - - -class BaseFailurePolicy(ABC): - """错误处理策略抽象基类""" - - @abstractmethod - async def handle_failure( - self, node: Any, exception: BaseException, step_input: StepInput, context: Any - ) -> PolicyResult: - pass - - -class AbortPolicy(BaseFailurePolicy): - """直接中断策略""" - - async def handle_failure( - self, node: Any, exception: BaseException, step_input: StepInput, context: Any - ) -> PolicyResult: - return PolicyResult(action=PolicyAction.ABORT) - - -class SkipPolicy(BaseFailurePolicy): - """跳过并继续策略""" - - async def handle_failure( - self, node: Any, exception: BaseException, step_input: StepInput, context: Any - ) -> PolicyResult: - return PolicyResult(action=PolicyAction.CONTINUE) - - -class RetryPolicy(BaseFailurePolicy): - """退避重试策略""" - - def __init__(self, max_retries: int = 3, delay: float = 1.0): - """ - 初始化退避重试策略。 - - 参数: - max_retries: 最大允许重试的次数限制,默认 3。 - delay: 每次重试前需要等待和睡眠的秒数,默认 1.0。 - """ - self.max_retries = max_retries - self.delay = delay - - async def handle_failure( - self, node: Any, exception: BaseException, step_input: StepInput, context: Any - ) -> PolicyResult: - counts = context.state.setdefault("__retry_counts__", {}) - key = f"{node.name}_{id(self)}" - counts[key] = counts.get(key, 0) + 1 - - if counts[key] <= self.max_retries: - return PolicyResult(action=PolicyAction.RETRY, delay=self.delay) - return PolicyResult(action=PolicyAction.ABORT) - - -class FallbackPolicy(BaseFailurePolicy): - """降级路由策略""" - - def __init__(self, fallback_node: Any): - """ - 初始化降级路由策略。 - - 参数: - fallback_node: 当主节点发生致命故障时,直接转入执行的备用降级节点。 - """ - self.fallback_node = fallback_node - - async def handle_failure( - self, node: Any, exception: BaseException, step_input: StepInput, context: Any - ) -> PolicyResult: - return PolicyResult( - action=PolicyAction.FALLBACK, fallback_node=self.fallback_node - ) - - -class SelfHealingPolicy(BaseFailurePolicy): - """大模型高级自愈策略""" - - def __init__(self, healer_model: str, max_retries: int = 2): - """ - 初始化大模型高级自愈策略。 - - 参数: - healer_model: 用于分析错误原因并智能修复入参的大模型名称。 - max_retries: 最大尝试自愈修复的次数,默认 2。 - """ - self.healer_model = healer_model - self.max_retries = max_retries - - async def handle_failure( - self, node: Any, exception: BaseException, step_input: StepInput, context: Any - ) -> PolicyResult: - counts = context.state.setdefault("__heal_counts__", {}) - key = f"{node.name}_{id(self)}" - counts[key] = counts.get(key, 0) + 1 - - if counts[key] > self.max_retries: - from zhenxun.services.log import logger - - logger.warning(f"节点 '{node.name}' 自愈次数达上限,宣告失败。") - return PolicyResult(action=PolicyAction.ABORT) - - import copy - - from zhenxun.services.ai.llm.api import generate_structured - from zhenxun.services.log import logger - - class HealedInput(BaseModel): - """自愈后输入结构""" - - fixed_input: str = Field( - description="""修复后的输入参数 必须是完全合法的数据结构""" - ) - - prompt = f"""# Self-Healing Task - -请修复节点 `{node.name}` 的参数错误。 - -## Original Input -{step_input.input} - -## Exception -{exception} - -## Requirements -- 分析错误原因 -- 将输入修复为可被程序正确解析的格式 -- 只输出修复后的结果,不要输出额外解释 -""" - - try: - logger.info(f"🩹 触发 AI 自愈分析 (节点: {node.name})...") - res = await generate_structured( - prompt, response_model=HealedInput, model=self.healer_model - ) - - new_input = copy.copy(step_input) - new_input.input = res.fixed_input - - return PolicyResult( - action=PolicyAction.RETRY, - new_input=new_input, - healer_agent_name=self.healer_model, - ) - - except Exception as e: - logger.error(f"自愈过程发生大模型调用异常: {e}") - return PolicyResult(action=PolicyAction.ABORT) diff --git a/zhenxun/services/ai/llm/adapters/base.py b/zhenxun/services/ai/llm/adapters/base.py index 2bc74a85..715a127e 100644 --- a/zhenxun/services/ai/llm/adapters/base.py +++ b/zhenxun/services/ai/llm/adapters/base.py @@ -64,20 +64,30 @@ class RequestData(BaseModel): """标准化的请求载体,用于向上层 HTTP 客户端传递请求参数。""" method: str = "POST" + """请求的 HTTP 方法,默认 'POST'""" url: str + """请求的目标 HTTP URL""" headers: dict[str, str] + """请求的 HTTP 头部键值对""" body: dict[str, Any] + """请求的 HTTP 载荷体 JSON 字典""" files: dict[str, Any] | list[tuple[str, Any]] | None = None + """要上传的多媒体或二进制文件字典""" class ResponseData(BaseModel): """标准化的响应载体,统一承接文本、多模态与附加元数据。""" content_parts: list[LLMContentPart] = Field(default_factory=list) + """大模型生成的结构化内容片段列表(如文本、图片、工具调用)""" usage_info: dict[str, Any] | None = None + """底层 API Token 消耗使用统计字典""" raw_response: dict[str, Any] | None = None + """接口返回的原始 JSON 响应字典""" grounding_metadata: Any | None = None + """Gemini 等模型特有的 Grounding 搜索依据元数据""" cache_info: Any | None = None + """接口缓存的命中与生成情况等元数据""" @property def text(self) -> str: diff --git a/zhenxun/services/ai/llm/api.py b/zhenxun/services/ai/llm/api.py index 1d7cfe9f..7dd2416c 100644 --- a/zhenxun/services/ai/llm/api.py +++ b/zhenxun/services/ai/llm/api.py @@ -8,7 +8,9 @@ from typing import Any, Literal, TypeVar, overload from pydantic import BaseModel from zhenxun.services.ai.core.exceptions import ( + ControlFlowExit, LLMException, + ModelRetry, UpstreamServerException, get_user_friendly_error_message, ) @@ -336,7 +338,7 @@ async def generate_structured( raise LLMException("结构化输出失败:中间件未返回解析后的对象。") return response.parsed_obj - except LLMException as e: + except (LLMException, ModelRetry, ControlFlowExit) as e: raise e.with_traceback(None) from None except Exception as e: friendly_msg = get_user_friendly_error_message(e) @@ -406,7 +408,7 @@ async def generate( return await LLMOrchestrator.invoke( request, model_name=model, task="chat", override_config=resolved_config ) - except LLMException as e: + except (LLMException, ModelRetry, ControlFlowExit) as e: raise e.with_traceback(None) from None except Exception as e: friendly_msg = get_user_friendly_error_message(e) diff --git a/zhenxun/services/ai/run/__init__.py b/zhenxun/services/ai/run/__init__.py index 542e60c5..0c17380f 100644 --- a/zhenxun/services/ai/run/__init__.py +++ b/zhenxun/services/ai/run/__init__.py @@ -12,8 +12,8 @@ from .hitl import HITLController from .hooks import Hooks from .models import ( AgentRunResult, + AgentTask, StreamedRunResult, - Task, ) from .session import session_manager from .ui import UIController @@ -21,6 +21,7 @@ from .ui import UIController __all__ = [ "GLOBAL_CAPABILITIES", "AgentRunResult", + "AgentTask", "BlackboardManager", "CancellationToken", "HITLController", @@ -30,7 +31,6 @@ __all__ = [ "NoneBotDeps", "RunContext", "StreamedRunResult", - "Task", "UIController", "get_current_run_context", "register_global_capability", diff --git a/zhenxun/services/ai/run/context.py b/zhenxun/services/ai/run/context.py index e40341e6..77718b7a 100644 --- a/zhenxun/services/ai/run/context.py +++ b/zhenxun/services/ai/run/context.py @@ -383,6 +383,7 @@ def set_run_context(ctx: RunContext): def _is_run_context_type(annotation: Any) -> bool: + """判断类型标注是否为 RunContext 类型或其子类。""" if annotation is RunContext: return True origin = get_origin(annotation) diff --git a/zhenxun/services/ai/run/hooks.py b/zhenxun/services/ai/run/hooks.py index 3148d33b..5aa69801 100644 --- a/zhenxun/services/ai/run/hooks.py +++ b/zhenxun/services/ai/run/hooks.py @@ -1,6 +1,6 @@ from collections.abc import Awaitable, Callable from dataclasses import dataclass -from typing import Any, Generic, Protocol, TypeVar +from typing import Any, Generic, Protocol, TypeVar, overload import anyio from nonebot.utils import is_coroutine_callable @@ -25,6 +25,7 @@ class HookTimeoutError(TimeoutError): """当 Hook 函数执行超过配置的时间时抛出此异常。""" def __init__(self, hook_name: str, func_name: str, timeout: float): + """初始化 HookTimeoutError 异常实例。""" self.hook_name = hook_name self.func_name = func_name self.timeout = timeout @@ -49,28 +50,38 @@ class _ToolHookEntry(_HookEntry[_FuncT]): class BeforeRunHookFunc(Protocol): + """运行前钩子:在 Agent 启动任何流转前触发。""" + def __call__(self, ctx: RunContext[Any], /) -> None | Awaitable[None]: ... class AfterRunHookFunc(Protocol): + """运行后钩子:在 Agent 获取最终结果后触发,可用于修改最终输出。""" + def __call__( self, ctx: RunContext[Any], /, *, result: AgentRunResult[Any] ) -> AgentRunResult[Any] | Awaitable[AgentRunResult[Any]]: ... class WrapRunHookFunc(Protocol): + """包裹运行钩子:以洋葱模型接管整个 Agent 运行生命周期。""" + def __call__( self, ctx: RunContext[Any], /, *, handler: WrapRunHandler ) -> AgentRunResult[Any] | Awaitable[AgentRunResult[Any]]: ... class OnRunErrorHookFunc(Protocol): + """运行异常钩子:捕获 Agent 级别的致命错误。""" + def __call__( self, ctx: RunContext[Any], /, *, error: BaseException ) -> AgentRunResult[Any] | Awaitable[AgentRunResult[Any]]: ... class BeforeModelRequestHookFunc(Protocol): + """大模型请求前钩子:在发送网络请求前触发,可篡改 Messages 上下文。""" + def __call__( self, ctx: RunContext[Any], @@ -83,6 +94,8 @@ class BeforeModelRequestHookFunc(Protocol): class AfterModelRequestHookFunc(Protocol): + """大模型请求后钩子:在接收到网络响应后触发,可验证或篡改原始 Response。""" + def __call__( self, ctx: RunContext[Any], @@ -94,6 +107,8 @@ class AfterModelRequestHookFunc(Protocol): class WrapModelRequestHookFunc(Protocol): + """包裹大模型请求钩子:以洋葱模型接管 LLM 网络请求过程(可用于实现缓存)。""" + def __call__( self, ctx: RunContext[Any], @@ -105,6 +120,8 @@ class WrapModelRequestHookFunc(Protocol): class OnModelRequestErrorHookFunc(Protocol): + """大模型请求异常钩子:捕获超时或网络等异常,可用于发起重试。""" + def __call__( self, ctx: RunContext[Any], @@ -116,24 +133,32 @@ class OnModelRequestErrorHookFunc(Protocol): class PrepareToolsHookFunc(Protocol): + """准备工具集钩子:向大模型渲染 JSON Schema 前触发,可动态增删工具。""" + def __call__( self, ctx: RunContext[Any], tool_defs: list[Any], / ) -> list[Any] | Awaitable[list[Any]]: ... class BeforeToolValidateHookFunc(Protocol): + """工具验证前钩子:在反序列化前触发,可篡改原始 JSON 字符串或字典。""" + def __call__( self, ctx: RunContext[Any], /, *, tool_name: str, args: str | dict[str, Any] ) -> str | dict[str, Any] | Awaitable[str | dict[str, Any]]: ... class AfterToolValidateHookFunc(Protocol): + """工具验证后钩子:在 Schema 校验通过后触发,接收并可篡改强类型参数字典。""" + def __call__( self, ctx: RunContext[Any], /, *, tool_name: str, args: dict[str, Any] ) -> dict[str, Any] | Awaitable[dict[str, Any]]: ... class WrapToolValidateHookFunc(Protocol): + """包裹工具验证钩子:以洋葱模型接管工具参数校验过程(可用于交互式拦截)。""" + def __call__( self, ctx: RunContext[Any], @@ -146,6 +171,8 @@ class WrapToolValidateHookFunc(Protocol): class OnToolValidateErrorHookFunc(Protocol): + """工具验证异常钩子:捕获校验失败的异常,可将其转化为大模型自愈提示。""" + def __call__( self, ctx: RunContext[Any], @@ -158,12 +185,16 @@ class OnToolValidateErrorHookFunc(Protocol): class BeforeToolExecuteHookFunc(Protocol): + """工具执行前钩子:在工具物理执行前触发,可用于最后的权限判定或参数篡改。""" + def __call__( self, ctx: RunContext[Any], /, *, tool_name: str, arguments: dict[str, Any] ) -> dict[str, Any] | Awaitable[dict[str, Any]]: ... class AfterToolExecuteHookFunc(Protocol): + """工具执行后钩子:工具执行完毕后触发,可篡改最终发往大模型的返回结果。""" + def __call__( self, ctx: RunContext[Any], @@ -176,6 +207,8 @@ class AfterToolExecuteHookFunc(Protocol): class WrapToolExecuteHookFunc(Protocol): + """包裹工具执行钩子:以洋葱模型接管特定工具的物理执行逻辑。""" + def __call__( self, ctx: RunContext[Any], @@ -188,12 +221,15 @@ class WrapToolExecuteHookFunc(Protocol): class OnToolExecuteErrorHookFunc(Protocol): + """工具执行异常钩子:捕获工具执行时的崩溃异常,可用于兜底返回或自愈。""" + def __call__( self, ctx: RunContext[Any], /, *, tool_name: str, error: Exception ) -> Any | Awaitable[Any]: ... async def _call_func(func: Callable[..., Any], *args: Any, **kwargs: Any) -> Any: + """统一执行同步或异步的可调用函数。""" if is_coroutine_callable(func): return await func(*args, **kwargs) return func(*args, **kwargs) @@ -276,224 +312,139 @@ def _tool_bare_or_parameterized( return decorator +class BoundHookPoint(Generic[_FuncT]): + """绑定到具体实例的 Hook 注册点。""" + + def __init__(self, key: str, registry_holder: Any): + """初始化绑定的 Hook 注册点。""" + self.key = key + self.registry_holder = registry_holder + + @overload + def __call__(self, func: _FuncT, /) -> _FuncT: ... + + @overload + def __call__( + self, *, timeout: float | None = None + ) -> Callable[[_FuncT], _FuncT]: ... + + def __call__( + self, func: _FuncT | None = None, *, timeout: float | None = None + ) -> Any: + """支持装饰器语法注册钩子函数。""" + return _bare_or_parameterized( + self.registry_holder._r, self.key, func, timeout=timeout + ) + + +class HookPoint(Generic[_FuncT]): + """描述符形式的 Hook 注册点。""" + + def __init__(self, key: str): + """初始化 Hook 描述符。""" + self.key = key + + def __get__(self, instance: Any, owner: Any) -> BoundHookPoint[_FuncT]: + """通过描述符协议绑定 Hook 实例。""" + return BoundHookPoint(self.key, instance) + + +class BoundToolHookPoint(Generic[_FuncT]): + """绑定到具体实例的工具级别 Hook 注册点。""" + + def __init__(self, key: str, registry_holder: Any): + """初始化绑定的工具 Hook 注册点。""" + self.key = key + self.registry_holder = registry_holder + + @overload + def __call__(self, func: _FuncT, /) -> _FuncT: ... + + @overload + def __call__( + self, *, tools: list[str] | None = None, timeout: float | None = None + ) -> Callable[[_FuncT], _FuncT]: ... + + def __call__( + self, + func: _FuncT | None = None, + *, + tools: list[str] | None = None, + timeout: float | None = None, + ) -> Any: + """支持装饰器语法注册工具钩子函数。""" + return _tool_bare_or_parameterized( + self.registry_holder._r, self.key, func, tools=tools, timeout=timeout + ) + + +class ToolHookPoint(Generic[_FuncT]): + """描述符形式的工具级别 Hook 注册点。""" + + def __init__(self, key: str): + """初始化工具 Hook 描述符。""" + self.key = key + + def __get__(self, instance: Any, owner: Any) -> BoundToolHookPoint[_FuncT]: + """通过描述符协议绑定工具 Hook 实例。""" + return BoundToolHookPoint(self.key, instance) + + class _HookRegistration: - """ - Hooks 的装饰器命名空间。 - 利用 @overload 提供完美的 IDE 强类型补全和文档提示。 - """ + """Hooks 装饰器注册辅助类,用于各阶段钩子的声明式注册""" def __init__(self, hooks: "Hooks"): + """初始化钩子注册辅助实例。""" self._hooks = hooks @property def _r(self) -> dict[str, list[_HookEntry[Any]]]: + """获取底层注册的钩子字典。""" return self._hooks._registry - from typing import overload + before_run = HookPoint[BeforeRunHookFunc]("before_run") + """注册运行前钩子。在 Agent 启动任何流转前触发。""" + after_run = HookPoint[AfterRunHookFunc]("after_run") + """注册运行后钩子。在 Agent 获取最终结果后触发,可修改结果。""" + wrap_run = HookPoint[WrapRunHookFunc]("wrap_run") + """注册运行包裹钩子。以洋葱模型接管整个 Agent 运行过程。""" + on_run_error = HookPoint[OnRunErrorHookFunc]("on_run_error") + """注册运行异常钩子。捕获 Agent 级别的致命错误。""" - @overload - def before_run(self, func: BeforeRunHookFunc, /) -> BeforeRunHookFunc: ... - @overload - def before_run( - self, *, timeout: float | None = None - ) -> Callable[[BeforeRunHookFunc], BeforeRunHookFunc]: ... - def before_run( - self, func: BeforeRunHookFunc | None = None, *, timeout: float | None = None - ) -> Any: - """注册运行前钩子。在 Agent 启动任何流转前触发。""" - return _bare_or_parameterized(self._r, "before_run", func, timeout=timeout) + before_model_request = HookPoint[BeforeModelRequestHookFunc]("before_model_request") + """注册大模型请求前钩子。可在此修改发送给 LLM 的 Messages 等上下文。""" + after_model_request = HookPoint[AfterModelRequestHookFunc]("after_model_request") + """注册大模型请求后钩子。可在此验证或修改 LLM 的原始 Response。""" + wrap_model_request = HookPoint[WrapModelRequestHookFunc]("wrap_model_request") + """注册大模型请求包裹钩子。以洋葱模型接管 LLM 的网络请求过程。""" + on_model_request_error = HookPoint[OnModelRequestErrorHookFunc]( + "on_model_request_error" + ) + """注册大模型请求异常钩子。捕获超时或网络等异常。""" - @overload - def after_run(self, func: AfterRunHookFunc, /) -> AfterRunHookFunc: ... - @overload - def after_run( - self, *, timeout: float | None = None - ) -> Callable[[AfterRunHookFunc], AfterRunHookFunc]: ... - def after_run( - self, func: AfterRunHookFunc | None = None, *, timeout: float | None = None - ) -> Any: - """注册运行后钩子。在 Agent 获取最终结果后触发,可修改结果。""" - return _bare_or_parameterized(self._r, "after_run", func, timeout=timeout) + before_tool_validate = HookPoint[BeforeToolValidateHookFunc]("before_tool_validate") + """注册工具验证前钩子。在工具输入参数校验前触发,可篡改或校验参数。""" + after_tool_validate = HookPoint[AfterToolValidateHookFunc]("after_tool_validate") + """注册工具验证后钩子。在工具输入参数校验成功后触发,可篡改最终参数。""" + wrap_tool_validate = HookPoint[WrapToolValidateHookFunc]("wrap_tool_validate") + """注册工具验证包裹钩子。以洋葱模型接管工具参数校验过程。""" + on_tool_validate_error = HookPoint[OnToolValidateErrorHookFunc]( + "on_tool_validate_error" + ) + """注册工具验证异常钩子。捕获并处理校验阶段抛出的异常。""" - @overload - def wrap_run(self, func: WrapRunHookFunc, /) -> WrapRunHookFunc: ... - @overload - def wrap_run( - self, *, timeout: float | None = None - ) -> Callable[[WrapRunHookFunc], WrapRunHookFunc]: ... - def wrap_run( - self, func: WrapRunHookFunc | None = None, *, timeout: float | None = None - ) -> Any: - """注册运行包裹钩子。以洋葱模型接管整个 Agent 运行过程。""" - return _bare_or_parameterized(self._r, "wrap_run", func, timeout=timeout) - - @overload - def on_run_error(self, func: OnRunErrorHookFunc, /) -> OnRunErrorHookFunc: ... - @overload - def on_run_error( - self, *, timeout: float | None = None - ) -> Callable[[OnRunErrorHookFunc], OnRunErrorHookFunc]: ... - def on_run_error( - self, func: OnRunErrorHookFunc | None = None, *, timeout: float | None = None - ) -> Any: - """注册运行异常钩子。捕获 Agent 级别的致命错误。""" - return _bare_or_parameterized(self._r, "on_run_error", func, timeout=timeout) - - @overload - def before_model_request( - self, func: BeforeModelRequestHookFunc, / - ) -> BeforeModelRequestHookFunc: ... - @overload - def before_model_request( - self, *, timeout: float | None = None - ) -> Callable[[BeforeModelRequestHookFunc], BeforeModelRequestHookFunc]: ... - def before_model_request( - self, - func: BeforeModelRequestHookFunc | None = None, - *, - timeout: float | None = None, - ) -> Any: - """注册大模型请求前钩子。可在此修改发送给 LLM 的 Messages 等上下文。""" - return _bare_or_parameterized( - self._r, "before_model_request", func, timeout=timeout - ) - - @overload - def after_model_request( - self, func: AfterModelRequestHookFunc, / - ) -> AfterModelRequestHookFunc: ... - @overload - def after_model_request( - self, *, timeout: float | None = None - ) -> Callable[[AfterModelRequestHookFunc], AfterModelRequestHookFunc]: ... - def after_model_request( - self, - func: AfterModelRequestHookFunc | None = None, - *, - timeout: float | None = None, - ) -> Any: - """注册大模型请求后钩子。可在此验证或修改 LLM 的原始 Response。""" - return _bare_or_parameterized( - self._r, "after_model_request", func, timeout=timeout - ) - - @overload - def wrap_model_request( - self, func: WrapModelRequestHookFunc, / - ) -> WrapModelRequestHookFunc: ... - @overload - def wrap_model_request( - self, *, timeout: float | None = None - ) -> Callable[[WrapModelRequestHookFunc], WrapModelRequestHookFunc]: ... - def wrap_model_request( - self, - func: WrapModelRequestHookFunc | None = None, - *, - timeout: float | None = None, - ) -> Any: - """注册大模型请求包裹钩子。以洋葱模型接管 LLM 的网络请求过程。""" - return _bare_or_parameterized( - self._r, "wrap_model_request", func, timeout=timeout - ) - - @overload - def on_model_request_error( - self, func: OnModelRequestErrorHookFunc, / - ) -> OnModelRequestErrorHookFunc: ... - @overload - def on_model_request_error( - self, *, timeout: float | None = None - ) -> Callable[[OnModelRequestErrorHookFunc], OnModelRequestErrorHookFunc]: ... - def on_model_request_error( - self, - func: OnModelRequestErrorHookFunc | None = None, - *, - timeout: float | None = None, - ) -> Any: - """注册大模型请求异常钩子。捕获超时或网络等异常。""" - return _bare_or_parameterized( - self._r, "on_model_request_error", func, timeout=timeout - ) - - @overload - def before_tool_execute( - self, func: BeforeToolExecuteHookFunc, / - ) -> BeforeToolExecuteHookFunc: ... - @overload - def before_tool_execute( - self, *, tools: list[str] | None = None, timeout: float | None = None - ) -> Callable[[BeforeToolExecuteHookFunc], BeforeToolExecuteHookFunc]: ... - def before_tool_execute( - self, - func: BeforeToolExecuteHookFunc | None = None, - *, - tools: list[str] | None = None, - timeout: float | None = None, - ) -> Any: - """注册工具执行前钩子。可通过 tools 参数指定拦截特定工具,可篡改传入参数。""" - return _tool_bare_or_parameterized( - self._r, "before_tool_execute", func, tools=tools, timeout=timeout - ) - - @overload - def after_tool_execute( - self, func: AfterToolExecuteHookFunc, / - ) -> AfterToolExecuteHookFunc: ... - @overload - def after_tool_execute( - self, *, tools: list[str] | None = None, timeout: float | None = None - ) -> Callable[[AfterToolExecuteHookFunc], AfterToolExecuteHookFunc]: ... - def after_tool_execute( - self, - func: AfterToolExecuteHookFunc | None = None, - *, - tools: list[str] | None = None, - timeout: float | None = None, - ) -> Any: - """注册工具执行后钩子。可通过 tools 参数指定拦截特定工具,可篡改返回结果。""" - return _tool_bare_or_parameterized( - self._r, "after_tool_execute", func, tools=tools, timeout=timeout - ) - - @overload - def wrap_tool_execute( - self, func: WrapToolExecuteHookFunc, / - ) -> WrapToolExecuteHookFunc: ... - @overload - def wrap_tool_execute( - self, *, tools: list[str] | None = None, timeout: float | None = None - ) -> Callable[[WrapToolExecuteHookFunc], WrapToolExecuteHookFunc]: ... - def wrap_tool_execute( - self, - func: WrapToolExecuteHookFunc | None = None, - *, - tools: list[str] | None = None, - timeout: float | None = None, - ) -> Any: - """注册工具执行包裹钩子。以洋葱模型接管特定工具的执行逻辑。""" - return _tool_bare_or_parameterized( - self._r, "wrap_tool_execute", func, tools=tools, timeout=timeout - ) - - @overload - def on_tool_execute_error( - self, func: OnToolExecuteErrorHookFunc, / - ) -> OnToolExecuteErrorHookFunc: ... - @overload - def on_tool_execute_error( - self, *, tools: list[str] | None = None, timeout: float | None = None - ) -> Callable[[OnToolExecuteErrorHookFunc], OnToolExecuteErrorHookFunc]: ... - def on_tool_execute_error( - self, - func: OnToolExecuteErrorHookFunc | None = None, - *, - tools: list[str] | None = None, - timeout: float | None = None, - ) -> Any: - """注册工具执行异常钩子。捕获特定工具的崩溃异常,可用于自愈重试。""" - return _tool_bare_or_parameterized( - self._r, "on_tool_execute_error", func, tools=tools, timeout=timeout - ) + before_tool_execute = ToolHookPoint[BeforeToolExecuteHookFunc]( + "before_tool_execute" + ) + """注册工具执行前钩子。可通过 tools 参数指定拦截特定工具,可篡改传入参数。""" + after_tool_execute = ToolHookPoint[AfterToolExecuteHookFunc]("after_tool_execute") + """注册工具执行后钩子。可通过 tools 参数指定拦截特定工具,可篡改返回结果。""" + wrap_tool_execute = ToolHookPoint[WrapToolExecuteHookFunc]("wrap_tool_execute") + """注册工具执行包裹钩子。以洋葱模型接管特定工具的执行逻辑。""" + on_tool_execute_error = ToolHookPoint[OnToolExecuteErrorHookFunc]( + "on_tool_execute_error" + ) + """注册工具执行异常钩子。捕获特定工具的崩溃异常,可用于自愈重试。""" class Hooks(AbstractCapability): @@ -503,93 +454,111 @@ class Hooks(AbstractCapability): """ def __init__(self): + """初始化 Hooks 拦截器实例。""" self._registry: dict[str, list[_HookEntry[Any]]] = {} self.on = _HookRegistration(self) - async def wrap_run( - self, context: RunContext, handler: WrapRunHandler - ) -> AgentRunResult[Any]: - """派发:洋葱模型组装与全周期执行接管""" - for entry in self._registry.get("before_run", []): - await _call_entry(entry, "before_run", context) + async def _dispatch_pipeline( + self, + hook_prefix: str, + context: RunContext, + handler: Callable, + get_entries: Callable[[str], list[_HookEntry[Any]]], + do_before: Callable[[_HookEntry[Any]], Awaitable[Any]], + get_wrap_kwargs: Callable[[Callable], dict[str, Any]], + invoke_chain: Callable[[Callable], Awaitable[Any]], + do_error: Callable[[_HookEntry[Any], BaseException], Awaitable[Any]], + do_after: Callable[[_HookEntry[Any], Any], Awaitable[Any]], + ) -> Any: + """核心泛型流水线引擎,按阶段分发并执行注册的钩子链。""" + for entry in get_entries(f"before_{hook_prefix}"): + await do_before(entry) - entries = self._registry.get("wrap_run", []) chain = handler - if entries: - for entry in reversed(entries): + wrap_entries = get_entries(f"wrap_{hook_prefix}") + if wrap_entries: + for entry in reversed(wrap_entries): - def _wrap( - e: _HookEntry[Any], h: Callable[..., Any] - ) -> Callable[..., Any]: - async def _wrapped() -> Any: - return await _call_entry(e, "wrap_run", context, h) + def _wrap(e: _HookEntry[Any], h: Callable) -> Callable: + async def _wrapped(*args, **kwargs) -> Any: + return await _call_entry( + e, f"wrap_{hook_prefix}", context, **get_wrap_kwargs(h) + ) return _wrapped chain = _wrap(entry, chain) try: - result = await chain() + result = await invoke_chain(chain) except BaseException as error: - for err_entry in reversed(self._registry.get("on_run_error", [])): + err_entries = get_entries(f"on_{hook_prefix}_error") + for err_entry in reversed(err_entries): try: - return await _call_entry(err_entry, "on_run_error", context, error) + return await do_error(err_entry, error) except BaseException as new_err: error = new_err raise error - for after_entry in reversed(self._registry.get("after_run", [])): - result = await _call_entry(after_entry, "after_run", context, result) + for after_entry in reversed(get_entries(f"after_{hook_prefix}")): + res = await do_after(after_entry, result) + if res is not None: + result = res return result + async def wrap_run( + self, context: RunContext, handler: WrapRunHandler + ) -> AgentRunResult[Any]: + """接管并包裹 Agent 运行生命周期的执行。""" + return await self._dispatch_pipeline( + "run", + context, + handler, + get_entries=lambda p: self._registry.get(p, []), + do_before=lambda e: _call_entry(e, "before_run", context), + get_wrap_kwargs=lambda h: {"handler": h}, + invoke_chain=lambda c: c(), + do_error=lambda e, err: _call_entry(e, "on_run_error", context, error=err), + do_after=lambda e, res: _call_entry(e, "after_run", context, result=res), + ) + async def wrap_model_request( self, context: RunContext, llm_context: LLMContext[ChatRequest, ChatResponse], handler: WrapModelRequestHandler, ) -> ChatResponse: - """派发:洋葱模型接管网络请求及前后生命周期""" - for entry in self._registry.get("before_model_request", []): - llm_context = await _call_entry( - entry, "before_model_request", context, llm_context - ) + """接管并包裹 LLM 大模型请求的网络交互过程。""" - entries = self._registry.get("wrap_model_request", []) - chain = handler - if entries: - for entry in reversed(entries): + async def do_before(e): + nonlocal llm_context + res = await _call_entry(e, "before_model_request", context, llm_context) + if res is not None: + llm_context = res - def _make_chain( - e: _HookEntry[Any], h: Callable[..., Any] - ) -> Callable[..., Any]: - async def _wrapped( - ctx_inner: LLMContext[ChatRequest, ChatResponse], - ) -> Any: - return await _call_entry( - e, "wrap_model_request", context, ctx_inner, h - ) - - return _wrapped - - chain = _make_chain(entry, chain) - - try: - response = await chain(llm_context) - except Exception as error: - for err_entry in reversed(self._registry.get("on_model_request_error", [])): - try: - return await _call_entry( - err_entry, "on_model_request_error", context, llm_context, error - ) - except Exception as new_err: - error = new_err - raise error - - for after_entry in reversed(self._registry.get("after_model_request", [])): - response = await _call_entry( - after_entry, "after_model_request", context, llm_context, response - ) - return response + return await self._dispatch_pipeline( + "model_request", + context, + handler, + get_entries=lambda p: self._registry.get(p, []), + do_before=do_before, + get_wrap_kwargs=lambda h: {"request_context": llm_context, "handler": h}, + invoke_chain=lambda c: c(llm_context), + do_error=lambda e, err: _call_entry( + e, + "on_model_request_error", + context, + request_context=llm_context, + error=err, + ), + do_after=lambda e, res: _call_entry( + e, + "after_model_request", + context, + request_context=llm_context, + response=res, + ), + ) async def wrap_tool_validate( self, @@ -598,51 +567,40 @@ class Hooks(AbstractCapability): args: str | dict[str, Any], handler: WrapToolValidateHandler, ) -> dict[str, Any]: - """派发:洋葱模型接管工具参数校验及生命周期""" - for entry in self._registry.get("before_tool_validate", []): - args = await _call_entry( - entry, "before_tool_validate", context, tool_name, args + """接管并包裹工具参数的校验过程。""" + + async def do_before(e): + nonlocal args + res = await _call_entry( + e, "before_tool_validate", context, tool_name=tool_name, args=args ) + if res is not None: + args = res - entries = self._registry.get("wrap_tool_validate", []) - chain = handler - if entries: - for entry in reversed(entries): - - def _wrap( - e: _HookEntry[Any], h: Callable[..., Any] - ) -> Callable[..., Any]: - async def _wrapped(args_inner: str | dict[str, Any]) -> Any: - return await _call_entry( - e, "wrap_tool_validate", context, tool_name, args_inner, h - ) - - return _wrapped - - chain = _wrap(entry, chain) - - try: - validated_args = await chain(args) - except Exception as error: - for err_entry in reversed(self._registry.get("on_tool_validate_error", [])): - try: - return await _call_entry( - err_entry, - "on_tool_validate_error", - context, - tool_name, - args, - error, - ) - except Exception as new_err: - error = new_err - raise error - - for after_entry in reversed(self._registry.get("after_tool_validate", [])): - validated_args = await _call_entry( - after_entry, "after_tool_validate", context, tool_name, validated_args - ) - return validated_args + return await self._dispatch_pipeline( + "tool_validate", + context, + handler, + get_entries=lambda p: self._registry.get(p, []), + do_before=do_before, + get_wrap_kwargs=lambda h: { + "tool_name": tool_name, + "args": args, + "handler": h, + }, + invoke_chain=lambda c: c(args), + do_error=lambda e, err: _call_entry( + e, + "on_tool_validate_error", + context, + tool_name=tool_name, + args=args, + error=err, + ), + do_after=lambda e, res: _call_entry( + e, "after_tool_validate", context, tool_name=tool_name, args=res + ), + ) async def wrap_tool_execute( self, @@ -651,55 +609,43 @@ class Hooks(AbstractCapability): arguments: dict[str, Any], handler: WrapToolExecuteHandler, ) -> Any: - """派发:洋葱模型接管工具执行(支持匹配名称)及生命周期""" - for entry in _filter_tool_entries( - self._registry.get("before_tool_execute", []), tool_name=tool_name - ): - arguments = await _call_entry( - entry, "before_tool_execute", context, tool_name, arguments - ) + """接管并包裹特定工具的物理执行逻辑。""" - entries = _filter_tool_entries( - self._registry.get("wrap_tool_execute", []), tool_name=tool_name + async def do_before(e): + nonlocal arguments + res = await _call_entry( + e, + "before_tool_execute", + context, + tool_name=tool_name, + arguments=arguments, + ) + if res is not None: + arguments = res + + return await self._dispatch_pipeline( + "tool_execute", + context, + handler, + get_entries=lambda p: _filter_tool_entries( + self._registry.get(p, []), tool_name=tool_name + ), + do_before=do_before, + get_wrap_kwargs=lambda h: { + "tool_name": tool_name, + "arguments": arguments, + "handler": h, + }, + invoke_chain=lambda c: c(arguments), + do_error=lambda e, err: _call_entry( + e, "on_tool_execute_error", context, tool_name=tool_name, error=err + ), + do_after=lambda e, res: _call_entry( + e, + "after_tool_execute", + context, + tool_name=tool_name, + arguments=arguments, + result=res, + ), ) - chain = handler - if entries: - for entry in reversed(entries): - - def _wrap( - e: _HookEntry[Any], h: Callable[..., Any] - ) -> Callable[..., Any]: - async def _wrapped(args_inner: dict[str, Any]) -> Any: - return await _call_entry( - e, "wrap_tool_execute", context, tool_name, args_inner, h - ) - - return _wrapped - - chain = _wrap(entry, chain) - - try: - result = await chain(arguments) - except Exception as error: - for err_entry in reversed( - _filter_tool_entries( - self._registry.get("on_tool_execute_error", []), tool_name=tool_name - ) - ): - try: - return await _call_entry( - err_entry, "on_tool_execute_error", context, tool_name, error - ) - except Exception as new_err: - error = new_err - raise error - - for after_entry in reversed( - _filter_tool_entries( - self._registry.get("after_tool_execute", []), tool_name=tool_name - ) - ): - result = await _call_entry( - after_entry, "after_tool_execute", context, tool_name, arguments, result - ) - return result diff --git a/zhenxun/services/ai/run/models.py b/zhenxun/services/ai/run/models.py index 29468c37..a8b42986 100644 --- a/zhenxun/services/ai/run/models.py +++ b/zhenxun/services/ai/run/models.py @@ -232,7 +232,7 @@ class TaskResult(BaseModel): model_config = ConfigDict(arbitrary_types_allowed=True) -class Task(BaseModel): +class AgentTask(BaseModel): """标准化数据契约(意图载体 Payload),定义大模型需要做什么及产出什么格式""" id: str = Field(default_factory=lambda: __import__("uuid").uuid4().hex) @@ -263,7 +263,7 @@ class Task(BaseModel): model_config = ConfigDict(arbitrary_types_allowed=True) @model_validator(mode="after") - def _parse_and_set_guardrails(self) -> "Task": + def _parse_and_set_guardrails(self) -> "AgentTask": from zhenxun.services.ai.guardrails import parse_guardrails self._parsed_guardrails = parse_guardrails(self.guardrails) @@ -272,8 +272,8 @@ class Task(BaseModel): __all__ = [ "AgentRunResult", + "AgentTask", "OutputDataT", "StreamedRunResult", - "Task", "TaskResult", ] diff --git a/zhenxun/services/ai/run/session.py b/zhenxun/services/ai/run/session.py index eabed282..154427bf 100644 --- a/zhenxun/services/ai/run/session.py +++ b/zhenxun/services/ai/run/session.py @@ -71,7 +71,9 @@ class LockContext(BaseModel): model_config = ConfigDict(arbitrary_types_allowed=True) active_task: Any | None = None + """当前持锁运行的异步任务""" cancel_token: CancellationToken | None = None + """当前任务关联的取消令牌,以便由抢占者随时下发取消指令""" class AgentSessionManager: @@ -115,32 +117,38 @@ class AgentSessionManager: return count def _get_lock(self, session_id: str) -> asyncio.Lock: + """获取或创建指定会话的内部同步锁""" if session_id not in self._locks: self._locks[session_id] = asyncio.Lock() return self._locks[session_id] def get_exec_lock(self, session_id: str) -> asyncio.Lock: + """获取或创建指定会话的任务排队执行锁""" if session_id not in self._exec_locks: self._exec_locks[session_id] = asyncio.Lock() return self._exec_locks[session_id] async def get_or_create(self, session_id: str) -> SessionInfo: + """获取或创建指定会话的信息,不存在则自动初始化""" async with self._get_lock(session_id): if session_id not in self._sessions: self._sessions[session_id] = SessionInfo(session_id=session_id) return self._sessions[session_id] async def get(self, session_id: str) -> SessionInfo | None: + """获取指定会话的信息,不存在则返回 None""" async with self._get_lock(session_id): return self._sessions.get(session_id) async def update_state(self, session_id: str, new_state: dict[str, Any]): + """用新字典更新指定会话的状态载荷""" async with self._get_lock(session_id): if session_id in self._sessions: self._sessions[session_id].state.update(new_state) self._sessions[session_id].updated_at = time.time() async def delete(self, session_id: str): + """删除指定会话及其对应的长期内存上下文""" async with self._get_lock(session_id): self._sessions.pop(session_id, None) from zhenxun.services.ai.context.memory.manager import memory_manager diff --git a/zhenxun/services/ai/run/subscribers.py b/zhenxun/services/ai/run/subscribers.py index 361ee59d..72a52519 100644 --- a/zhenxun/services/ai/run/subscribers.py +++ b/zhenxun/services/ai/run/subscribers.py @@ -22,10 +22,12 @@ class TelemetrySubscriber: """纯粹的数据观察者:默默记录时间戳与 Token 消耗""" def __init__(self): + """初始化遥测数据观察者,准备记录统计摘要""" self.summary = AgentRunSummary() self._start_times: dict[str, float] = {} def attach(self, bus: EventBus): + """注册监听的各种智能体生命周期和核心流程事件""" bus.subscribe(AgentRunStart, self.on_run_start) bus.subscribe(AgentRunEnd, self.on_run_end) bus.subscribe(LLMStartEvent, self.on_llm_start) @@ -34,19 +36,23 @@ class TelemetrySubscriber: bus.subscribe(ToolCallEndEvent, self.on_tool_end) async def on_run_start(self, event: AgentRunStart): + """处理智能体运行开始事件,记录运行起始时间""" self._start_times["run"] = time.monotonic() logger.debug(f"🚀 [Telemetry] 智能体 {event.agent_name} 开始运行") async def on_run_end(self, event: AgentRunEnd): + """处理智能体运行结束事件,统计总耗时并填充至结果""" dur = (time.monotonic() - self._start_times.get("run", time.monotonic())) * 1000 self.summary.total_latency_ms = dur event.result.telemetry = self.summary logger.debug(f"🏁 [Telemetry] 智能体运行结束 (总耗时: {dur:.2f}ms)") async def on_llm_start(self, event: LLMStartEvent): + """处理模型调用开始事件,记录单次大模型交互的起始时间""" self._start_times["llm"] = time.monotonic() async def on_llm_end(self, event: LLMEndEvent): + """处理模型调用结束事件,累计大模型交互次数和延迟,并记录停止原因""" dur = (time.monotonic() - self._start_times.pop("llm", time.monotonic())) * 1000 self.summary.chats.total += 1 self.summary.chats.total_latency_ms += dur @@ -60,9 +66,11 @@ class TelemetrySubscriber: logger.debug(f"🧠 [Telemetry] 模型调用完成 (耗时: {dur:.2f}ms)") async def on_tool_start(self, event: ToolCallStartEvent): + """处理工具调用开始事件,记录指定工具执行的起始时间""" self._start_times[f"tool_{event.tool_name}"] = time.monotonic() async def on_tool_end(self, event: ToolCallEndEvent): + """处理工具调用结束事件,统计工具执行耗时,累计成功和失败次数""" dur = ( time.monotonic() - self._start_times.pop(f"tool_{event.tool_name}", time.monotonic()) @@ -93,6 +101,14 @@ class DefaultUISubscriber: def __init__( self, context: RunContext, reply_to: bool = False, verbose: bool = False ): + """ + 初始化默认的 UI 观察者,绑定运行上下文并获取 bot 与事件实例。 + + 参数: + context: 智能体单次运行的上下文对象,用于提取 bot、事件及依赖。 + reply_to: 在发送平台消息时是否对用户发起的消息进行回复(引用/艾特),默认 False。 + verbose: 是否开启冗长模式,若开启则会将工具流等过程事件也输出到平台,默认 False。 + """ # noqa: E501 self.context = context self.reply_to = reply_to self.verbose = verbose @@ -100,10 +116,12 @@ class DefaultUISubscriber: self.event = context.get_event() def attach(self, bus: EventBus): + """注册订阅相关的 UI 交互事件(如工具输出流、用户自定义事件等)""" bus.subscribe(ToolStreamChunkEvent, self.on_tool_stream) bus.subscribe(UserCustomEvent, self.on_custom_event) async def _send_to_platform(self, display: Any): + """将通用的显示内容或富文本消息部件渲染并发送至具体的聊天平台""" if not self.bot or not self.event or not display: return @@ -131,8 +149,10 @@ class DefaultUISubscriber: await self.bot.send(self.event, str(display)) async def on_tool_stream(self, event: ToolStreamChunkEvent): + """处理工具流式数据块事件,仅在 verbose 开启时发送给平台""" if self.verbose and event.content: await self._send_to_platform(event.content) async def on_custom_event(self, event: UserCustomEvent): + """处理用户自定义事件,将展示内容发送至平台""" await self._send_to_platform(event.display) diff --git a/zhenxun/services/ai/sandbox/addons/base.py b/zhenxun/services/ai/sandbox/addons/base.py index 5f70aa65..d5056496 100644 --- a/zhenxun/services/ai/sandbox/addons/base.py +++ b/zhenxun/services/ai/sandbox/addons/base.py @@ -10,22 +10,30 @@ if TYPE_CHECKING: class BaseSandboxExtension(ABC): + """沙箱功能能力扩展基类""" + def __init__(self, session: "BaseSandboxSession"): + """初始化沙箱扩展实例,绑定当前沙箱会话""" self.session = session @property @abstractmethod def extension_name(self) -> str: + """获取扩展能力的唯一名称""" pass async def on_mount(self) -> None: + """在扩展挂载到会话时触发的钩子函数""" logger.debug(f"[SandboxExtension] 扩展 '{self.extension_name}' 已挂载。") async def on_unmount(self) -> None: + """在扩展从会话卸载时触发的钩子函数""" logger.debug(f"[SandboxExtension] 扩展 '{self.extension_name}' 已卸载。") class BaseMcpProxyExtension(BaseSandboxExtension): + """沙箱 MCP 代理扩展基类""" + @abstractmethod @asynccontextmanager async def connect_mcp( @@ -34,4 +42,5 @@ class BaseMcpProxyExtension(BaseSandboxExtension): args: list[str], env: dict[str, str] | None = None, ) -> AsyncGenerator[tuple[Any, Any], None]: + """建立与沙箱内部 MCP 服务的代理连接并返回输入输出管道""" yield None, None diff --git a/zhenxun/services/ai/sandbox/addons/mcp_proxy.py b/zhenxun/services/ai/sandbox/addons/mcp_proxy.py index 860af3df..06834f35 100644 --- a/zhenxun/services/ai/sandbox/addons/mcp_proxy.py +++ b/zhenxun/services/ai/sandbox/addons/mcp_proxy.py @@ -15,14 +15,18 @@ from zhenxun.utils.pydantic_compat import model_dump_json, model_validate class UniversalMcpExtension(BaseMcpProxyExtension): + """通用 MCP 代理扩展类,用于在沙箱内连接 MCP 服务""" + @property def extension_name(self) -> str: + """获取 MCP 代理扩展的唯一名称""" return "universal_mcp" @asynccontextmanager async def connect_mcp( self, command: str, args: list[str], env: dict[str, str] | None = None ) -> AsyncGenerator[tuple[Any, Any], None]: + """启动沙箱内的 MCP 服务器,并建立与之进行 JSON-RPC 通信的双向内存流管道""" if not isinstance(self.session, SupportsStreamExecution): raise RuntimeError( "当前沙箱驱动不支持流式后台进程执行 (SupportsStreamExecution)," diff --git a/zhenxun/services/ai/sandbox/drivers/base.py b/zhenxun/services/ai/sandbox/drivers/base.py index 94c736e1..fc663ac5 100644 --- a/zhenxun/services/ai/sandbox/drivers/base.py +++ b/zhenxun/services/ai/sandbox/drivers/base.py @@ -17,6 +17,7 @@ class BaseSandboxSession(SandboxChannel): """沙箱会话接口,持有 Client 分配的具体资源,提供统一的标准操作""" def __init__(self, state: SandboxSessionState): + """初始化沙箱会话实例并设置工作空间""" self.state = state self.last_active_time: float = time.time() self.loaded_skills: set[str] = set() @@ -27,6 +28,7 @@ class BaseSandboxSession(SandboxChannel): @property def session_id(self) -> str: + """获取当前会话的 ID""" return self.state.session_id def get_meta(self, key: str, default: Any = None) -> Any: @@ -103,6 +105,7 @@ class BaseSandboxSession(SandboxChannel): @abstractmethod async def close(self) -> None: + """关闭当前沙箱会话并释放关联资源""" pass @abstractmethod @@ -114,37 +117,46 @@ class BaseSandboxSession(SandboxChannel): env: dict[str, str] | None = None, on_output: Any = None, ) -> Any: + """在沙箱内异步执行子进程""" pass @abstractmethod async def read(self, path: str) -> bytes: + """读取沙箱内指定路径的文件内容""" pass @abstractmethod async def write(self, path: str, data: bytes) -> bool: + """向沙箱内指定路径写入文件数据""" pass @abstractmethod async def rm(self, path: str, recursive: bool = False) -> bool: + """在沙箱内删除指定路径的文件或目录""" pass @abstractmethod async def mkdir(self, path: str, parents: bool = False) -> bool: + """在沙箱内创建目录""" pass async def write_raw_file(self, path: str, content: str) -> bool: + """向沙箱中写入纯文本文件内容""" return await self.write(path, content.encode("utf-8")) async def read_raw_file(self, path: str) -> str: + """读取并以 utf-8 解码沙箱中的文本文件内容""" data = await self.read(path) return data.decode("utf-8", errors="replace") async def delete_raw_file(self, path: str) -> bool: + """删除沙箱中的纯文本文件""" return await self.rm(path) async def upload_raw_dir( self, local_dir_path: str, sandbox_target_path: str ) -> bool: + """将本地文件夹上传并映射至沙箱内指定位置""" return True @@ -159,14 +171,17 @@ class BaseSandboxClient(ABC): session_id: str, blueprint: SandboxBlueprint | None = None, ) -> BaseSandboxSession: + """创建并启动一个全新的沙箱会话环境""" pass @abstractmethod async def resume(self, state: SandboxSessionState) -> BaseSandboxSession: + """恢复一个已有的沙箱会话环境""" pass @abstractmethod async def delete(self, session: BaseSandboxSession) -> None: + """销毁指定的沙箱环境并清理容器或实例""" pass diff --git a/zhenxun/services/ai/sandbox/drivers/docker.py b/zhenxun/services/ai/sandbox/drivers/docker.py index 1824cecc..a8f10f2d 100644 --- a/zhenxun/services/ai/sandbox/drivers/docker.py +++ b/zhenxun/services/ai/sandbox/drivers/docker.py @@ -10,6 +10,7 @@ from typing import Any, ClassVar import aiodocker import anyio +from zhenxun.configs.config import BotConfig from zhenxun.services.ai.config import get_llm_config from zhenxun.services.ai.core.exceptions import SandboxPathEscapeError, WorkspaceIOError from zhenxun.services.ai.sandbox.models import ( @@ -32,6 +33,7 @@ class DockerInteractiveTerminalSession(InteractiveTerminalSession): """PTY 交互式会话:接管 Docker Stream,带有防死循环 Token 截断机制""" def __init__(self, session: "DockerSandboxSession"): + """初始化 Docker PTY 交互式会话实例""" self.session = session self.exec_stream = None self.buffer = "" @@ -39,6 +41,7 @@ class DockerInteractiveTerminalSession(InteractiveTerminalSession): self.ansi_escape = re.compile(r"(?:\x1B[@-_]|[\x80-\x9F])[0-?]*[ -/]*[@-~]") async def start(self, cmd: str, env: dict[str, str] | None = None) -> None: + """在容器中异步开启 PTY 终端执行指定命令""" if not self.session.container: raise RuntimeError("沙箱容器未启动") @@ -62,6 +65,7 @@ class DockerInteractiveTerminalSession(InteractiveTerminalSession): self._read_task = asyncio.create_task(self._read_loop()) async def _read_loop(self): + """异步循环读取容器执行输出流,并写入本地缓冲区(包含防超长截断)""" try: while True: if not self.exec_stream: @@ -82,18 +86,22 @@ class DockerInteractiveTerminalSession(InteractiveTerminalSession): logger.debug(f"[PTY] 流读取结束: {e}") async def send_input(self, text: str) -> None: + """向容器的 PTY 终端输入标准输入数据""" if self.exec_stream: await self.exec_stream.write_in(text.encode("utf-8")) async def read_output(self, timeout: int = 5) -> str: # noqa: ASYNC109 + """获取缓冲区最新的 50 行终端输出内容""" lines = self.buffer.split("\n") return "\n".join(lines[-50:]).strip() async def interrupt(self) -> None: + """发送 Ctrl+C 中断信号以打断当前执行进程""" if self.exec_stream: await self.exec_stream.write_in(b"\x03") async def close(self) -> None: + """关闭 PTY 读取循环任务和 exec 流资源""" if self._read_task: self._read_task.cancel() if self.exec_stream: @@ -105,33 +113,41 @@ class DockerSandboxProcessStream(SandboxProcessStream): """包装 aiodocker 的流,使其符合 SandboxProcessStream 协议""" def __init__(self, docker_stream): + """初始化封装的 Docker 流程管道流""" self.stream = docker_stream async def read(self) -> ProcessStreamMessage | None: + """异步读取管道流中的数据块并转换为 ProcessStreamMessage""" msg = await self.stream.read_out() if msg is None: return None return ProcessStreamMessage(stream_type=msg.stream, data=msg.data) async def write(self, data: bytes) -> None: + """向管道中异步写入数据字节""" await self.stream.write_in(data) async def close(self) -> None: + """关闭 Docker 管道流资源""" await self.stream.close() class DockerSandboxSession(BaseSandboxSession): + """Docker 驱动底层的具体沙箱会话通道实现""" + def __init__( self, state: SandboxSessionState, container: Any, ): + """初始化 Docker 沙箱会话并传入 Docker 容器句柄""" super().__init__(state) self.container = container self._vfs_helper_installed = False self._workspace_created = False async def is_alive(self) -> bool: + """查询底层 Docker 容器是否正处于 Running 状态""" if not self.container: return False try: @@ -141,9 +157,11 @@ class DockerSandboxSession(BaseSandboxSession): return False async def create_pty_session(self) -> InteractiveTerminalSession: + """创建交互式 Docker 终端会话实例""" return DockerInteractiveTerminalSession(self) async def _ensure_workspace(self): + """确保在容器内成功创建该会话的工作空间目录""" if not self._workspace_created and self.container: exec_inst = await self.container.exec( cmd=["/bin/sh", "-c", f"mkdir -p '{self.workspace_path}'"] @@ -162,6 +180,7 @@ class DockerSandboxSession(BaseSandboxSession): cwd: str | None = None, env: dict[str, str] | None = None, ) -> AsyncGenerator[SandboxProcessStream, None]: + """在指定工作目录下创建一个流式交互的进程通道""" cmd_list = ["/bin/sh", "-c", command] if isinstance(command, str) else command env_list = [f"{k}={v}" for k, v in env.items()] if env else None @@ -178,7 +197,7 @@ class DockerSandboxSession(BaseSandboxSession): yield DockerSandboxProcessStream(raw_stream) async def _ensure_vfs_helper(self): - """预置路径安全探针""" + """确保在沙箱容器内装有路径安全分析二进制文件""" if self._vfs_helper_installed: return check = await self.run_process(f"test -x {RESOLVE_PATH_HELPER.install_path}") @@ -191,7 +210,7 @@ class DockerSandboxSession(BaseSandboxSession): async def _validate_remote_path( self, path: str | Path, for_write: bool = False, base_dir: str | None = None ) -> Path: - """基于沙箱内真实环境的防软链接逃逸解析""" + """使用 VFS 探针对给定的沙箱路径进行安全性防逃逸校验""" base_dir = base_dir or self.workspace_path target_posix = coerce_posix_path(path).as_posix() is_write = "1" if for_write else "0" @@ -227,6 +246,7 @@ class DockerSandboxSession(BaseSandboxSession): env: dict[str, str] | None = None, on_output: Any = None, ) -> SandboxExecutionResult: + """在容器中指定目录下执行非交互式进程并等待其运行结果""" from zhenxun.services.ai.core.exceptions import SandboxFatalError self.touch() @@ -288,6 +308,7 @@ class DockerSandboxSession(BaseSandboxSession): return SandboxExecutionResult(exit_code=-1, error=str(e)) async def read(self, path: str | Path) -> bytes: + """通过 Docker Tar 归档接口读取容器内的指定文件内容""" self.touch() secure_path = await self._validate_remote_path(path, for_write=False) try: @@ -305,6 +326,7 @@ class DockerSandboxSession(BaseSandboxSession): raise WorkspaceIOError(str(path), f"读取文件异常: {e}") async def write(self, path: str | Path, data: bytes) -> bool: + """通过 Docker put_archive 接口将文件写入容器的指定路径""" self.touch() secure_path = await self._validate_remote_path(path, for_write=True) @@ -328,12 +350,14 @@ class DockerSandboxSession(BaseSandboxSession): return False async def rm(self, path: str | Path, recursive: bool = False) -> bool: + """在容器内执行 rm 命令移除指定文件或目录""" secure_path = await self._validate_remote_path(path, for_write=True) flag = "-rf" if recursive else "-f" res = await self.run_process(f"rm {flag} '{secure_path.as_posix()}'") return res.exit_code == 0 async def mkdir(self, path: str | Path, parents: bool = False) -> bool: + """在容器内执行 mkdir 命令创建目录""" secure_path = await self._validate_remote_path(path, for_write=True) flag = "-p" if parents else "" res = await self.run_process(f"mkdir {flag} '{secure_path.as_posix()}'") @@ -342,6 +366,7 @@ class DockerSandboxSession(BaseSandboxSession): async def upload_raw_dir( self, local_dir_path: str, sandbox_target_path: str ) -> bool: + """通过打包 tar 归档将宿主机本地目录上传至容器内指定路径""" aio_path = anyio.Path(local_dir_path) if not await aio_path.exists() or not await aio_path.is_dir(): return False @@ -369,7 +394,7 @@ class DockerSandboxSession(BaseSandboxSession): return False async def close(self) -> None: - """关闭会话:仅清理工作区,不销毁共享容器""" + """关闭会话并在容器内清理本会话对应的工作空间目录""" try: if await self.is_alive(): await self.rm(self.workspace_path, recursive=True) @@ -378,6 +403,8 @@ class DockerSandboxSession(BaseSandboxSession): class DockerSandboxClient(BaseSandboxClient): + """基于 Docker 实现的沙箱物理资源管理器类""" + backend_id = "docker" _global_docker_client: ClassVar[Any] = None _containers: ClassVar[dict[str, Any]] = {} @@ -390,10 +417,25 @@ class DockerSandboxClient(BaseSandboxClient): session_id: str, blueprint: SandboxBlueprint | None = None, ) -> BaseSandboxSession: + """建立或复用物理容器,并为该会话初始化专属的工作目录和 Python 虚拟环境""" bp = blueprint or SandboxBlueprint() eff_image = bp.image or get_llm_config().sandbox.docker_image eff_cname = bp.container_name + proxy_envs = [] + if BotConfig.system_proxy: + sandbox_proxy = BotConfig.system_proxy.replace( + "127.0.0.1", "host.docker.internal" + ).replace("localhost", "host.docker.internal") + + proxy_envs = [ + f"HTTP_PROXY={sandbox_proxy}", + f"HTTPS_PROXY={sandbox_proxy}", + f"http_proxy={sandbox_proxy}", + f"https_proxy={sandbox_proxy}", + f"ALL_PROXY={sandbox_proxy}", + ] + async with self._init_lock: if DockerSandboxClient._global_docker_client is None: import aiodocker @@ -465,6 +507,7 @@ class DockerSandboxClient(BaseSandboxClient): "HF_HOME=/tmp/hf_cache", "JUPYTER_RUNTIME_DIR=/tmp/jupyter_runtime", "JUPYTER_DATA_DIR=/tmp/jupyter_data", + *proxy_envs, ], "Cmd": [ "/bin/sh", @@ -476,6 +519,7 @@ class DockerSandboxClient(BaseSandboxClient): "HostConfig": { "PortBindings": port_bindings, "Binds": binds, + "ExtraHosts": ["host.docker.internal:host-gateway"], }, "Labels": { "zhenxun_component": "sandbox", @@ -549,10 +593,11 @@ class DockerSandboxClient(BaseSandboxClient): return session async def resume(self, state: SandboxSessionState) -> BaseSandboxSession: + """暂不支持通过还原状态重建 Docker 沙箱会话""" raise NotImplementedError("Docker Driver 不支持无状态重建恢复。") async def delete(self, session: BaseSandboxSession) -> None: - """触发会话销毁及引用计数物理回收""" + """清理会话工作区,并当物理容器处于长闲置时触发物理销毁回收""" await session.close() from zhenxun.services.ai.sandbox.manager import sandbox_manager @@ -584,6 +629,7 @@ class DockerSandboxClient(BaseSandboxClient): @classmethod async def close_env(cls): + """清理释放全部管理的 Docker 容器并关闭 Docker 客户端连接""" for cname, container in cls._containers.items(): try: await container.delete(force=True) @@ -602,7 +648,7 @@ class DockerSandboxClient(BaseSandboxClient): @classmethod async def silent_prune_orphans(cls): - """供 manager 后台延迟调用的静默清理函数""" + """在系统启动时静默搜寻并强力删除带有残留标记的孤儿容器""" try: async with aiodocker.Docker() as docker: await asyncio.wait_for(docker.system.info(), timeout=2.0) diff --git a/zhenxun/services/ai/sandbox/environments.py b/zhenxun/services/ai/sandbox/environments.py index 62f1c0e9..193c4aa1 100644 --- a/zhenxun/services/ai/sandbox/environments.py +++ b/zhenxun/services/ai/sandbox/environments.py @@ -37,6 +37,7 @@ class ProvisionerRegistry: @classmethod def register(cls, provisioner: BaseProvisioner) -> None: + """向注册中心注册一个环境配置器""" if provisioner.name in cls._provisioners: logger.warning( f"[ProvisionerRegistry] 覆盖已存在的配置器: {provisioner.name}" @@ -46,10 +47,12 @@ class ProvisionerRegistry: @classmethod def get(cls, name: str) -> BaseProvisioner | None: + """获取指定名称的环境配置器""" return cls._provisioners.get(name) @classmethod def get_all(cls) -> dict[str, BaseProvisioner]: + """获取所有已注册的环境配置器列表""" return cls._provisioners.copy() @@ -61,11 +64,13 @@ class UnifiedManifestProvisioner(BaseProvisioner): @property def name(self) -> str: + """返回统一环境清单装配器的名称""" return "unified_manifest" async def install( self, session: "BaseSandboxSession", blueprint: "SandboxBlueprint" ) -> bool: + """根据 SandboxBlueprint 环境清单指纹装配沙箱依赖环境""" target_hash = blueprint.calculate_hash() if session.get_meta("env_hash") == target_hash: @@ -96,6 +101,7 @@ class UnifiedManifestProvisioner(BaseProvisioner): async def scan_and_setup_workspace( self, session: "BaseSandboxSession", workspace_dir: str ) -> bool: + """扫描项目工作区,根据 requirements.txt 或 package.json 自动安装所需依赖""" check = await session.run_process(f"test -f {workspace_dir}/requirements.txt") if check.exit_code == 0: check_uv = await session.run_process("command -v uv") diff --git a/zhenxun/services/ai/sandbox/manager.py b/zhenxun/services/ai/sandbox/manager.py index 8a708cae..6b41494c 100644 --- a/zhenxun/services/ai/sandbox/manager.py +++ b/zhenxun/services/ai/sandbox/manager.py @@ -4,9 +4,7 @@ from typing import Any, cast import nonebot from zhenxun.services.ai.config import get_llm_config -from zhenxun.services.ai.sandbox.models import ( - SandboxBlueprint, -) +from zhenxun.services.ai.sandbox.models import SandboxBlueprint from zhenxun.services.log import logger from zhenxun.utils.lifespan import LifespanManager @@ -23,11 +21,13 @@ class SandboxManager: """ def __init__(self): + """初始化沙箱环境管理器,配置活跃会话、会话锁及生存期管理器""" self._active_sessions: dict[str, BaseSandboxSession] = {} self._session_locks: dict[str, asyncio.Lock] = {} self.lifespan_manager = LifespanManager() def _get_lock(self, session_id: str) -> asyncio.Lock: + """获取或创建指定会话的 asyncio 互斥锁,确保并发操作的安全性""" if session_id not in self._session_locks: self._session_locks[session_id] = asyncio.Lock() return self._session_locks[session_id] @@ -36,6 +36,7 @@ class SandboxManager: self, blueprint: SandboxBlueprint, ) -> BaseSandboxClient: + """根据全局配置和蓝图参数决定并获取具体的沙箱驱动客户端实例""" global_type = get_llm_config().sandbox.sandbox_type effective_type = ( @@ -148,6 +149,7 @@ class SandboxManager: return True async def release_resource(self, resource_id: str): + """释放指定会话的沙箱资源并销毁底层物理容器""" if resource_id in self._active_sessions: session = self._active_sessions.pop(resource_id) try: @@ -161,10 +163,12 @@ class SandboxManager: ) async def close_session(self, session_id: str) -> None: + """从生存期管理器中注销并销毁指定会话的沙箱环境""" await self.lifespan_manager.unregister(session_id) await self.release_resource(session_id) async def shutdown_all(self) -> None: + """关闭所有当前活跃的沙箱环境并停止生存期监测""" keys = list(self._active_sessions.keys()) for sid in keys: await self.close_session(sid) @@ -179,6 +183,7 @@ driver = nonebot.get_driver() @driver.on_startup async def _startup_sandboxes(): + """异步初始化沙箱,触发后台自动清理孤儿容器""" if not get_llm_config().sandbox.enable_sandbox: return clients = SandboxRegistry.get_all_clients() @@ -199,6 +204,7 @@ async def _startup_sandboxes(): @driver.on_shutdown async def _shutdown_sandboxes(): + """在系统关闭时优雅关闭所有沙箱并清理 Docker 环境""" if not get_llm_config().sandbox.enable_sandbox: return await sandbox_manager.shutdown_all() diff --git a/zhenxun/services/ai/sandbox/registry.py b/zhenxun/services/ai/sandbox/registry.py index aed483c6..209959e6 100644 --- a/zhenxun/services/ai/sandbox/registry.py +++ b/zhenxun/services/ai/sandbox/registry.py @@ -8,11 +8,14 @@ if TYPE_CHECKING: class SandboxRegistry: + """沙箱组件与能力扩展注册中心""" + _clients: ClassVar[dict[str, type["BaseSandboxClient"]]] = {} _extensions: ClassVar[dict[str, type["BaseSandboxExtension"]]] = {} @classmethod def register_client(cls, name: str, client_cls: type["BaseSandboxClient"]) -> None: + """注册一个沙箱驱动客户端类""" if name in cls._clients: logger.warning(f"[SandboxRegistry] 覆盖已存在的沙箱客户端: {name}") cls._clients[name] = client_cls @@ -20,16 +23,19 @@ class SandboxRegistry: @classmethod def get_client_cls(cls, name: str) -> type["BaseSandboxClient"]: + """获取指定名称的沙箱驱动客户端类""" if name not in cls._clients: raise ValueError(f"未找到名为 '{name}' 的沙箱客户端。") return cls._clients[name] @classmethod def get_all_clients(cls) -> dict[str, type["BaseSandboxClient"]]: + """获取所有已注册的沙箱驱动客户端类""" return cls._clients.copy() @classmethod def register_extension(cls, extension_cls: type["BaseSandboxExtension"]) -> None: + """注册一个沙箱功能能力扩展类""" if ( isinstance(getattr(extension_cls, "extension_name", None), property) and extension_cls.extension_name.fget @@ -42,4 +48,5 @@ class SandboxRegistry: @classmethod def get_extension_cls(cls, name: str) -> type["BaseSandboxExtension"] | None: + """获取指定名称的沙箱能力扩展类""" return cls._extensions.get(name) diff --git a/zhenxun/services/ai/sandbox/runtimes.py b/zhenxun/services/ai/sandbox/runtimes.py index e1fefe9c..13b70d14 100644 --- a/zhenxun/services/ai/sandbox/runtimes.py +++ b/zhenxun/services/ai/sandbox/runtimes.py @@ -17,6 +17,7 @@ from zhenxun.utils.utils import infer_plugin_namespace def parse_shebang(script_path: str | Path) -> str | None: + """解析脚本首行的 Shebang,获取解释器名称""" path = Path(script_path) if not path.is_file(): return None @@ -46,6 +47,7 @@ def parse_shebang(script_path: str | Path) -> str | None: def get_execution_command( script_path: str | Path, args: list[str] | None = None ) -> str: + """根据 Shebang 或文件后缀生成对应的脚本执行命令""" path = Path(script_path) interpreter = parse_shebang(path) if not interpreter: @@ -66,6 +68,8 @@ def get_execution_command( class JupyterWSClient: + """管理与单个 Jupyter Kernel WebSocket 通道通信的客户端""" + def __init__( self, http_session: aiohttp.ClientSession, @@ -73,6 +77,7 @@ class JupyterWSClient: ws_url: str, kernel_id: str, ): + """初始化 Jupyter WebSocket 客户端并绑定内核ID""" self.http_session = http_session self.base_url = base_url self.ws_url = ws_url @@ -80,6 +85,7 @@ class JupyterWSClient: self.ws: aiohttp.ClientWebSocketResponse | None = None async def _connect_ws(self) -> aiohttp.ClientWebSocketResponse: + """建立并返回与 Jupyter Kernel 活跃的 WebSocket 连接""" if self.ws is None or self.ws.closed: self.ws = await self.http_session.ws_connect( f"{self.ws_url}/api/kernels/{self.kernel_id}/channels" @@ -88,6 +94,7 @@ class JupyterWSClient: return self.ws async def interrupt(self): + """向 Jupyter Kernel 发送 HTTP POST 中断正在执行的进程""" try: async with self.http_session.post( f"{self.base_url}/api/kernels/{self.kernel_id}/interrupt" @@ -100,6 +107,7 @@ class JupyterWSClient: logger.warning(f"[JupyterWSClient] 中断 Kernel 失败: {e}") async def execute(self, code: str, timeout: int = 30, on_output=None): + """在 Jupyter 核心中执行代码,并通过 WebSocket 接收标准输出、标准错误和图像""" background_tasks: set[asyncio.Task[None]] = set() try: @@ -226,6 +234,7 @@ class JupyterWSClient: ) async def close(self): + """关闭与 Jupyter Kernel 的 WebSocket 通道连接""" if self.ws and not self.ws.closed: await self.ws.close() self.ws = None @@ -235,6 +244,7 @@ class JupyterServerManager: """管理沙箱内的 Jupyter 引擎生命周期及长连接""" def __init__(self, session: BaseSandboxSession): + """初始化 Jupyter 服务生命周期及会话连接管理器""" self.session = session self._http_session: aiohttp.ClientSession | None = None self.base_url = "" @@ -243,6 +253,7 @@ class JupyterServerManager: self._clients: dict[str, JupyterWSClient] = {} async def ensure_started(self, env_vars: dict[str, str] | None = None): + """确保沙箱环境内部已拉起并运行着 Jupyter Server""" if self._is_started: return @@ -258,6 +269,14 @@ class JupyterServerManager: if check_jupyter.exit_code != 0: raise RuntimeError("沙箱内未安装 jupyter-server,请检查 Blueprint") + await self.session.run_process( + "mkdir -p /tmp/jupyter_runtime /tmp/jupyter_data && " + "chmod 777 /tmp/jupyter_runtime /tmp/jupyter_data" + ) + + await self.session.run_process("pkill -9 -f jupyter-server || true") + await asyncio.sleep(0.5) + env_str = " ".join([f"{k}={v}" for k, v in (env_vars or {}).items()]) start_cmd = ( f"nohup env {env_str} jupyter-server " @@ -300,6 +319,7 @@ class JupyterServerManager: ) async def get_client(self, kernel_name: str) -> JupyterWSClient: + """获取并连接到指定内核的 Jupyter WebSocket 客户端""" await self.ensure_started() if kernel_name in self._clients: return self._clients[kernel_name] @@ -321,6 +341,7 @@ class JupyterServerManager: return client async def close(self): + """关闭并清理所有 Jupyter WS 连接及 HTTP 会话资源""" for client in self._clients.values(): await client.close() self._clients.clear() @@ -334,6 +355,7 @@ class BaseCodeExecutor(ABC): """代码执行器抽象基类""" def __init__(self, session: BaseSandboxSession): + """初始化代码执行器,绑定沙箱会话""" self.session = session @abstractmethod @@ -344,6 +366,7 @@ class BaseCodeExecutor(ABC): injected_code: str | None = None, on_output: Callable[[str, bytes], Awaitable[None]] | None = None, ) -> SandboxExecutionResult: + """在沙箱会话中执行指定代码并返回执行结果""" pass @@ -351,6 +374,7 @@ class GenericCLIExecutor(BaseCodeExecutor): """模板驱动的通用代码执行器""" def __init__(self, session, profile: "LanguageProfile"): + """初始化通用命令行代码执行器,传入语言配置模板""" super().__init__(session) self.profile = profile @@ -361,6 +385,7 @@ class GenericCLIExecutor(BaseCodeExecutor): injected_code: str | None = None, on_output: Callable[[str, bytes], Awaitable[None]] | None = None, ) -> SandboxExecutionResult: + """将代码写入临时文件,必要时编译并执行,返回执行结果""" env = None if injected_code: await self.session.write( @@ -403,6 +428,7 @@ class CodeExecutorRegistry: @classmethod def _normalize_lang(cls, language: str) -> str: + """规范化语言名称,转换为统一小写格式并应用别名""" lang_lower = language.lower().strip() return cls._aliases.get(lang_lower, lang_lower) @@ -414,6 +440,7 @@ class CodeExecutorRegistry: is_stateful: bool = False, scope: str | None = None, ) -> None: + """注册一个语言对应的代码执行器构造工厂""" ns = scope if scope is not None else infer_plugin_namespace() lang_norm = cls._normalize_lang(language) @@ -435,7 +462,7 @@ class CodeExecutorRegistry: kernel_name: str, scope: str | None = None, ) -> None: - """语法糖:快速注册一门基于 Jupyter Kernel 的有状态语言""" + """快速注册一个基于 Jupyter 内核的有状态运行语言""" cls.register( language, lambda session: GenericJupyterExecutor(session, kernel_name=kernel_name), @@ -445,6 +472,7 @@ class CodeExecutorRegistry: @classmethod def register_profile(cls, profile: "LanguageProfile") -> None: + """注册一个基于命令行的无状态运行语言配置模板""" cls._profiles[profile.language.lower()] = profile for alias in profile.aliases: cls._profiles[alias.lower()] = profile @@ -454,6 +482,7 @@ class CodeExecutorRegistry: def create_executor( cls, language: str, needs_state: bool, session: Any, namespace: str = "global" ) -> BaseCodeExecutor: + """根据语言和有无状态需求为指定会话创建具体的执行器实例""" lang_norm = cls._normalize_lang(language) for target_ns in [namespace, "global"]: @@ -482,6 +511,7 @@ class CodeExecutorRegistry: @classmethod def get_supported_languages(cls, namespace: str = "global") -> list[str]: + """获取所有目前已注册支持的编程语言列表""" langs = set(cls._executors.get("global", {}).keys()) if namespace in cls._executors: langs.update(cls._executors[namespace].keys()) @@ -493,6 +523,7 @@ class GenericJupyterExecutor(BaseCodeExecutor): """基于 Jupyter 协议的泛化有状态执行器。支持多语言 REPL。""" def __init__(self, session, kernel_name: str = "python3"): + """初始化泛用 Jupyter 有状态执行器,设置默认内核""" super().__init__(session) self.kernel_name = kernel_name self.manager = JupyterServerManager(session) @@ -504,6 +535,7 @@ class GenericJupyterExecutor(BaseCodeExecutor): injected_code: str | None = None, on_output: Callable[[str, bytes], Awaitable[None]] | None = None, ) -> SandboxExecutionResult: + """利用 Jupyter 内核长连接异步执行代码并返回运行结果""" try: await self.manager.ensure_started() except Exception as e: @@ -524,6 +556,7 @@ class GenericJupyterExecutor(BaseCodeExecutor): return result async def close(self): + """关闭 Jupyter 客户端及后台运行环境""" await self.manager.close() diff --git a/zhenxun/services/ai/tools/__init__.py b/zhenxun/services/ai/tools/__init__.py index e356cc5a..7c7fafaa 100644 --- a/zhenxun/services/ai/tools/__init__.py +++ b/zhenxun/services/ai/tools/__init__.py @@ -1,7 +1,7 @@ import nonebot -from .bridges.matcher_bridge import MatcherTool, bind_matcher -from .core.decorators import Rules, tool +from .bridges.matcher_bridge import bind_matcher +from .core.decorators import Rules, tool, toolkit from .core.toolkit import BaseToolkit from .engine.registry import tool_provider_manager from .models import ( @@ -22,11 +22,11 @@ async def _shutdown_mcp_provider(): __all__ = [ "BaseToolkit", - "MatcherTool", "Native", "Rules", "ToolOptions", "ToolResult", "bind_matcher", "tool", + "toolkit", ] diff --git a/zhenxun/services/ai/tools/bridges/delegate.py b/zhenxun/services/ai/tools/bridges/delegate.py index 0dc9763e..90e37aaa 100644 --- a/zhenxun/services/ai/tools/bridges/delegate.py +++ b/zhenxun/services/ai/tools/bridges/delegate.py @@ -41,6 +41,14 @@ class DelegateTool(BaseTool): name: str | None = None, description: str | None = None, ): + """ + 初始化委派工具,将可运行实体包装为子例程工具。 + + 参数: + runnable: 被包装的可运行实体,可以是 Agent、Team 或 Workflow 等。 + name: 自定义工具名称,若为 None 则默认推导为 runnable 实体名。 + description: 工具描述信息,用于指导大模型何时代用此工具。 + """ resolved_name = name or getattr(runnable, "name", "SubRunnable") resolved_desc = description or getattr( runnable, "description", f"将子任务委派给 {resolved_name} 执行" diff --git a/zhenxun/services/ai/tools/bridges/handoff.py b/zhenxun/services/ai/tools/bridges/handoff.py index 20a13a83..0ac799a5 100644 --- a/zhenxun/services/ai/tools/bridges/handoff.py +++ b/zhenxun/services/ai/tools/bridges/handoff.py @@ -2,7 +2,7 @@ from typing import Any from pydantic import BaseModel, Field, create_model -from zhenxun.services.ai.run import RunContext +from zhenxun.services.ai.run.context import RunContext from zhenxun.services.ai.tools.core.tool import BaseTool from zhenxun.services.ai.tools.models import HandoffResult, ToolResult @@ -19,6 +19,14 @@ class HandoffTool(BaseTool): target_description: str, input_schema: type[BaseModel] | Any | None = None, ): + """ + 初始化移交工具,为模型赋予转移对话控制权到指定实体的能力。 + + 参数: + target_name: 被转移的目标接收者(Agent 或负责人)的唯一标识名称。 + target_description: 目标接收者的职责或专长说明,供模型决策是否移交。 + input_schema: 自定义移交数据结构,指定转移时所需携带的结构化参数。 + """ super().__init__( name=f"transfer_to_{target_name}", description=( diff --git a/zhenxun/services/ai/tools/bridges/matcher_bridge.py b/zhenxun/services/ai/tools/bridges/matcher_bridge.py index 538e0651..83571e58 100644 --- a/zhenxun/services/ai/tools/bridges/matcher_bridge.py +++ b/zhenxun/services/ai/tools/bridges/matcher_bridge.py @@ -20,7 +20,7 @@ from pydantic import BaseModel from zhenxun.services.ai.core.exceptions import ToolRetryError from zhenxun.services.ai.core.models import ToolDefinition -from zhenxun.services.ai.run import RunContext +from zhenxun.services.ai.run.context import RunContext from zhenxun.services.ai.tools.core.schema import build_schema_hint from zhenxun.services.ai.tools.core.tool import BaseTool from zhenxun.services.ai.tools.models import EndRunResult, ToolOptions, ToolResult @@ -37,6 +37,14 @@ class MatcherAdapter(ABC): args_schema: type[BaseModel] | None, command_formatter: Callable[[Any], list[Any] | str] | None = None, ): + """ + 初始化 Matcher 桥接适配器。 + + 参数: + matcher: 目标 Matcher 类,负责实际接收 Fake State 并执行命令。 + args_schema: 描述 LLM 需要填写的参数 Pydantic Schema,可以为 None。 + command_formatter: 可选的自定义命令行格式化器,将参数序列化为命令行参数。 + """ self.matcher = matcher self.args_schema = args_schema self.command_formatter = command_formatter @@ -314,6 +322,17 @@ class MatcherTool(BaseTool): args_schema: type[BaseModel] | None, terminal: bool = True, ): + """ + 初始化通用状态机穿透工具。 + + 参数: + matcher: 目标 Matcher 类,由 nonebot 装饰器创建的指令匹配器。 + adapter: 对应的 MatcherAdapter 桥接适配器,用于构建 Fake State。 + name: 工具的唯一标识名称(仅限英文字母、数字及下划线)。 + description: 工具的职责说明,指导大模型选用此工具。 + args_schema: 描述 LLM 需要填写的参数 Pydantic Schema,可以为 None。 + terminal: 是否为终端工具,若为 True 则执行完毕后立即中断大模型循环,默认 True。 + """ # noqa: E501 self.matcher = matcher self.adapter = adapter self.terminal = terminal diff --git a/zhenxun/services/ai/tools/core/__init__.py b/zhenxun/services/ai/tools/core/__init__.py index 5da5882c..683d0116 100644 --- a/zhenxun/services/ai/tools/core/__init__.py +++ b/zhenxun/services/ai/tools/core/__init__.py @@ -1,4 +1,4 @@ -from .decorators import Rules, tool +from .decorators import Rules, tool, toolkit from .schema import ( FieldPermission, RequireAdminLevel, @@ -20,4 +20,5 @@ __all__ = [ "RequireSuperUser", "Rules", "tool", + "toolkit", ] diff --git a/zhenxun/services/ai/tools/core/capabilities.py b/zhenxun/services/ai/tools/core/capabilities.py index 7be4905d..e2570a17 100644 --- a/zhenxun/services/ai/tools/core/capabilities.py +++ b/zhenxun/services/ai/tools/core/capabilities.py @@ -1,16 +1,24 @@ import asyncio from collections.abc import Callable +import json from typing import Any from aiocache import SimpleMemoryCache +from zhenxun.configs.config import Config from zhenxun.services.ai.capabilities import ( AbstractCapability, WrapToolExecuteHandler, WrapToolValidateHandler, ) -from zhenxun.services.ai.run import RunContext +from zhenxun.services.ai.core.exceptions import ( + NeedsInputException, + ToolFatalError, + ToolRetryError, +) +from zhenxun.services.ai.run.context import RunContext from zhenxun.services.ai.run.di import DependencyInjector +from zhenxun.services.ai.run.hitl import HITLController from zhenxun.services.ai.tools.models import ToolResult from zhenxun.services.ai.utils import PermissionUtils from zhenxun.services.log import logger @@ -99,8 +107,6 @@ class ApprovalCapability(AbstractCapability): if tool is None: return await handler(arguments) - import json - args_str = json.dumps(arguments, ensure_ascii=False, indent=2) confirm_msg = f"即将在本地执行高危工具 [{tool_name}]\n参数:\n{args_str}" @@ -110,8 +116,6 @@ class ApprovalCapability(AbstractCapability): await hitl_lock.acquire() try: - from zhenxun.services.ai.run.hitl import HITLController - hitl = HITLController(context) await hitl.ask_confirm(f"⚠️ **安全交互审批**\n\n{confirm_msg}", timeout=60.0) logger.info(f"🛡️ [HITL] 工具 {tool_name} 审批通过。") @@ -275,8 +279,6 @@ class ConfigDependencyCapability(AbstractCapability): async def prepare_tools( self, context: RunContext, tool_defs: list[Any] ) -> list[Any]: - from zhenxun.configs.config import Config - if Config.get_config(self.module, self.key) == self.expected_value: return tool_defs return [] @@ -292,12 +294,6 @@ class InteractiveCapability(AbstractCapability): args: str | dict[str, Any], handler: WrapToolValidateHandler, ) -> dict[str, Any]: - from zhenxun.services.ai.core.exceptions import ( - NeedsInputException, - ToolFatalError, - ToolRetryError, - ) - tool = context.call.current_tool current_kwargs = dict(args) if isinstance(args, dict) else args @@ -322,8 +318,6 @@ class InteractiveCapability(AbstractCapability): "请发送文本补充,或回复“取消”中止。" ) try: - from zhenxun.services.ai.run.hitl import HITLController - hitl = HITLController(context) user_input = await hitl.ask_text(prompt_msg, timeout=60.0) except ToolFatalError: diff --git a/zhenxun/services/ai/tools/core/decorators.py b/zhenxun/services/ai/tools/core/decorators.py index d08c6ae1..3ea6c4d3 100644 --- a/zhenxun/services/ai/tools/core/decorators.py +++ b/zhenxun/services/ai/tools/core/decorators.py @@ -4,53 +4,99 @@ from collections.abc import Callable from typing import Any from zhenxun.services.ai.capabilities import AbstractCapability -from zhenxun.services.ai.run import RunContext +from zhenxun.services.ai.run.context import RunContext +from zhenxun.services.ai.tools.core.capabilities import ( + AdminLevelCapability, + ApprovalCapability, + CacheCapability, + FallbackCapability, + GroupOnlyCapability, + InteractiveCapability, + LifecycleCapability, + SuperuserCapability, +) from zhenxun.services.ai.tools.models import ToolOptions +from zhenxun.utils.pydantic_compat import model_copy -def _update_settings(func: Callable, **kwargs) -> Callable: - """叠加协议底层:创建或更新函数的 __tool_settings__""" - if not hasattr(func, "__tool_settings__"): - setattr(func, "__tool_settings__", ToolOptions()) - settings: ToolOptions = getattr(func, "__tool_settings__") - for k, v in kwargs.items(): - if k == "capabilities": - settings.capabilities.extend(v) - elif k == "metadata": - settings.metadata.update(v) +def toolkit( + rules: list[ToolOptions] | ToolOptions | None = None, + prefix: str = "", + instructions: str | None = None, +): + """ + 类级别的工具箱装饰器,用于向内部所有 @tool 统一下发配置规则、名称前缀以及工具箱级系统提示词说明。 + + 参数: + rules: 声明式规则集合或单个规则,应用于工具箱内所有工具。 + prefix: 工具箱内所有工具的名称前缀,通常以下划线结尾。 + instructions: 工具箱级别的系统提示词补充说明,大模型可见。 + + 返回: + Callable: 装饰器函数,接收一个类并返回该类。 + """ # noqa: E501 + + def decorator(cls): + merged_options = ToolOptions() + if rules: + rule_list = rules if isinstance(rules, list) else [rules] + for r in rule_list: + if isinstance(r, ToolOptions): + merged_options = merged_options.merge(r) + + from zhenxun.services.ai.tools.models import ToolkitConfig + + base_config = getattr(cls, "_default_config", None) or ToolkitConfig() + new_config = model_copy(base_config, deep=True) + + if prefix: + new_config.prefix = prefix + + if new_config.shared_options: + new_config.shared_options = new_config.shared_options.merge(merged_options) else: - setattr(settings, k, v) - return func + new_config.shared_options = merged_options + + cls._default_config = new_config + + if instructions is not None: + cls.default_instructions = instructions + + return cls + + return decorator def require_sandbox( python_packages: list[str] | None = None, node_packages: list[str] | None = None, system_packages: list[str] | None = None, -): +) -> ToolOptions: """显式声明该工具所需安装的沙箱环境依赖""" - - def decorator(func): - setattr( - func, - "__sandbox_requirements__", - { - "python": python_packages or [], - "node": node_packages or [], - "system": system_packages or [], - }, - ) - return func - - return decorator + return ToolOptions( + sandbox_requirements={ + "python": python_packages or [], + "node": node_packages or [], + "system": system_packages or [], + } + ) class ToolkitMethodDescriptor: """用于 Toolkit 类方法的描述符,确保在实例化时绑定正确的 self 并生成 FunctionTool""" def __init__( - self, func: Callable, name: str, description: str, settings: ToolOptions + self, func: Callable, name: str, description: str | None, settings: ToolOptions ): + """ + 初始化 Toolkit 类方法描述符。 + + 参数: + func: 底层的 Python 可执行函数对象。 + name: 大模型识别与调用的工具名。 + description: 大模型阅读的工具功能描述说明。 + settings: 工具的声明式高阶配置项 (ToolOptions)。 + """ self.func = func self.name = name self.description = description @@ -64,7 +110,6 @@ class ToolkitMethodDescriptor: import types from zhenxun.services.ai.tools.core.tool import FunctionTool - from zhenxun.utils.pydantic_compat import model_copy bound_func = types.MethodType(self.func, instance) return FunctionTool( @@ -82,6 +127,7 @@ class ToolkitMethodDescriptor: def tool( name: str | None = None, description: str | None = None, + rules: list[ToolOptions] | ToolOptions | None = None, settings: ToolOptions | None = None, tags: list[str] | None = None, auto_register: bool = False, @@ -91,23 +137,26 @@ def tool( 将普通函数或类方法注册为 LLM 工具的统一装饰器大一统。 参数: - name: 工具的名称(英文字母及下划线)。 - 大模型将看到此名称。如果为空则默认使用函数名。 - description: 工具描述。 - 大模型将基于此决定何时、如何使用该工具。如果为空,将读取函数 docstring。 - settings: 工具的高阶配置对象 (ToolOptions)。 - 用于控制极速缓存、人工审批拦截、静默执行等扩展能力。 - tags: 工具的标签列表。 - 用于被 Agent 的智能字符串路由识别并进行能力注入。 - auto_register: 是否自动注册到当前命名空间的工具注册表。 - 默认为 False,开发者需要显式传递实例。若为 True 方可通过字符串调用。 - require_prefix: 是否自动添加插件命名空间前缀(仅对游离函数生效)。默认为 False。 + name: 工具的名称(英文字母及下划线),大模型将看到此名称,如果为空则默认使用函数名。 + description: 工具描述,大模型将基于此决定何时、如何使用该工具,如果为空,将读取函数 docstring。 + rules: 声明式规则集合预设列表,用于接收通过 Rules.xxx() 生成的策略载荷并组合。 + settings: 工具的高阶配置对象 (ToolOptions),用于控制极速缓存、人工审批拦截、静默执行等扩展能力。 + tags: 工具的标签列表,用于被 Agent 的智能字符串路由识别并进行能力注入。 + auto_register: 是否自动注册到当前命名空间的工具注册表,默认为 False,若为 True 方可通过字符串调用。 + require_prefix: 是否自动添加插件命名空间前缀(仅对游离函数生效),默认为 False。 返回: Callable | FunctionTool: 包装后的函数或方法。 - """ + """ # noqa: E501 def decorator(func: Callable): base_settings = settings or ToolOptions() + + if rules: + rule_list = rules if isinstance(rules, list) else [rules] + for r in rule_list: + if isinstance(r, ToolOptions): + base_settings = base_settings.merge(r) + if tags: base_settings.tags = list(set(base_settings.tags + tags)) existing_settings = getattr(func, "__tool_settings__", None) @@ -116,10 +165,13 @@ def tool( reqs = getattr(func, "__sandbox_requirements__", None) if reqs: - base_settings.sandbox_requirements = reqs + if not base_settings.sandbox_requirements: + base_settings.sandbox_requirements = reqs + else: + base_settings.sandbox_requirements.update(reqs) tool_name = name or func.__name__ - tool_desc = description or func.__doc__ or "未提供描述" + tool_desc = description import inspect @@ -172,30 +224,19 @@ def tool( return decorator -def with_cache(ttl: int = 3600, cache_function: Callable | None = None): +def with_cache(ttl: int = 3600, cache_function: Callable | None = None) -> ToolOptions: """开启极速缓存,阻止参数相同的重复请求发往底层。""" - - from zhenxun.services.ai.tools.core.capabilities import CacheCapability - - def decorator(func: Callable): - return _update_settings( - func, - capabilities=[CacheCapability(ttl=ttl, cache_function=cache_function)], - ) - - return decorator + return ToolOptions( + capabilities=[CacheCapability(ttl=ttl, cache_function=cache_function)] + ) -def silent(): +def silent() -> ToolOptions: """静默执行。工具执行过程与结果不会作为界面流渲染给用户,仅作大模型内部参考。""" - - def decorator(func: Callable): - return _update_settings(func, silent=True) - - return decorator + return ToolOptions(silent=True) -def direct_reply(): +def direct_reply() -> ToolOptions: """直出模式。工具执行完毕后强制中断大模型的思考循环,将工具输出作为最终回答返回。""" class DirectReplyCapability(AbstractCapability): @@ -209,139 +250,70 @@ def direct_reply(): return EndRunResult(output=result.output) return EndRunResult(output=result) - def decorator(func: Callable): - return _update_settings(func, capabilities=[DirectReplyCapability()]) - - return decorator + return ToolOptions(capabilities=[DirectReplyCapability()]) -def interactive(): +def interactive() -> ToolOptions: """开启交互式参数补全。如果参数缺失或校验失败, 会主动通过 Bot 向用户提问要求补全。""" - from zhenxun.services.ai.tools.core.capabilities import InteractiveCapability - - def decorator(func: Callable): - return _update_settings(func, capabilities=[InteractiveCapability()]) - - return decorator + return ToolOptions(capabilities=[InteractiveCapability()]) -def require_approval(): +def require_approval() -> ToolOptions: """高危操作标记。调用前将拦截并发送至群组要求超管人工审核。""" - from zhenxun.services.ai.tools.core.capabilities import ApprovalCapability - - def decorator(func: Callable): - return _update_settings(func, capabilities=[ApprovalCapability()]) - - return decorator + return ToolOptions(capabilities=[ApprovalCapability()]) -def fallback(tool_name: str): +def fallback(tool_name: str) -> ToolOptions: """降级备用工具。主工具执行失败时自动路由至备用工具。""" - from zhenxun.services.ai.tools.core.capabilities import FallbackCapability - - def decorator(func: Callable): - return _update_settings( - func, - capabilities=[FallbackCapability(fallback_tool_name=tool_name)], - ) - - return decorator + return ToolOptions(capabilities=[FallbackCapability(fallback_tool_name=tool_name)]) -def before_execute(hook_func: Callable): +def before_execute(hook_func: Callable) -> ToolOptions: """生命周期拦截:在工具即将执行前触发,可在此篡改全局状态或通过依赖注入访问当前参数。""" - from zhenxun.services.ai.tools.core.capabilities import LifecycleCapability - - def decorator(func: Callable): - return _update_settings( - func, capabilities=[LifecycleCapability(before_execute=hook_func)] - ) - - return decorator + return ToolOptions(capabilities=[LifecycleCapability(before_execute=hook_func)]) -def after_execute(hook_func: Callable): +def after_execute(hook_func: Callable) -> ToolOptions: """生命周期拦截:在工具执行完毕后触发,可在此修改将发往大模型的返回结果。""" - from zhenxun.services.ai.tools.core.capabilities import LifecycleCapability - - def decorator(func: Callable): - return _update_settings( - func, capabilities=[LifecycleCapability(after_execute=hook_func)] - ) - - return decorator + return ToolOptions(capabilities=[LifecycleCapability(after_execute=hook_func)]) -def validate_args(hook_func: Callable): +def validate_args(hook_func: Callable) -> ToolOptions: """生命周期拦截:在工具反序列化校验前触发。如果校验失败直接抛出异常,将触发大模型重试机制。""" - from zhenxun.services.ai.tools.core.capabilities import LifecycleCapability - - def decorator(func: Callable): - return _update_settings( - func, capabilities=[LifecycleCapability(validate_args=hook_func)] - ) - - return decorator + return ToolOptions(capabilities=[LifecycleCapability(validate_args=hook_func)]) -def prepare_tool(hook_func: Callable): +def prepare_tool(hook_func: Callable) -> ToolOptions: """生命周期拦截:在向大模型渲染 JSON Schema 前触发, 可在此动态修改该工具的定义或返回 None 对大模型隐藏。""" - from zhenxun.services.ai.tools.core.capabilities import LifecycleCapability - - def decorator(func: Callable): - return _update_settings( - func, capabilities=[LifecycleCapability(prepare_tool=hook_func)] - ) - - return decorator + return ToolOptions(capabilities=[LifecycleCapability(prepare_tool=hook_func)]) -def require_superuser(): +def require_superuser() -> ToolOptions: """仅限超级管理员使用。若调用者无权限,该工具将直接对大模型隐藏。""" - from zhenxun.services.ai.tools.core.capabilities import SuperuserCapability - - def decorator(func: Callable): - return _update_settings(func, capabilities=[SuperuserCapability()]) - - return decorator + return ToolOptions(capabilities=[SuperuserCapability()]) -def require_admin_level(min_level: int = 1): +def require_admin_level(min_level: int = 1) -> ToolOptions: """仅限满足指定群聊权限等级的用户使用。""" - from zhenxun.services.ai.tools.core.capabilities import AdminLevelCapability - - def decorator(func: Callable): - return _update_settings( - func, - capabilities=[AdminLevelCapability(min_level)], - metadata={"admin_level": min_level}, - ) - - return decorator + return ToolOptions( + capabilities=[AdminLevelCapability(min_level)], + metadata={"admin_level": min_level}, + ) -def require_group(): +def require_group() -> ToolOptions: """限制该工具仅能在群聊环境中被大模型调用。""" - from zhenxun.services.ai.tools.core.capabilities import GroupOnlyCapability - - def decorator(func: Callable): - return _update_settings(func, capabilities=[GroupOnlyCapability()]) - - return decorator + return ToolOptions(capabilities=[GroupOnlyCapability()]) -def require_minimum_gold(amount: int): +def require_minimum_gold(amount: int) -> ToolOptions: """需满足一定的金币余额才可调用,同时执行后会自动扣除该金币。""" - - def decorator(func: Callable): - return _update_settings(func, metadata={"cost_gold": amount}) - - return decorator + return ToolOptions(metadata={"cost_gold": amount}) -def require_session_state(key: str, expected_value: Any = None): +def require_session_state(key: str, expected_value: Any = None) -> ToolOptions: """需要当前执行上下文中包含指定状态变量,否则隐藏工具""" class StateDependencyCapability(AbstractCapability): @@ -356,10 +328,7 @@ def require_session_state(key: str, expected_value: Any = None): return tool_defs return [] - def decorator(func: Callable): - return _update_settings(func, capabilities=[StateDependencyCapability()]) - - return decorator + return ToolOptions(capabilities=[StateDependencyCapability()]) class Rules: @@ -368,6 +337,15 @@ class Rules: 用于控制工具的极速缓存、沙箱要求、前端静默以及人机交互审批等。 """ + @staticmethod + def combine(*rules: ToolOptions) -> ToolOptions: + """打包组合多个规则预设载荷""" + base = ToolOptions() + for r in rules: + if isinstance(r, ToolOptions): + base = base.merge(r) + return base + sandbox = staticmethod(require_sandbox) """环境:声明执行该工具所需的物理沙箱及 pip/npm 包依赖""" diff --git a/zhenxun/services/ai/tools/core/tool.py b/zhenxun/services/ai/tools/core/tool.py index 8ebae12b..06b4ace7 100644 --- a/zhenxun/services/ai/tools/core/tool.py +++ b/zhenxun/services/ai/tools/core/tool.py @@ -1,4 +1,5 @@ from collections.abc import Callable +import copy import hashlib import inspect import json @@ -6,18 +7,22 @@ from typing import Any, Literal from pydantic import BaseModel, ValidationError +from zhenxun.services.ai.capabilities import CombinedCapability from zhenxun.services.ai.core.exceptions import ( + ControlFlowExit, NeedsInputException, ToolFatalError, ToolRetryError, ) from zhenxun.services.ai.core.models import ToolDefinition -from zhenxun.services.ai.run import RunContext +from zhenxun.services.ai.run.context import RunContext +from zhenxun.services.ai.tools.core.capabilities import InteractiveCapability from zhenxun.services.ai.tools.models import ( ResolvedToolPayload, ToolOptions, ToolResult, ) +from zhenxun.services.ai.utils.utils import wrap_to_async from zhenxun.services.log import logger from zhenxun.utils.pydantic_compat import model_dump, model_json_schema, model_validate @@ -112,8 +117,6 @@ class BaseTool: def clone_with_options(self, override: Any) -> "BaseTool": """创建当前工具的浅拷贝,并覆盖 ToolOptions 和基本属性""" - import copy - new_tool = copy.copy(self) if hasattr(override, "to_tool_options"): @@ -154,8 +157,6 @@ class BaseTool: ) schema_hint = build_schema_hint(validation_model) - from zhenxun.services.ai.tools.core.capabilities import InteractiveCapability - if any( isinstance(c, InteractiveCapability) for c in self.settings.capabilities ): @@ -225,10 +226,6 @@ class BaseTool: ) if context and self.settings.capabilities: - from zhenxun.services.ai.capabilities import ( - CombinedCapability, - ) - combined_cap = CombinedCapability(self.settings.capabilities) defs = await combined_cap.prepare_tools(context, [tool_def]) if not defs: @@ -239,8 +236,6 @@ class BaseTool: async def resolve(self, context: RunContext | None = None) -> ResolvedToolPayload: """实现 ToolResolvable 协议,将自身解析为标准 Payload""" - import copy - definition = await self.get_definition(context) if definition is None: return ResolvedToolPayload() @@ -288,8 +283,6 @@ class BaseTool: try: return await self._core_execution(context_to_pass, **kwargs) except Exception as e: - from zhenxun.services.ai.core.exceptions import ControlFlowExit - if not isinstance(e, ControlFlowExit): logger.error(f"工具 {self.name} 执行抛出异常,将交由底层引擎处理: {e}") raise @@ -381,8 +374,6 @@ class FunctionTool(BaseTool): super().__init__(name=name, description=description, settings=settings) self._original_func = func - from zhenxun.services.ai.utils.utils import wrap_to_async - self.__name__ = self.name self._func = wrap_to_async(func) self._schema_built = False diff --git a/zhenxun/services/ai/tools/core/toolkit.py b/zhenxun/services/ai/tools/core/toolkit.py index 9cc2769a..9d53066f 100644 --- a/zhenxun/services/ai/tools/core/toolkit.py +++ b/zhenxun/services/ai/tools/core/toolkit.py @@ -1,9 +1,10 @@ from collections.abc import Callable import copy +import inspect from typing import Any, ClassVar from zhenxun.services.ai.core.protocols.tool import ToolExecutable -from zhenxun.services.ai.run import RunContext +from zhenxun.services.ai.run.context import RunContext from zhenxun.services.ai.run.di import DependencyInjector from zhenxun.services.ai.tools.models import ( ResolvedToolPayload, @@ -51,19 +52,27 @@ class BaseToolkit: tools: 除了带有 @tool 装饰器的方法外,要额外动态注入的独立工具列表。 instructions: 当前工具箱实例的系统提示词补充说明。 """ + base_config = getattr(self.__class__, "_default_config", ToolkitConfig()) + if config is not None: merged_dict = { - **model_dump(self._default_config, exclude_unset=True), + **model_dump(base_config, exclude_unset=True), **model_dump(config, exclude_unset=True), } self.config = ToolkitConfig(**merged_dict) else: - self.config = ToolkitConfig( - prefix=prefix, - include=include, - exclude=exclude, - shared_options=shared_options, - ) + overrides = {} + if prefix is not None: + overrides["prefix"] = prefix + if include is not None: + overrides["include"] = include + if exclude is not None: + overrides["exclude"] = exclude + if shared_options is not None: + overrides["shared_options"] = shared_options + + merged_dict = {**model_dump(base_config, exclude_unset=True), **overrides} + self.config = ToolkitConfig(**merged_dict) if self.config.prefix is None: self.config.prefix = "" @@ -101,6 +110,7 @@ class BaseToolkit: def _apply_global_settings( self, original_name: str, settings: ToolOptions ) -> ToolOptions: + """合并工具箱全局配置到单个工具的配置项中""" if self.config.shared_options: settings = self.config.shared_options.merge(settings) @@ -143,16 +153,18 @@ class BaseToolkit: ) async def get_tools(self, context: RunContext | None = None) -> dict[str, BaseTool]: + """获取该工具箱中所有合法注册的工具实例映射表""" if self._cached_tools is not None: return self._cached_tools tools_dict: dict[str, BaseTool] = {} for t in self._injected_tools: t.parent_toolkit = self + if self.config.prefix and not t.name.startswith(self.config.prefix): + t.name = f"{self.config.prefix}{t.name}" + t.settings = self._apply_global_settings(t.name, t.settings) tools_dict[t.name] = t - import inspect - for name, member in inspect.getmembers(self.__class__): if hasattr(member, "__toolkit_tool__"): bound_tool = getattr(self, name) @@ -201,6 +213,7 @@ class BaseToolkit: return f"<{tag_name}>\n{text}\n" def prefixed(self, prefix: str) -> "BaseToolkit": + """克隆工具箱并为其中所有工具追加统一的前缀""" new_tk = copy.copy(self) new_tk.config = model_copy(self.config, deep=True) current_prefix = new_tk.config.prefix or "" @@ -209,6 +222,7 @@ class BaseToolkit: return new_tk def filtered(self, filter_func: Callable[[BaseTool], bool]) -> "BaseToolkit": + """克隆工具箱并通过自定义过滤器筛选其中的工具""" new_tk = copy.copy(self) new_tk.config = model_copy(self.config, deep=True) old_filter = self._instance_filter diff --git a/zhenxun/services/ai/tools/engine/executor.py b/zhenxun/services/ai/tools/engine/executor.py index ee1a27b9..ec4f3efd 100644 --- a/zhenxun/services/ai/tools/engine/executor.py +++ b/zhenxun/services/ai/tools/engine/executor.py @@ -13,6 +13,7 @@ from nonebot.adapters import Message as PlatformMessage from zhenxun.services.ai.capabilities import CombinedCapability from zhenxun.services.ai.core.exceptions import ( + ControlFlowExit, ToolFatalError, ToolRetryError, ) @@ -24,11 +25,14 @@ from zhenxun.services.ai.core.stream_events import ( ToolStreamChunkEvent, UserCustomEvent, ) +from zhenxun.services.ai.message_builder import MessageBuilder from zhenxun.services.ai.run.context import RunContext from zhenxun.services.ai.run.di import DependencyInjector from zhenxun.services.ai.tools.models import ( + StateSyncResult, ToolOptions, ToolResult, + ToolResultChunk, ValidatedToolCall, ) from zhenxun.services.log import logger @@ -144,8 +148,6 @@ class ToolExecutor: available_tools: "ToolCollection | dict[str, Any] | None" = None, ) -> RunContext: """准备/克隆工具调用所使用的隔离 RunContext""" - from zhenxun.services.ai.run import RunContext - safe_context = ( context.clone_for_tool_call(tool_call_id, tool_name) if context @@ -256,8 +258,6 @@ class ToolExecutor: available_tools = {} tool_name = validated.call.tool_name if not validated.args_valid or validated.tool is None: - from zhenxun.services.ai.core.exceptions import ControlFlowExit - if isinstance(validated.validation_error, ControlFlowExit): raise validated.validation_error @@ -303,16 +303,12 @@ class ToolExecutor: if not isinstance(result, ToolResult): result = ToolResult(output=result) - from zhenxun.services.ai.tools.models import StateSyncResult - if isinstance(result, StateSyncResult) and result.state_notice: if safe_context: safe_context.run.add_system_prompt( f"[系统通知(状态同步)]:{result.state_notice}" ) except BaseException as e: - from zhenxun.services.ai.core.exceptions import ControlFlowExit - if isinstance(e, ControlFlowExit): raise e if isinstance(e, asyncio.CancelledError): @@ -394,8 +390,6 @@ class ToolExecutor: func_name = original_call.tool_name if isinstance(result_pair, BaseException): - from zhenxun.services.ai.core.exceptions import ControlFlowExit - if isinstance(result_pair, ControlFlowExit): raise result_pair if isinstance(result_pair, asyncio.CancelledError): @@ -438,8 +432,6 @@ class ToolExecutionPolicy: """ 计算当前工具的绝对最大重试次数。 优先使用工具级配置 (ToolOptions.max_retries),如果未设置,则使用全局配置。 - 由于重试机制是保证 Agent 稳定性的防线, - 即使全局为 0,底层默认也会给予至少 1 次的机会。 """ tool_retries = getattr(self.settings, "max_retries", None) if tool_retries is not None: @@ -496,8 +488,6 @@ class NativeToolRunner(ToolRunner): if isinstance(chunk, ToolResult): res = chunk else: - from zhenxun.services.ai.tools.models import ToolResultChunk - chunk_obj = ( chunk if isinstance(chunk, ToolResultChunk) @@ -525,8 +515,6 @@ class NativeToolRunner(ToolRunner): final_result = res else: if str(type(res)).find("Message") != -1: - from zhenxun.services.ai.message_builder import MessageBuilder - uni_msg = ( MessageBuilder.message_to_unimessage(res) if isinstance(res, PlatformMessage) diff --git a/zhenxun/services/ai/tools/engine/registry.py b/zhenxun/services/ai/tools/engine/registry.py index 0b7c7b58..5521db79 100644 --- a/zhenxun/services/ai/tools/engine/registry.py +++ b/zhenxun/services/ai/tools/engine/registry.py @@ -26,6 +26,7 @@ class ToolCollection(list[T], Generic[T]): """支持按索引和按名称获取的工具集合 (List + Dict)""" def __init__(self, iterable: Iterable[T] | None = None): + """初始化工具集合。""" super().__init__(iterable or []) self._name_cache: dict[str, T] = {} self._build_name_cache() @@ -79,9 +80,11 @@ class ToolCollection(list[T], Generic[T]): self._name_cache[value.name.lower()] = value def get(self, key: str, default: Any = None) -> T | Any: + """通过名称获取工具,若不存在则返回默认值。""" return self._name_cache.get(key.lower(), default) def append(self, object: T) -> None: + """向集合中添加工具,并更新名称缓存。""" name = object.name if name.lower() in self._name_cache: old_tool = self._name_cache[name.lower()] @@ -95,16 +98,19 @@ class ToolCollection(list[T], Generic[T]): self._name_cache[name.lower()] = object def extend(self, iterable: Iterable[T]) -> None: + """批量添加工具到集合中。""" for t in iterable: self.append(t) def remove(self, value: T) -> None: + """从集合中移除指定工具,并同步更新缓存。""" super().remove(value) name = getattr(value, "name", None) if name and name.lower() in self._name_cache: del self._name_cache[name.lower()] def pop(self, index: SupportsIndex = -1) -> T: + """弹出指定位置的工具,并从缓存中移除。""" tool = super().pop(index) name = getattr(tool, "name", None) if name and name.lower() in self._name_cache: @@ -112,6 +118,7 @@ class ToolCollection(list[T], Generic[T]): return tool def filter_by_names(self, names: list[str] | None = None) -> "ToolCollection[T]": + """根据名称列表筛选并返回新的工具子集合。""" if names is None: return self return ToolCollection( @@ -123,16 +130,20 @@ class ToolCollection(list[T], Generic[T]): ) def clear(self) -> None: + """清空集合及所有名称缓存。""" super().clear() self._name_cache.clear() def keys(self): + """获取所有工具名称缓存的键。""" return self._name_cache.keys() def values(self): + """获取所有已缓存的工具实例。""" return self._name_cache.values() def items(self): + """获取所有工具名称与实例的键值对。""" return self._name_cache.items() @@ -145,11 +156,13 @@ class _StringResolver: def __init__( self, name: str, manager: "ToolProviderManager", default_namespace: str ): + """初始化字符串路由解析器。""" self.name = name self.manager = manager self.default_namespace = default_namespace async def resolve(self, context: RunContext | None = None) -> ResolvedToolPayload: + """解析字符串路由并返回匹配的工具载荷。""" if self.name in self.manager._macro_resolvers: resolver = self.manager._macro_resolvers[self.name] resolved = ( @@ -192,11 +205,13 @@ class _QueryResolver: def __init__( self, query: Query, manager: "ToolProviderManager", default_namespace: str ): + """初始化查询对象路由解析器。""" self.query = query self.manager = manager self.default_namespace = default_namespace async def resolve(self, context: RunContext | None = None) -> ResolvedToolPayload: + """执行查询以解析并返回匹配的工具载荷。""" payload = ResolvedToolPayload() namespaces_to_search = [] @@ -235,10 +250,12 @@ class _CallableResolver: """ def __init__(self, func: Callable, manager: "ToolProviderManager"): + """初始化可调用对象路由解析器。""" self.func = func self.manager = manager async def resolve(self, context: RunContext | None = None) -> ResolvedToolPayload: + """将可调用对象转换为函数工具并返回其解析载荷。""" for candidate in ( getattr(self.func, "__tool_name__", None), getattr(self.func, "__name__", None), @@ -263,12 +280,14 @@ class _TypeAdapterResolver: manager: "ToolProviderManager", default_namespace: str, ): + """初始化类型适配器解析器。""" self.item = item self.resolver_func = resolver_func self.manager = manager self.default_namespace = default_namespace async def resolve(self, context: RunContext | None = None) -> ResolvedToolPayload: + """执行类型适配器函数并解析返回对应的工具载荷。""" resolved = ( await self.resolver_func(self.item) if is_coroutine_callable(self.resolver_func) @@ -285,11 +304,13 @@ class ToolProviderManager: _instance: "ToolProviderManager | None" = None def __new__(cls) -> Self: + """单例模式的实例创建方法。""" if cls._instance is None: cls._instance = super().__new__(cls) return cast(Self, cls._instance) def __init__(self): + """初始化工具提供者管理器。""" if hasattr(self, "_initialized") and self._initialized: return @@ -304,11 +325,13 @@ class ToolProviderManager: self._type_resolvers: dict[type, Callable] = {} def register_macro_resolver(self, macro_str: str, resolver_func: Callable) -> None: + """注册宏解析器函数。""" self._macro_resolvers[macro_str] = resolver_func def register_type_resolver( self, target_type: type, resolver_func: Callable ) -> None: + """注册特定类型的工具解析函数。""" self._type_resolvers[target_type] = resolver_func def register(self, provider: ToolProvider): @@ -445,6 +468,7 @@ class ToolProviderManager: excluded_servers: list[str] | None = None, namespaces: list[str] | None = None, ) -> ToolCollection: + """获取已解析完成的所有可用工具集合。""" has_filters = ( allowed_servers is not None or excluded_servers is not None @@ -462,11 +486,13 @@ class ToolProviderManager: return tools async def resolve_specific_tools(self, tool_names: list[str]) -> ToolCollection: + """根据名称列表检索并返回特定的工具集合。""" return await self._query_engine(names=tool_names, include_providers=True) async def get_function_tools( self, names: list[str] | None = None ) -> ToolCollection: + """获取本地注册的所有函数工具集合。""" return await self._query_engine(names=names, include_providers=False) def _normalize_to_resolver(self, item: Any, default_ns: str) -> Any: diff --git a/zhenxun/services/ai/tools/providers/builtin/hitl.py b/zhenxun/services/ai/tools/providers/builtin/hitl.py index 9af187fb..63c2b83c 100644 --- a/zhenxun/services/ai/tools/providers/builtin/hitl.py +++ b/zhenxun/services/ai/tools/providers/builtin/hitl.py @@ -1,6 +1,7 @@ from zhenxun.services.ai.core.exceptions import AbortException, ToolFatalError from zhenxun.services.ai.core.stream_events import ToolStreamChunkEvent -from zhenxun.services.ai.run import Inject, RunContext +from zhenxun.services.ai.run.context import RunContext +from zhenxun.services.ai.run.di import Inject from zhenxun.services.ai.tools.core.decorators import tool from zhenxun.services.ai.tools.core.toolkit import BaseToolkit from zhenxun.services.ai.tools.models import ToolResult diff --git a/zhenxun/services/ai/tools/providers/builtin/rest_api.py b/zhenxun/services/ai/tools/providers/builtin/rest_api.py index 7e0c7f04..b46da774 100644 --- a/zhenxun/services/ai/tools/providers/builtin/rest_api.py +++ b/zhenxun/services/ai/tools/providers/builtin/rest_api.py @@ -4,7 +4,7 @@ from typing import Any, Literal import httpx from zhenxun.services.ai.core.stream_events import ToolStreamChunkEvent -from zhenxun.services.ai.run import RunContext +from zhenxun.services.ai.run.context import RunContext from zhenxun.services.ai.tools.core.decorators import tool from zhenxun.services.ai.tools.core.toolkit import BaseToolkit from zhenxun.services.ai.tools.models import ToolResult diff --git a/zhenxun/services/ai/tools/providers/builtin/sandbox.py b/zhenxun/services/ai/tools/providers/builtin/sandbox.py index 9c2f9b20..c2169e07 100644 --- a/zhenxun/services/ai/tools/providers/builtin/sandbox.py +++ b/zhenxun/services/ai/tools/providers/builtin/sandbox.py @@ -1,31 +1,19 @@ import asyncio -from typing import Any, Protocol +from typing import Any from zhenxun.services.ai.core.stream_events import ToolStreamChunkEvent, UserCustomEvent -from zhenxun.services.ai.run import Inject, RunContext +from zhenxun.services.ai.run.context import RunContext +from zhenxun.services.ai.run.di import Inject from zhenxun.services.ai.sandbox.models import ( SandboxBlueprint, - SandboxExecutionResult, ) -from zhenxun.services.ai.tools.core.decorators import silent, tool +from zhenxun.services.ai.tools.core.decorators import Rules, tool from zhenxun.services.ai.tools.core.toolkit import BaseToolkit from zhenxun.services.ai.tools.models import ToolResult from zhenxun.services.log import logger from zhenxun.utils.pydantic_compat import model_copy -class PythonPluginProtocol(Protocol): - @property - def supports_state(self) -> bool: ... - - async def execute( - self, - code: str, - timeout: int = 30, - injected_code: str | None = None, - ) -> SandboxExecutionResult: ... - - class SandboxToolkit(BaseToolkit): default_prefix = "" @@ -383,8 +371,8 @@ class SandboxToolkit(BaseToolkit): @tool( name="write_sandbox_file", description="将文本内容写入沙箱文件系统中,支持保存大块数据或配置,避免超过对话上下文。", + rules=[Rules.silent()], ) - @silent() async def write_sandbox_file( self, path: str, content: str, context: RunContext, sandbox: Inject.Sandbox ) -> ToolResult: @@ -405,8 +393,8 @@ class SandboxToolkit(BaseToolkit): @tool( name="read_sandbox_file", description="从沙箱文件系统中读取指定文件的文本内容。", + rules=[Rules.silent()], ) - @silent() async def read_sandbox_file( self, path: str, context: RunContext, sandbox: Inject.Sandbox ) -> ToolResult: diff --git a/zhenxun/services/ai/tools/providers/mcp/toolkit.py b/zhenxun/services/ai/tools/providers/mcp/toolkit.py index 2e286c6e..2d557673 100644 --- a/zhenxun/services/ai/tools/providers/mcp/toolkit.py +++ b/zhenxun/services/ai/tools/providers/mcp/toolkit.py @@ -17,7 +17,7 @@ from pydantic import ValidationError from zhenxun.services.ai.core.models import ToolDefinition from zhenxun.services.ai.core.stream_events import UserCustomEvent -from zhenxun.services.ai.run import RunContext +from zhenxun.services.ai.run.context import RunContext from zhenxun.services.ai.sandbox.addons.base import BaseMcpProxyExtension from zhenxun.services.ai.sandbox.models import SandboxBlueprint from zhenxun.services.ai.tools.core.tool import BaseTool @@ -120,6 +120,16 @@ class MCPRemoteTool(BaseTool): parameters: dict, toolkit: "MCPToolkit", ): + """ + 初始化远端 MCP 工具在本地的代理对象。 + + 参数: + name: 工具在本地注册的唯一标识名称。 + original_tool_name: 该工具在远端 MCP 服务器上的原始名称。 + description: 工具的描述信息,用于指导大模型选用此工具。 + parameters: 工具参数的 JSON 模式描述。 + toolkit: 所属的 MCPToolkit 工具箱实例。 + """ super().__init__(name=name, description=description) self.original_tool_name = original_tool_name self.parameters = parameters @@ -132,7 +142,7 @@ class MCPRemoteTool(BaseTool): @property def effective_ttl(self) -> float: - """动态获取当前最新的 TTL 配置,支持无缝热重载""" + """获取当前工具关联的有效生存期 TTL""" try: from zhenxun.services.ai.config import get_llm_config @@ -272,6 +282,28 @@ class MCPToolkit(BaseToolkit): sandbox_session_id: str | None = None, sandbox_blueprint: SandboxBlueprint | None = None, ): + """ + 初始化模型上下文协议 (MCP) 的工具箱封装。 + + 参数: + server_name: MCP 服务器的唯一标识名称。 + prefix: 工具名称的前缀,防止工具命名冲突。 + transport: 连接 MCP 服务器的通信传输协议,可选 "stdio", "sse", "streamable-http", "sandbox_proxy"。 + command: 用于 stdio 或 sandbox_proxy 模式启动服务器的可执行命令。 + args: 启动服务器时附加的命令行参数。 + url: 用于 sse 或 streamable-http 模式的服务器连接 URL。 + env: 启动服务器时的进程环境变量。 + cwd: 启动服务器时的进程工作目录。 + install_command: 首次启动前执行的环境热装配/依赖安装命令。 + isolation: 环境隔离模式,可选 "shared" (共享模式) 或 "per_session" (每个会话独立)。 + timeout: 初始化和请求网络接口时的超时秒数,默认 30。 + admin_level: 调用工具所需的管理权限等级,默认 0 (无限制)。 + header_provider: 动态生成 HTTP 头部信息的工厂函数。 + env_provider: 动态生成进程环境变量的工厂函数。 + ttl: 闲置清理的生存时间 (秒),默认 600。 + sandbox_session_id: 当 transport 为 sandbox_proxy 时指定的沙箱会话 ID。 + sandbox_blueprint: 用于沙箱自动装配的环境蓝图配置。 + """ # noqa: E501 super().__init__( config=ToolkitConfig(prefix=prefix) if prefix is not None else None ) @@ -304,7 +336,7 @@ class MCPToolkit(BaseToolkit): @property def effective_ttl(self) -> float: - """动态获取当前最新的 TTL 配置,支持无缝热重载""" + """获取当前工具箱的有效闲置生存期 TTL""" try: from zhenxun.services.ai.config import get_llm_config @@ -314,7 +346,7 @@ class MCPToolkit(BaseToolkit): return 31536000.0 if self.ttl <= 0 else float(self.ttl) async def _run_install_command(self, force: bool = False): - """执行环境预安装(热加载依赖)""" + """在本地运行预热依赖安装命令以准备 MCP 服务器环境""" if not self.install_command or not self.cwd: return @@ -366,6 +398,7 @@ class MCPToolkit(BaseToolkit): async def get_session( self, context: RunContext | None = None ) -> ClientSession | None: + """获取或建立与 MCP 服务器的会话连接,并刷新其生存状态""" await self.lifespan_manager.register( self.server_name, ttl=self.effective_ttl, @@ -376,8 +409,7 @@ class MCPToolkit(BaseToolkit): return self._shared_session async def _spawn_session_task(self, dynamic_headers: dict, dynamic_env: dict): - """生成后台连接任务""" - + """在后台异步拉起 MCP 传输协议通道并进行服务初始化""" max_attempts = 2 if (self.transport == "stdio" and self.install_command) else 1 try: @@ -546,6 +578,7 @@ class MCPToolkit(BaseToolkit): self._ready_event.set() async def initialize(self, context: RunContext | None = None): + """清空当前工具集并异步执行 MCP 服务端的全套初始化流程""" if self._is_initialized: return @@ -584,6 +617,7 @@ class MCPToolkit(BaseToolkit): await self.close() async def close(self): + """彻底关闭 MCP 工具箱,注销生存期并释放底层 WebSocket/进程通道""" current_task = asyncio.current_task() if ( diff --git a/zhenxun/services/ai/tools/providers/skills/capabilities.py b/zhenxun/services/ai/tools/providers/skills/capabilities.py index bae1e26e..4550477d 100644 --- a/zhenxun/services/ai/tools/providers/skills/capabilities.py +++ b/zhenxun/services/ai/tools/providers/skills/capabilities.py @@ -3,7 +3,7 @@ from pathlib import Path from typing import Any from zhenxun.services.ai.capabilities import AbstractCapability -from zhenxun.services.ai.run import RunContext +from zhenxun.services.ai.run.context import RunContext from zhenxun.services.ai.tools.providers.skills.manager import skill_manager from zhenxun.services.ai.tools.providers.skills.models import Skill, SkillSource from zhenxun.services.ai.tools.providers.skills.toolkit import SkillMetaToolkit @@ -18,6 +18,13 @@ class SkillCapability(AbstractCapability): skills: Sequence[str | Path | Skill | SkillSource] | None = None, namespace: str | None = None, ): + """ + 初始化技能库挂载能力组件。 + + 参数: + skills: 需要挂载到该环境下的技能列表,支持名称、路径、Skill 实例或 SkillSource。 + namespace: 该技能库所处的作用域命名空间,若不指定则自动推导。 + """ # noqa: E501 self.skills = skills or [] self.namespace = namespace or infer_plugin_namespace() diff --git a/zhenxun/services/ai/tools/providers/skills/toolkit.py b/zhenxun/services/ai/tools/providers/skills/toolkit.py index 87c799f9..877f7ee8 100644 --- a/zhenxun/services/ai/tools/providers/skills/toolkit.py +++ b/zhenxun/services/ai/tools/providers/skills/toolkit.py @@ -1,12 +1,13 @@ from typing import Any, cast -from zhenxun.services.ai.run import Inject, RunContext +from zhenxun.services.ai.run.context import RunContext +from zhenxun.services.ai.run.di import Inject from zhenxun.services.ai.sandbox.models import SandboxBlueprint from zhenxun.services.ai.sandbox.protocols import ( SupportsCommandExecution, SupportsFileSystem, ) -from zhenxun.services.ai.tools.core.decorators import silent, tool +from zhenxun.services.ai.tools.core.decorators import Rules, tool from zhenxun.services.ai.tools.core.toolkit import BaseToolkit from zhenxun.services.ai.tools.models import ResolvedToolPayload, ToolResult from zhenxun.services.ai.tools.providers.skills.manager import ( @@ -211,8 +212,8 @@ class SkillMetaToolkit(BaseToolkit, SkillSandboxExecutionMixin): @tool( name="read_skill_instructions", description="加载指定技能的完整使用说明与可用资源清单。返回值为 XML 结构。", + rules=[Rules.silent()], ) - @silent() async def read_skill_instructions(self, skill_name: str) -> ToolResult: skill = await self._get_skill(skill_name) if not skill: @@ -261,8 +262,8 @@ class SkillMetaToolkit(BaseToolkit, SkillSandboxExecutionMixin): "安全读取指定技能目录下的附加文件" "(如 references/ 里的参考文档或 scripts/ 里的代码)。" ), + rules=[Rules.silent()], ) - @silent() async def read_skill_file( self, skill_name: str, diff --git a/zhenxun/services/ai/utils/runtime.py b/zhenxun/services/ai/utils/runtime.py index dfa3414e..bd60b691 100644 --- a/zhenxun/services/ai/utils/runtime.py +++ b/zhenxun/services/ai/utils/runtime.py @@ -8,11 +8,11 @@ from zhenxun.services.ai.utils.scope import ScopeBuilder class ContextUtils: """ 从底层依赖容器 (deps) 中提取运行环境信息的纯静态工具类。 - """ @staticmethod def extract_user_id(deps: Any) -> str | None: + """从依赖容器中提取当前用户的 ID""" if not deps: return None if hasattr(deps, "user_id") and getattr(deps, "user_id") is not None: @@ -31,6 +31,7 @@ class ContextUtils: @staticmethod def extract_group_id(deps: Any) -> str | None: + """从依赖容器中提取当前群聊的 ID""" if not deps: return None if hasattr(deps, "group_id") and getattr(deps, "group_id") is not None: @@ -42,6 +43,7 @@ class ContextUtils: @staticmethod def extract_platform(deps: Any) -> str: + """从依赖容器的 Bot 实例中提取当前聊天平台名称""" if not deps: return "unknown" if hasattr(deps, "platform") and getattr(deps, "platform") is not None: @@ -57,6 +59,7 @@ class ContextUtils: def extract_concurrency_lock_id( context: Any, scope: Any, default_session_id: str ) -> str: + """根据并发隔离范围 scope 动态计算并返回当前会话的并发锁 ID""" from zhenxun.services.ai.flow.base import ConcurrencyScope scope = scope or ConcurrencyScope.GROUP @@ -136,6 +139,7 @@ class PermissionUtils: @staticmethod async def check_superuser(context: Any) -> bool: + """异步校验当前运行上下文中的用户是否为超级用户""" bot = context.get_bot() event = context.get_event() if bot and event: @@ -146,6 +150,7 @@ class PermissionUtils: @staticmethod async def check_admin_level(context: Any, min_level: int) -> bool: + """异步校验当前用户在全局或当前群聊中的管理权限等级是否达到最低要求""" if await PermissionUtils.check_superuser(context): return True