diff --git a/zhenxun/builtin_plugins/scheduler_admin/data_source.py b/zhenxun/builtin_plugins/scheduler_admin/data_source.py index bdb4b659..f7a2e629 100644 --- a/zhenxun/builtin_plugins/scheduler_admin/data_source.py +++ b/zhenxun/builtin_plugins/scheduler_admin/data_source.py @@ -6,7 +6,14 @@ from zhenxun.models.level_user import LevelUser from zhenxun.models.scheduled_job import ScheduledJob from zhenxun.services import scheduler_manager from zhenxun.services.log import logger +from zhenxun.services.scheduler.registry import scheduler_registry from zhenxun.services.scheduler.repository import ScheduleRepository +from zhenxun.services.scheduler.types import ( + BaseTrigger, + ExecutionOptions, + JobConfig, + TargetType, +) from zhenxun.utils.pydantic_compat import model_dump, model_validate from . import presenters @@ -58,7 +65,7 @@ class SchedulerAdminService: targets: list[str], creator_permission_level: int, plugin_name: str, - trigger_info: tuple[str, dict], + trigger: BaseTrigger, job_kwargs: dict, permission: int, bot_id: str, @@ -69,7 +76,6 @@ class SchedulerAdminService: created_by: str, ) -> str: """创建或更新一个定时任务""" - trigger_type, trigger_config = trigger_info success_targets = [] failed_targets = [] permission_denied_targets = [] @@ -103,7 +109,7 @@ class SchedulerAdminService: ) continue - if target_type in ["TAG", "ALL_GROUPS"]: + if target_type in [TargetType.TAG.value, TargetType.ALL_GROUPS.value]: logger.debug( f"检测到多目标任务 (类型: {target_type})," f"将所需权限强制提升至超级用户级别。" @@ -111,18 +117,26 @@ class SchedulerAdminService: permission = 9 try: + exec_opts = ( + model_validate(ExecutionOptions, execution_options) + if execution_options + else model_validate(ExecutionOptions, {}) + ) + + config = JobConfig( + trigger=trigger, + job_kwargs=job_kwargs, + bot_id=bot_id, + name=job_name, + created_by=created_by, + required_permission=permission, + execution_options=exec_opts, + ) schedule = await scheduler_manager.add_schedule( plugin_name=plugin_name, target_type=target_type, target_identifier=target_id, - trigger_type=trigger_type, - trigger_config=trigger_config, - job_kwargs=job_kwargs, - bot_id=bot_id, - required_permission=permission, - name=job_name, - created_by=created_by, - execution_options=execution_options if execution_options else None, + config=config, ) if schedule: success_targets.append((target_desc, schedule.id)) @@ -150,7 +164,10 @@ class SchedulerAdminService: permission_denied = False if all_flag or global_flag: permission_denied = True - elif targeter._filters.get("target_type") in ["TAG", "ALL_GROUPS"]: + elif targeter._filters.get("target_type") in [ + TargetType.TAG.value, + TargetType.ALL_GROUPS.value, + ]: permission_denied = True if permission_denied: @@ -202,11 +219,16 @@ class SchedulerAdminService: ) async def update_schedule( - self, schedule: ScheduledJob, trigger_info: tuple | None, kwargs_str: str | None + self, + schedule: ScheduledJob, + trigger: BaseTrigger | None, + kwargs_str: str | None, ) -> str: """更新一个任务的配置""" - trigger_type = trigger_info[0] if trigger_info else None - trigger_config = trigger_info[1] if trigger_info else None + trigger_type = trigger.trigger_type if trigger else None + trigger_config = ( + model_dump(trigger, exclude={"trigger_type"}) if trigger else None + ) job_kwargs = await self._parse_and_validate_kwargs_for_update( schedule.plugin_name, kwargs_str ) @@ -243,7 +265,7 @@ class SchedulerAdminService: def _generate_view_title(self, filters: dict) -> str: title = "定时任务" - if filters.get("target_type") == "ALL_GROUPS": + if filters.get("target_type") == TargetType.ALL_GROUPS.value: title = "全局定时任务" elif "target_identifier" in filters: title = f"群 {filters['target_identifier']} 的定时任务" @@ -253,12 +275,12 @@ class SchedulerAdminService: def _resolve_target_descriptor(self, target_desc: str) -> tuple[str, str]: if target_desc == scheduler_manager.ALL_GROUPS: - return "ALL_GROUPS", scheduler_manager.ALL_GROUPS + return TargetType.ALL_GROUPS.value, scheduler_manager.ALL_GROUPS if target_desc.startswith("tag:"): - return "TAG", target_desc[4:] + return TargetType.TAG.value, target_desc[4:] if target_desc.isdigit(): - return "GROUP", target_desc - return "USER", target_desc + return TargetType.GROUP.value, target_desc + return TargetType.USER.value, target_desc def _format_set_result_message( self, targets: list, success: list, failed: list, permission_denied: list @@ -286,7 +308,7 @@ class SchedulerAdminService: if not kwargs_str: return {} - task_meta = scheduler_manager._registered_tasks.get(plugin_name) + task_meta = scheduler_registry.tasks.get(plugin_name) if not task_meta: raise ValueError(f"插件 '{plugin_name}' 未注册。") diff --git a/zhenxun/builtin_plugins/scheduler_admin/dependencies.py b/zhenxun/builtin_plugins/scheduler_admin/dependencies.py index 1f1e1ac5..eceecf04 100644 --- a/zhenxun/builtin_plugins/scheduler_admin/dependencies.py +++ b/zhenxun/builtin_plugins/scheduler_admin/dependencies.py @@ -22,6 +22,8 @@ from zhenxun.configs.config import Config from zhenxun.models.level_user import LevelUser from zhenxun.models.scheduled_job import ScheduledJob from zhenxun.services import scheduler_manager +from zhenxun.services.scheduler.registry import scheduler_registry +from zhenxun.services.scheduler.types import BaseTrigger, TargetType, Trigger from zhenxun.utils.time_utils import TimeUtils @@ -103,7 +105,7 @@ def parse_daily_time(time_str: str) -> dict: raise ValueError("时间格式错误,请使用 'HH:MM' 或 'HH:MM:SS' 格式。") -def _parse_trigger_from_arparma(arp: Arparma) -> tuple[str, dict] | None: +def _parse_trigger_from_arparma(arp: Arparma) -> BaseTrigger | None: """从 Arparma 中解析时间触发器配置""" subcommand_name = next(iter(arp.subcommands.keys()), None) if not subcommand_name: @@ -111,19 +113,22 @@ def _parse_trigger_from_arparma(arp: Arparma) -> tuple[str, dict] | None: try: if cron_expr := arp.query[str](f"{subcommand_name}.cron.cron_expr", None): - return "cron", dict( - zip( - ["minute", "hour", "day", "month", "day_of_week"], cron_expr.split() + return Trigger.cron( + **dict( + zip( + ["minute", "hour", "day", "month", "day_of_week"], + cron_expr.split(), + ) ) ) if interval_expr := arp.query[str]( f"{subcommand_name}.interval.interval_expr", None ): - return "interval", TimeUtils.parse_interval_to_dict(interval_expr) + return Trigger.interval(**TimeUtils.parse_interval_to_dict(interval_expr)) if date_expr := arp.query[str](f"{subcommand_name}.date.date_expr", None): - return "date", {"run_date": datetime.fromisoformat(date_expr)} + return Trigger.date(run_date=datetime.fromisoformat(date_expr)) if daily_expr := arp.query[str](f"{subcommand_name}.daily.daily_expr", None): - return "cron", parse_daily_time(daily_expr) + return Trigger.cron(**parse_daily_time(daily_expr)) except ValueError as e: raise ValueError(f"时间参数解析错误: {e}") from e return None @@ -132,12 +137,12 @@ def _parse_trigger_from_arparma(arp: Arparma) -> tuple[str, dict] | None: async def GetTriggerInfo( matcher: AlconnaMatcher, arp: Arparma = AlconnaMatches(), -) -> tuple[str, dict]: +) -> BaseTrigger: """依赖注入函数:解析并验证时间触发器""" try: - trigger_info = _parse_trigger_from_arparma(arp) - if trigger_info: - return trigger_info + trigger = _parse_trigger_from_arparma(arp) + if trigger: + return trigger except ValueError as e: await matcher.finish(f"时间参数解析错误: {e}") @@ -199,24 +204,24 @@ async def GetTargeter( filters["plugin_name"] = plugin_name.result if global_flag: - filters["target_type"] = "ALL_GROUPS" + filters["target_type"] = TargetType.ALL_GROUPS.value filters["target_identifier"] = scheduler_manager.ALL_GROUPS elif user_id.available: - filters["target_type"] = "USER" + filters["target_type"] = TargetType.USER.value filters["target_identifier"] = user_id.result elif all_enabled: pass elif tag_name.available: - filters["target_type"] = "TAG" + filters["target_type"] = TargetType.TAG.value filters["target_identifier"] = tag_name.result elif group_ids.available: gids = [str(gid) for gid in group_ids.result] - filters["target_type"] = "GROUP" + filters["target_type"] = TargetType.GROUP.value filters["target_identifier__in"] = gids else: current_group_id = getattr(event, "group_id", None) if current_group_id: - filters["target_type"] = "GROUP" + filters["target_type"] = TargetType.GROUP.value filters["target_identifier"] = str(current_group_id) return scheduler_manager.target(**filters) @@ -230,7 +235,7 @@ async def GetValidatedJobKwargs( ) -> dict: """依赖注入函数:解析、合并和验证任务的关键字参数""" p_name = plugin_name.result - task_meta = scheduler_manager._registered_tasks.get(p_name) + task_meta = scheduler_registry.tasks.get(p_name) if not task_meta: await matcher.finish(f"插件 '{p_name}' 未注册可定时执行的任务。") @@ -311,7 +316,7 @@ async def GetFinalPermission( else: base_permission = effective_user_level - task_meta = scheduler_manager._registered_tasks.get(plugin_name.result) + task_meta = scheduler_registry.tasks.get(plugin_name.result) if task_meta and "default_permission" in task_meta: default_perm = task_meta.get("default_permission") if isinstance(default_perm, int): diff --git a/zhenxun/builtin_plugins/scheduler_admin/handlers.py b/zhenxun/builtin_plugins/scheduler_admin/handlers.py index 26b91a88..f00f0d40 100644 --- a/zhenxun/builtin_plugins/scheduler_admin/handlers.py +++ b/zhenxun/builtin_plugins/scheduler_admin/handlers.py @@ -16,6 +16,8 @@ from nonebot_plugin_uninfo import Uninfo from zhenxun.configs.config import Config from zhenxun.models.scheduled_job import ScheduledJob from zhenxun.services import scheduler_manager +from zhenxun.services.scheduler.registry import scheduler_registry +from zhenxun.services.scheduler.types import BaseTrigger from zhenxun.utils.message import MessageUtils from .commands import schedule_cmd @@ -66,7 +68,7 @@ async def handle_set( interval: Match[int] = AlconnaMatch("interval_seconds"), job_name: Match[str] = AlconnaMatch("job_name"), bot_id_to_operate: str = Depends(GetBotId), - trigger_info: tuple[str, dict] = Depends(GetTriggerInfo), + trigger: BaseTrigger = Depends(GetTriggerInfo), job_kwargs: dict = Depends(GetValidatedJobKwargs), creator_permission_level: int = Depends(GetCreatorPermissionLevel), final_permission: int = Depends(GetFinalPermission), @@ -86,7 +88,7 @@ async def handle_set( ) if is_multi_target: - task_meta = scheduler_manager._registered_tasks.get(p_name) + task_meta = scheduler_registry.tasks.get(p_name) if jitter_val is None: if task_meta and task_meta.get("default_jitter") is not None: jitter_val = cast(int | None, task_meta["default_jitter"]) @@ -114,7 +116,7 @@ async def handle_set( targets=target_groups, creator_permission_level=creator_permission_level, plugin_name=p_name, - trigger_info=trigger_info, + trigger=trigger, job_kwargs=job_kwargs, permission=final_permission, bot_id=bot_id_to_operate, @@ -210,14 +212,14 @@ async def handle_update( kwargs_str: Match[str] = AlconnaMatch("kwargs_str"), ): """处理 '更新' 子命令""" - trigger_info = _parse_trigger_from_arparma(arp) - if not trigger_info and not kwargs_str.available: + trigger = _parse_trigger_from_arparma(arp) + if not trigger and not kwargs_str.available: await schedule_cmd.finish( "请提供需要更新的时间 (--cron/--interval/--date/--daily) 或参数 (--kwargs)" ) result_message = await scheduler_admin_service.update_schedule( - schedule, trigger_info, kwargs_str.result if kwargs_str.available else None + schedule, trigger, kwargs_str.result if kwargs_str.available else None ) await schedule_cmd.finish(result_message) diff --git a/zhenxun/builtin_plugins/scheduler_admin/presenters.py b/zhenxun/builtin_plugins/scheduler_admin/presenters.py index 58973931..a6f653f2 100644 --- a/zhenxun/builtin_plugins/scheduler_admin/presenters.py +++ b/zhenxun/builtin_plugins/scheduler_admin/presenters.py @@ -3,6 +3,8 @@ from typing import Any from zhenxun import ui from zhenxun.models.scheduled_job import ScheduledJob from zhenxun.services import scheduler_manager +from zhenxun.services.scheduler.registry import scheduler_registry +from zhenxun.services.scheduler.types import TargetType from zhenxun.ui.models import StatusBadgeCell, TextCell from zhenxun.utils.pydantic_compat import model_json_schema @@ -173,15 +175,15 @@ async def format_schedule_list_as_image( def format_target_info(target_type: str, target_identifier: str) -> str: """格式化目标信息以供显示""" - if target_type == "GLOBAL": + if target_type == TargetType.GLOBAL.value: return "全局" - elif target_type == "ALL_GROUPS": + elif target_type == TargetType.ALL_GROUPS.value: return "所有群组" - elif target_type == "TAG": + elif target_type == TargetType.TAG.value: return f"标签: {target_identifier}" - elif target_type == "GROUP": + elif target_type == TargetType.GROUP.value: return f"群: {target_identifier}" - elif target_type == "USER": + elif target_type == TargetType.USER.value: return f"用户: {target_identifier}" else: return f"{target_type}: {target_identifier}" @@ -215,7 +217,7 @@ async def format_plugins_list() -> str: message_parts = ["📋 已注册的定时任务插件:"] for i, plugin_name in enumerate(registered_plugins, 1): - task_meta = scheduler_manager._registered_tasks[plugin_name] + task_meta = scheduler_registry.tasks[plugin_name] params_model = task_meta.get("model") param_info_str = "无参数" diff --git a/zhenxun/services/ai/capabilities/__init__.py b/zhenxun/services/ai/capabilities/__init__.py index a8af4713..1fe4464b 100644 --- a/zhenxun/services/ai/capabilities/__init__.py +++ b/zhenxun/services/ai/capabilities/__init__.py @@ -5,10 +5,13 @@ from .base import ( WrapToolExecuteHandler, WrapToolValidateHandler, ) +from .manager import CapabilityQuery, CapabilitySource, capability from .wrappers import CombinedCapability, DynamicCapability, WrapperCapability __all__ = [ "AbstractCapability", + "CapabilityQuery", + "CapabilitySource", "CombinedCapability", "DynamicCapability", "WrapModelRequestHandler", @@ -16,4 +19,5 @@ __all__ = [ "WrapToolExecuteHandler", "WrapToolValidateHandler", "WrapperCapability", + "capability", ] diff --git a/zhenxun/services/ai/capabilities/base.py b/zhenxun/services/ai/capabilities/base.py index f6f93d49..52ee887d 100644 --- a/zhenxun/services/ai/capabilities/base.py +++ b/zhenxun/services/ai/capabilities/base.py @@ -2,20 +2,20 @@ from __future__ import annotations from collections.abc import Awaitable, Callable, Sequence from dataclasses import dataclass -from typing import TYPE_CHECKING, Any, ClassVar, Literal, Union +from typing import TYPE_CHECKING, Any, Literal, Union from zhenxun.services.ai.core.messages import ChatRequest, ChatResponse +from zhenxun.services.ai.core.models import LLMContext from zhenxun.services.ai.core.options import GenerationConfig if TYPE_CHECKING: - from zhenxun.services.ai.core.models import LLMContext from zhenxun.services.ai.run import AgentRunResult, RunContext WrapRunHandler = Callable[[], Awaitable["AgentRunResult[Any]"]] """整个 Agent 运行过程包裹的处理函数类型""" WrapModelRequestHandler = Callable[ - ["LLMContext[ChatRequest, ChatResponse]"], Awaitable[ChatResponse] + [LLMContext[ChatRequest, ChatResponse]], Awaitable[ChatResponse] ] """单次大模型 API 请求包裹的处理函数类型""" @@ -52,34 +52,14 @@ class CapabilityOrdering: class AbstractCapability: """ Agent 能力组件基类协议。 - 所有业务逻辑拦截(限流、权限、动态 Prompt)请在此实现。 - 底层网络重试、并发控制等请勿在此处理。 """ - @classmethod - def get_serialization_name(cls) -> str | None: - """用于 YAML/JSON 反序列化的注册标识符""" - return cls.__name__ - - @classmethod - def from_spec(cls, **kwargs) -> "AbstractCapability": - """从 Spec 的 kwargs 中实例化对象""" - return cls(**kwargs) - - def __init_subclass__(cls, **kwargs): - """自动将继承此类的所有拦截器注册到中心表""" - super().__init_subclass__(**kwargs) - CapabilityRegistry.register(cls) - def get_ordering(self) -> CapabilityOrdering | None: """获取该拦截器的拓扑排序约束。子类可重写此方法以锁定执行顺序。""" return None async def for_run(self, context: RunContext) -> "AbstractCapability": - """获取专用于单次运行的实例。 - 默认返回自身(无状态)。 - 若需要记录单次运行的上下文状态,请返回深/浅拷贝(如 return copy.copy(self))。 - """ + """获取专用于单次运行的实例,默认返回自身(无状态)。""" return self async def get_generation_config( @@ -89,9 +69,11 @@ class AbstractCapability: return None async def get_system_prompts(self, context: RunContext) -> list[str]: + """获取该能力提供的系统提示词列表。""" return [] async def get_tools(self, context: RunContext) -> list[Any]: + """获取该能力附带的工具列表。""" return [] async def prepare_tools( @@ -135,19 +117,3 @@ class AbstractCapability: ) -> Any: """包裹单一工具的执行 (洋葱模型)。""" return await handler(arguments) - - -class CapabilityRegistry: - """Capability 序列化注册表""" - - _registry: ClassVar[dict[str, type[AbstractCapability]]] = {} - - @classmethod - def register(cls, cap_cls: type[AbstractCapability]): - name = cap_cls.get_serialization_name() - if name: - cls._registry[name] = cap_cls - - @classmethod - def get(cls, name: str) -> type[AbstractCapability] | None: - return cls._registry.get(name) diff --git a/zhenxun/services/ai/capabilities/builtin.py b/zhenxun/services/ai/capabilities/builtin.py index 92a4276d..f6b4eaf3 100644 --- a/zhenxun/services/ai/capabilities/builtin.py +++ b/zhenxun/services/ai/capabilities/builtin.py @@ -5,11 +5,6 @@ import json from typing import Any, cast from zhenxun.models.user_console import UserConsole -from zhenxun.services.ai.capabilities import ( - AbstractCapability, - WrapModelRequestHandler, - WrapToolExecuteHandler, -) from zhenxun.services.ai.config import get_llm_config from zhenxun.services.ai.core.exceptions import ( AbortException, @@ -33,10 +28,17 @@ from zhenxun.services.ai.core.messages import ( from zhenxun.services.ai.core.models import LLMContext from zhenxun.services.ai.run.context import RunContext from zhenxun.services.ai.utils import PermissionUtils -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_capability as logger from zhenxun.utils.enum import GoldHandle from zhenxun.utils.exception import InsufficientGold +from .base import ( + AbstractCapability, + WrapModelRequestHandler, + WrapToolExecuteHandler, +) +from .manager import capability + def _get_tool_meta(tool: Any, key: str, default: Any = None) -> Any: """辅助方法:安全地提取工具元数据中指定的键值""" @@ -47,6 +49,7 @@ def _get_tool_meta(tool: Any, key: str, default: Any = None) -> Any: return meta.get(key, default) +@capability(namespace="global", auto_apply=True) class StuckDetectionCapability(AbstractCapability): """死循环检测:使用前置请求拦截防止 LLM 陷入无限重试""" @@ -120,6 +123,7 @@ class StuckDetectionCapability(AbstractCapability): return await handler(llm_context) +@capability(namespace="global", auto_apply=True) class GlobalCycleLimitCapability(AbstractCapability): """全局防死循环检测中间件:跨 Agent 追踪大模型调用总次数""" @@ -152,6 +156,7 @@ class GlobalCycleLimitCapability(AbstractCapability): return await handler(llm_context) +@capability(namespace="global", auto_apply=True) class PermissionCapability(AbstractCapability): """权限校验中间件:在执行前根据确定参数进行动态鉴权""" @@ -182,6 +187,7 @@ class PermissionCapability(AbstractCapability): return await handler(arguments) +@capability(namespace="global", auto_apply=True) class BillingCapability(AbstractCapability): """经济系统中间件:执行前扣除金币""" @@ -221,6 +227,7 @@ class BillingCapability(AbstractCapability): return await handler(arguments) +@capability(namespace="global", auto_apply=True) class ToolRetryAndReflectionCapability(AbstractCapability): """ 重试与自愈反思中间件。 @@ -267,6 +274,7 @@ class ToolRetryAndReflectionCapability(AbstractCapability): return ToolResult(output=f"执行发生异常: {e}").as_error() +@capability(namespace="global", auto_apply=True) class ReflexionCapability(AbstractCapability): """自愈反思与验证引擎 (Reflexion Engine)。 统一处理结构化解析失败和语义护栏拦截。""" diff --git a/zhenxun/services/ai/capabilities/manager.py b/zhenxun/services/ai/capabilities/manager.py new file mode 100644 index 00000000..ce84382b --- /dev/null +++ b/zhenxun/services/ai/capabilities/manager.py @@ -0,0 +1,206 @@ +from collections.abc import Callable +from dataclasses import dataclass, field +import fnmatch +from typing import Any, cast +from typing_extensions import Self + +from pydantic import BaseModel, Field + +from zhenxun.services.ai.utils.logger import log_capability as logger +from zhenxun.services.ai.utils.utils import parse_routing_string +from zhenxun.utils.utils import infer_plugin_namespace + +from .base import AbstractCapability +from .wrappers import DynamicCapability + + +class CapabilityQuery(BaseModel): + """ + 拦截器/能力组件的声明式查询对象。 + 用于在 Agent 中精确或批量筛选加载特定命名空间、特定标签的能力。 + """ + + name: str | list[str] | None = Field(default=None) + """如果提供,则能力的名称必须等于该字符串或在列表中。支持 * / ? 通配符。""" + tags: list[str] | None = Field(default=None) + """如果提供,则能力必须包含这里列出的所有标签 (交集/AND匹配)。""" + exclude_tags: list[str] | None = Field(default=None) + """如果提供,则能力不能包含这里列出的任何标签 (排斥过滤)。""" + namespace: str | None = Field(default=None) + """限制搜索的插件命名空间。如果不指定,将自动推导为调用者所在的插件; + 'global' 将跨全插件搜索,'*' 代表所有插件。""" + + +CapabilitySource = ( + str | Callable | AbstractCapability | type[AbstractCapability] | CapabilityQuery +) +"""能力/拦截器来源 + +支持字符串别名/标签、普通函数、Capability类或实例,以及声明式 Query 对象。 +""" + + +@dataclass +class CapabilityEntry: + """能力组件元数据载体""" + + cls: type[AbstractCapability] + """能力类""" + name: str + """能力名称""" + namespace: str + """能力所在的命名空间""" + tags: list[str] = field(default_factory=list) + """能力标签列表""" + auto_apply: bool = False + """是否自动挂载该能力""" + + +class CapabilityManager: + """能力组件全局注册与发现中心 (单例)""" + + _instance: "CapabilityManager | None" = None + _entries: list[CapabilityEntry] + + def __new__(cls) -> Self: + """单例模式:获取或创建全局唯一的能力管理器实例""" + if cls._instance is None: + cls._instance = super().__new__(cls) + cls._instance._entries = [] + return cast(Self, cls._instance) + + def register( + self, + cls: type[AbstractCapability], + name: str, + namespace: str, + tags: list[str], + auto_apply: bool, + ) -> None: + """注册一个能力组件到管理器中""" + self._entries.append( + CapabilityEntry( + cls=cls, + name=name, + namespace=namespace, + tags=tags, + auto_apply=auto_apply, + ) + ) + tag_str = f" | Tags: {tags}" if tags else "" + logger.debug( + f"已注册 Capability: '{name}' -> Namespace: '{namespace}'{tag_str})" + ) + + def get_auto_apply_capabilities(self, namespace: str) -> list[AbstractCapability]: + """获取指定命名空间及其它全局命名空间下自动挂载的能力实例""" + instances = [] + for entry in self._entries: + if entry.auto_apply and entry.namespace in ("global", namespace): + try: + instances.append(entry.cls()) + except Exception as e: + logger.error(f"实例化自动装配能力 {entry.name} 失败: {e}") + return instances + + def query_capabilities( + self, query: CapabilityQuery, default_namespace: str + ) -> list[AbstractCapability]: + """根据声明式查询条件筛选并实例化匹配的能力组件""" + matched = [] + ns = query.namespace or default_namespace + + for entry in self._entries: + if ns != "*" and entry.namespace != ns: + continue + + if query.name: + names = [query.name] if isinstance(query.name, str) else query.name + name_matched = False + for pattern in names: + if fnmatch.fnmatch(entry.name, pattern): + name_matched = True + break + if not name_matched: + continue + + if query.tags: + if not all(tag in entry.tags for tag in query.tags): + continue + + if query.exclude_tags: + if any(tag in entry.tags for tag in query.exclude_tags): + continue + + try: + matched.append(entry.cls()) + except Exception as e: + logger.error(f"实例化能力 {entry.name} 失败: {e}") + + return matched + + def resolve_capabilities( + self, sources: list[Any], default_namespace: str + ) -> list[AbstractCapability]: + """解析多种类型的能力来源并实例化为能力组件列表""" + resolved = [] + for source in sources: + if isinstance(source, AbstractCapability): + resolved.append(source) + elif callable(source) and not isinstance(source, type): + resolved.append(DynamicCapability(source)) + elif isinstance(source, CapabilityQuery): + resolved.extend(self.query_capabilities(source, default_namespace)) + elif isinstance(source, str): + s = cast(str, source) + parsed_args = parse_routing_string(s, default_namespace) + q = CapabilityQuery(**parsed_args) + resolved.extend(self.query_capabilities(q, default_namespace)) + elif isinstance(source, type) and issubclass(source, AbstractCapability): + try: + resolved.append(source()) + except Exception as e: + logger.error(f"实例化能力 {source.__name__} 失败: {e}") + else: + raise TypeError(f"不支持的 Capability 来源: {type(source)}") + return resolved + + +capability_manager = CapabilityManager() + + +def capability( + name: str | None = None, + tags: list[str] | None = None, + auto_apply: bool = False, + namespace: str | None = None, +) -> Callable: + """ + 类装饰器:声明式地注册一个 Capability 到全局能力池中。 + 允许第三方插件通过字符串别名或标签进行引用,彻底解耦模块依赖。 + + 参数: + name: 能力的名称,如果为None则默认使用类名。 + tags: 能力的标签列表,用于分类或批量筛选。 + auto_apply: 是否自动应用挂载该能力。 + namespace: 能力的命名空间,如果为None则自动推导为调用者所在的插件。 + + 返回: + Callable: 装饰器函数,用于包装 AbstractCapability 类。 + """ + + def decorator(cls: type[AbstractCapability]): + """装饰器内部函数,实现类注册""" + final_name = name or cls.__name__ + final_tags = tags or [] + ns = ( + namespace + if namespace is not None + else infer_plugin_namespace(default="global") + ) + capability_manager.register( + cls, name=final_name, namespace=ns, tags=final_tags, auto_apply=auto_apply + ) + return cls + + return decorator diff --git a/zhenxun/services/ai/capabilities/wrappers.py b/zhenxun/services/ai/capabilities/wrappers.py index e8d27926..82b1238f 100644 --- a/zhenxun/services/ai/capabilities/wrappers.py +++ b/zhenxun/services/ai/capabilities/wrappers.py @@ -8,6 +8,7 @@ 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.models import LLMContext from zhenxun.services.ai.core.options import GenerationConfig from .base import ( @@ -20,7 +21,6 @@ from .base import ( ) if TYPE_CHECKING: - from zhenxun.services.ai.core.models import LLMContext from zhenxun.services.ai.run import AgentRunResult, RunContext @@ -266,11 +266,6 @@ class DynamicCapability(AbstractCapability): """初始化动态能力""" self.capability_func = capability_func - @classmethod - def get_serialization_name(cls) -> str | None: - """获取反序列化标识""" - return None - async def for_run(self, context: RunContext) -> "AbstractCapability": """在运行时基于当前上下文动态实例化并执行真正的 Capability""" if is_coroutine_callable(self.capability_func): @@ -292,11 +287,6 @@ class WrapperCapability(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) diff --git a/zhenxun/services/ai/config/manager.py b/zhenxun/services/ai/config/manager.py index e82f1fe5..3012db38 100644 --- a/zhenxun/services/ai/config/manager.py +++ b/zhenxun/services/ai/config/manager.py @@ -24,6 +24,7 @@ def get_default_providers() -> list[dict[str, Any]]: "models": [ { "model_name": "deepseek-v4-pro", + "reasoning_effort": "high", }, { "model_name": "deepseek-v4-flash", diff --git a/zhenxun/services/ai/config/models.py b/zhenxun/services/ai/config/models.py index cd890b1d..7f9de232 100644 --- a/zhenxun/services/ai/config/models.py +++ b/zhenxun/services/ai/config/models.py @@ -107,11 +107,9 @@ class ProviderConfig(BaseModel): """API 基础 URL 路径""" api_type: str = "openai" """API 协议类型 (openai/gemini/zhipu/etc.)""" - openai_compat: bool = False - """是否强制使用 OpenAI 兼容模式""" - temperature: float | None = 0.7 + temperature: float | None = None """该提供商下模型的默认温度""" - generation_max_tokens: int | None = None + max_output_tokens: int | None = None """该提供商下模型的默认最大输出限制""" models: list[ModelDetail] """该提供商提供的具体模型列表""" diff --git a/zhenxun/services/ai/context/knowledge/filesystem.py b/zhenxun/services/ai/context/knowledge/filesystem.py index c6fa791b..2ed5b95e 100644 --- a/zhenxun/services/ai/context/knowledge/filesystem.py +++ b/zhenxun/services/ai/context/knowledge/filesystem.py @@ -5,7 +5,7 @@ from typing import Any from zhenxun.services.ai.tools.core.decorators import tool from zhenxun.services.ai.tools.models import ToolResult -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_knowledge as logger from .base import BaseKnowledge diff --git a/zhenxun/services/ai/context/knowledge/readers.py b/zhenxun/services/ai/context/knowledge/readers.py index 403fd728..5116cc86 100644 --- a/zhenxun/services/ai/context/knowledge/readers.py +++ b/zhenxun/services/ai/context/knowledge/readers.py @@ -3,7 +3,7 @@ import csv from pathlib import Path from zhenxun.services.ai.context.rag.models import BaseRecord -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_knowledge as logger class BaseReader: diff --git a/zhenxun/services/ai/context/knowledge/vector.py b/zhenxun/services/ai/context/knowledge/vector.py index 429f9c62..cbeb5015 100644 --- a/zhenxun/services/ai/context/knowledge/vector.py +++ b/zhenxun/services/ai/context/knowledge/vector.py @@ -5,12 +5,6 @@ import anyio from nonebot.adapters import Bot, Event from pydantic import BaseModel, Field -from zhenxun.services.ai.context.knowledge.base import BaseKnowledge -from zhenxun.services.ai.context.knowledge.readers import ( - BaseReader, - CSVReader, - TextReader, -) from zhenxun.services.ai.context.rag.engine import ScopedRAGClient from zhenxun.services.ai.context.rag.models import BaseRecord from zhenxun.services.ai.core.messages import LLMMessage @@ -18,7 +12,14 @@ from zhenxun.services.ai.llm.api import generate_structured from zhenxun.services.ai.run import RunContext from zhenxun.services.ai.tools.core.decorators import tool from zhenxun.services.ai.tools.models import ToolkitConfig, ToolResult -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_knowledge as logger + +from .base import BaseKnowledge +from .readers import ( + BaseReader, + CSVReader, + TextReader, +) class QueryAnalysis(BaseModel): @@ -312,7 +313,7 @@ class VectorKnowledge(BaseKnowledge): for result in results: doc_name = result.record.metadata.get("name", "未命名文档") formatted_results.append( - f"📄 来源: {doc_name}\n" f"片段内容:\n{result.record.content}" + f"📄 来源: {doc_name}\n片段内容:\n{result.record.content}" ) final_text = "\n\n======\n\n".join(formatted_results) diff --git a/zhenxun/services/ai/context/memory/builder.py b/zhenxun/services/ai/context/memory/builder.py index 52be7af5..cc7092d8 100644 --- a/zhenxun/services/ai/context/memory/builder.py +++ b/zhenxun/services/ai/context/memory/builder.py @@ -1,33 +1,31 @@ from __future__ import annotations from collections.abc import Callable -from typing import TYPE_CHECKING, Any +from typing import Any from typing_extensions import Self from pydantic import BaseModel -if TYPE_CHECKING: - from zhenxun.services.ai.context.memory.models import MemorySlot - from zhenxun.services.ai.context.memory.storage.interfaces import ( - BaseChatContext, - BaseMemoryIngestionMiddleware, - BaseSlotContext, - ) - from zhenxun.services.ai.context.rag.backends import Embedder, StorageBackend - from zhenxun.services.ai.context.rag.engine import ScopedRAGClient - from zhenxun.services.ai.context.memory.compression import MemoryPolicy from zhenxun.services.ai.context.memory.models import ( ContextCompressionConfig, IngestionConfig, LongTermConfig, MemoryConfig, + MemorySlot, ShortTermConfig, SlotMemoryConfig, ) +from zhenxun.services.ai.context.memory.storage.interfaces import ( + BaseChatContext, + BaseMemoryIngestionMiddleware, + BaseSlotContext, +) from zhenxun.services.ai.context.memory.types import ( AutoRecallPolicy, ) +from zhenxun.services.ai.context.rag.backends import Embedder, StorageBackend +from zhenxun.services.ai.context.rag.engine import ScopedRAGClient from zhenxun.services.ai.utils.scope import ScopeBuilder @@ -51,7 +49,7 @@ class MemoryBuilder: ) @classmethod - def auto(cls) -> "MemoryBuilder": + def auto(cls) -> MemoryBuilder: """ 创建一个开箱即用的默认记忆配置构建器。 @@ -61,7 +59,7 @@ class MemoryBuilder: @classmethod def resolve( - cls, memory: bool | MemoryConfig | "MemoryBuilder" | None + cls, memory: bool | MemoryConfig | MemoryBuilder | None ) -> MemoryConfig: if isinstance(memory, MemoryConfig): return memory @@ -85,7 +83,7 @@ class MemoryBuilder: self, enable: bool = True, isolation: ScopeBuilder | None = None, - backend: "str | BaseChatContext | None" = None, + backend: str | BaseChatContext | None = None, ) -> Self: """ 配置短期对话历史记忆。 @@ -107,8 +105,8 @@ class MemoryBuilder: self, enable: bool = True, scopes: dict[str, ScopeBuilder] | None = None, - default_slots: list["MemorySlot"] | None = None, - backend: "str | BaseSlotContext | None" = None, + default_slots: list[MemorySlot] | None = None, + backend: str | BaseSlotContext | None = None, instructions: str | None = None, ) -> Self: """ @@ -136,9 +134,9 @@ class MemoryBuilder: self, enable: bool = True, scopes: dict[str, ScopeBuilder] | None = None, - engine: "ScopedRAGClient | None" = None, - backend: "str | StorageBackend | None" = None, - embedder: "Embedder | str | None" = None, + engine: ScopedRAGClient | None = None, + backend: str | StorageBackend | None = None, + embedder: Embedder | str | None = None, agentic: bool = True, auto_recall: AutoRecallPolicy = False, instructions: str | None = None, @@ -250,7 +248,7 @@ class MemoryBuilder: return self def with_ingestion_middlewares( - self, *middlewares: "BaseMemoryIngestionMiddleware" + self, *middlewares: BaseMemoryIngestionMiddleware ) -> Self: """ 配置记忆入库管线中间件。 diff --git a/zhenxun/services/ai/context/memory/compression.py b/zhenxun/services/ai/context/memory/compression.py index 9aec2f8c..05dd7f0d 100644 --- a/zhenxun/services/ai/context/memory/compression.py +++ b/zhenxun/services/ai/context/memory/compression.py @@ -4,9 +4,6 @@ from typing import Any, Generic, TypeVar from pydantic import BaseModel, Field -from zhenxun.services.ai.context.memory.storage.interfaces import ( - BaseMemoryReducer, -) from zhenxun.services.ai.core.engine.token_counter import token_counter from zhenxun.services.ai.core.messages import ( AudioPart, @@ -18,9 +15,13 @@ from zhenxun.services.ai.core.messages import ( VideoPart, ) from zhenxun.services.ai.llm.manager import get_default_model -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_memory as logger from zhenxun.utils.pydantic_compat import model_copy +from .storage.interfaces import ( + BaseMemoryReducer, +) + class MultimodalPlaceholderReducer(BaseMemoryReducer): """视觉媒体降级:将超过一定轮数的老图片/视频替换为 <图片> 占位符文本""" @@ -126,7 +127,7 @@ class MessageDropper(BaseMemoryReducer): return messages, False, current_tokens logger.info( - "✂️ [MemoryCompression] 触发硬截断丢弃策略 | 原因: " + "✂️ 触发硬截断丢弃策略 | 原因: " f"当前 Token 预估 ({current_tokens}) 仍超过硬性上限 ({self.trigger_tokens})," # noqa: E501 "开始丢弃最旧的历史对话..." ) @@ -198,9 +199,7 @@ class ToolPrunerReducer(BaseMemoryReducer): if is_turn_exceeded: reasons.append(f"工具调用轮数超限 ({tool_turns} > {self.max_turns})") - logger.info( - f"✂️ [MemoryCompression] 触发工具结果修剪策略 | 原因: {' 且 '.join(reasons)}" - ) + logger.info(f"✂️ 触发工具结果修剪策略 | 原因: {' 且 '.join(reasons)}") from zhenxun.services.ai.core.messages import ToolReturnPart diff --git a/zhenxun/services/ai/context/memory/engine.py b/zhenxun/services/ai/context/memory/engine.py index 42047f95..150e73bd 100644 --- a/zhenxun/services/ai/context/memory/engine.py +++ b/zhenxun/services/ai/context/memory/engine.py @@ -1,17 +1,18 @@ from collections.abc import Sequence from typing import Any, cast -from zhenxun.services.ai.context.memory.compression import ( - CondenserPipeline, -) -from zhenxun.services.ai.context.memory.manager import memory_manager -from zhenxun.services.ai.context.memory.models import MemoryConfig -from zhenxun.services.ai.context.memory.types import SessionMetadata from zhenxun.services.ai.core.engine.context_renderer import ContextConverter from zhenxun.services.ai.core.messages import AgentMessage -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_memory as logger from zhenxun.utils.pydantic_compat import model_copy +from .compression import ( + CondenserPipeline, +) +from .manager import memory_manager +from .models import MemoryConfig +from .types import SessionMetadata + class MemoryReader: """ diff --git a/zhenxun/services/ai/context/memory/facades.py b/zhenxun/services/ai/context/memory/facades.py index 539d1c45..1fb07cfe 100644 --- a/zhenxun/services/ai/context/memory/facades.py +++ b/zhenxun/services/ai/context/memory/facades.py @@ -1,29 +1,30 @@ -from collections.abc import Sequence -from typing import TYPE_CHECKING, Literal +from __future__ import annotations -from zhenxun.services.ai.context.memory.types import ( +from collections.abc import Sequence +from typing import Literal + +from zhenxun.services.ai.core.messages import AgentMessage, LLMMessage + +from .manager import GlobalMemoryManager +from .storage.interfaces import ( + BaseChatContext, + BaseSlotContext, +) +from .types import ( MemorySlot, SessionMetadata, ) -from zhenxun.services.ai.core.messages import AgentMessage, LLMMessage - -if TYPE_CHECKING: - from zhenxun.services.ai.context.memory.manager import GlobalMemoryManager - from zhenxun.services.ai.context.memory.storage.interfaces import ( - BaseChatContext, - BaseSlotContext, - ) class ChatHistoryFacade: """短期对话历史门面""" - def __init__(self, manager: "GlobalMemoryManager", session_meta: SessionMetadata): + def __init__(self, manager: GlobalMemoryManager, session_meta: SessionMetadata): self.manager = manager self.session_meta = session_meta @property - def _backend(self) -> "BaseChatContext | None": + def _backend(self) -> BaseChatContext | None: return self.manager.get_chat_context( None, self.session_meta.namespace or "global" ) @@ -56,12 +57,12 @@ class ChatHistoryFacade: class SlotFacade: """中期记忆槽门面""" - def __init__(self, manager: "GlobalMemoryManager", session_meta: SessionMetadata): + def __init__(self, manager: GlobalMemoryManager, session_meta: SessionMetadata): self.manager = manager self.session_meta = session_meta @property - def _backend(self) -> "BaseSlotContext | None": + def _backend(self) -> BaseSlotContext | None: """获取底层槽位存储后端""" return self.manager.get_slot_context( None, self.session_meta.namespace or "global" @@ -115,7 +116,7 @@ class AgentSessionFacade: 提供给第三方开发者的会话记忆访问聚合门面 (Facade)。 """ - def __init__(self, manager: "GlobalMemoryManager", session_meta: SessionMetadata): + def __init__(self, manager: GlobalMemoryManager, session_meta: SessionMetadata): self.manager = manager self.session_meta = session_meta self.history = ChatHistoryFacade(manager, session_meta) diff --git a/zhenxun/services/ai/context/memory/manager.py b/zhenxun/services/ai/context/memory/manager.py index 64f28ff9..0b9ba851 100644 --- a/zhenxun/services/ai/context/memory/manager.py +++ b/zhenxun/services/ai/context/memory/manager.py @@ -1,18 +1,20 @@ from collections.abc import Callable from typing import Any, cast -from zhenxun.services.ai.context.memory.models import MemoryConfig -from zhenxun.services.ai.context.memory.storage.backends import ( +from zhenxun.services.ai.context.rag.backends import Embedder, StorageBackend +from zhenxun.services.ai.utils.logger import log_memory as logger +from zhenxun.services.ai.utils.scope import BaseScopeBuilder +from zhenxun.utils.utils import infer_plugin_namespace + +from .models import MemoryConfig +from .storage.backends import ( InMemoryChatContext, MemoryScope, ) -from zhenxun.services.ai.context.memory.storage.interfaces import ( +from .storage.interfaces import ( BaseChatContext, BaseSlotContext, ) -from zhenxun.services.ai.context.rag.backends import Embedder, StorageBackend -from zhenxun.services.ai.utils.scope import BaseScopeBuilder -from zhenxun.utils.utils import infer_plugin_namespace class MemoryCleaner(BaseScopeBuilder["MemoryCleaner"]): @@ -64,14 +66,11 @@ class MemoryCleaner(BaseScopeBuilder["MemoryCleaner"]): async def clear_all(self): """一键清理指定范围下的所有生命周期记忆(对话、槽位、RAG)""" - from zhenxun.services.log import logger - await self.clear_short_term() await self.clear_slots() await self.clear_long_term() logger.info( - f"🧹 [MemoryCleaner] 成功清理作用域 '{self._selector.scope_prefix}'" - "下的所有记忆痕迹!" + f"🧹 成功清理作用域 '{self._selector.scope_prefix}'下的所有记忆痕迹!" ) @@ -189,7 +188,7 @@ class GlobalMemoryManager: if embedder: builder.with_embedder(embedder) - from zhenxun.services.ai.context.memory.models import MemoryScoringConfig + from .models import MemoryScoringConfig scoring_cfg = MemoryScoringConfig() diff --git a/zhenxun/services/ai/context/memory/models.py b/zhenxun/services/ai/context/memory/models.py index e2f8ccb2..d5f97b05 100644 --- a/zhenxun/services/ai/context/memory/models.py +++ b/zhenxun/services/ai/context/memory/models.py @@ -6,21 +6,22 @@ from typing import Any from pydantic import BaseModel, ConfigDict, Field -from zhenxun.services.ai.context.memory.storage.interfaces import ( +from zhenxun.services.ai.context.rag.backends import Embedder, StorageBackend +from zhenxun.services.ai.context.rag.engine import ScopedRAGClient +from zhenxun.services.ai.utils.scope import ScopeBuilder + +from .storage.interfaces import ( BaseChatContext, BaseMemoryIngestionMiddleware, BaseMemoryReducer, BaseSlotContext, ) -from zhenxun.services.ai.context.memory.types import ( +from .types import ( AutoRecallPolicy, Isolation, MemorySlot, SessionMetadata, ) -from zhenxun.services.ai.context.rag.backends import Embedder, StorageBackend -from zhenxun.services.ai.context.rag.engine import ScopedRAGClient -from zhenxun.services.ai.utils.scope import ScopeBuilder class SlotMemoryConfig(BaseModel): diff --git a/zhenxun/services/ai/context/memory/storage/backends.py b/zhenxun/services/ai/context/memory/storage/backends.py index ce81a79a..4e7da1d0 100644 --- a/zhenxun/services/ai/context/memory/storage/backends.py +++ b/zhenxun/services/ai/context/memory/storage/backends.py @@ -1,26 +1,22 @@ +from __future__ import annotations + import asyncio import base64 from collections.abc import Callable import datetime from pathlib import Path import time -from typing import TYPE_CHECKING, Any, cast - -if TYPE_CHECKING: - from zhenxun.services.ai.context.rag.engine import ScopedRAGClient +from typing import Any, cast from nonebot.utils import is_coroutine_callable from tortoise import fields from tortoise.timezone import now -from zhenxun.services.ai.context.memory.storage.interfaces import ( - BaseChatContext, - BaseSlotContext, -) from zhenxun.services.ai.context.memory.types import ( MemorySlot, SessionMetadata, ) +from zhenxun.services.ai.context.rag.engine import ScopedRAGClient from zhenxun.services.ai.context.rag.models import BaseRecord, SearchResult from zhenxun.services.ai.core.messages import ( AssistantMessage, @@ -35,6 +31,11 @@ from zhenxun.services.ai.utils.scope import ScopeSelector from zhenxun.services.db_context import Model from zhenxun.utils.pydantic_compat import TypeAdapter, model_dump +from .interfaces import ( + BaseChatContext, + BaseSlotContext, +) + class AbstractMemoryRecord(Model): """Tortoise ORM 短期记忆持久化基类 (Mixin)。""" @@ -129,7 +130,7 @@ class MemoryScope: def __init__( self, - rag_client: "ScopedRAGClient", + rag_client: ScopedRAGClient, ): """初始化长期记忆作用域与 RAG 客户端""" self.rag_client = rag_client diff --git a/zhenxun/services/ai/context/rag/backends/embedders.py b/zhenxun/services/ai/context/rag/backends/embedders.py index 246c3628..6391fef6 100644 --- a/zhenxun/services/ai/context/rag/backends/embedders.py +++ b/zhenxun/services/ai/context/rag/backends/embedders.py @@ -6,7 +6,7 @@ from typing import Any, Literal, Protocol, runtime_checkable from zhenxun.services.ai.core.messages import EmbedBatch from zhenxun.services.ai.llm.api import embed as api_embed from zhenxun.services.ai.message_builder import MessageBuilder -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_rag as logger EmbedTaskType = Literal[ "general", "query", "document", "similarity", "classification", "clustering" diff --git a/zhenxun/services/ai/context/rag/backends/storages.py b/zhenxun/services/ai/context/rag/backends/storages.py index af3a4c43..18c21d17 100644 --- a/zhenxun/services/ai/context/rag/backends/storages.py +++ b/zhenxun/services/ai/context/rag/backends/storages.py @@ -11,6 +11,7 @@ from zhenxun.services.ai.context.rag.models import ( SearchResult, ) from zhenxun.services.ai.context.rag.retrieval import FilterEvaluator +from zhenxun.services.ai.utils.logger import log_rag as logger from zhenxun.services.ai.utils.scope import ScopeSelector from zhenxun.services.db_context import Model @@ -124,8 +125,6 @@ class DictStorageBackend(StorageBackend): SearchResult(record=self._records[r_id], score=float(score)) ) except ValueError as e: - from zhenxun.services.log import logger - logger.warning( "⚠️ DictStorage 中缓存的向量维度与当前查询维度不匹配," f"跳过向量检索。原因: {e}" @@ -284,8 +283,6 @@ class TortoiseStorageBackend(StorageBackend): ) ) except ValueError as e: - from zhenxun.services.log import logger - logger.warning( "⚠️ 数据库中缓存的向量维度与当前模型查询维度不匹配," f"已安全跳过向量检索(降级为稀疏匹配)。原因: {e}" @@ -579,8 +576,6 @@ class LanceDBStorageBackend(StorageBackend): .to_list() ) except Exception as e: - from zhenxun.services.log import logger - logger.warning(f"LanceDB FTS 检索失败(可能是由于尚未创建FTS索引): {e}") return [] else: diff --git a/zhenxun/services/ai/context/rag/builder.py b/zhenxun/services/ai/context/rag/builder.py index 883a66b0..753d5c96 100644 --- a/zhenxun/services/ai/context/rag/builder.py +++ b/zhenxun/services/ai/context/rag/builder.py @@ -1,9 +1,11 @@ from typing import Any -from zhenxun.services.ai.context.rag.backends import StorageBackend -from zhenxun.services.ai.context.rag.configs import RAGConfig -from zhenxun.services.ai.context.rag.engine import ScopedRAGClient -from zhenxun.services.ai.context.rag.ingestion import ( +from zhenxun.services.ai.utils.logger import log_rag as logger + +from .backends import StorageBackend +from .configs import RAGConfig +from .engine import ScopedRAGClient +from .ingestion import ( ChunkingStrategy, DedupNode, DocumentChunking, @@ -12,7 +14,7 @@ from zhenxun.services.ai.context.rag.ingestion import ( IndexPipeline, StorageCommitNode, ) -from zhenxun.services.ai.context.rag.retrieval import ( +from .retrieval import ( BaseRetriever, DatabaseSparseRetriever, HybridRetriever, @@ -23,7 +25,6 @@ from zhenxun.services.ai.context.rag.retrieval import ( RerankRetriever, VectorDBRetriever, ) -from zhenxun.services.log import logger class RAGBuilder: @@ -201,9 +202,10 @@ class RAGBuilder: embedder = cfg.embedder if not embedder: - from zhenxun.services.ai.context.rag.backends import DefaultEmbedder from zhenxun.services.ai.llm.manager import get_default_model + from .backends import DefaultEmbedder + embedder = DefaultEmbedder(model_name=get_default_model("embedding")) logger.debug("RAGBuilder: 未指定 Embedder,已使用系统默认 Embedder。") diff --git a/zhenxun/services/ai/context/rag/configs.py b/zhenxun/services/ai/context/rag/configs.py index d2938a61..b2f64d68 100644 --- a/zhenxun/services/ai/context/rag/configs.py +++ b/zhenxun/services/ai/context/rag/configs.py @@ -2,9 +2,9 @@ from typing import Any from pydantic import BaseModel, ConfigDict, Field -from zhenxun.services.ai.context.rag.backends import Embedder, StorageBackend -from zhenxun.services.ai.context.rag.ingestion import ChunkingStrategy -from zhenxun.services.ai.context.rag.retrieval import ( +from .backends import Embedder, StorageBackend +from .ingestion import ChunkingStrategy +from .retrieval import ( BaseRetriever, PostProcessor, PreProcessor, diff --git a/zhenxun/services/ai/context/rag/engine.py b/zhenxun/services/ai/context/rag/engine.py index b6f870a0..bad85aed 100644 --- a/zhenxun/services/ai/context/rag/engine.py +++ b/zhenxun/services/ai/context/rag/engine.py @@ -1,20 +1,21 @@ from typing import Any -from zhenxun.services.ai.context.rag.backends import ( +from zhenxun.services.ai.utils.scope import normalize_scope_path + +from .backends import ( StorageBackend, ) -from zhenxun.services.ai.context.rag.configs import RAGConfig -from zhenxun.services.ai.context.rag.ingestion import ( +from .configs import RAGConfig +from .ingestion import ( IndexPipeline, ) -from zhenxun.services.ai.context.rag.models import ( +from .models import ( BaseRecord, SearchResult, ) -from zhenxun.services.ai.context.rag.retrieval import ( +from .retrieval import ( BaseRetriever, ) -from zhenxun.services.ai.utils.scope import normalize_scope_path class ScopedRAGClient: diff --git a/zhenxun/services/ai/context/rag/ingestion.py b/zhenxun/services/ai/context/rag/ingestion.py index 5f4e96d6..0974894d 100644 --- a/zhenxun/services/ai/context/rag/ingestion.py +++ b/zhenxun/services/ai/context/rag/ingestion.py @@ -2,9 +2,10 @@ from abc import ABC, abstractmethod import asyncio import re -from zhenxun.services.ai.context.rag.models import BaseRecord -from zhenxun.services.ai.context.rag.utils import cosine_similarity -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_rag as logger + +from .models import BaseRecord +from .utils import cosine_similarity class ChunkingStrategy(ABC): diff --git a/zhenxun/services/ai/context/rag/retrieval.py b/zhenxun/services/ai/context/rag/retrieval.py index b306ca15..fff2e05b 100644 --- a/zhenxun/services/ai/context/rag/retrieval.py +++ b/zhenxun/services/ai/context/rag/retrieval.py @@ -3,11 +3,12 @@ import asyncio import time from typing import TYPE_CHECKING, Any, Protocol, cast, runtime_checkable -from zhenxun.services.ai.context.rag.models import QueryRequest, SearchResult -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_rag as logger + +from .models import QueryRequest, SearchResult if TYPE_CHECKING: - from zhenxun.services.ai.context.rag.backends.storages import StorageBackend + from .backends.storages import StorageBackend def normalize_query_text(query: Any) -> str: diff --git a/zhenxun/services/ai/core/engine/append_only.py b/zhenxun/services/ai/core/engine/append_only.py index 1a68cd86..307f1b1a 100644 --- a/zhenxun/services/ai/core/engine/append_only.py +++ b/zhenxun/services/ai/core/engine/append_only.py @@ -22,19 +22,23 @@ class StablePrefix: """ def __init__(self): + """初始化稳定前缀实例。""" self._snapshot: StablePrefixSnapshot | None = None self._version = 0 @property def fingerprint(self) -> str: + """获取当前快照的唯一指纹。""" return self._snapshot.fingerprint if self._snapshot else "" @property def version(self) -> int: + """获取当前前缀的版本号。""" return self._version @property def built(self) -> bool: + """判断当前是否已构建快照。""" return self._snapshot is not None def build(self, system_prompt: list[str], tools: list[Any]) -> bool: @@ -50,9 +54,11 @@ class StablePrefix: return True def invalidate(self): + """使当前前缀快照失效。""" self._snapshot = None def to_context(self) -> tuple[list[str], list[Any]]: + """将前缀快照导出为上下文元组。""" if not self._snapshot: raise RuntimeError("StablePrefix.to_context() called before build()") return self._snapshot.system_prompt, self._snapshot.tools @@ -60,6 +66,7 @@ class StablePrefix: def _take_snapshot( self, system_prompt: list[str], tools: list[Any] ) -> StablePrefixSnapshot: + """为系统提示词与工具列表生成指纹快照。""" parsed_tools = [] for t in tools: if hasattr(t, "name"): @@ -83,19 +90,24 @@ class AppendOnlyLog: """追加写入模式 of 对话日志管理器""" def __init__(self): + """初始化追加对话日志。""" self._entries: list[Any] = [] @property def length(self) -> int: + """获取当前对话日志的条数。""" return len(self._entries) def append(self, message: Any): + """向对话日志中追加单条消息。""" self._entries.append(message) def extend(self, messages: list[Any]): + """批量追加多条消息到对话日志。""" self._entries.extend(messages) def clear(self): + """清空所有对话日志。""" self._entries.clear() def to_messages(self) -> list[Any]: @@ -110,6 +122,7 @@ class AppendOnlyContextManager: """ def __init__(self): + """初始化上下文管理器。""" self.prefix = StablePrefix() self.log = AppendOnlyLog() self._last_sync_count = 0 @@ -158,7 +171,7 @@ class AppendOnlyContextManager: def _compute_digest(self, messages: list[Any]) -> int: """核心:计算消息列表 of 指纹,用于识别内容篡改。包含 role 与 content。""" - from zhenxun.services.ai.core.engine.context_renderer import ContextConverter + from .context_renderer import ContextConverter payloads = [] flattened = ContextConverter.flatten_to_llm_messages(messages) diff --git a/zhenxun/services/ai/core/engine/context_renderer.py b/zhenxun/services/ai/core/engine/context_renderer.py index 87304f0a..7c15e267 100644 --- a/zhenxun/services/ai/core/engine/context_renderer.py +++ b/zhenxun/services/ai/core/engine/context_renderer.py @@ -2,7 +2,7 @@ from collections.abc import Sequence from typing import Any from zhenxun.services.ai.core.messages import AgentEvent, AgentMessage, LLMMessage -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_core as logger class ContextConverter: diff --git a/zhenxun/services/ai/core/engine/structured_parser.py b/zhenxun/services/ai/core/engine/structured_parser.py index 21997fce..7c40ff16 100644 --- a/zhenxun/services/ai/core/engine/structured_parser.py +++ b/zhenxun/services/ai/core/engine/structured_parser.py @@ -14,7 +14,7 @@ from zhenxun.services.ai.core.models import ToolDefinition from zhenxun.services.ai.run.models import OutputDataT from zhenxun.services.ai.tools.core.tool import BaseTool from zhenxun.services.ai.tools.models import StructuredSubmissionResult, ToolResult -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_core as logger from zhenxun.utils.pydantic_compat import model_json_schema, model_validate DEFAULT_IVR_TEMPLATE = ( diff --git a/zhenxun/services/ai/core/messages/__init__.py b/zhenxun/services/ai/core/messages/__init__.py index 4fc4a55a..a3c7333f 100644 --- a/zhenxun/services/ai/core/messages/__init__.py +++ b/zhenxun/services/ai/core/messages/__init__.py @@ -2,7 +2,7 @@ 消息与响应域类型定义 - 统一导出门面 """ -from nonebot.compat import PYDANTIC_V2 +from zhenxun.utils.pydantic_compat import model_rebuild from .context_events import ( AgentEvent, @@ -59,30 +59,25 @@ from .shared import ( from .types import ( AgentMessage, AnyLLMMessage, - AssistantContentUnion, ContentT, PromptInput, RoleT, - SystemContentUnion, - ToolContentUnion, UserContentUnion, ) -if PYDANTIC_V2: - ChatResponse.model_rebuild() - LLMMessage.model_rebuild() - SystemMessage.model_rebuild() - UserMessage.model_rebuild() - AssistantMessage.model_rebuild() - ToolMessage.model_rebuild() - RerankResponse.model_rebuild() +model_rebuild(ChatResponse) +model_rebuild(LLMMessage) +model_rebuild(SystemMessage) +model_rebuild(UserMessage) +model_rebuild(AssistantMessage) +model_rebuild(ToolMessage) +model_rebuild(RerankResponse) __all__ = [ "AgentEvent", "AgentMessage", "AnyLLMMessage", - "AssistantContentUnion", "AssistantMessage", "AudioPart", "AudioResponse", @@ -112,7 +107,6 @@ __all__ = [ "RerankResult", "RoleT", "SpeechRequest", - "SystemContentUnion", "SystemMessage", "TaskLifecycleEvent", "TextDeltaPart", @@ -121,7 +115,6 @@ __all__ = [ "ThoughtPart", "ToolCallDeltaPart", "ToolCallPart", - "ToolContentUnion", "ToolMessage", "ToolReturnPart", "UsageInfo", diff --git a/zhenxun/services/ai/core/messages/parts.py b/zhenxun/services/ai/core/messages/parts.py index f7030370..25b46d09 100644 --- a/zhenxun/services/ai/core/messages/parts.py +++ b/zhenxun/services/ai/core/messages/parts.py @@ -7,6 +7,7 @@ from typing_extensions import Self from pydantic import BaseModel, Field +from zhenxun.services.ai.utils.logger import log_core as logger from zhenxun.utils.pydantic_compat import model_validator @@ -309,7 +310,6 @@ class EmbedBatch(BaseModel): def to_text_only(self, context_name: str) -> list[str]: """降级工具:将多模态的批量向量安全剔除图片等内容,回退为纯文本数组""" - from zhenxun.services.log import logger texts = [] for payload in self.payloads: diff --git a/zhenxun/services/ai/core/messages/types.py b/zhenxun/services/ai/core/messages/types.py index 5d4dfe52..c1676f8f 100644 --- a/zhenxun/services/ai/core/messages/types.py +++ b/zhenxun/services/ai/core/messages/types.py @@ -12,33 +12,12 @@ from .parts import ( ImagePart, LLMContentPart, TextPart, - ThoughtPart, - ToolCallPart, - ToolReturnPart, VideoPart, ) -SystemContentUnion = TextPart -"""系统消息允许的内容片段联合类型""" - UserContentUnion = TextPart | ImagePart | AudioPart | VideoPart | FilePart """用户消息允许的多模态内容片段联合类型""" -AssistantContentUnion = ( - TextPart - | ThoughtPart - | ToolCallPart - | ToolReturnPart - | ImagePart - | AudioPart - | VideoPart - | FilePart -) -"""助手回复允许的内容片段联合类型""" - -ToolContentUnion = ToolReturnPart | ImagePart | AudioPart | VideoPart | FilePart -"""工具消息允许的内容片段联合类型""" - RoleT = TypeVar("RoleT", default=str, covariant=True) """泛型:消息参与者角色类型变量""" @@ -72,12 +51,9 @@ AgentMessage = LLMMessage | AgentEvent __all__ = [ "AgentMessage", "AnyLLMMessage", - "AssistantContentUnion", "ContentT", "LLMContentPart", "PromptInput", "RoleT", - "SystemContentUnion", - "ToolContentUnion", "UserContentUnion", ] diff --git a/zhenxun/services/ai/core/models.py b/zhenxun/services/ai/core/models.py index d61da8c6..1afdefed 100644 --- a/zhenxun/services/ai/core/models.py +++ b/zhenxun/services/ai/core/models.py @@ -9,7 +9,7 @@ from typing import Any, Generic, Literal, TypeVar from pydantic import BaseModel, ConfigDict, Field -from zhenxun.services.ai.core.options import GenerationConfig +from .options import GenerationConfig ModelName = str | None @@ -131,6 +131,8 @@ class ModelCapabilities(BaseModel): """模型支持的输出模态集合。""" supports_tool_calling: bool = False """是否支持工具调用能力。""" + supports_thinking_toggle: bool = False + """是否支持通过 {"thinking": {"type": "enabled/disabled"}} 显式控制思考模式。""" is_embedding_model: bool = False """是否为嵌入模型。""" is_rerank_model: bool = False @@ -187,7 +189,7 @@ class ModelDetail(BaseModel): """模型是否可用。""" temperature: float | None = None """采样温度参数。""" - generation_max_tokens: int | None = None + max_output_tokens: int | None = None """单次生成最大 Token 数。""" api_type: str | None = None """API 类型标识。""" @@ -197,6 +199,10 @@ class ModelDetail(BaseModel): """显式声明的主任务类型 (如 'image_generation')。""" path_prefix: str | None = Field(default=None) """中转路由前缀,例如 '/cogvideox' 或 '/minimax'。""" + max_input_tokens: int | None = None + """最大输入上下文窗口(用于控制记忆压缩策略)""" + reasoning_effort: str | None = None + """该模型的默认思考/推理等级(如 'high', 'low', 'none')""" __all__ = [ diff --git a/zhenxun/services/ai/core/options.py b/zhenxun/services/ai/core/options.py index e465b4db..06275646 100644 --- a/zhenxun/services/ai/core/options.py +++ b/zhenxun/services/ai/core/options.py @@ -212,13 +212,6 @@ class GeminiOptions(BaseProviderOption): """检索定位配置,如 LBS 经纬度信息,配合 Google Maps 工具使用""" -class DeepSeekOptions(BaseProviderOption): - """DeepSeek 专属特权参数""" - - thinking: bool | None = Field(default=None) - """是否强制开启或关闭 R1 模型的思维链过程""" - - class OpenAITTSOptions(BaseProviderOption): """OpenAI TTS 专属特权参数""" @@ -331,8 +324,6 @@ class GenerationConfig(BaseModel): """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/core/protocols/tool.py b/zhenxun/services/ai/core/protocols/tool.py index 385a819d..989c1115 100644 --- a/zhenxun/services/ai/core/protocols/tool.py +++ b/zhenxun/services/ai/core/protocols/tool.py @@ -67,8 +67,6 @@ class ToolProvider(Protocol): async def discover_tools( self, - allowed_servers: list[str] | None = None, - excluded_servers: list[str] | None = None, ) -> dict[str, ToolExecutable]: """ 异步发现此提供者提供的所有工具。 diff --git a/zhenxun/services/ai/core/stream_events.py b/zhenxun/services/ai/core/stream_events.py index 9c208196..1bf8da23 100644 --- a/zhenxun/services/ai/core/stream_events.py +++ b/zhenxun/services/ai/core/stream_events.py @@ -8,7 +8,7 @@ from typing import Any, TypeVar from nonebot.utils import is_coroutine_callable from pydantic import BaseModel, ConfigDict -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_core as logger class AgentStreamEvent(BaseModel): diff --git a/zhenxun/services/ai/core/templates.py b/zhenxun/services/ai/core/templates.py index c27222e6..4a010d17 100644 --- a/zhenxun/services/ai/core/templates.py +++ b/zhenxun/services/ai/core/templates.py @@ -4,7 +4,7 @@ from typing_extensions import TypeVar from jinja2 import Environment -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_core as logger from zhenxun.utils.pydantic_compat import model_dump if TYPE_CHECKING: diff --git a/zhenxun/services/ai/flow/agent/agent.py b/zhenxun/services/ai/flow/agent/agent.py index fcabfc94..598f87f6 100644 --- a/zhenxun/services/ai/flow/agent/agent.py +++ b/zhenxun/services/ai/flow/agent/agent.py @@ -5,9 +5,8 @@ from pathlib import Path from typing import Any, Generic, cast from zhenxun.services.ai.capabilities import ( - AbstractCapability, + CapabilitySource, CombinedCapability, - DynamicCapability, ) from zhenxun.services.ai.config import get_llm_config from zhenxun.services.ai.context.knowledge.base import BaseKnowledge @@ -67,7 +66,7 @@ from zhenxun.services.ai.tools.providers.skills.capabilities import ( ) 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.services.ai.utils.logger import log_agent as logger from zhenxun.utils.pydantic_compat import ( model_construct, model_copy, @@ -78,7 +77,9 @@ from zhenxun.utils.utils import infer_plugin_namespace from .engine.builders import ( AgentProfileResolver, + CapabilityBuilder, ContextBuilder, + SessionBuilder, ToolBuilder, ) from .engine.executor import BaseAgentExecutor, StandardAgentExecutor @@ -94,9 +95,6 @@ ToolSource = ( ) """任何可以作为工具提供给大模型的实体对象(函数、基础工具类、字典定义、工具名、工具箱、声明式查询对象)""" -CapabilitySource = Callable | AbstractCapability -"""能力/拦截器来源(函数或 AbstractCapability 实例)""" - class AgentBuilder(Generic[AgentDepsT, OutputDataT]): """ @@ -148,7 +146,7 @@ class AgentBuilder(Generic[AgentDepsT, OutputDataT]): return self def with_tools( - self, *tools: ToolSource | list[ToolSource] + self, *tools: ToolSource | Sequence[ToolSource] ) -> "AgentBuilder[AgentDepsT, OutputDataT]": """ 配置可供智能体调用的工具列表。 @@ -158,7 +156,7 @@ class AgentBuilder(Generic[AgentDepsT, OutputDataT]): """ current_tools = self._kwargs.setdefault("tools", []) for t in tools: - if isinstance(t, list): + if isinstance(t, Sequence) and not isinstance(t, str): current_tools.extend(t) else: current_tools.append(t) @@ -353,7 +351,7 @@ class Agent( description: str | None = None, persona: Persona | dict | None = None, model: str | Callable[[], str] | None = None, - tools: list[ToolSource] | None = None, + tools: Sequence[ToolSource] | None = None, skills: Sequence[str | Path | Skill | SkillSource] | None = None, generation_config: GenerationConfig | IntentBuilder | dict | None = None, response_model: BaseOutputDefinition | type[OutputDataT] | None = None, @@ -458,14 +456,14 @@ class Agent( def _assemble_plugins(self, tools, knowledge, capabilities, skills): """私有方法:集中处理各类能力、知识与技能的挂载,消解冗余样板代码""" - self.tool_definitions = tools or [] + self.tool_definitions = list(tools) if tools else [] if knowledge: if not isinstance(knowledge, list): knowledge = [knowledge] self.tool_definitions.extend(knowledge) - self.capabilities: list[AbstractCapability] = [] + self.capabilities: list[CapabilitySource] = [] if self.memory_config.long_term.enable and self.memory_config.long_term.agentic: self.capabilities.append( @@ -478,11 +476,7 @@ class Agent( ) if capabilities: - for cap in capabilities: - if isinstance(cap, AbstractCapability): - self.capabilities.append(cap) - elif callable(cap): - self.capabilities.append(DynamicCapability(cap)) + self.capabilities.extend(capabilities) if self.config.enable_hitl: self.tool_definitions.append(HITLToolkit()) @@ -789,11 +783,6 @@ class Agent( **kwargs: Any, ) -> tuple[AgentState, AgentRunResources]: """解析任务意图,初始化隔离域与基础状态载体""" - from zhenxun.services.ai.flow.agent.engine.builders import ( - AgentProfileResolver, - CapabilityBuilder, - SessionBuilder, - ) if context is None: raise ValueError("RunContext 不能为空") @@ -900,7 +889,6 @@ class Agent( toolset_funcs=getattr(self, "toolset_funcs", []), system_tools=[], namespace=self.namespace or "unknown", - tool_filter=resources.config.tool_filter, run_context=context, run_scoped_cap=run_scoped_cap, ) diff --git a/zhenxun/services/ai/flow/agent/bridge.py b/zhenxun/services/ai/flow/agent/bridge.py index 61ba4a32..6fd2ad2a 100644 --- a/zhenxun/services/ai/flow/agent/bridge.py +++ b/zhenxun/services/ai/flow/agent/bridge.py @@ -1,6 +1,7 @@ import asyncio from typing import Any, Generic, cast from typing_extensions import TypeVar +import uuid from nonebot.adapters import Bot, Event from nonebot_plugin_alconna.uniseg import UniMessage @@ -16,8 +17,9 @@ from zhenxun.services.ai.flow.base import BaseRunnable from zhenxun.services.ai.run import AgentRunResult, RunContext from zhenxun.services.ai.run.models import AgentRunEnd, AgentRunError from zhenxun.services.ai.run.ui import UIController -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_agent as logger from zhenxun.utils.message import MessageUtils +from zhenxun.utils.platform import PlatformUtils T_Deps = TypeVar("T_Deps", default=Any) T_Out = TypeVar("T_Out", default=str) @@ -46,8 +48,6 @@ class AgentRunner(Generic[T_Out]): ) if is_stateless and self.context.session_id: if not self.context.session_id.startswith("stateless_"): - import uuid - self.context.session_id = ( f"stateless_{self.context.session_id}_{uuid.uuid4().hex[:8]}" ) @@ -126,15 +126,24 @@ class AgentRunner(Generic[T_Out]): await MessageUtils.build_message(f"❌ 运行发生错误: {e}").send() raise e - if final_result and final_result.output and self._bot and self._event: - if isinstance(final_result.output, UniMessage): - await final_result.output.send( - self._event, bot=self._bot, reply_to=reply_to - ) - final_result.output = final_result.output.extract_plain_text() + if final_result and final_result.output and self._bot: + msg_to_send = ( + final_result.output + if isinstance(final_result.output, UniMessage) + else MessageUtils.build_message(str(final_result.output)) + ) + if self._event: + await msg_to_send.send(self._event, bot=self._bot, reply_to=reply_to) else: - final_msg = str(final_result.output) - await MessageUtils.build_message(final_msg).send() + target = PlatformUtils.get_target( + user_id=self.context.get_user_id(), + group_id=self.context.get_group_id(), + ) + if target: + await msg_to_send.send(target=target, bot=self._bot) + + if isinstance(final_result.output, UniMessage): + final_result.output = final_result.output.extract_plain_text() if final_result is None: raise RuntimeError("智能体运行流异常结束:未返回最终结果。") diff --git a/zhenxun/services/ai/flow/agent/capabilities.py b/zhenxun/services/ai/flow/agent/capabilities.py index 0b5e0f22..756c1836 100644 --- a/zhenxun/services/ai/flow/agent/capabilities.py +++ b/zhenxun/services/ai/flow/agent/capabilities.py @@ -12,7 +12,7 @@ from zhenxun.services.ai.core.messages import TaskLifecycleEvent from zhenxun.services.ai.core.options import BaseOutputDefinition, ToolOutput from zhenxun.services.ai.guardrails import parse_guardrails from zhenxun.services.ai.run import AgentRunResult, AgentTask, RunContext -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_agent as logger class OutputValidationCapability(AbstractCapability): diff --git a/zhenxun/services/ai/flow/agent/engine/builders.py b/zhenxun/services/ai/flow/agent/engine/builders.py index db56d883..86a8b393 100644 --- a/zhenxun/services/ai/flow/agent/engine/builders.py +++ b/zhenxun/services/ai/flow/agent/engine/builders.py @@ -6,9 +6,7 @@ from typing import Any, cast from nonebot.utils import is_coroutine_callable 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 @@ -22,13 +20,13 @@ from zhenxun.services.ai.flow.agent.capabilities import ( TaskTrackingCapability, ) from zhenxun.services.ai.flow.agent.models import Persona -from zhenxun.services.ai.run import GLOBAL_CAPABILITIES, RunContext +from zhenxun.services.ai.run import 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.tools.models import ResolvedToolPayload from zhenxun.services.ai.utils.scope import ScopeSelector from zhenxun.utils.pydantic_compat import model_copy @@ -135,22 +133,21 @@ class CapabilityBuilder: if task_obj: dynamic_caps.append(TaskTrackingCapability(task_obj, agent_name)) - run_level_caps = [] - if profile_capabilities: - for cap in profile_capabilities: - if isinstance(cap, AbstractCapability): - run_level_caps.append(cap) - elif callable(cap): - run_level_caps.append(DynamicCapability(cap)) + from zhenxun.services.ai.capabilities.manager import capability_manager - base_caps = GLOBAL_CAPABILITIES.get("global", []).copy() - if namespace != "global" and namespace in GLOBAL_CAPABILITIES: - base_caps.extend(GLOBAL_CAPABILITIES[namespace]) + run_level_caps = capability_manager.resolve_capabilities( + profile_capabilities or [], namespace + ) + agent_level_caps = capability_manager.resolve_capabilities( + agent_capabilities or [], namespace + ) + + auto_caps = capability_manager.get_auto_apply_capabilities(namespace) combined_cap = CombinedCapability( - base_caps + auto_caps + getattr(context, "capabilities", []) - + agent_capabilities + + agent_level_caps + run_level_caps + dynamic_caps ) @@ -298,7 +295,6 @@ class ToolBuilder: toolset_funcs: list[Any], system_tools: list[Any], namespace: str, - tool_filter: GlobalToolFilter | None, run_context: RunContext, run_scoped_cap: CombinedCapability, ) -> ResolvedToolPayload: @@ -309,9 +305,8 @@ class ToolBuilder: 参数: tool_definitions: 静态工具或工具集合的定义列表。 toolset_funcs: 待依赖注入解析的工具集生成函数列表。 - system_tools: 系统默认强制集成的工具定义列表。 + system_tools: 系统默认强制集成的工具定义列表. namespace: 会话命名空间。 - tool_filter: 全局的工具过滤条件,此处为占位。 run_context: 运行上下文对象实例. run_scoped_cap: 运行域下的合并能力中间件,用于提供特定的能力工具。 @@ -411,14 +406,14 @@ class ToolBuilder: current_tool_defs = list(_cap_res) final_defs_map = {d.name.lower(): d for d in current_tool_defs if d} - final_effective_tools = ToolCollection() + final_effective_tools_list = [] for t_exec in effective_tools: t_name = getattr(t_exec, "name", "unknown") if t_name.lower() in final_defs_map: cloned_tool = copy.copy(t_exec) cloned_tool._dynamic_def = final_defs_map[t_name.lower()] - final_effective_tools.append(cloned_tool) - return final_effective_tools + final_effective_tools_list.append(cloned_tool) + return ToolCollection(final_effective_tools_list) class SessionBuilder: diff --git a/zhenxun/services/ai/flow/agent/engine/directive.py b/zhenxun/services/ai/flow/agent/engine/directive.py index cbf49a9e..0e8bbecf 100644 --- a/zhenxun/services/ai/flow/agent/engine/directive.py +++ b/zhenxun/services/ai/flow/agent/engine/directive.py @@ -4,7 +4,7 @@ from collections.abc import Awaitable, Callable from zhenxun.services.ai.flow.agent.models import AgentRunResources, AgentState from zhenxun.services.ai.run.models import AgentRunResult, HandoffPayload from zhenxun.services.ai.tools.models import ToolResult -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_agent as logger from zhenxun.utils.pydantic_compat import model_construct from zhenxun.utils.utils import infer_plugin_namespace @@ -63,7 +63,7 @@ async def handle_submit_structured( parsed_obj = ( tool_res.directive.payload.get("parsed_obj") if tool_res.directive else None ) - logger.info("✅ 拦截到结构化结果提交,结束循环。") + logger.debug("✅ 拦截到结构化结果提交,结束循环。") state.is_finished = True state.final_result = model_construct( AgentRunResult, diff --git a/zhenxun/services/ai/flow/agent/engine/executor.py b/zhenxun/services/ai/flow/agent/engine/executor.py index f6f8af40..37417bdd 100644 --- a/zhenxun/services/ai/flow/agent/engine/executor.py +++ b/zhenxun/services/ai/flow/agent/engine/executor.py @@ -1,7 +1,7 @@ from abc import ABC, abstractmethod import asyncio import json -from typing import Any, cast +from typing import Any from zhenxun.services.ai.capabilities import CombinedCapability from zhenxun.services.ai.core.engine.context_renderer import ContextConverter @@ -15,7 +15,6 @@ from zhenxun.services.ai.core.exceptions import ( ) from zhenxun.services.ai.core.messages import ( AgentMessage, - AssistantContentUnion, AssistantMessage, AudioPart, ChatRequest, @@ -35,19 +34,20 @@ from zhenxun.services.ai.core.stream_events import ( LLMStartEvent, ToolStreamChunkEvent, ) -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 +from zhenxun.services.ai.utils.logger import log_agent as logger from zhenxun.utils.pydantic_compat import dump_json_safely, model_construct +from .directive import ( + DirectiveHandlerFunc, + directive_manager, +) + class BaseAgentExecutor(ABC): """ @@ -219,13 +219,12 @@ class StandardAgentExecutor(BaseAgentExecutor): messages, run_context ) - if run_context.run.event_bus: - await run_context.run.event_bus.emit( - LLMStartEvent( - model_name=resources.model_name or "unknown", - messages=flattened_messages, - ) + await run_context.run.emit( + LLMStartEvent( + model_name=resources.model_name or "unknown", + messages=flattened_messages, ) + ) response = await self._execute_model_request( model_name=resources.model_name, @@ -238,8 +237,7 @@ class StandardAgentExecutor(BaseAgentExecutor): cancellation_token=cancellation_token, ) - if run_context.run.event_bus: - await run_context.run.event_bus.emit(LLMEndEvent(response=response)) + await run_context.run.emit(LLMEndEvent(response=response)) assistant_content = ( response.content_parts if response.content_parts else response.text @@ -252,9 +250,7 @@ class StandardAgentExecutor(BaseAgentExecutor): part.metadata["thought_signature"] = response.thought_signature break - assistant_message = AssistantMessage( - content=cast(list[AssistantContentUnion], response.content_parts) - ) + assistant_message = AssistantMessage(content=response.content_parts) if hasattr(response, "parsed_obj") and response.parsed_obj is not None: if not isinstance(response.parsed_obj, str): @@ -376,7 +372,7 @@ class StandardAgentExecutor(BaseAgentExecutor): state.messages, resources.model_name or "", base_overhead=0 ) logger.debug( - f"[TokenTracker] (Iter {state.current_cycle + 1}) " + f"(Iter {state.current_cycle + 1}) " f"预估将消耗 {est_tokens} Token " f"(Model: {resources.model_name or 'Unknown'})" ) @@ -444,7 +440,7 @@ class StandardAgentExecutor(BaseAgentExecutor): return if not response.tool_calls: - logger.debug("✅ AgentExecutor:模型未请求工具调用,推理循环结束。") + logger.debug("✅ 模型未请求工具调用,推理循环结束。") state.is_finished = True state.final_result = model_construct( AgentRunResult, @@ -476,8 +472,7 @@ class StandardAgentExecutor(BaseAgentExecutor): if is_server_side: logger.debug( - "☁️ [AgentExecutor] 检测到云端工具调用: " - f"{call.tool_name},已跳过本地执行。" + f"☁️ 检测到云端工具调用: {call.tool_name},已跳过本地执行。" ) async with self.tool_executor._tool_stream_scope( event_bus, @@ -500,7 +495,7 @@ class StandardAgentExecutor(BaseAgentExecutor): client_tool_calls.append(call) if not client_tool_calls: - logger.info("✅ AgentExecutor:无本地客户端工具需执行,推理循环平滑结束。") + logger.info("✅ 无本地客户端工具需执行,推理循环平滑结束。") state.is_finished = True state.final_result = model_construct( @@ -638,7 +633,6 @@ class StandardAgentExecutor(BaseAgentExecutor): self, state: AgentState, resources: AgentRunResources ) -> AgentRunResult[Any]: run_context = resources.run_context - event_bus = run_context.run.event_bus if not resources.config.enable_fallback_summary: raise UpstreamServerException( @@ -646,17 +640,15 @@ class StandardAgentExecutor(BaseAgentExecutor): ) logger.warning( - f"AgentExecutor 达到最大循环次数 ({resources.config.max_cycles})," - "触发兜底总结机制。" + f"达到最大循环次数 ({resources.config.max_cycles}),触发兜底总结机制。" ) - if event_bus: - await event_bus.emit( - ToolStreamChunkEvent( - tool_name="System", - content="⏳ 思考过程过于复杂,正在强制生成最终总结...", - ) + await run_context.run.emit( + ToolStreamChunkEvent( + tool_name="System", + content="⏳ 思考过程过于复杂,正在强制生成最终总结...", ) + ) fallback_msg = LLMMessage.user( "### 🚨 [系统强制指令]\n" diff --git a/zhenxun/services/ai/flow/agent/models.py b/zhenxun/services/ai/flow/agent/models.py index 7cf0cc3d..43b0a020 100644 --- a/zhenxun/services/ai/flow/agent/models.py +++ b/zhenxun/services/ai/flow/agent/models.py @@ -1,12 +1,15 @@ -""" -Agent 相关静态声明类型定义 -""" +from __future__ import annotations from collections.abc import Sequence +from pathlib import Path from typing import Any from pydantic import BaseModel, ConfigDict, Field +from zhenxun.services.ai.capabilities import CapabilitySource, CombinedCapability +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.context.memory.types import SessionMetadata from zhenxun.services.ai.core.messages import ( AgentMessage, @@ -18,8 +21,9 @@ from zhenxun.services.ai.core.messages import ( from zhenxun.services.ai.core.options import GenerationConfig from zhenxun.services.ai.flow.base import BaseRuntimeConfig from zhenxun.services.ai.run import RunContext +from zhenxun.services.ai.run.models import AgentRunResult, AgentTask, HandoffPayload +from zhenxun.services.ai.tools.core.toolkit import BaseToolkit from zhenxun.services.ai.tools.engine.registry import ToolCollection -from zhenxun.services.ai.tools.models import GlobalToolFilter from zhenxun.utils.pydantic_compat import model_copy @@ -60,15 +64,14 @@ class AgentConfig(BaseRuntimeConfig): message_history: Sequence[AgentMessage] | None = Field(default=None) """初始化的底层对话历史记录。""" - tool_filter: GlobalToolFilter | None = Field(default=None) - """全局工具过滤器,限制本次运行可用的工具池。""" - memory: Any | None = Field(default=None) + + memory: MemoryConfig | MemoryBuilder | bool | None = Field(default=None) """单次运行级别的记忆门面覆盖 (支持 bool, MemoryConfig, MemoryBuilder)。""" generation_config: GenerationConfig | None = Field(default=None) """单次运行覆盖的大模型生成配置。""" - capabilities: list[Any] | None = Field(default=None) + capabilities: list[CapabilitySource] | None = Field(default=None) """仅针对本次运行动态注入的临时拦截器/能力组件列表。""" - skills: Sequence[Any] | None = Field(default=None) + skills: Sequence[str | Path | Any] | None = Field(default=None) """仅针对本次运行动态注入的临时技能集合。""" executor: Any | None = Field(default=None) """单次运行覆盖的核心执行引擎策略 (BaseAgentExecutor)。""" @@ -128,11 +131,11 @@ class AgentState(BaseModel): """拦截到的早期终止输出结果""" should_terminate: bool = False """标记是否应提前终止循环""" - handoff_triggered: Any | None = None + handoff_triggered: HandoffPayload | None = None """标记是否触发了移交""" is_finished: bool = False """标记大模型循环是否彻底结束""" - final_result: Any | None = None + final_result: AgentRunResult[Any] | None = None """最终的运行结果 (AgentRunResult)""" origin_msg_len: int = 0 """初始进入循环时的消息历史长度 (用于增量保存记忆)""" @@ -160,15 +163,15 @@ class AgentRunResources(BaseModel): """保留依赖注入(DI)与黑板引用的全局运行时上下文""" session_meta: SessionMetadata | None = None """隔离会话的元信息(Session ID, 命名空间, 权限等)""" - memory_reader: Any | None = None + memory_reader: MemoryReader | None = None """用于读取短/中/长期上下文记忆的读取器""" - memory_writer: Any | None = None + memory_writer: MemoryWriter | None = None """用于将对话历史安全落盘的写入器""" - run_scoped_cap: Any | None = None + run_scoped_cap: CombinedCapability | None = None """聚合了 Agent/AgentTask/全局 的复合能力拦截器 (CombinedCapability)""" - task_obj: Any | None = None + task_obj: AgentTask | None = None """(如有) 解析后的结构化数据任务契约""" - toolkits: list[Any] = Field(default_factory=list) + toolkits: list[BaseToolkit] = Field(default_factory=list) """当前轮次生效的工具箱列表 (需要执行生命周期挂载)""" config: AgentConfig = Field(default_factory=AgentConfig) """Agent 全局与运行时的统一策略配置""" diff --git a/zhenxun/services/ai/flow/base.py b/zhenxun/services/ai/flow/base.py index 8775a1b0..8f53e2ed 100644 --- a/zhenxun/services/ai/flow/base.py +++ b/zhenxun/services/ai/flow/base.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from abc import ABC, abstractmethod import asyncio from collections.abc import AsyncIterator @@ -7,12 +9,14 @@ from typing import TYPE_CHECKING, Any, Generic, TypeVar, cast from pydantic import BaseModel, Field +from zhenxun.services.ai.core.exceptions import ControlFlowExit from zhenxun.services.ai.run.context import RunContext +from zhenxun.services.ai.run.models import StreamedRunResult from zhenxun.services.ai.run.ui import UIController +from zhenxun.services.ai.utils.logger import log_flow as logger if TYPE_CHECKING: - from zhenxun.services.ai.flow.agent.models import Persona - from zhenxun.services.ai.run.models import StreamedRunResult + from .agent.models import Persona from zhenxun.services.ai.core.messages import PromptInput @@ -92,7 +96,7 @@ class BaseRunnable(ABC, Generic[T_RunResult]): """DI 注入语法糖:返回 Depends,自动绑定当前上下文""" from nonebot.params import Depends - from zhenxun.services.ai.flow.agent.bridge import AgentRunner + from .agent.bridge import AgentRunner async def _dependency() -> AgentRunner[Any]: return AgentRunner[Any](self, **kwargs) @@ -108,7 +112,7 @@ class BaseRunnable(ABC, Generic[T_RunResult]): **kwargs: Any, ) -> T_RunResult: """交互执行语法糖,自动渲染流式进度并最终将结果回复给终端用户""" - from zhenxun.services.ai.flow.agent.bridge import AgentRunner + from .agent.bridge import AgentRunner runner = AgentRunner(self, context=context, **kwargs) return cast(T_RunResult, await runner.reply(prompt=prompt, reply_to=reply_to)) @@ -121,8 +125,6 @@ class BaseRunnable(ABC, Generic[T_RunResult]): **kwargs: Any, ) -> T_RunResult: """阻塞式核心运行入口,安全捕获内部抛出的静默退出信号""" - from zhenxun.services.ai.core.exceptions import ControlFlowExit - from zhenxun.services.log import logger try: async with self.run_stream( @@ -144,6 +146,6 @@ class BaseRunnable(ABC, Generic[T_RunResult]): *, context: RunContext | None = None, **kwargs: Any, - ) -> "AsyncIterator[StreamedRunResult[Any]]": + ) -> AsyncIterator[StreamedRunResult[Any]]: """流式运行入口,返回上下文管理器,用于消费底层执行流事件 (StreamedRunResult)""" yield cast(Any, None) diff --git a/zhenxun/services/ai/flow/concurrency.py b/zhenxun/services/ai/flow/concurrency.py index bfdffded..0af17685 100644 --- a/zhenxun/services/ai/flow/concurrency.py +++ b/zhenxun/services/ai/flow/concurrency.py @@ -9,7 +9,7 @@ from zhenxun.services.ai.core.exceptions import ( from zhenxun.services.ai.core.models import CancellationToken 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 zhenxun.services.ai.utils.logger import log_flow as logger from .base import ConcurrencyPolicy, InterventionPolicy diff --git a/zhenxun/services/ai/flow/team/capabilities.py b/zhenxun/services/ai/flow/team/capabilities.py index 699dd6a2..66dbfec5 100644 --- a/zhenxun/services/ai/flow/team/capabilities.py +++ b/zhenxun/services/ai/flow/team/capabilities.py @@ -6,8 +6,11 @@ from nonebot.utils import is_coroutine_callable from zhenxun.services.ai.capabilities import AbstractCapability from zhenxun.services.ai.run import RunContext +from zhenxun.services.ai.run.di import DependencyInjector from zhenxun.services.ai.tools.bridges.handoff import HandoffTool +from .models import Transition + class TeamRoutingCapability(AbstractCapability): """团队路由能力组件:动态向所有团队成员""" @@ -17,10 +20,12 @@ class TeamRoutingCapability(AbstractCapability): team_name: str, members: list[Any], state_flow: Mapping[str, Sequence[Any]] | Callable | None = None, + max_handoffs: int = 3, ): self.team_name = team_name self.members = members self.state_flow = state_flow + self.max_handoffs = max_handoffs async def _get_allowed_transitions(self, context: RunContext) -> list[Any] | None: """核心FSM解析:解析静态字典或动态执行函数获取允许的 Transition 列表""" @@ -34,15 +39,12 @@ class TeamRoutingCapability(AbstractCapability): current_speaker, [m.name for m in self.members if m.name != current_speaker], ) - from zhenxun.services.ai.flow.team.models import Transition return [ Transition(target=t) if isinstance(t, str) else t for t in raw_targets ] if callable(self.state_flow): - from zhenxun.services.ai.run.di import DependencyInjector - sig = inspect.signature(self.state_flow) kwargs = await DependencyInjector.resolve_all( sig, call_kwargs={}, context=context @@ -55,7 +57,6 @@ class TeamRoutingCapability(AbstractCapability): if result is None: return None - from zhenxun.services.ai.flow.team.models import Transition return [Transition(target=t) if isinstance(t, str) else t for t in result] @@ -97,6 +98,7 @@ class TeamRoutingCapability(AbstractCapability): target_name=m.name, target_description=desc, input_schema=input_schema, + max_handoffs=self.max_handoffs, ) ) return tools diff --git a/zhenxun/services/ai/flow/team/models.py b/zhenxun/services/ai/flow/team/models.py index 3ed1ac35..2afa8dae 100644 --- a/zhenxun/services/ai/flow/team/models.py +++ b/zhenxun/services/ai/flow/team/models.py @@ -1,3 +1,5 @@ +from __future__ import annotations + from collections.abc import Callable, Sequence from enum import Enum from typing import Any @@ -7,7 +9,8 @@ from pydantic import BaseModel, ConfigDict, Field from zhenxun.services.ai.core.messages import AgentMessage from zhenxun.services.ai.core.options import BaseOutputDefinition -from zhenxun.services.ai.flow.base import BaseRuntimeConfig +from zhenxun.services.ai.flow.base import BaseRunnable, BaseRuntimeConfig +from zhenxun.services.ai.run import AgentTask class TeamRuntimeConfig(BaseRuntimeConfig): @@ -45,7 +48,7 @@ class Transition(BaseModel): trigger_regex: str | None = None """(可选) 正则表达式。 如果用户的输入匹配此正则,将触发极速硬路由,跳过大模型思考。""" - trigger_func: Callable[..., Any] | None = None + trigger_func: Callable[..., bool | str | None] | None = None """(可选) 自定义校验函数。返回 True 或目标名称时触发硬路由。支持依赖注入。""" model_config = ConfigDict(arbitrary_types_allowed=True) @@ -62,9 +65,9 @@ class CallAction(TeamAction): 调度动作:呼叫指定的 Agent 执行任务 """ - agent: str | Any + agent: str | BaseRunnable[Any] """目标 Agent 的名称(字符串)或动态生成的 Agent 实例""" - task: str | Any + task: str | AgentTask """派发给该 Agent 的具体任务或提示词""" history: Sequence[AgentMessage] | None = None """需要传递给该 Agent 的上下文历史记录(可选)""" diff --git a/zhenxun/services/ai/flow/team/router.py b/zhenxun/services/ai/flow/team/router.py index b2cb3e9b..f1d0055b 100644 --- a/zhenxun/services/ai/flow/team/router.py +++ b/zhenxun/services/ai/flow/team/router.py @@ -8,10 +8,11 @@ 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 AgentTask, RunContext from zhenxun.services.ai.run.di import DependencyInjector -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_team as logger + +from .models import RouteDecision, Transition class BaseRouter(ABC): @@ -148,6 +149,7 @@ class LLMRouter(BaseRouter): runtime_config: Any = None, custom_prompt: str | None = None, allowed_transitions: list[Transition] | None = None, + max_handoffs: int = 3, ): """ 初始化基于大模型的意图路由器。 @@ -161,6 +163,7 @@ class LLMRouter(BaseRouter): runtime_config: 团队级别的运行时全局配置。 custom_prompt: 自定义的系统提示词模板,用以覆盖默认的路由系统指令。 allowed_transitions: 允许的状态移交规则与前置条件列表。 + max_handoffs: 同一会话中允许连续移交的最大次数。 """ self.team_name = team_name self.members = members @@ -170,6 +173,7 @@ class LLMRouter(BaseRouter): self.runtime_config = runtime_config self.custom_prompt = custom_prompt self.allowed_transitions = allowed_transitions + self.max_handoffs = max_handoffs async def route( self, @@ -179,7 +183,8 @@ class LLMRouter(BaseRouter): ) -> RouteDecision | None: from zhenxun.services.ai.flow.agent.agent import Agent from zhenxun.services.ai.flow.agent.models import AgentConfig - from zhenxun.services.ai.flow.team.capabilities import TeamRoutingCapability + + from .capabilities import TeamRoutingCapability default_system_prompt = """## 角色与目标 你是一个高级任务路由器 (所在团队: {{ team_name }})。 @@ -200,7 +205,10 @@ class LLMRouter(BaseRouter): route_prompt = PromptTemplate(template).render(team_name=self.team_name) routing_cap = TeamRoutingCapability( - team_name=self.team_name, members=self.members, state_flow=self.state_flow + team_name=self.team_name, + members=self.members, + state_flow=self.state_flow, + max_handoffs=self.max_handoffs, ) leader_config = AgentConfig( diff --git a/zhenxun/services/ai/flow/team/runner.py b/zhenxun/services/ai/flow/team/runner.py index fadb6c0a..e23ef5d4 100644 --- a/zhenxun/services/ai/flow/team/runner.py +++ b/zhenxun/services/ai/flow/team/runner.py @@ -8,17 +8,19 @@ from zhenxun.services.ai.core.exceptions import ( LLMException, ) from zhenxun.services.ai.core.messages import UsageInfo -from zhenxun.services.ai.flow.team.capabilities import TeamRoutingCapability -from zhenxun.services.ai.flow.team.models import ( +from zhenxun.services.ai.flow.agent.models import AgentConfig +from zhenxun.services.ai.run import AgentRunResult, RunContext +from zhenxun.services.ai.run.models import AgentRunEnd +from zhenxun.services.ai.utils.logger import log_team as logger +from zhenxun.utils.pydantic_compat import model_construct + +from .capabilities import TeamRoutingCapability +from .models import ( CallAction, ConcurrentCallAction, FinishAction, ) -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 +from .strategy import BaseTeamStrategy, RouteStrategy class TeamRunner: @@ -47,7 +49,7 @@ class TeamRunner: (m for m in self.team.members if m.name == action.agent), None ) if not target_agent: - logger.error(f"❌ [TeamRunner] 找不到团队成员: {action.agent}") + logger.error(f"❌ 找不到团队成员: {action.agent}") await queue.put( ( "result", @@ -67,13 +69,12 @@ class TeamRunner: sub_context = context.clone_for_member(target_agent.name) sub_context.capabilities = list(sub_context.capabilities) - from zhenxun.services.ai.flow.team.strategy import RouteStrategy - if isinstance(self.strategy, RouteStrategy): routing_cap = TeamRoutingCapability( team_name=self.team.name, members=self.team.members, state_flow=getattr(self.strategy, "state_flow", None), + max_handoffs=getattr(self.strategy, "max_handoffs", 3), ) sub_context.capabilities.append(routing_cap) @@ -82,8 +83,6 @@ class TeamRunner: agent_res = None try: - from zhenxun.services.ai.flow.agent.models import AgentConfig - async with target_agent.run_stream( prompt=action.task, context=sub_context, diff --git a/zhenxun/services/ai/flow/team/strategy.py b/zhenxun/services/ai/flow/team/strategy.py index be099734..3c5fd665 100644 --- a/zhenxun/services/ai/flow/team/strategy.py +++ b/zhenxun/services/ai/flow/team/strategy.py @@ -10,19 +10,26 @@ from zhenxun.services.ai.core.stream_events import ToolStreamChunkEvent from zhenxun.services.ai.core.templates import PromptTemplate from zhenxun.services.ai.flow.agent.agent import Agent, ToolSource from zhenxun.services.ai.flow.agent.models import AgentConfig -from zhenxun.services.ai.flow.team.models import ( +from zhenxun.services.ai.run import AgentTask, RunContext +from zhenxun.services.ai.run.blackboard import BlackboardManager +from zhenxun.services.ai.tools.bridges.delegate import DelegateTool +from zhenxun.services.ai.tools.providers.builtin.blackboard import BlackboardToolkit +from zhenxun.services.ai.utils.logger import log_team as logger + +from .models import ( CallAction, ConcurrentCallAction, FinishAction, + TaskBoardState, + TaskNodeStatus, TeamAction, + Transition, ) -from zhenxun.services.ai.flow.team.router import BaseRouter -from zhenxun.services.ai.run import AgentTask, RunContext -from zhenxun.services.ai.tools.bridges.delegate import DelegateTool -from zhenxun.services.log import logger +from .router import BaseRouter +from .task_tools import TaskPlanningToolkit if TYPE_CHECKING: - from zhenxun.services.ai.flow.team.team import Team + from .team import Team class BaseTeamStrategy(ABC): @@ -106,6 +113,7 @@ class RouteStrategy(BaseTeamStrategy): leader_model: str | None = None, leader_tools: list[ToolSource] | None = None, custom_prompt: str | None = None, + max_handoffs: int = 3, ): """ 路由策略初始化,基于挂载的 Router 进行最合适的专家分发。 @@ -117,16 +125,16 @@ class RouteStrategy(BaseTeamStrategy): leader_model: 路由节点 (Leader) 使用的大模型名称,若为空则默认继承全局。 leader_tools: 挂载给路由节点 (Leader) 的专属工具列表。 custom_prompt: 自定义系统提示词,用于覆盖默认的路由系统提示词。 + max_handoffs: 同一会话中允许连续移交的最大次数。 """ super().__init__(custom_prompt=custom_prompt) self.selector_func = selector_func self.router = router self.leader_model = leader_model self.leader_tools = leader_tools or [] + self.max_handoffs = max_handoffs if isinstance(state_flow, dict): - from zhenxun.services.ai.flow.team.models import Transition - normalized_flow = {} for k, targets in state_flow.items(): normalized_targets = [] @@ -163,6 +171,7 @@ class RouteStrategy(BaseTeamStrategy): state_flow=self.state_flow, runtime_config=getattr(team, "runtime_config", None), custom_prompt=self.custom_prompt, + max_handoffs=self.max_handoffs, ) ) router = ChainRouter(routers) @@ -171,13 +180,11 @@ class RouteStrategy(BaseTeamStrategy): exec_config = kwargs.get("config") max_cycles = getattr(exec_config, "max_cycles", 15) if exec_config else 15 - logger.info(f"🛣️ [RouteStrategy] '{team.name}' 正在获取初始路由决策...") + logger.debug(f"🛣️ '{team.name}' 正在获取初始路由决策...") decision = await router.route(context, [], prompt) if not decision: - logger.warning( - f"🚨 [RouteStrategy] Team '{team.name}' 的所有路由策略未能命中目标。" - ) + logger.warning(f"🚨 Team '{team.name}' 的所有路由策略未能命中目标。") raise AbortException( reason=f"Team '{team.name}' 无法找到合适的路由节点处理该任务", display="🚨 团队协作失败,无法分配任务。", @@ -191,7 +198,7 @@ class RouteStrategy(BaseTeamStrategy): cycle_count += 1 if cycle_count > max_cycles: logger.error( - f"🚨 [RouteStrategy] Team '{team.name}' 路由陷入死循环!" + f"🚨 Team '{team.name}' 路由陷入死循环!" f"已达到最大限制 {max_cycles} 次。" ) raise AbortException( @@ -229,7 +236,7 @@ class RouteStrategy(BaseTeamStrategy): run_result = yield CallAction( agent=current_target, - task=prompt, + task=prompt or "", history=handoff_history_messages, kwargs=kwargs, ) @@ -268,7 +275,7 @@ class RouteStrategy(BaseTeamStrategy): pass if fast_routed: - logger.info( + logger.debug( f"🛣️ **路由决策**: 委派给专员 👨💼`{current_target}`" "(系统拦截:正则/函数状态流发生转移)" ) @@ -291,18 +298,21 @@ class CoordinateStrategy(BaseTeamStrategy): leader_model: str | None = None, leader_tools: list[ToolSource] | None = None, custom_prompt: str | None = None, + max_delegations: int = 3, ): """ 协作策略初始化,Leader 主动拆解任务,委派给 Sub-Agents 并汇总结果。 参数: leader_model: 协调节点 (Leader) 使用的大模型名称,若为空则默认继承全局。 - leader_tools: 挂载给协调节点 (Leader) 的专属工具列表。 + leader_tools: 挂载给协调节点 (Leader) 的专属工具列表. custom_prompt: 自定义系统提示词,用于覆盖默认的协调系统提示词。 + max_delegations: 允许向同一个专员连续委派失败的最大重试次数。 """ super().__init__(custom_prompt=custom_prompt) self.leader_model = leader_model self.leader_tools = leader_tools or [] + self.max_delegations = max_delegations async def generate_plan( self, @@ -323,6 +333,7 @@ class CoordinateStrategy(BaseTeamStrategy): runnable=m, name=f"delegate_to_{m.name}", description=f"将子任务委派给专员 [{m.name}] 处理。专长:{desc}", + max_delegations=self.max_delegations, ) ) @@ -338,17 +349,16 @@ class CoordinateStrategy(BaseTeamStrategy): logger.debug(f"✨ **团队 [{team.name}] Leader** 正在汇总各方报告...") - if context.run.event_bus: - await context.run.event_bus.emit( - ToolStreamChunkEvent( - tool_name="Team Leader", - content="✨ 团队 Leader 正在汇总各方报告...", - ) + await context.run.emit( + ToolStreamChunkEvent( + tool_name="Team Leader", + content="✨ 团队 Leader 正在汇总各方报告...", ) + ) logger.debug(f"👨💼 [CoordinateStrategy] '{team.name}' 正在启动协调推理循环...") - leader_res = yield CallAction(agent=leader_agent, task=prompt) + leader_res = yield CallAction(agent=leader_agent, task=prompt or "") yield FinishAction(result=leader_res) @@ -389,24 +399,22 @@ class BroadcastStrategy(BaseTeamStrategy): prompt.description if isinstance(prompt, AgentTask) else (prompt or "") ) - if context.run.event_bus: - await context.run.event_bus.emit( - ToolStreamChunkEvent( - tool_name="Team Broadcaster", - content=f"🚀 正在并发广播任务给 {len(team.members)} 位专家...", - ) + await context.run.emit( + ToolStreamChunkEvent( + tool_name="Team Broadcaster", + content=f"🚀 正在并发广播任务给 {len(team.members)} 位专家...", ) + ) actions = [CallAction(agent=m.name, task=task_desc_str) for m in team.members] results = yield ConcurrentCallAction(actions=actions) - if context.run.event_bus: - await context.run.event_bus.emit( - ToolStreamChunkEvent( - tool_name="Team Leader", - content="✨ 所有专家汇报完毕,Leader 正在融合各方观点...", - ) + await context.run.emit( + ToolStreamChunkEvent( + tool_name="Team Leader", + content="✨ 所有专家汇报完毕,Leader 正在融合各方观点...", ) + ) logger.debug(f"✨ **团队 [{team.name}] Leader** 正在汇总各方报告...") @@ -451,8 +459,7 @@ class TaskStrategy(BaseTeamStrategy): leader_model: str | None = None, leader_tools: list[ToolSource] | None = None, max_iterations: int = 15, - blackboard_schema: type[BaseModel] | None = None, - initial_blackboard_state: BaseModel | None = None, + blackboard: type[BaseModel] | BaseModel | None = None, custom_prompt: str | None = None, ): """ @@ -463,10 +470,9 @@ class TaskStrategy(BaseTeamStrategy): leader_model: 规划节点 (Leader) 使用的大模型名称,若为空则默认继承全局。 leader_tools: 挂载给规划节点 (Leader) 的专属附加工具列表。 max_iterations: 引擎驱动的状态机最大迭代/循环次数,防止死循环。 - blackboard_schema: 团队共享黑板的数据结构类型 (Pydantic Model 类)。 - initial_blackboard_state: 共享黑板的初始数据状态实例。 + blackboard: (可选) 团队共享黑板。可传入 Schema 类型类,或直接传入带有初始数据的 Schema 实例对象。 custom_prompt: 自定义系统提示词,用于覆盖默认的规划系统提示词。 - """ + """ # noqa: E501 super().__init__(custom_prompt=custom_prompt) self.leader_model = leader_model self.leader_tools = leader_tools or [] @@ -474,14 +480,21 @@ class TaskStrategy(BaseTeamStrategy): self.blackboard = None self.bb_toolkit = None - if blackboard_schema is not None: - from zhenxun.services.ai.run.blackboard import BlackboardManager - from zhenxun.services.ai.tools.providers.builtin.blackboard import ( - BlackboardToolkit, - ) + if blackboard is not None: + schema = None + initial_state = None + if isinstance(blackboard, type) and issubclass(blackboard, BaseModel): + schema = blackboard + elif isinstance(blackboard, BaseModel): + schema = type(blackboard) + initial_state = blackboard + else: + raise ValueError( + "blackboard 参数必须是 Pydantic BaseModel 的子类(类型)或其实例" + ) self.blackboard = BlackboardManager( - schema=blackboard_schema, initial_state=initial_blackboard_state + schema=schema, initial_state=initial_state ) self.bb_toolkit = BlackboardToolkit(self.blackboard) self.leader_tools.append(self.bb_toolkit) @@ -493,9 +506,6 @@ class TaskStrategy(BaseTeamStrategy): 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 - if self.blackboard is not None: context.session.blackboard = self.blackboard @@ -541,9 +551,7 @@ class TaskStrategy(BaseTeamStrategy): ) if "__task_board__" not in context.session.shared_state: - context.session.shared_state["__task_board__"] = ( - self.blackboard._state if self.blackboard else TaskBoardState() - ) + context.session.shared_state["__task_board__"] = TaskBoardState() board = cast(TaskBoardState, context.session.shared_state["__task_board__"]) max_iterations = self.max_iterations @@ -576,15 +584,17 @@ class TaskStrategy(BaseTeamStrategy): 请检查是否有 failed 的任务需要修复重新指派?或者如果所有任务均已 completed, 请立刻调用 `mark_all_complete` 汇报总结。""" - logger.info(f"🧠 [TaskStrategy] 唤醒 Planner (Iter: {iteration})") - leader_res = yield CallAction(agent=leader_agent, task=planner_prompt) + logger.debug(f"🧠 [TaskStrategy] 唤醒 Planner (Iter: {iteration})") + leader_res = yield CallAction( + agent=leader_agent, task=planner_prompt or "" + ) if board.is_goal_complete: yield FinishAction(result=board.final_summary or leader_res.output) return continue - logger.info( + logger.debug( f"🚀 [TaskStrategy] 引擎接管:并发执行 {len(available_tasks)} 个任务..." ) actions = [] diff --git a/zhenxun/services/ai/flow/team/task_tools.py b/zhenxun/services/ai/flow/team/task_tools.py index f2188c5f..8d9a3a83 100644 --- a/zhenxun/services/ai/flow/team/task_tools.py +++ b/zhenxun/services/ai/flow/team/task_tools.py @@ -6,8 +6,8 @@ from zhenxun.services.ai.flow.base import BaseRunnable 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 -from zhenxun.services.log import logger +from zhenxun.services.ai.tools.models import EndRunResult, ToolResult +from zhenxun.services.ai.utils.logger import log_team as logger from .models import TaskBoardState, TaskNodeStatus @@ -188,8 +188,6 @@ class TaskPlanningToolkit(BaseToolkit): ], context: RunContext, ) -> ToolResult: - from zhenxun.services.ai.tools.models import EndRunResult - board = self._get_board(context) board.is_goal_complete = True board.final_summary = summary diff --git a/zhenxun/services/ai/flow/team/team.py b/zhenxun/services/ai/flow/team/team.py index c96d51c4..efa952bc 100644 --- a/zhenxun/services/ai/flow/team/team.py +++ b/zhenxun/services/ai/flow/team/team.py @@ -1,27 +1,42 @@ import asyncio from collections.abc import Callable, Mapping, Sequence +import contextlib from pathlib import Path from typing import Any from typing_extensions import Self from pydantic import BaseModel -from zhenxun.services.ai.capabilities import AbstractCapability, DynamicCapability +from zhenxun.services.ai.capabilities import ( + AbstractCapability, + CapabilitySource, + DynamicCapability, +) from zhenxun.services.ai.core.exceptions import ConcurrencyInterruptException from zhenxun.services.ai.core.messages import PromptInput from zhenxun.services.ai.core.models import CancellationToken from zhenxun.services.ai.core.stream_events import EventBus -from zhenxun.services.ai.flow.agent.agent import CapabilitySource, ToolSource +from zhenxun.services.ai.flow.agent.agent import ToolSource from zhenxun.services.ai.flow.agent.models import Persona -from zhenxun.services.ai.flow.base import BaseRunnable -from zhenxun.services.ai.flow.team.models import TeamRuntimeConfig, Transition -from zhenxun.services.ai.flow.team.router import BaseRouter -from zhenxun.services.ai.flow.team.strategy import ( +from zhenxun.services.ai.flow.base import BaseRunnable, ConcurrencyPolicy +from zhenxun.services.ai.flow.concurrency import apply_concurrency_policy +from zhenxun.services.ai.run import ( + AgentRunResult, + AgentTask, + RunContext, + StreamedRunResult, +) +from zhenxun.services.ai.run.models import AgentRunError +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.utils.utils import infer_plugin_namespace + +from .models import TeamRuntimeConfig, Transition +from .router import BaseRouter +from .strategy import ( BaseTeamStrategy, ) -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 class Team(BaseRunnable[AgentRunResult[Any]]): @@ -115,6 +130,7 @@ class Team(BaseRunnable[AgentRunResult[Any]]): leader_model: str | None = None, leader_tools: list[ToolSource] | None = None, custom_prompt: str | None = None, + max_handoffs: int = 3, ) -> Self: """ 应用路由策略,基于挂载的 Router 进行最合适的专家动态分发。 @@ -128,8 +144,9 @@ class Team(BaseRunnable[AgentRunResult[Any]]): leader_model: 路由节点 (Leader) 使用的大模型名称,若为空则默认继承全局。 leader_tools: 挂载给路由节点 (Leader) 的专属工具列表。 custom_prompt: 自定义系统提示词,用于覆盖默认的路由系统提示词。 + max_handoffs: 同一会话中允许连续移交的最大次数,防止无限踢皮球。 """ - from zhenxun.services.ai.flow.team.strategy import RouteStrategy + from .strategy import RouteStrategy self.strategy = RouteStrategy( state_flow=state_flow, @@ -138,6 +155,7 @@ class Team(BaseRunnable[AgentRunResult[Any]]): leader_model=leader_model, leader_tools=leader_tools, custom_prompt=custom_prompt, + max_handoffs=max_handoffs, ) self.selector_func = selector_func return self @@ -147,6 +165,7 @@ class Team(BaseRunnable[AgentRunResult[Any]]): leader_model: str | None = None, leader_tools: list[ToolSource] | None = None, custom_prompt: str | None = None, + max_delegations: int = 3, ) -> Self: """ 应用协作策略,Leader 自主规划并主动将子任务委派给 Sub-Agents,最后汇总结果。 @@ -158,13 +177,15 @@ class Team(BaseRunnable[AgentRunResult[Any]]): leader_model: 协调节点 (Leader) 使用的大模型名称,若为空则默认继承全局。 leader_tools: 挂载给协调节点 (Leader) 的专属附加工具列表。 custom_prompt: 自定义系统提示词,用于覆盖默认的协调系统提示词。 + max_delegations: 允许向同一个专员连续委派失败的最大重试次数。 """ - from zhenxun.services.ai.flow.team.strategy import CoordinateStrategy + from .strategy import CoordinateStrategy self.strategy = CoordinateStrategy( leader_model=leader_model, leader_tools=leader_tools, custom_prompt=custom_prompt, + max_delegations=max_delegations, ) return self @@ -184,7 +205,7 @@ class Team(BaseRunnable[AgentRunResult[Any]]): leader_tools: 挂载给总结节点 (Leader) 的专属附加工具列表。 custom_prompt: 自定义系统提示词,用于覆盖默认的广播总结系统提示词。 """ - from zhenxun.services.ai.flow.team.strategy import BroadcastStrategy + from .strategy import BroadcastStrategy self.strategy = BroadcastStrategy( leader_model=leader_model, @@ -198,8 +219,7 @@ class Team(BaseRunnable[AgentRunResult[Any]]): leader_model: str | None = None, leader_tools: list[ToolSource] | None = None, max_iterations: int = 15, - blackboard_schema: type[BaseModel] | None = None, - initial_blackboard_state: BaseModel | None = None, + blackboard: type[BaseModel] | BaseModel | None = None, custom_prompt: str | None = None, ) -> Self: """ @@ -213,18 +233,16 @@ class Team(BaseRunnable[AgentRunResult[Any]]): leader_model: 规划节点 (Leader) 使用的大模型名称,若为空则默认继承全局。 leader_tools: 挂载给规划节点 (Leader) 的专属附加工具列表。 max_iterations: 引擎驱动的状态机最大迭代/循环次数,防止死循环。 - blackboard_schema: 团队共享黑板的数据结构类型 (Pydantic Model 类)。 - initial_blackboard_state: 共享黑板的初始数据状态实例。 + blackboard: (可选) 团队共享黑板。可传入 Schema 类型类,或直接传入带有初始数据的 Schema 实例对象。 custom_prompt: 自定义系统提示词,用于覆盖默认的规划系统提示词。 - """ - from zhenxun.services.ai.flow.team.strategy import TaskStrategy + """ # noqa: E501 + from .strategy import TaskStrategy self.strategy = TaskStrategy( leader_model=leader_model, leader_tools=leader_tools, max_iterations=max_iterations, - blackboard_schema=blackboard_schema, - initial_blackboard_state=initial_blackboard_state, + blackboard=blackboard, custom_prompt=custom_prompt, ) return self @@ -260,10 +278,6 @@ class Team(BaseRunnable[AgentRunResult[Any]]): """ self._ensure_strategy() if skills: - from zhenxun.services.ai.tools.providers.skills.capabilities import ( - SkillCapability, - ) - capabilities = list(capabilities) if capabilities else [] capabilities.append( SkillCapability(skills=skills, namespace=self.namespace) @@ -273,8 +287,6 @@ class Team(BaseRunnable[AgentRunResult[Any]]): prompt=prompt, context=context, capabilities=capabilities, **kwargs ) - import contextlib - @contextlib.asynccontextmanager async def run_stream( self, @@ -297,10 +309,6 @@ class Team(BaseRunnable[AgentRunResult[Any]]): context.capabilities.extend(self.capabilities) if skills: - from zhenxun.services.ai.tools.providers.skills.capabilities import ( - SkillCapability, - ) - capabilities = list(capabilities) if capabilities else [] capabilities.append( SkillCapability(skills=skills, namespace=self.namespace) @@ -313,8 +321,7 @@ class Team(BaseRunnable[AgentRunResult[Any]]): elif callable(cap): context.capabilities.append(DynamicCapability(cap)) - from zhenxun.services.ai.flow.team.runner import TeamRunner - from zhenxun.services.ai.run import StreamedRunResult + from .runner import TeamRunner event_bus = EventBus() context.run.event_bus = event_bus @@ -323,8 +330,6 @@ class Team(BaseRunnable[AgentRunResult[Any]]): policy = getattr(self.runtime_config, "concurrency_policy", None) if policy is None: - from zhenxun.services.ai.flow.base import ConcurrencyPolicy - policy = ( ConcurrencyPolicy.ALLOW if getattr(self.runtime_config, "stateless", True) @@ -333,8 +338,6 @@ class Team(BaseRunnable[AgentRunResult[Any]]): intervention_policy = getattr(self.runtime_config, "intervention_policy", None) - from zhenxun.services.ai.utils import ContextUtils - lock_id = ContextUtils.extract_concurrency_lock_id( context, getattr(self.runtime_config, "concurrency_scope", None), @@ -342,8 +345,6 @@ class Team(BaseRunnable[AgentRunResult[Any]]): ) async def _execution_task(): - from zhenxun.services.ai.flow.concurrency import apply_concurrency_policy - cancel_token = context.run.cancellation_token or CancellationToken() context.run.cancellation_token = cancel_token @@ -359,8 +360,6 @@ class Team(BaseRunnable[AgentRunResult[Any]]): async for event in runner.run_stream(prompt, context, **kwargs): await event_bus.emit(event) except BaseException as e: - from zhenxun.services.ai.run.models import AgentRunError - if isinstance(e, asyncio.CancelledError): e = ConcurrencyInterruptException("团队执行已被新请求打断并接管") await event_bus.emit(AgentRunError(error=e)) diff --git a/zhenxun/services/ai/flow/workflow/base.py b/zhenxun/services/ai/flow/workflow/base.py index 8651387d..f1f3c197 100644 --- a/zhenxun/services/ai/flow/workflow/base.py +++ b/zhenxun/services/ai/flow/workflow/base.py @@ -8,18 +8,19 @@ from zhenxun.services.ai.core.exceptions import ( ControlFlowExit, ToolFatalError, ) -from zhenxun.services.ai.flow.workflow.policies import ( +from zhenxun.services.ai.run import RunContext +from zhenxun.services.ai.utils.logger import log_flow as logger + +from .policies import ( AbortPolicy, BaseFailurePolicy, PolicyAction, ) -from zhenxun.services.ai.flow.workflow.types import ( +from .types import ( StepInput, StepOutput, StepType, ) -from zhenxun.services.ai.run import RunContext -from zhenxun.services.log import logger class BaseNode(ABC): diff --git a/zhenxun/services/ai/flow/workflow/engine.py b/zhenxun/services/ai/flow/workflow/engine.py index 664b904b..3c29ccbb 100644 --- a/zhenxun/services/ai/flow/workflow/engine.py +++ b/zhenxun/services/ai/flow/workflow/engine.py @@ -5,20 +5,16 @@ from typing import TYPE_CHECKING, Any import uuid from nonebot.params import Depends +from pydantic import BaseModel if TYPE_CHECKING: - from zhenxun.services.ai.flow.workflow.nodes import NodeSource + from .nodes import NodeSource from zhenxun.services.ai.core.exceptions import ControlFlowExit, ToolRetryError from zhenxun.services.ai.core.messages import PromptInput, UsageInfo from zhenxun.services.ai.core.stream_events import EventBus from zhenxun.services.ai.flow.base import BaseRunnable, BaseRuntimeConfig -from zhenxun.services.ai.flow.workflow.nodes import Steps -from zhenxun.services.ai.flow.workflow.types import ( - StepInput, - StepOutput, - WorkflowRunResult, -) +from zhenxun.services.ai.run.blackboard import BlackboardManager from zhenxun.services.ai.run.context import RunContext from zhenxun.services.ai.run.models import ( AgentRunEnd, @@ -27,9 +23,16 @@ from zhenxun.services.ai.run.models import ( StreamedRunResult, ) from zhenxun.services.ai.tools.core.tool import FunctionTool -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_flow as logger from zhenxun.utils.message import MessageUtils +from .nodes import Steps +from .types import ( + StepInput, + StepOutput, + WorkflowRunResult, +) + class Workflow(BaseRunnable[WorkflowRunResult]): """ @@ -37,7 +40,13 @@ class Workflow(BaseRunnable[WorkflowRunResult]): 继承自 BaseRunnable,支持被作为节点嵌套在 Team 或 其他工作流中。 """ - def __init__(self, name: str, steps: list["NodeSource"], description: str = ""): + def __init__( + self, + name: str, + steps: list["NodeSource"], + description: str = "", + blackboard: type[BaseModel] | BaseModel | None = None, + ): """ 静态图元工作流容器初始化。 @@ -45,7 +54,8 @@ class Workflow(BaseRunnable[WorkflowRunResult]): name: 工作流的名称标识。 steps: 工作流的节点列表(按列表顺序构成串行或嵌套结构)。 description: 工作流的说明描述,用于被 Agent 调用时理解其功能。 - """ + blackboard: (可选) 结构化黑板。可传入 Schema 类型类,或直接传入带有初始数据的 Schema 实例对象。 + """ # noqa: E501 self.name = name self.description = description self.id = uuid.uuid4().hex @@ -54,6 +64,19 @@ class Workflow(BaseRunnable[WorkflowRunResult]): self.runtime_config = BaseRuntimeConfig(stateless=True) self.persona = None + self.blackboard_schema = None + self.initial_blackboard_state = None + if blackboard is not None: + if isinstance(blackboard, type) and issubclass(blackboard, BaseModel): + self.blackboard_schema = blackboard + elif isinstance(blackboard, BaseModel): + self.blackboard_schema = type(blackboard) + self.initial_blackboard_state = blackboard + else: + raise ValueError( + "blackboard 参数必须是 Pydantic BaseModel 的子类(类型)或其实例" + ) + def _build_result( self, initial_input: StepInput, @@ -94,7 +117,12 @@ class Workflow(BaseRunnable[WorkflowRunResult]): return Depends(_dependency) async def reply( - self, prompt: PromptInput | None = None, reply_to: bool = False, **kwargs: Any + self, + prompt: PromptInput | None = None, + reply_to: bool = False, + *, + context: RunContext | None = None, + **kwargs: Any, ) -> WorkflowRunResult: """ 工作流交互执行语法糖,隐式提取上下文并自动将最终流水线产出发送回复给用户。 @@ -102,12 +130,13 @@ class Workflow(BaseRunnable[WorkflowRunResult]): 参数: prompt: 传入工作流入口根节点的初始参数或指令。 reply_to: 是否将结果作为回复消息发送 (at用户或引用原消息)。 + context: 显式传入的会话与运行上下文。 kwargs: 追加的工作流附带参数 (additional_data)。 返回: WorkflowRunResult: 包含执行状态、断点快照、各节点产出的全量工作流结果对象。 """ - ctx = RunContext() + ctx = context or RunContext() bot = ctx.get_bot() event = ctx.get_event() @@ -152,6 +181,12 @@ class Workflow(BaseRunnable[WorkflowRunResult]): ) safe_context = context or RunContext(session_id=session_id) + if self.blackboard_schema and not safe_context.session.blackboard: + safe_context.session.blackboard = BlackboardManager( + schema=self.blackboard_schema, + initial_state=self.initial_blackboard_state, + ) + logger.debug(f"🏭 **工作流 [{self.name}] 启动**") initial_input = StepInput(input=prompt) @@ -214,6 +249,12 @@ class Workflow(BaseRunnable[WorkflowRunResult]): ) safe_context = context or RunContext(session_id=session_id) + if self.blackboard_schema and not safe_context.session.blackboard: + safe_context.session.blackboard = BlackboardManager( + schema=self.blackboard_schema, + initial_state=self.initial_blackboard_state, + ) + logger.debug(f"🏭 **工作流 [{self.name}] 启动**") initial_input = StepInput(input=prompt) diff --git a/zhenxun/services/ai/flow/workflow/nodes.py b/zhenxun/services/ai/flow/workflow/nodes.py index b81d56d4..755cfe24 100644 --- a/zhenxun/services/ai/flow/workflow/nodes.py +++ b/zhenxun/services/ai/flow/workflow/nodes.py @@ -5,17 +5,18 @@ from typing import Any, cast 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 ( +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.ai.utils.logger import log_flow as logger + +from .base import BaseNode +from .policies import BaseFailurePolicy +from .types import ( StepInput, StepOutput, StepType, ) -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 """工作流节点来源,可以是图元、可执行引擎或原生函数""" diff --git a/zhenxun/services/ai/flow/workflow/policies.py b/zhenxun/services/ai/flow/workflow/policies.py index ea48902d..cf0c9864 100644 --- a/zhenxun/services/ai/flow/workflow/policies.py +++ b/zhenxun/services/ai/flow/workflow/policies.py @@ -4,9 +4,10 @@ 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 +from zhenxun.services.ai.utils.logger import log_flow as logger + +from .types import StepInput class PolicyAction(str, Enum): diff --git a/zhenxun/services/ai/llm/adapters/__init__.py b/zhenxun/services/ai/llm/adapters/__init__.py index ac57a21c..6fa41181 100644 --- a/zhenxun/services/ai/llm/adapters/__init__.py +++ b/zhenxun/services/ai/llm/adapters/__init__.py @@ -13,7 +13,7 @@ from .glm import GLMAdapter from .jina import JinaAdapter from .mimo import MiMoAdapter from .minimax import MiniMaxAdapter -from .openai import OpenAIAdapter, OpenAICompatAdapter +from .openai import OpenAIAdapter, OpenAICompatAdapter, OpenAIResponsesAdapter from .openrouter import OpenRouterAdapter LLMAdapterFactory.initialize() @@ -30,6 +30,7 @@ __all__ = [ "MiniMaxAdapter", "OpenAIAdapter", "OpenAICompatAdapter", + "OpenAIResponsesAdapter", "OpenRouterAdapter", "RequestData", "ResponseData", diff --git a/zhenxun/services/ai/llm/adapters/base.py b/zhenxun/services/ai/llm/adapters/base.py index 715a127e..b97670bf 100644 --- a/zhenxun/services/ai/llm/adapters/base.py +++ b/zhenxun/services/ai/llm/adapters/base.py @@ -47,11 +47,11 @@ from zhenxun.services.ai.core.messages import ( ThoughtPart, ) from zhenxun.services.ai.core.models import ModelIdentity -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_llm as logger from zhenxun.utils.log_sanitizer import sanitize_for_logging if TYPE_CHECKING: - from zhenxun.services.ai.llm.adapters.handlers.base import ( + from .handlers.base import ( BaseAudioHandler, BaseEmbeddingHandler, BaseImageHandler, @@ -496,8 +496,7 @@ class BaseAdapter(ABC): async def parse_speech_response( self, identity: ModelIdentity, raw_response: httpx.Response ) -> AudioResponse: - """解析语音响应并委派给 `audio_handler`。 - 注意传入的是 httpx.Response 的 raw 对象""" + """解析语音响应并委派给 `audio_handler`""" if self.audio_handler: return await self.audio_handler.parse_speech_response( adapter=self, identity=identity, raw_response=raw_response diff --git a/zhenxun/services/ai/llm/adapters/deepseek.py b/zhenxun/services/ai/llm/adapters/deepseek.py index 8be7207b..3bcfdd01 100644 --- a/zhenxun/services/ai/llm/adapters/deepseek.py +++ b/zhenxun/services/ai/llm/adapters/deepseek.py @@ -6,12 +6,13 @@ from zhenxun.services.ai.core.models import ( ModelIdentity, ) from zhenxun.services.ai.core.options import GenerationConfig -from zhenxun.services.ai.llm.adapters.handlers.openai_handlers import ( + +from .handlers.openai_handlers import ( OpenAIConfigMapper, OpenAITextHandler, OpenAIToolSerializer, ) -from zhenxun.services.ai.llm.adapters.openai import OpenAICompatAdapter +from .openai import OpenAICompatAdapter class DeepSeekToolSerializer(OpenAIToolSerializer): @@ -84,21 +85,6 @@ class DeepSeekConfigMapper(OpenAIConfigMapper): if isinstance(rf, dict) and rf.get("type") == "json_schema": params["response_format"] = {"type": "json_object"} - if config.common.reasoning_effort: - effort = str(config.common.reasoning_effort).lower() - if effort == "none": - params["thinking"] = {"type": "disabled"} - else: - params["thinking"] = {"type": "enabled"} - elif ( - hasattr(config, "deepseek_options") - and config.deepseek_options.thinking is not None - ): - if config.deepseek_options.thinking is True: - params["thinking"] = {"type": "enabled"} - elif config.deepseek_options.thinking is False: - params["thinking"] = {"type": "disabled"} - return params diff --git a/zhenxun/services/ai/llm/adapters/factory.py b/zhenxun/services/ai/llm/adapters/factory.py index 4c9b2b83..43d329f3 100644 --- a/zhenxun/services/ai/llm/adapters/factory.py +++ b/zhenxun/services/ai/llm/adapters/factory.py @@ -34,10 +34,11 @@ class LLMAdapterFactory: from .jina import JinaAdapter from .mimo import MiMoAdapter from .minimax import MiniMaxAdapter - from .openai import OpenAIAdapter + from .openai import OpenAIAdapter, OpenAIResponsesAdapter from .openrouter import OpenRouterAdapter cls.register_adapter(OpenAIAdapter()) + cls.register_adapter(OpenAIResponsesAdapter()) cls.register_adapter(OpenRouterAdapter()) cls.register_adapter(DeepSeekAdapter()) cls.register_adapter(JinaAdapter()) diff --git a/zhenxun/services/ai/llm/adapters/glm.py b/zhenxun/services/ai/llm/adapters/glm.py index 16ed579f..edb8dbc1 100644 --- a/zhenxun/services/ai/llm/adapters/glm.py +++ b/zhenxun/services/ai/llm/adapters/glm.py @@ -1,13 +1,14 @@ from zhenxun.services.ai.core.messages import RerankRequest from zhenxun.services.ai.core.models import ModelIdentity -from zhenxun.services.ai.llm.adapters.base import BaseAdapter, RequestData -from zhenxun.services.ai.llm.adapters.handlers.openai_handlers import ( + +from .base import BaseAdapter, RequestData +from .handlers.openai_handlers import ( OpenAIConfigMapper, OpenAIEmbeddingHandler, OpenAIRerankHandler, OpenAITextHandler, ) -from zhenxun.services.ai.llm.adapters.openai import OpenAICompatAdapter +from .openai import OpenAICompatAdapter class GLMRerankHandler(OpenAIRerankHandler): diff --git a/zhenxun/services/ai/llm/adapters/handlers/gemini_handlers.py b/zhenxun/services/ai/llm/adapters/handlers/gemini_handlers.py index e820ba5c..319c4032 100644 --- a/zhenxun/services/ai/llm/adapters/handlers/gemini_handlers.py +++ b/zhenxun/services/ai/llm/adapters/handlers/gemini_handlers.py @@ -56,7 +56,9 @@ from zhenxun.services.ai.llm.adapters.base import ( ResponseData, process_image_data, ) -from zhenxun.services.ai.llm.adapters.handlers.base import ( +from zhenxun.services.ai.utils.logger import log_llm as logger + +from .base import ( BaseAudioHandler, BaseEmbeddingHandler, BaseImageHandler, @@ -66,7 +68,6 @@ from zhenxun.services.ai.llm.adapters.handlers.base import ( ResponseParser, ToolSerializer, ) -from zhenxun.services.log import logger class GeminiConfigMapper(ConfigMapper): @@ -121,22 +122,19 @@ class GeminiConfigMapper(ConfigMapper): fc_config["allowedFunctionNames"] = user_funcs params["toolConfig"] = {"functionCallingConfig": fc_config} - has_effort = bool( - config.common.reasoning_effort - and str(config.common.reasoning_effort).lower() != "none" - ) - if ( - has_effort - and capabilities - and capabilities.reasoning_mode == ReasoningMode.LEVEL - ): - thinking_config = params.setdefault("thinkingConfig", {}) + if config.common.reasoning_effort and capabilities: effort = str(config.common.reasoning_effort).lower() - if capabilities.reasoning_effort_map: effort = capabilities.reasoning_effort_map.get(effort, effort) - thinking_config["thinkingLevel"] = effort + if capabilities.reasoning_mode == ReasoningMode.LEVEL and effort != "none": + thinking_config = params.setdefault("thinkingConfig", {}) + thinking_config["thinkingLevel"] = effort + elif ( + capabilities.reasoning_mode == ReasoningMode.BUDGET and effort == "none" + ): + thinking_config = params.setdefault("thinkingConfig", {}) + thinking_config["thinkingBudget"] = 0 if config.gemini_options.include_thoughts is not None: thinking_config = params.setdefault("thinkingConfig", {}) @@ -829,10 +827,6 @@ class GeminiEmbeddingHandler(BaseEmbeddingHandler): url = f"{base_url}/v1beta/{api_model_name}:batchEmbedContents" headers = adapter.get_base_headers(api_key) - from zhenxun.services.ai.llm.adapters.handlers.gemini_handlers import ( - GeminiMessageConverter, - ) - converter = GeminiMessageConverter() requests_payload = [] diff --git a/zhenxun/services/ai/llm/adapters/handlers/mimo_handlers.py b/zhenxun/services/ai/llm/adapters/handlers/mimo_handlers.py index 3e209737..83055e7f 100644 --- a/zhenxun/services/ai/llm/adapters/handlers/mimo_handlers.py +++ b/zhenxun/services/ai/llm/adapters/handlers/mimo_handlers.py @@ -16,13 +16,13 @@ from zhenxun.services.ai.core.messages import ( ) from zhenxun.services.ai.core.models import ( ModelCapabilities, - ModelDetail, ModelIdentity, ) -from zhenxun.services.ai.core.options import GenerationConfig, TTSConfig +from zhenxun.services.ai.core.options import TTSConfig from zhenxun.services.ai.llm.adapters.base import BaseAdapter, RequestData -from zhenxun.services.ai.llm.adapters.handlers.base import BaseAudioHandler -from zhenxun.services.ai.llm.adapters.handlers.openai_handlers import ( + +from .base import BaseAudioHandler +from .openai_handlers import ( OpenAIConfigMapper, OpenAIMessageConverter, OpenAITextHandler, @@ -53,35 +53,6 @@ class MiMoToolSerializer(OpenAIToolSerializer): return res -class MiMoConfigMapper(OpenAIConfigMapper): - """MiMo 配置映射器,处理深度思考参数差异""" - - def map_config( - self, - config: GenerationConfig, - model_detail: ModelDetail | None = None, - capabilities: ModelCapabilities | None = None, - ) -> dict[str, Any]: - params = super().map_config(config, model_detail, capabilities) - - if config.common.reasoning_effort: - effort = str(config.common.reasoning_effort).lower() - if effort == "none": - params["thinking"] = {"type": "disabled"} - else: - params["thinking"] = {"type": "enabled"} - elif ( - hasattr(config, "deepseek_options") - and config.deepseek_options.thinking is not None - ): - if config.deepseek_options.thinking is True: - params["thinking"] = {"type": "enabled"} - elif config.deepseek_options.thinking is False: - params["thinking"] = {"type": "disabled"} - - return params - - class MiMoMessageConverter(OpenAIMessageConverter): """MiMo 消息转换器,拦截处理特有的音视频多模态结构""" @@ -141,7 +112,7 @@ class MiMoTextHandler(OpenAITextHandler): super().__init__(api_type=api_type) self.converter = MiMoMessageConverter(api_type=api_type) self.serializer = MiMoToolSerializer(api_type=api_type) - self.mapper = MiMoConfigMapper(api_type=api_type) + self.mapper = OpenAIConfigMapper(api_type=api_type) class MiMoAudioHandler(BaseAudioHandler): diff --git a/zhenxun/services/ai/llm/adapters/handlers/openai_handlers.py b/zhenxun/services/ai/llm/adapters/handlers/openai_handlers.py index eb9e221a..3f502f53 100644 --- a/zhenxun/services/ai/llm/adapters/handlers/openai_handlers.py +++ b/zhenxun/services/ai/llm/adapters/handlers/openai_handlers.py @@ -53,7 +53,9 @@ from zhenxun.services.ai.llm.adapters.base import ( ResponseData, process_image_data, ) -from zhenxun.services.ai.llm.adapters.handlers.base import ( +from zhenxun.services.ai.utils.logger import log_llm as logger + +from .base import ( BaseAudioHandler, BaseEmbeddingHandler, BaseImageHandler, @@ -64,7 +66,6 @@ from zhenxun.services.ai.llm.adapters.handlers.base import ( ResponseParser, ToolSerializer, ) -from zhenxun.services.log import logger class OpenAIConfigMapper(ConfigMapper): @@ -114,7 +115,13 @@ class OpenAIConfigMapper(ConfigMapper): if capabilities and capabilities.reasoning_effort_map: effort = capabilities.reasoning_effort_map.get(effort, effort) - if effort != "none": + if ( + effort == "none" + and capabilities + and capabilities.supports_thinking_toggle + ): + params["thinking"] = {"type": "disabled"} + else: params["reasoning_effort"] = effort if isinstance(config.output.response_format, dict): @@ -1165,45 +1172,3 @@ class OpenAIAudioHandler(BaseAudioHandler): usage=UsageInfo(), model_name=identity.model_name, ) - - -class CompositeOpenAITextHandler(BaseTextHandler): - """ - OpenAI 复合文本对话处理器 (Composite Pattern)。 - 内部包装标准协议与 responses 协议 Handler,根据模型配置在请求时动态路由 - """ - - def __init__(self, api_type: str = "openai"): - self.api_type = api_type - self._standard_handler = OpenAITextHandler(api_type=api_type) - self._responses_handler = OpenAIResponsesTextHandler( - api_type="openai_responses" - ) - - def _get_active_handler(self, identity: ModelIdentity) -> BaseTextHandler: - current_api_type = identity.api_type - if current_api_type == "openai_responses": - return self._responses_handler - return self._standard_handler - - async def prepare_text_request( - self, - adapter: BaseAdapter, - identity: ModelIdentity, - api_key: str, - request: ChatRequest, - ) -> RequestData: - handler = self._get_active_handler(identity) - return await handler.prepare_text_request(adapter, identity, api_key, request) - - def parse_text_response( - self, - adapter: BaseAdapter, - identity: ModelIdentity, - response_json: dict[str, Any], - is_advanced: bool = False, - ) -> ResponseData: - handler = self._get_active_handler(identity) - return handler.parse_text_response( - adapter, identity, response_json, is_advanced - ) diff --git a/zhenxun/services/ai/llm/adapters/jina.py b/zhenxun/services/ai/llm/adapters/jina.py index 42dc2989..f7ba41ee 100644 --- a/zhenxun/services/ai/llm/adapters/jina.py +++ b/zhenxun/services/ai/llm/adapters/jina.py @@ -1,12 +1,14 @@ from zhenxun.services.ai.core.messages import EmbeddingRequest from zhenxun.services.ai.core.models import ModelIdentity from zhenxun.services.ai.core.options import LLMEmbeddingConfig -from zhenxun.services.ai.llm.adapters.base import BaseAdapter, RequestData -from zhenxun.services.ai.llm.adapters.handlers.openai_handlers import ( +from zhenxun.services.ai.utils.logger import log_llm as logger + +from .base import BaseAdapter, RequestData +from .handlers.openai_handlers import ( OpenAIEmbeddingHandler, OpenAIRerankHandler, ) -from zhenxun.services.ai.llm.adapters.openai import OpenAICompatAdapter +from .openai import OpenAICompatAdapter class JinaEmbeddingHandler(OpenAIEmbeddingHandler): @@ -36,7 +38,6 @@ class JinaEmbeddingHandler(OpenAIEmbeddingHandler): TextPart, VideoPart, ) - from zhenxun.services.log import logger for payload in batch.payloads: jina_content = [] diff --git a/zhenxun/services/ai/llm/adapters/mimo.py b/zhenxun/services/ai/llm/adapters/mimo.py index 031bcdd3..313de6ba 100644 --- a/zhenxun/services/ai/llm/adapters/mimo.py +++ b/zhenxun/services/ai/llm/adapters/mimo.py @@ -1,12 +1,13 @@ from zhenxun.services.ai.core.models import ModelIdentity -from zhenxun.services.ai.llm.adapters.handlers.mimo_handlers import ( + +from .handlers.mimo_handlers import ( MiMoAudioHandler, MiMoTextHandler, ) -from zhenxun.services.ai.llm.adapters.handlers.openai_handlers import ( +from .handlers.openai_handlers import ( OpenAIImageHandler, ) -from zhenxun.services.ai.llm.adapters.openai import OpenAICompatAdapter +from .openai import OpenAICompatAdapter class MiMoAdapter(OpenAICompatAdapter): diff --git a/zhenxun/services/ai/llm/adapters/minimax.py b/zhenxun/services/ai/llm/adapters/minimax.py index b4b7fd20..b1f2967d 100644 --- a/zhenxun/services/ai/llm/adapters/minimax.py +++ b/zhenxun/services/ai/llm/adapters/minimax.py @@ -19,9 +19,9 @@ from zhenxun.services.ai.core.models import ( ModelIdentity, ) from zhenxun.services.ai.core.options import GenerationConfig, TTSConfig -from zhenxun.services.ai.llm.adapters.base import BaseAdapter, RequestData -from zhenxun.services.ai.llm.adapters.handlers.base import BaseAudioHandler +from .base import BaseAdapter, RequestData +from .handlers.base import BaseAudioHandler from .handlers.openai_handlers import ( OpenAIConfigMapper, OpenAIMessageConverter, diff --git a/zhenxun/services/ai/llm/adapters/openai.py b/zhenxun/services/ai/llm/adapters/openai.py index 60377377..f7697deb 100644 --- a/zhenxun/services/ai/llm/adapters/openai.py +++ b/zhenxun/services/ai/llm/adapters/openai.py @@ -15,11 +15,12 @@ from .base import ( RequestData, ) from .handlers.openai_handlers import ( - CompositeOpenAITextHandler, OpenAIAudioHandler, OpenAIEmbeddingHandler, OpenAIImageHandler, OpenAIRerankHandler, + OpenAIResponsesTextHandler, + OpenAITextHandler, ) @@ -83,12 +84,12 @@ class OpenAICompatAdapter(BaseAdapter): class OpenAIAdapter(OpenAICompatAdapter): - """OpenAI 系列适配器,统一装配文本/图像/嵌入/重排处理链。""" + """标准 OpenAI 系列适配器,统一装配文本/图像/嵌入/重排处理链。""" def __init__(self): - """初始化并挂载复合文本处理器与通用多模态处理器。""" + """初始化并挂载标准文本处理器与通用多模态处理器。""" super().__init__() - self.text_handler = CompositeOpenAITextHandler(api_type=self.api_type) + self.text_handler = OpenAITextHandler(api_type=self.api_type) self.image_handler = OpenAIImageHandler() self.embedding_handler = OpenAIEmbeddingHandler() self.rerank_handler = OpenAIRerankHandler() @@ -102,21 +103,35 @@ class OpenAIAdapter(OpenAICompatAdapter): @property def supported_api_types(self) -> list[str]: """支持的 API 类型及别名。""" - return [ - "openai", - "openai_responses", - ] + return ["openai"] def get_chat_endpoint(self, identity: ModelIdentity) -> str: """返回聊天完成端点""" - current_api_type = identity.api_type - - if current_api_type == "openai_responses": - return "/v1/responses" - if current_api_type == "doubao": - return "/api/v3/chat/completions" return "/v1/chat/completions" def get_embedding_endpoint(self, identity: ModelIdentity) -> str: """返回嵌入端点。""" return "/v1/embeddings" + + +class OpenAIResponsesAdapter(OpenAICompatAdapter): + """OpenAI v1/responses 协议专用适配器""" + + def __init__(self): + super().__init__() + self.text_handler = OpenAIResponsesTextHandler(api_type=self.api_type) + self.image_handler = OpenAIImageHandler() + self.embedding_handler = OpenAIEmbeddingHandler() + self.rerank_handler = OpenAIRerankHandler() + self.audio_handler = OpenAIAudioHandler() + + @property + def api_type(self) -> str: + return "openai_responses" + + @property + def supported_api_types(self) -> list[str]: + return ["openai_responses"] + + def get_chat_endpoint(self, identity: ModelIdentity) -> str: + return "/v1/responses" diff --git a/zhenxun/services/ai/llm/adapters/openrouter.py b/zhenxun/services/ai/llm/adapters/openrouter.py index 64f4554c..ce97872d 100644 --- a/zhenxun/services/ai/llm/adapters/openrouter.py +++ b/zhenxun/services/ai/llm/adapters/openrouter.py @@ -10,18 +10,19 @@ from zhenxun.services.ai.core.messages import ( ThoughtPart, ) from zhenxun.services.ai.core.models import ModelIdentity -from zhenxun.services.ai.llm.adapters.base import ( + +from .base import ( BaseAdapter, RequestData, ResponseData, process_image_data, ) -from zhenxun.services.ai.llm.adapters.handlers.base import BaseImageHandler -from zhenxun.services.ai.llm.adapters.handlers.openai_handlers import ( - CompositeOpenAITextHandler, +from .handlers.base import BaseImageHandler +from .handlers.openai_handlers import ( OpenAIMessageConverter, + OpenAITextHandler, ) -from zhenxun.services.ai.llm.adapters.openai import OpenAIAdapter +from .openai import OpenAIAdapter class OpenRouterMessageConverter(OpenAIMessageConverter): @@ -60,12 +61,12 @@ class OpenRouterMessageConverter(OpenAIMessageConverter): return openai_messages -class OpenRouterTextHandler(CompositeOpenAITextHandler): +class OpenRouterTextHandler(OpenAITextHandler): """OpenRouter 专有文本处理器,挂载专有 Converter""" def __init__(self, api_type: str = "openrouter"): super().__init__(api_type=api_type) - self._standard_handler.converter = OpenRouterMessageConverter(api_type=api_type) + self.converter = OpenRouterMessageConverter(api_type=api_type) class OpenRouterImageHandler(BaseImageHandler): diff --git a/zhenxun/services/ai/llm/api.py b/zhenxun/services/ai/llm/api.py index 7dd2416c..0f5fe3f5 100644 --- a/zhenxun/services/ai/llm/api.py +++ b/zhenxun/services/ai/llm/api.py @@ -35,10 +35,10 @@ from zhenxun.services.ai.core.options import ( TTSConfig, ) from zhenxun.services.ai.guardrails import GuardrailSource -from zhenxun.services.ai.llm.engine.router import LLMOrchestrator -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_llm as logger from .builder import IntentBuilder +from .engine.router import LLMOrchestrator T = TypeVar("T", bound=BaseModel) diff --git a/zhenxun/services/ai/llm/builder.py b/zhenxun/services/ai/llm/builder.py index d1bc81fb..4593d7c6 100644 --- a/zhenxun/services/ai/llm/builder.py +++ b/zhenxun/services/ai/llm/builder.py @@ -10,7 +10,7 @@ from zhenxun.services.ai.core.options import ( GenerationConfig, ResponseFormat, ) -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_llm as logger from zhenxun.utils.pydantic_compat import model_json_schema, model_validate @@ -43,18 +43,6 @@ class OpenAIIntentNamespace: return self._builder -class DeepSeekIntentNamespace: - """DeepSeek 专属高级参数构建域""" - - def __init__(self, builder: "IntentBuilder"): - self._builder = builder - - def disable_thinking(self) -> "IntentBuilder": - """显式关闭 DeepSeek 的思维链""" - self._builder._config.deepseek_options.thinking = False - return self._builder - - class IntentBuilder: """ 基于能力意图声明的构建器 (Intent-Driven Builder)。 @@ -72,10 +60,6 @@ class IntentBuilder: def openai(self) -> OpenAIIntentNamespace: return OpenAIIntentNamespace(self) - @property - def deepseek(self) -> DeepSeekIntentNamespace: - return DeepSeekIntentNamespace(self) - def with_reasoning(self, level: str | None = None) -> Self: """ 跨厂商统一的思考/推理等级意图声明。 @@ -85,9 +69,6 @@ class IntentBuilder: self._config.common.reasoning_effort = level if level.lower() != "none": self._config.gemini_options.include_thoughts = True - self._config.deepseek_options.thinking = True - else: - self._config.deepseek_options.thinking = False return self def with_local_cache(self, ttl: int = 3600) -> Self: diff --git a/zhenxun/services/ai/llm/engine/middlewares.py b/zhenxun/services/ai/llm/engine/middlewares.py index 0b0980ed..ef8e01c4 100644 --- a/zhenxun/services/ai/llm/engine/middlewares.py +++ b/zhenxun/services/ai/llm/engine/middlewares.py @@ -51,7 +51,7 @@ from zhenxun.services.ai.llm.system.network import ( HealthManager, LLMHttpClient, ) -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_llm as logger from zhenxun.utils.http_utils import AsyncHttpx from zhenxun.utils.log_sanitizer import sanitize_for_logging from zhenxun.utils.pydantic_compat import ( @@ -136,7 +136,7 @@ class LLMCacheMiddleware: if cached_data is not None: logger.debug( - f"⚡ [LLMCache] 命中本地极速缓存 - " + f"命中本地极速缓存 - " f"model: {self.model_name}, type: {type(context.request).__name__}" ) diff --git a/zhenxun/services/ai/llm/engine/router.py b/zhenxun/services/ai/llm/engine/router.py index e862b32d..81a9e64e 100644 --- a/zhenxun/services/ai/llm/engine/router.py +++ b/zhenxun/services/ai/llm/engine/router.py @@ -16,7 +16,7 @@ from zhenxun.services.ai.llm.manager import ( ) from zhenxun.services.ai.llm.system.capabilities import get_model_capabilities from zhenxun.services.ai.llm.system.network import health_manager -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_llm as logger class BaseModelRouter(ABC): @@ -57,7 +57,7 @@ class FallbackRouter(BaseModelRouter): m_name = model_names[idx] if not health_manager.is_route_healthy(m_name, strict_mode=is_routed_call): - logger.debug(f"👉 [Orchestrator] 节点 '{m_name}' 熔断中,已跳过") + logger.debug(f"👉 节点 '{m_name}' 熔断中,已跳过") errors.append(f"{m_name}(熔断中)") continue @@ -70,7 +70,7 @@ class FallbackRouter(BaseModelRouter): try: if len(model_names) > 1: if idx != start_idx: - logger.debug(f"🔄 [Orchestrator] 切换至备用节点: '{m_name}'...") + logger.debug(f"🔄 切换至备用节点: '{m_name}'...") async with await get_model_instance( m_name, override_config, task=task @@ -83,26 +83,23 @@ class FallbackRouter(BaseModelRouter): except LLMException as e: if not e.should_failover: logger.warning( - f"🚫 [Orchestrator] 节点 '{m_name}' " + f"🚫 节点 '{m_name}' " f"返回不可恢复错误 ({e.__class__.__name__}),停止故障转移。" ) raise e logger.warning( - f"⚠️ [Orchestrator] 节点 '{m_name}' " + f"⚠️ 节点 '{m_name}' " f"错误 ({e.__class__.__name__}),触发故障转移..." ) errors.append(f"{m_name}({e.__class__.__name__})") except Exception as e: - logger.warning( - f"⚠️ [Orchestrator] 节点 '{m_name}' 发生未知异常,触发故障转移: {e}" - ) + logger.warning(f"⚠️ 节点 '{m_name}' 发生未知异常,触发故障转移: {e}") errors.append(f"{m_name}(Error)") if all_nodes_bypassed and len(model_names) > 1: fallback_model = health_manager.get_best_fallback_route(model_names) logger.warning( - f"⚠️ [Orchestrator] 路由组所有节点均已宕机!" - f"强制放行 '{fallback_model}' 探活..." + f"⚠️ 路由组所有节点均已宕机!" f"强制放行 '{fallback_model}' 探活..." ) try: async with await get_model_instance( diff --git a/zhenxun/services/ai/llm/engine/service.py b/zhenxun/services/ai/llm/engine/service.py index 0a831251..c431dd17 100644 --- a/zhenxun/services/ai/llm/engine/service.py +++ b/zhenxun/services/ai/llm/engine/service.py @@ -43,9 +43,22 @@ from zhenxun.services.ai.core.protocols.llm import ( SupportsTextEmbedding, ) from zhenxun.services.ai.core.protocols.middleware import LLMMiddleware +from zhenxun.services.ai.llm.adapters.factory import get_adapter_for_api_type from zhenxun.services.ai.llm.system.models import RetryConfig from zhenxun.services.ai.llm.system.network import HealthManager, LLMHttpClient -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_llm as logger + +from .middlewares import ( + ConfigMergeMiddleware, + FailoverAndRetryMiddleware, + HttpExecutionMiddleware, + LLMCacheMiddleware, + LoggingMiddleware, + MiddlewarePipeline, + ModalityFilterMiddleware, + OutputValidationMiddleware, + ResponseRescueMiddleware, +) T = TypeVar("T", bound=BaseModel) @@ -86,7 +99,7 @@ class LLMModel( ) self.model_name = model_detail.model_name self.temperature = model_detail.temperature - self.generation_max_tokens = model_detail.generation_max_tokens + self.max_output_tokens = model_detail.max_output_tokens self._is_closed = False self._ref_count = 0 @@ -100,8 +113,6 @@ class LLMModel( generation_config=self._generation_config, ) - from zhenxun.services.ai.llm.engine.middlewares import MiddlewarePipeline - self.pipeline = MiddlewarePipeline() self._setup_default_pipeline() @@ -110,17 +121,6 @@ class LLMModel( self.pipeline.add_middleware(middleware) def _setup_default_pipeline(self) -> None: - from zhenxun.services.ai.llm.adapters.factory import get_adapter_for_api_type - from zhenxun.services.ai.llm.engine.middlewares import ( - ConfigMergeMiddleware, - FailoverAndRetryMiddleware, - LLMCacheMiddleware, - LoggingMiddleware, - ModalityFilterMiddleware, - OutputValidationMiddleware, - ResponseRescueMiddleware, - ) - client_settings = get_llm_config().client_settings retry_config = RetryConfig( max_retries=client_settings.max_retries, @@ -216,10 +216,6 @@ class LLMModel( request=request, cancellation_token=cancellation_token, ) - - from zhenxun.services.ai.llm.adapters.factory import get_adapter_for_api_type - from zhenxun.services.ai.llm.engine.middlewares import HttpExecutionMiddleware - adapter = get_adapter_for_api_type(self.api_type) execution_middleware = HttpExecutionMiddleware( http_client=self.http_client, diff --git a/zhenxun/services/ai/llm/manager.py b/zhenxun/services/ai/llm/manager.py index e31271c8..4b5d8287 100644 --- a/zhenxun/services/ai/llm/manager.py +++ b/zhenxun/services/ai/llm/manager.py @@ -13,12 +13,13 @@ from zhenxun.services.ai.config import ( from zhenxun.services.ai.core.exceptions import ConfigurationException from zhenxun.services.ai.core.models import ModelDetail from zhenxun.services.ai.core.options import GenerationConfig -from zhenxun.services.ai.llm.system.capabilities import get_model_capabilities -from zhenxun.services.ai.llm.system.network import health_manager -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_llm as logger from zhenxun.utils.manager.priority_manager import PriorityLifecycle from zhenxun.utils.pydantic_compat import model_dump +from .system.capabilities import get_model_capabilities +from .system.network import health_manager + _RESOLVED_GROUP_CACHE: dict[str, list[str]] = {} """路由组解析缓存,避免每次调用重复打印剔除警告并提升性能""" @@ -276,7 +277,7 @@ async def get_model_instance( provider_config_found, model_detail_found = config_tuple_found - from zhenxun.services.ai.llm.system.cache import get_or_create_model + from .system.cache import get_or_create_model return await get_or_create_model( provider_config_found, model_detail_found, override_config @@ -287,7 +288,7 @@ def clear_all_cache() -> None: """ 清空模型实例缓存与路由组解析缓存。 """ - from zhenxun.services.ai.llm.system.cache import clear_model_cache + from .system.cache import clear_model_cache clear_model_cache() clear_resolved_group_cache() @@ -300,9 +301,10 @@ async def _init_llm_config_on_startup(): logger.info("正在初始化 LLM 配置并加载遥测状态...") try: from zhenxun.services.ai.config import get_llm_config - from zhenxun.services.ai.llm.system.network import health_manager from zhenxun.services.ai.tools.engine.registry import tool_provider_manager + from .system.network import health_manager + get_llm_config() await health_manager.initialize() await tool_provider_manager.initialize() diff --git a/zhenxun/services/ai/llm/system/cache.py b/zhenxun/services/ai/llm/system/cache.py index ffd1e342..906f0692 100644 --- a/zhenxun/services/ai/llm/system/cache.py +++ b/zhenxun/services/ai/llm/system/cache.py @@ -8,11 +8,12 @@ from zhenxun.services.ai.core.models import ModelDetail, ModelModality from zhenxun.services.ai.core.options import GenerationConfig from zhenxun.services.ai.llm.builder import validate_override_params from zhenxun.services.ai.llm.engine.service import LLMModel -from zhenxun.services.ai.llm.system.capabilities import get_model_capabilities -from zhenxun.services.ai.llm.system.network import health_manager, http_client_manager -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_llm as logger from zhenxun.utils.pydantic_compat import dump_json_safely, model_copy, model_dump +from .capabilities import get_model_capabilities +from .network import health_manager, http_client_manager + _model_cache: dict[str, tuple[LLMModel, float]] = {} _cache_ttl = 3600 _max_cache_size = 10 @@ -111,6 +112,22 @@ async def get_or_create_model( capabilities = get_model_capabilities(model_detail_found.model_name) capabilities = model_copy(capabilities, deep=True) + if model_detail_found.max_input_tokens is not None: + capabilities.max_input_tokens = model_detail_found.max_input_tokens + + base_gen_config = GenerationConfig() + if provider_config_found.temperature is not None: + base_gen_config.common.temperature = provider_config_found.temperature + if provider_config_found.max_output_tokens is not None: + base_gen_config.common.max_tokens = provider_config_found.max_output_tokens + + if model_detail_found.temperature is not None: + base_gen_config.common.temperature = model_detail_found.temperature + if model_detail_found.max_output_tokens is not None: + base_gen_config.common.max_tokens = model_detail_found.max_output_tokens + if model_detail_found.reasoning_effort is not None: + base_gen_config.common.reasoning_effort = model_detail_found.reasoning_effort + if model_detail_found.task_type == "image_generation": capabilities.output_modalities.add(ModelModality.IMAGE) capabilities.supports_tool_calling = False @@ -130,9 +147,8 @@ async def get_or_create_model( timeout=default_timeout, api_base=provider_config_found.api_base, api_type=provider_config_found.api_type, - openai_compat=provider_config_found.openai_compat, temperature=provider_config_found.temperature, - generation_max_tokens=provider_config_found.generation_max_tokens, + max_output_tokens=provider_config_found.max_output_tokens, ) shared_http_client = await http_client_manager.get_client(config_for_http_client) @@ -144,16 +160,21 @@ async def get_or_create_model( health_manager=health_manager, http_client=shared_http_client, capabilities=capabilities, + config_override=base_gen_config, ) if override_config: validated_override_params = validate_override_params(override_config) - model_instance._generation_config = validated_override_params - model_instance.identity.generation_config = validated_override_params + final_config = base_gen_config.merge_with(validated_override_params) + model_instance._generation_config = final_config + model_instance.identity.generation_config = final_config logger.debug( f"为新模型 {prov_name_str}/{mod_name_str} 应用配置覆盖: " f"{_get_clean_log_config(validated_override_params)}" ) + else: + model_instance._generation_config = base_gen_config + model_instance.identity.generation_config = base_gen_config _cache_model(cache_key, model_instance) logger.debug( diff --git a/zhenxun/services/ai/llm/system/capabilities.py b/zhenxun/services/ai/llm/system/capabilities.py index 557edd8d..45c15307 100644 --- a/zhenxun/services/ai/llm/system/capabilities.py +++ b/zhenxun/services/ai/llm/system/capabilities.py @@ -125,6 +125,7 @@ CAP_DEEPSEEK_V4 = ModelCapabilities( input_modalities={ModelModality.TEXT}, output_modalities={ModelModality.TEXT}, supports_tool_calling=True, + supports_thinking_toggle=True, reasoning_mode=ReasoningMode.EFFORT, reasoning_visibility="visible", reasoning_effort_map={"minimal": "low"}, @@ -133,6 +134,7 @@ CAP_MINIMAX_REASONING = ModelCapabilities( input_modalities={ModelModality.TEXT}, output_modalities={ModelModality.TEXT}, supports_tool_calling=True, + supports_thinking_toggle=True, reasoning_mode=ReasoningMode.EFFORT, reasoning_visibility="visible", ) @@ -146,6 +148,15 @@ CAP_GLM_MULTIMODAL = ModelCapabilities( }, output_modalities={ModelModality.TEXT}, supports_tool_calling=True, + supports_thinking_toggle=True, +) + +CAP_GLM_REASONING = ModelCapabilities( + input_modalities={ModelModality.TEXT}, + output_modalities={ModelModality.TEXT}, + supports_tool_calling=True, + supports_thinking_toggle=True, + reasoning_mode=ReasoningMode.EFFORT, ) CAP_MINIMAX_MULTIMODAL = ModelCapabilities( @@ -158,6 +169,7 @@ CAP_MIMO_TEXT = ModelCapabilities( input_modalities={ModelModality.TEXT}, output_modalities={ModelModality.TEXT}, supports_tool_calling=True, + supports_thinking_toggle=True, supported_native_tools={"web_search"}, reasoning_effort_map={"max": "high", "xhigh": "high", "minimal": "low"}, ) @@ -171,6 +183,7 @@ CAP_MIMO_MULTIMODAL = ModelCapabilities( }, output_modalities={ModelModality.TEXT, ModelModality.AUDIO}, supports_tool_calling=True, + supports_thinking_toggle=True, supported_native_tools={"web_search"}, reasoning_effort_map={"max": "high", "xhigh": "high", "minimal": "low"}, ) @@ -289,7 +302,7 @@ _ROUTING_TABLE: list[tuple[list[str], ModelCapabilities, int]] = [ CTX_256K, ), (["glm-5v*"], CAP_GLM_MULTIMODAL, CTX_200K), - (["glm-5*", "glm-4.7*", "glm-4.6*"], STANDARD_TEXT_TOOL_CAPABILITIES, CTX_200K), + (["glm-5*", "glm-4.7*", "glm-4.6*"], CAP_GLM_REASONING, CTX_200K), (["*MiniMax-M2*", "*minimax-m2*"], CAP_MINIMAX_REASONING, CTX_200K), (["gpt-4*", "gpt-3.5*", "gpt-*"], CAP_OPENAI_MULTIMODAL, CTX_128K), (["o1-*", "o3-*"], CAP_OPENAI_REASONING, CTX_128K), diff --git a/zhenxun/services/ai/llm/system/network.py b/zhenxun/services/ai/llm/system/network.py index f11d742d..34b4c208 100644 --- a/zhenxun/services/ai/llm/system/network.py +++ b/zhenxun/services/ai/llm/system/network.py @@ -25,7 +25,7 @@ from zhenxun.services.ai.core.exceptions import ( QuotaExceededException, RateLimitException, ) -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_llm as logger from zhenxun.utils.pydantic_compat import model_dump, parse_as from zhenxun.utils.user_agent import get_user_agent diff --git a/zhenxun/services/ai/message_builder.py b/zhenxun/services/ai/message_builder.py index 19e9a3d4..5940b7c3 100644 --- a/zhenxun/services/ai/message_builder.py +++ b/zhenxun/services/ai/message_builder.py @@ -1,23 +1,31 @@ from collections import defaultdict from collections.abc import Awaitable, Callable +import inspect from io import BytesIO import mimetypes from pathlib import Path from typing import Any, ClassVar, TypeVar, cast +import anyio from nonebot.adapters import Bot, Event from nonebot.adapters import Message as PlatformMessage +from nonebot.matcher import current_bot, current_event, current_matcher +from nonebot_plugin_alconna import UniMessage from nonebot_plugin_alconna.uniseg import ( + At, + AtAll, Audio, Image, + Reply, Segment, Text, - UniMessage, Video, Voice, ) +from nonebot_plugin_alconna.uniseg.tools import image_fetch, reply_fetch from PIL.Image import Image as PILImageType +from zhenxun.services.ai.core.engine.context_renderer import ContextConverter from zhenxun.services.ai.core.messages import ( AgentEvent, AudioPart, @@ -35,7 +43,9 @@ from zhenxun.services.ai.core.messages import ( VideoPart, ) from zhenxun.services.ai.core.options import LLMEmbeddingConfig -from zhenxun.services.log import logger +from zhenxun.services.ai.run import get_current_run_context +from zhenxun.services.ai.utils.logger import log_core as logger +from zhenxun.utils.http_utils import AsyncHttpx from zhenxun.utils.pydantic_compat import TypeAdapter, model_copy, model_dump from zhenxun.utils.utils import infer_plugin_namespace @@ -96,16 +106,12 @@ class MessageBuilder: ) -> LLMContentPart | None: """将本地路径读取为多态消息部件""" try: - import anyio - aio_path = anyio.Path(path_like) if not await aio_path.exists() or not await aio_path.is_file(): logger.warning(f"文件不存在或不是一个文件: {path_like}") return None - from pathlib import Path as StdPath - - std_path = StdPath(path_like) + std_path = Path(path_like) resolved_aio_path = await aio_path.absolute() _ = resolved_aio_path @@ -233,11 +239,6 @@ class MessageBuilder: """获取并解析引用消息的内容片段""" namespace = namespace or infer_plugin_namespace(default="global") try: - from nonebot.adapters import Message as PlatformMessage - from nonebot_plugin_alconna import UniMessage - from nonebot_plugin_alconna.uniseg import Reply - from nonebot_plugin_alconna.uniseg.tools import reply_fetch - orig_msg = await reply_fetch(event, bot) if not orig_msg or not orig_msg.msg: return None @@ -283,17 +284,11 @@ class MessageBuilder: reply_parts = [] try: - from nonebot.matcher import current_bot, current_event - from nonebot_plugin_alconna import UniMessage - from nonebot_plugin_alconna.uniseg import Reply - bot_inst = bot or current_bot.get(None) event_inst = event or current_event.get(None) if not bot_inst or not event_inst: try: - from zhenxun.services.ai.run import get_current_run_context - ctx = get_current_run_context() if ctx: bot_inst = bot_inst or ctx.get_bot() @@ -319,7 +314,6 @@ class MessageBuilder: converted_msgs: list[LLMMessage] = [] converted = False - import inspect for msg_type, converter_dict in cls._MESSAGE_CONVERTERS.items(): if isinstance(message, msg_type): @@ -348,10 +342,6 @@ class MessageBuilder: elif isinstance(message, str): converted_msgs = [LLMMessage.user(message)] elif isinstance(message, AgentEvent): - from zhenxun.services.ai.core.engine.context_renderer import ( - ContextConverter, - ) - converted_msgs = ContextConverter.flatten_to_llm_messages([message]) elif isinstance(message, list): parts = [] @@ -402,14 +392,6 @@ class MessageBuilder: elif config.multimodal is False: allowed_modalities = {"text"} - from zhenxun.services.ai.core.messages import ( - AudioPart, - FilePart, - ImagePart, - TextPart, - VideoPart, - ) - messages = await cls.normalize_to_llm_messages( item, bot=bot, @@ -437,9 +419,6 @@ class MessageBuilder: ) -> "EmbedBatch": """将任意输入标准化为嵌入向量批处理对象""" namespace = namespace or infer_plugin_namespace(default="global") - from nonebot_plugin_alconna import UniMessage - - from zhenxun.services.ai.core.messages import BaseContentPart, TextPart if isinstance(inputs, list) and not isinstance(inputs, UniMessage): if not inputs: @@ -510,9 +489,6 @@ def _extract_media_kwargs(seg: Segment, default_mime: str) -> dict | None: async def _handle_image(seg: Image) -> ImagePart | None: if not seg.raw and not getattr(seg, "path", None): try: - from nonebot.matcher import current_bot, current_event, current_matcher - from nonebot_plugin_alconna.uniseg.tools import image_fetch - bot = current_bot.get(None) event = current_event.get(None) matcher = current_matcher.get(None) @@ -526,8 +502,6 @@ async def _handle_image(seg: Image) -> ImagePart | None: logger.debug(f"底层静默水合下载图片失败: {e}") if not seg.raw and seg.url: - from zhenxun.utils.http_utils import AsyncHttpx - try: logger.debug(f"正在从临时 URL 物理固化图片: {seg.url[:50]}...") raw_bytes = await AsyncHttpx.get_content(seg.url) @@ -543,8 +517,6 @@ async def _handle_image(seg: Image) -> ImagePart | None: async def _process_audio_seg(seg: Audio | Voice) -> AudioPart | None: if not seg.raw and not getattr(seg, "path", None) and seg.url: - from zhenxun.utils.http_utils import AsyncHttpx - try: raw_bytes = await AsyncHttpx.get_content(seg.url) if raw_bytes: @@ -569,8 +541,6 @@ async def _handle_voice(seg: Voice) -> AudioPart | None: @MessageBuilder.register_segment_handler(Video, scope="global") async def _handle_video(seg: Video) -> VideoPart | None: if not seg.raw and not getattr(seg, "path", None) and seg.url: - from zhenxun.utils.http_utils import AsyncHttpx - try: raw_bytes = await AsyncHttpx.get_content(seg.url) if raw_bytes: @@ -582,21 +552,14 @@ async def _handle_video(seg: Video) -> VideoPart | None: return VideoPart(**kwargs) if kwargs else None -from nonebot_plugin_alconna.uniseg import At, AtAll, Reply - - @MessageBuilder.register_segment_handler(Reply, scope="global") async def _handle_reply(seg: Reply) -> list[LLMContentPart] | LLMContentPart | None: try: - from nonebot.matcher import current_bot, current_event - bot = current_bot.get(None) event = current_event.get(None) if not bot or not event: return None - from zhenxun.services.ai.run import get_current_run_context - ctx = get_current_run_context() ns = getattr(ctx.session, "namespace", "global") if ctx else "global" @@ -618,8 +581,6 @@ async def _handle_at(seg: At) -> TextPart: target_id = seg.target try: - from nonebot.matcher import current_bot, current_event - bot = current_bot.get(None) event = current_event.get(None) diff --git a/zhenxun/services/ai/run/__init__.py b/zhenxun/services/ai/run/__init__.py index 0c17380f..f76d1ef7 100644 --- a/zhenxun/services/ai/run/__init__.py +++ b/zhenxun/services/ai/run/__init__.py @@ -1,7 +1,6 @@ from zhenxun.services.ai.core.models import CancellationToken from .blackboard import BlackboardManager -from .capabilities import GLOBAL_CAPABILITIES, register_global_capability from .context import ( NoneBotDeps, RunContext, @@ -19,7 +18,6 @@ from .session import session_manager from .ui import UIController __all__ = [ - "GLOBAL_CAPABILITIES", "AgentRunResult", "AgentTask", "BlackboardManager", @@ -33,6 +31,5 @@ __all__ = [ "StreamedRunResult", "UIController", "get_current_run_context", - "register_global_capability", "session_manager", ] diff --git a/zhenxun/services/ai/run/capabilities.py b/zhenxun/services/ai/run/capabilities.py deleted file mode 100644 index 0a3215a9..00000000 --- a/zhenxun/services/ai/run/capabilities.py +++ /dev/null @@ -1,40 +0,0 @@ -from __future__ import annotations - -from collections import defaultdict - -from zhenxun.services.ai.capabilities import ( - AbstractCapability, -) -from zhenxun.services.log import logger -from zhenxun.utils.utils import infer_plugin_namespace - -GLOBAL_CAPABILITIES: dict[str, list[AbstractCapability]] = defaultdict(list) - -from zhenxun.services.ai.capabilities.builtin import ( - BillingCapability, - GlobalCycleLimitCapability, - PermissionCapability, - ReflexionCapability, - StuckDetectionCapability, - ToolRetryAndReflectionCapability, -) - -for _cap in [ - GlobalCycleLimitCapability(), - StuckDetectionCapability(), - PermissionCapability(), - BillingCapability(), - ToolRetryAndReflectionCapability(), - ReflexionCapability(), -]: - GLOBAL_CAPABILITIES["global"].append(_cap) - - -def register_global_capability( - capability: AbstractCapability, scope: str | None = None -) -> None: - ns = scope if scope is not None else infer_plugin_namespace() - GLOBAL_CAPABILITIES[ns].append(capability) - logger.debug( - f"已注册全局 Capability: {capability.__class__.__name__} -> Namespace: {ns}" - ) diff --git a/zhenxun/services/ai/run/context.py b/zhenxun/services/ai/run/context.py index 77718b7a..c46ea584 100644 --- a/zhenxun/services/ai/run/context.py +++ b/zhenxun/services/ai/run/context.py @@ -4,17 +4,31 @@ from contextlib import contextmanager from contextvars import ContextVar import dataclasses import inspect -from typing import Any, Generic, cast, get_origin +from typing import TYPE_CHECKING, Any, Generic, cast, get_origin from typing_extensions import TypeVar +import uuid from nonebot.adapters import Bot, Event -from nonebot.matcher import Matcher +from nonebot.matcher import Matcher, current_bot, current_event, current_matcher from pydantic import BaseModel, ConfigDict, Field +from zhenxun.services.ai.capabilities.base import AbstractCapability +from zhenxun.services.ai.core.engine.append_only import AppendOnlyContextManager from zhenxun.services.ai.core.messages import AgentEvent, AgentMessage +from zhenxun.services.ai.core.models import CancellationToken +from zhenxun.services.ai.core.protocols.tool import ToolExecutable +from zhenxun.services.ai.core.stream_events import AgentStreamEvent, EventBus from zhenxun.services.ai.utils import ContextUtils +from zhenxun.services.ai.utils.scope import ScopeSelector +from zhenxun.services.scheduler.types import ScheduleContext +from zhenxun.utils.platform import PlatformUtils from zhenxun.utils.utils import infer_plugin_namespace +from .blackboard import BlackboardManager + +if TYPE_CHECKING: + from zhenxun.services.ai.context.memory.facades import AgentSessionFacade + AgentDepsT = TypeVar("AgentDepsT", default=Any) """泛型类型变量:外部环境依赖对象 (Agent Dependencies)。""" ProviderFunc = Callable[["RunContext"], Any | Awaitable[Any]] @@ -40,19 +54,47 @@ class NoneBotDeps(BaseModel): @classmethod def get_current(cls) -> "NoneBotDeps | None": """利用 NoneBot 原生魔法,基于 ContextVars 隐式提取当前执行上下文""" - try: - from nonebot.matcher import current_bot, current_event, current_matcher + bot = current_bot.get(None) + event = current_event.get(None) + matcher = current_matcher.get(None) + + if bot or event: + return cls(bot=bot, event=event, matcher=matcher) - bot = current_bot.get(None) - event = current_event.get(None) - matcher = current_matcher.get(None) - if bot or event: - return cls(bot=bot, event=event, matcher=matcher) - except Exception: - pass return None +class ScheduledDeps(BaseModel): + """后台/定时任务环境下的标准依赖容器。""" + + model_config = ConfigDict(arbitrary_types_allowed=True) + + bot: Bot | None = Field(default=None) + """机器人实例""" + group_id: str | None = Field(default=None) + """群组/频道 ID""" + user_id: str | None = Field(default=None) + """用户 ID""" + schedule_id: int | str | None = Field(default=None) + """定时任务/调度任务 ID""" + platform: str | None = Field(default=None) + """适配器平台标识""" + + @classmethod + def from_schedule_context( + cls, bot: Bot, context: ScheduleContext + ) -> "ScheduledDeps": + """从调度器上下文中快速构造并提取依赖""" + return cls( + bot=bot, + group_id=getattr(context, "group_id", None), + user_id=getattr(context, "user_id", None), + schedule_id=getattr(context, "schedule_id", None), + platform=getattr(context, "platform_scope", None) + or PlatformUtils.get_platform_scope(bot), + ) + + @dataclasses.dataclass class SessionContext(Generic[AgentDepsT]): """ @@ -66,9 +108,7 @@ class SessionContext(Generic[AgentDepsT]): """强类型的外部依赖注入对象(如 Bot, Event),供跨工具共享。""" shared_state: dict[str, Any] = dataclasses.field(default_factory=dict) """共享状态字典:全局引用穿透,用于主智能体与嵌套子智能体之间的数据通信。""" - auth_tokens: dict[str, str] = dataclasses.field(default_factory=dict) - """授权凭证字典:保存用户针对各 Provider 的 OAuth 或 API Token。""" - blackboard: Any | None = None + blackboard: BlackboardManager | None = None """结构化黑板管理器,作为共享状态的高级替代方案,提供并发锁和强类型校验。""" namespace: str = "global" """触发事件的插件命名空间""" @@ -76,7 +116,7 @@ class SessionContext(Generic[AgentDepsT]): """用于大模型前缀缓存命中优化的追加写入管理器。""" @property - def memory(self) -> Any: + def memory(self) -> "AgentSessionFacade": """ 获取当前会话的持久化记忆访问门面 (AgentSessionFacade)。 提供极简的 history 和 slots 操作 API。 @@ -84,7 +124,6 @@ class SessionContext(Generic[AgentDepsT]): from zhenxun.services.ai.context.memory.facades import AgentSessionFacade from zhenxun.services.ai.context.memory.manager import memory_manager from zhenxun.services.ai.context.memory.types import SessionMetadata - from zhenxun.services.ai.utils.scope import ScopeSelector user_id = ContextUtils.extract_user_id(self.deps) group_id = ContextUtils.extract_group_id(self.deps) @@ -127,9 +166,9 @@ class AgentRunContext(Generic[AgentDepsT]): """子智能体委派深度标记,用于防范无限递归嵌套。""" tool_retries: dict[str, int] = dataclasses.field(default_factory=dict) """记录当前轮次内各个工具的累积失败重试次数,用于系统熔断。""" - cancellation_token: Any | None = None + cancellation_token: CancellationToken | None = None """全局级联取消令牌,用于跨 Agent 的协程挂起中断。""" - event_bus: Any | None = None + event_bus: EventBus | None = None """底层的事件流发射器 (EventBus),由执行引擎在运行时挂载。""" dynamic_prompts: dict[str, str] = dataclasses.field(default_factory=dict) @@ -146,6 +185,11 @@ class AgentRunContext(Generic[AgentDepsT]): """向当前运行上下文中安全追加业务事件""" self.messages.append(event) + async def emit(self, event: AgentStreamEvent) -> None: + """安全地向事件总线发射流式事件(内置判空逻辑,消解业务层的样板代码)""" + if self.event_bus: + await self.event_bus.emit(event) + @dataclasses.dataclass class ToolCallContext(Generic[AgentDepsT]): @@ -162,7 +206,7 @@ class ToolCallContext(Generic[AgentDepsT]): """本次调用的工具名称。""" retry_count: int = 0 """当前工具调用的重试序号 (第几次重试)。""" - current_tool: Any | None = None + current_tool: ToolExecutable | None = None """当前工具的可执行实例 (ToolExecutable)。""" @@ -188,7 +232,7 @@ class RunContext(Generic[AgentDepsT]): upstream_results: dict[str, Any] = dataclasses.field(default_factory=dict) """前置节点产出字典:标准化的数据流载荷契约,键为 Agent Name,值为输出内容""" - capabilities: list[Any] = dataclasses.field(default_factory=list) + capabilities: list[AbstractCapability] = dataclasses.field(default_factory=list) """当前上下文绑定的拦截器 (Capabilities) 链,用于生命周期拦截""" deps: AgentDepsT = dataclasses.field(default=cast(AgentDepsT, None)) @@ -211,6 +255,14 @@ class RunContext(Generic[AgentDepsT]): ) """标记 session_id 是否为框架隐式生成的。""" + @classmethod + def from_schedule( + cls, bot: Bot, context: ScheduleContext, **kwargs + ) -> "RunContext[ScheduledDeps]": + """极简构造语法糖:从定时调度任务上下文中直接生成 RunContext""" + deps = ScheduledDeps.from_schedule_context(bot, context) + return cast("RunContext[ScheduledDeps]", cls(deps=cast(Any, deps), **kwargs)) + def get_bot(self) -> Bot | None: """强类型安全地提取 Bot 实例""" if not self.deps: @@ -256,21 +308,15 @@ class RunContext(Generic[AgentDepsT]): bot = self.get_bot() event = self.get_event() - if bot and event: + if bot: meta = ContextUtils.generate_session_meta( - bot, event, scope_builder=Isolation.AGENT_USER() + bot=bot, + event=event, + deps=self.deps, + scope_builder=Isolation.AGENT_USER(), ) self.session_id = meta.session_id self._is_auto_session_id = True - else: - uid = self.get_user_id() - gid = self.get_group_id() - if uid and gid: - self.session_id = f"auto_{gid}_{uid}" - self._is_auto_session_id = True - elif uid: - self.session_id = f"auto_private_{uid}" - self._is_auto_session_id = True ns = "global" if self.deps: @@ -286,9 +332,6 @@ class RunContext(Generic[AgentDepsT]): shared_state=self.shared_state, namespace=ns, ) - from zhenxun.services.ai.core.engine.append_only import ( - AppendOnlyContextManager, - ) self.session.append_only_manager = AppendOnlyContextManager() @@ -351,8 +394,6 @@ class RunContext(Generic[AgentDepsT]): new_ctx.run.messages = [] new_ctx.capabilities = [] - import uuid - base_sid = self.session_id or "default" new_ctx.session_id = f"{base_sid}/sub_{member_name}_{uuid.uuid4().hex[:6]}" new_ctx.session.session_id = new_ctx.session_id @@ -403,6 +444,7 @@ __all__ = [ "AgentRunContext", "NoneBotDeps", "RunContext", + "ScheduledDeps", "SessionContext", "ToolCallContext", "ToolsPrepareFunc", diff --git a/zhenxun/services/ai/run/di.py b/zhenxun/services/ai/run/di.py index 7eeace6f..0e7965a6 100644 --- a/zhenxun/services/ai/run/di.py +++ b/zhenxun/services/ai/run/di.py @@ -9,9 +9,9 @@ from nonebot.utils import is_coroutine_callable from nonebot_plugin_session import EventSession, extract_session from zhenxun.services.ai.context.memory.facades import AgentSessionFacade -from zhenxun.services.ai.run.blackboard import BlackboardManager from zhenxun.utils.utils import infer_plugin_namespace +from .blackboard import BlackboardManager from .context import ProviderFunc, RunContext, _is_run_context_type from .hitl import HITLController from .ui import UIController @@ -105,6 +105,10 @@ class Inject: """ return _UpstreamResultMarker(step_name) + CurrentBlackboardState = Annotated[Any, Hidden(), _InjectMarker("blackboard_state")] + BlackboardState = CurrentBlackboardState + """自动注入:当前挂载的强类型黑板底层 Pydantic 实例(需配合 Annotated 使用)""" + UserId = CurrentUserId """自动注入:触发当前任务的用户 ID""" @@ -366,6 +370,19 @@ def _resolve_blackboard(ctx: RunContext): Inject.register_provider("blackboard", _resolve_blackboard, scope="global") +def _resolve_blackboard_state(ctx: RunContext): + bb = ctx.session.blackboard + if bb is None: + raise ValueError( + "Inject.BlackboardState 注入失败:" + "当前会话上下文中未挂载 BlackboardManager 实例。" + ) + return bb._state + + +Inject.register_provider("blackboard_state", _resolve_blackboard_state, scope="global") + + def _resolve_session(ctx): bot = ctx.get_bot() event = ctx.get_event() diff --git a/zhenxun/services/ai/run/hitl.py b/zhenxun/services/ai/run/hitl.py index 889f9062..d51875aa 100644 --- a/zhenxun/services/ai/run/hitl.py +++ b/zhenxun/services/ai/run/hitl.py @@ -1,17 +1,17 @@ +from __future__ import annotations + import asyncio from collections.abc import Callable import time -from typing import TYPE_CHECKING from nonebot.adapters import Event from nonebot.permission import SUPERUSER from nonebot_plugin_waiter import waiter from zhenxun.services.ai.core.exceptions import AbortException, ToolFatalError -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_agent as logger -if TYPE_CHECKING: - from zhenxun.services.ai.run.context import RunContext +from .context import RunContext CANCEL_WORDS = {"取消", "cancel", "0", "退出", "quit"} CONFIRM_WORDS = {"y", "yes", "是", "1", "ok", "确认"} @@ -24,7 +24,7 @@ class HITLController: 封装底层的物理环境隔离等待逻辑,供上层工具和中间件发起提问或审批。 """ - def __init__(self, context: "RunContext"): + def __init__(self, context: RunContext): self.context = context async def wait_event( @@ -42,7 +42,9 @@ class HITLController: if not bot or not event: logger.warning("HITLController: 当前环境无 Bot/Event 实例,无法发起交互。") - return None + raise ToolFatalError( + "当前处于自动化后台调度环境,无法发起人工交互。任务已强行中止。" + ) if prompt_msg: await bot.send(event, prompt_msg) diff --git a/zhenxun/services/ai/run/models.py b/zhenxun/services/ai/run/models.py index a8b42986..6e86afa2 100644 --- a/zhenxun/services/ai/run/models.py +++ b/zhenxun/services/ai/run/models.py @@ -1,8 +1,10 @@ +from __future__ import annotations + """ 运行时(Run)相关核心类型定义 """ -from collections.abc import AsyncIterator +from collections.abc import AsyncIterator, Callable import json from typing import Any, Generic, cast from typing_extensions import TypeVar @@ -10,12 +12,16 @@ from typing_extensions import TypeVar from pydantic import BaseModel, ConfigDict, Field, PrivateAttr from zhenxun.services.ai.core.messages import AgentMessage, LLMMessage, UsageInfo +from zhenxun.services.ai.core.options import BaseOutputDefinition +from zhenxun.services.ai.core.protocols.tool import ToolResolvable from zhenxun.services.ai.core.stream_events import ( AgentStreamEvent, EventBus, ToolCallStartEvent, ToolStreamChunkEvent, ) +from zhenxun.services.ai.guardrails import BaseGuardrail, GuardrailSource +from zhenxun.services.ai.tools.core.tool import BaseTool from zhenxun.utils.pydantic_compat import model_dump, model_validator @@ -37,7 +43,7 @@ class ToolSummary(BaseModel): """工具执行失败次数""" total_latency_ms: float = 0.0 """工具执行总耗时(毫秒)""" - by_name: dict[str, dict[str, Any]] = Field(default_factory=dict) + by_name: dict[str, dict[str, float | int]] = Field(default_factory=dict) """按工具名称细分的执行状态统计""" @@ -126,9 +132,8 @@ class StreamedRunResult(Generic[OutputDataT]): self.is_complete: bool = False self._result: AgentRunResult[OutputDataT] | None = None - async def stream_events(self) -> AsyncIterator[Any]: + async def stream_events(self) -> "AsyncIterator[AgentStreamEvent]": """获取底层的所有原始事件(包含工具调用过程等)""" - from zhenxun.services.ai.run.models import AgentRunEnd, AgentRunError async for event in self._event_bus: if isinstance(event, AgentRunEnd): @@ -214,24 +219,6 @@ class StreamedRunResult(Generic[OutputDataT]): return await self.get_run_result() -class TaskResult(BaseModel): - """单个数据契约任务的执行结果""" - - task_id: str - """关联的任务唯一 ID""" - - output: Any - """任务的实际产出(如果是结构化任务,则为解析后的 Pydantic 实例;否则为纯文本)""" - - raw_response: str | None = None - """大模型返回的原始纯文本内容""" - - usage: UsageInfo = Field(default_factory=UsageInfo) - """该任务执行期间的 Token 消耗统计""" - - model_config = ConfigDict(arbitrary_types_allowed=True) - - class AgentTask(BaseModel): """标准化数据契约(意图载体 Payload),定义大模型需要做什么及产出什么格式""" @@ -247,23 +234,25 @@ class AgentTask(BaseModel): expected_output: str """预期输出的自然语言描述(指导大模型如何组织最终答案)""" - response_model: Any | None = None + response_model: type[BaseModel] | BaseOutputDefinition | None = None """强制要求返回的强类型结构 (Pydantic Model) 或 OutputDefinition,为空则返回普通文本""" - tools: list[str | Any] | None = None + tools: list[str | Callable | dict[str, Any] | BaseTool | ToolResolvable] | None = ( + None + ) """针对此特定任务动态追加或覆盖的工具列表""" - guardrails: list[Any] | None = None + guardrails: list[GuardrailSource] | None = None """护栏验证列表。支持传入函数、BaseGuardrail 实例, 或直接传入自然语言字符串规则(自动转为 LLM 裁判)""" - _parsed_guardrails: list[Any] = PrivateAttr(default_factory=list) + _parsed_guardrails: list[BaseGuardrail] = PrivateAttr(default_factory=list) model_config = ConfigDict(arbitrary_types_allowed=True) @model_validator(mode="after") - def _parse_and_set_guardrails(self) -> "AgentTask": + def _parse_and_set_guardrails(self) -> AgentTask: from zhenxun.services.ai.guardrails import parse_guardrails self._parsed_guardrails = parse_guardrails(self.guardrails) @@ -275,5 +264,4 @@ __all__ = [ "AgentTask", "OutputDataT", "StreamedRunResult", - "TaskResult", ] diff --git a/zhenxun/services/ai/run/session.py b/zhenxun/services/ai/run/session.py index 154427bf..93ac365f 100644 --- a/zhenxun/services/ai/run/session.py +++ b/zhenxun/services/ai/run/session.py @@ -108,7 +108,7 @@ class AgentSessionManager: count += 1 if task and not task.done(): task.cancel() - from zhenxun.services.log import logger + from zhenxun.services.ai.utils.logger import log_agent as logger if count > 0: logger.info( diff --git a/zhenxun/services/ai/run/subscribers.py b/zhenxun/services/ai/run/subscribers.py index 72a52519..3b55cc4c 100644 --- a/zhenxun/services/ai/run/subscribers.py +++ b/zhenxun/services/ai/run/subscribers.py @@ -13,9 +13,12 @@ from zhenxun.services.ai.core.stream_events import ( ToolStreamChunkEvent, UserCustomEvent, ) -from zhenxun.services.ai.run.context import RunContext -from zhenxun.services.ai.run.models import AgentRunEnd, AgentRunStart, AgentRunSummary -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_agent as logger +from zhenxun.utils.message import MessageUtils +from zhenxun.utils.platform import PlatformUtils + +from .context import RunContext +from .models import AgentRunEnd, AgentRunStart, AgentRunSummary class TelemetrySubscriber: @@ -122,7 +125,7 @@ class DefaultUISubscriber: async def _send_to_platform(self, display: Any): """将通用的显示内容或富文本消息部件渲染并发送至具体的聊天平台""" - if not self.bot or not self.event or not display: + if not self.bot or not display: return if ( @@ -144,9 +147,27 @@ class DefaultUISubscriber: display = msg if isinstance(display, UniMessage): - await display.send(self.event, bot=self.bot, reply_to=self.reply_to) + if self.event: + await display.send(self.event, bot=self.bot, reply_to=self.reply_to) + else: + target = PlatformUtils.get_target( + user_id=self.context.get_user_id(), + group_id=self.context.get_group_id(), + ) + if target: + await display.send(target=target, bot=self.bot) else: - await self.bot.send(self.event, str(display)) + if self.event: + await self.bot.send(self.event, str(display)) + else: + target = PlatformUtils.get_target( + user_id=self.context.get_user_id(), + group_id=self.context.get_group_id(), + ) + if target: + await MessageUtils.build_message(str(display)).send( + target=target, bot=self.bot + ) async def on_tool_stream(self, event: ToolStreamChunkEvent): """处理工具流式数据块事件,仅在 verbose 开启时发送给平台""" diff --git a/zhenxun/services/ai/run/ui.py b/zhenxun/services/ai/run/ui.py index 6a08f668..566393b2 100644 --- a/zhenxun/services/ai/run/ui.py +++ b/zhenxun/services/ai/run/ui.py @@ -1,17 +1,20 @@ -from typing import TYPE_CHECKING, Any +from __future__ import annotations + +from typing import Any from nonebot_plugin_alconna import UniMessage from zhenxun.services.ai.core.stream_events import ToolStreamChunkEvent, UserCustomEvent +from zhenxun.utils.message import MessageUtils +from zhenxun.utils.platform import PlatformUtils -if TYPE_CHECKING: - from zhenxun.services.ai.run.context import RunContext +from .context import RunContext class UIController: """前端 UI 富交互流式控制器 (按需生成模式)""" - def __init__(self, context: "RunContext"): + def __init__(self, context: RunContext): self.context = context @property @@ -19,38 +22,31 @@ class UIController: """从上下文中动态获取当前调用的工具名""" return getattr(self.context.call, "tool_name", "UnknownTool") - @property - def _event_bus(self) -> Any | None: - """从运行时上下文中获取底层的事件发射器""" - return getattr(self.context.run, "event_bus", None) - async def send_text(self, text: str, status: str = "running") -> None: """向前端流式反馈执行进度文本""" - if self._event_bus: - await self._event_bus.emit( - ToolStreamChunkEvent( - tool_name=self.tool_name, content=text, metadata={"status": status} - ) + await self.context.run.emit( + ToolStreamChunkEvent( + tool_name=self.tool_name, content=text, metadata={"status": status} ) + ) async def send_image(self, image: bytes | str) -> None: """向前端发送富文本图片气泡(字节流或 URL)""" - if self._event_bus: - display_msg = UniMessage() - if isinstance(image, bytes): - display_msg = display_msg.image(raw=image) - else: - display_msg = display_msg.image(url=image) - await self._event_bus.emit(UserCustomEvent(display=display_msg)) + display_msg = UniMessage() + if isinstance(image, bytes): + display_msg = display_msg.image(raw=image) + else: + display_msg = display_msg.image(url=image) + await self.context.run.emit(UserCustomEvent(display=display_msg)) async def send_display(self, display: Any) -> None: """向前端发送任意展示载体""" - if self._event_bus and display is not None: - await self._event_bus.emit(UserCustomEvent(display=display)) + if display is not None: + await self.context.run.emit(UserCustomEvent(display=display)) @staticmethod async def handle_control_flow_exit_display( - e: BaseException, context: "RunContext | None", reply_to: bool = False + e: BaseException, context: RunContext | None, reply_to: bool = False ) -> None: """统一处理 ControlFlowExit 异常带来的 UI 反馈逻辑""" from zhenxun.services.ai.core.exceptions import ControlFlowExit @@ -63,19 +59,24 @@ class UIController: return try: - from nonebot_plugin_alconna import UniMessage - bot = context.get_bot() if context else None event = context.get_event() if context else None - if bot and event: - if isinstance(display_msg, UniMessage): - await display_msg.send(event, bot=bot, reply_to=reply_to) - else: - await bot.send(event, str(display_msg)) + if bot: + msg_to_send = ( + display_msg + if isinstance(display_msg, UniMessage) + else MessageUtils.build_message(str(display_msg)) + ) + if event: + await msg_to_send.send(event, bot=bot, reply_to=reply_to) + elif context: + target = PlatformUtils.get_target( + user_id=context.get_user_id(), group_id=context.get_group_id() + ) + if target: + await msg_to_send.send(target=target, bot=bot) else: - from zhenxun.utils.message import MessageUtils - await MessageUtils.build_message(str(display_msg)).send( reply_to=reply_to ) diff --git a/zhenxun/services/ai/sandbox/addons/base.py b/zhenxun/services/ai/sandbox/addons/base.py index d5056496..2b08a18d 100644 --- a/zhenxun/services/ai/sandbox/addons/base.py +++ b/zhenxun/services/ai/sandbox/addons/base.py @@ -3,7 +3,7 @@ from collections.abc import AsyncGenerator from contextlib import asynccontextmanager from typing import TYPE_CHECKING, Any -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_sandbox as logger if TYPE_CHECKING: from zhenxun.services.ai.sandbox.drivers.base import BaseSandboxSession diff --git a/zhenxun/services/ai/sandbox/addons/mcp_proxy.py b/zhenxun/services/ai/sandbox/addons/mcp_proxy.py index 06834f35..16dac911 100644 --- a/zhenxun/services/ai/sandbox/addons/mcp_proxy.py +++ b/zhenxun/services/ai/sandbox/addons/mcp_proxy.py @@ -7,12 +7,13 @@ from anyio import create_memory_object_stream, create_task_group from mcp.shared.message import SessionMessage from mcp.types import JSONRPCMessage -from zhenxun.services.ai.sandbox.addons.base import BaseMcpProxyExtension from zhenxun.services.ai.sandbox.protocols import SupportsStreamExecution from zhenxun.services.ai.sandbox.registry import SandboxRegistry -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_sandbox as logger from zhenxun.utils.pydantic_compat import model_dump_json, model_validate +from .base import BaseMcpProxyExtension + class UniversalMcpExtension(BaseMcpProxyExtension): """通用 MCP 代理扩展类,用于在沙箱内连接 MCP 服务""" diff --git a/zhenxun/services/ai/sandbox/drivers/docker.py b/zhenxun/services/ai/sandbox/drivers/docker.py index a8f10f2d..f9643b89 100644 --- a/zhenxun/services/ai/sandbox/drivers/docker.py +++ b/zhenxun/services/ai/sandbox/drivers/docker.py @@ -24,7 +24,7 @@ from zhenxun.services.ai.sandbox.protocols import ( SandboxProcessStream, ) from zhenxun.services.ai.sandbox.storage import RESOLVE_PATH_HELPER, coerce_posix_path -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_sandbox as logger from .base import BaseSandboxClient, BaseSandboxSession @@ -388,8 +388,6 @@ class DockerSandboxSession(BaseSandboxSession): await self.container.put_archive(sandbox_target_path, tar_bytes) return True except Exception as e: - from zhenxun.services.log import logger - logger.error(f"[Docker I/O] 上传目录失败: {e}") return False @@ -552,10 +550,7 @@ class DockerSandboxClient(BaseSandboxClient): DockerSandboxClient._containers[eff_cname] = container DockerSandboxClient._jupyter_ports[eff_cname] = jupyter_port - logger.info( - f"[DockerSandbox] 已启动物理隔离容器: {eff_cname} " - f"(镜像: {eff_image})" - ) + logger.info(f"已启动物理隔离容器: {eff_cname} (镜像: {eff_image})") state = SandboxSessionState( session_id=session_id, @@ -575,10 +570,7 @@ class DockerSandboxClient(BaseSandboxClient): "test -x /global_env/python_venv/bin/pip" ) if check_venv.exit_code != 0: - logger.info( - f"正在初始化/修复容器 [{eff_cname}] 的共享 Python 虚拟环境...", - command="SandboxManager", - ) + logger.info(f"正在初始化/修复容器 [{eff_cname}] 的共享 Python 虚拟环境...") init_res = await session.run_process( "rm -rf /global_env/python_venv && " "uv venv --seed --system-site-packages /global_env/python_venv || " @@ -586,8 +578,7 @@ class DockerSandboxClient(BaseSandboxClient): ) if init_res.exit_code != 0: logger.error( - f"初始化虚拟环境失败: {init_res.stderr or init_res.stdout}", - command="SandboxManager", + f"初始化虚拟环境失败: {init_res.stderr or init_res.stdout}" ) return session @@ -615,17 +606,14 @@ class DockerSandboxClient(BaseSandboxClient): await self._containers[cname].delete(force=True) self._containers.pop(cname, None) self._jupyter_ports.pop(cname, None) - logger.info( - f"[DockerSandbox] 物理容器 {cname} " - "已长时间闲置,已触发彻底销毁释放内存。" - ) + logger.info(f"物理容器 {cname} 已长时间闲置,已触发彻底销毁释放内存。") except Exception as e: self._containers.pop(cname, None) self._jupyter_ports.pop(cname, None) if getattr(e, "status", None) == 404 or "No such container" in str(e): - logger.debug(f"[DockerSandbox] 物理容器 {cname} 已不存在。") + logger.debug(f"物理容器 {cname} 已不存在。") else: - logger.error(f"[DockerSandbox] 闲置销毁物理容器 {cname} 失败: {e}") + logger.error(f"闲置销毁物理容器 {cname} 失败: {e}") @classmethod async def close_env(cls): @@ -663,8 +651,6 @@ class DockerSandboxClient(BaseSandboxClient): except Exception: pass if count > 0: - from zhenxun.services.log import logger - logger.info( f"静默清理:已成功回收 {count} 个上次异常退出遗留的沙箱容器。", command="SandboxManager", diff --git a/zhenxun/services/ai/sandbox/environments.py b/zhenxun/services/ai/sandbox/environments.py index 193c4aa1..51972958 100644 --- a/zhenxun/services/ai/sandbox/environments.py +++ b/zhenxun/services/ai/sandbox/environments.py @@ -1,11 +1,11 @@ from abc import ABC, abstractmethod from typing import TYPE_CHECKING, Any, ClassVar -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_sandbox as logger if TYPE_CHECKING: - from zhenxun.services.ai.sandbox.drivers.base import BaseSandboxSession - from zhenxun.services.ai.sandbox.models import SandboxBlueprint + from .drivers.base import BaseSandboxSession + from .models import SandboxBlueprint class BaseProvisioner(ABC): diff --git a/zhenxun/services/ai/sandbox/manager.py b/zhenxun/services/ai/sandbox/manager.py index 6b41494c..07ff2392 100644 --- a/zhenxun/services/ai/sandbox/manager.py +++ b/zhenxun/services/ai/sandbox/manager.py @@ -4,11 +4,11 @@ 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.log import logger +from zhenxun.services.ai.utils.logger import log_sandbox as logger from zhenxun.utils.lifespan import LifespanManager from .drivers.base import BaseSandboxClient, BaseSandboxSession +from .models import SandboxBlueprint from .registry import SandboxRegistry _startup_tasks = set() @@ -84,10 +84,7 @@ class SandboxManager: session = self._active_sessions[session_id] is_alive = await session.is_alive() if not is_alive: - logger.warning( - f"⚠️ Session '{session_id}' 容器已死亡,重建中。", - command="SandboxManager", - ) + logger.warning(f"⚠️ Session '{session_id}' 容器已死亡,重建中。") await self.release_resource(session_id) else: session.touch() @@ -98,9 +95,7 @@ class SandboxManager: await session.mount_extension(p) return session - logger.info( - f"为 Session '{session_id}' 创建沙箱环境...", command="SandboxManager" - ) + logger.info(f"为 Session '{session_id}' 创建沙箱环境...") client = self._get_client(blueprint) try: session = await client.create(session_id, blueprint) @@ -121,10 +116,7 @@ class SandboxManager: from zhenxun.services.ai.core.exceptions import SandboxFatalError err_msg = str(e) or type(e).__name__ - logger.error( - f"创建 Session 失败: {err_msg}\n{traceback.format_exc()}", - command="SandboxManager", - ) + logger.error(f"创建 Session 失败: {err_msg}\n{traceback.format_exc()}") if not isinstance(e, SandboxFatalError): raise SandboxFatalError(f"沙箱初始化异常: {err_msg}") from e raise e @@ -134,15 +126,12 @@ class SandboxManager: ) -> bool: """统一扫描指定工作区,通过 Provisioner 体系完成环境装配""" if session_id not in self._active_sessions: - logger.warning( - f"找不到活跃的 Session '{session_id}',无法配置环境。", - command="SandboxManager", - ) + logger.warning(f"找不到活跃的 Session '{session_id}',无法配置环境。") return False session = self._active_sessions[session_id] - from zhenxun.services.ai.sandbox.environments import ProvisionerRegistry + from .environments import ProvisionerRegistry for prov in ProvisionerRegistry.get_all().values(): await prov.scan_and_setup_workspace(cast(Any, session), workspace_dir) @@ -157,10 +146,7 @@ class SandboxManager: client = client_cls() await client.delete(cast(Any, session)) except Exception as e: - logger.error( - f"销毁沙箱环境失败 (Session: {resource_id}): {e}", - command="SandboxManager", - ) + logger.error(f"销毁沙箱环境失败 (Session: {resource_id}): {e}") async def close_session(self, session_id: str) -> None: """从生存期管理器中注销并销毁指定会话的沙箱环境""" diff --git a/zhenxun/services/ai/sandbox/models.py b/zhenxun/services/ai/sandbox/models.py index f09cec4d..92b70e9c 100644 --- a/zhenxun/services/ai/sandbox/models.py +++ b/zhenxun/services/ai/sandbox/models.py @@ -12,11 +12,11 @@ from typing import TYPE_CHECKING, Annotated, Any, Literal from pydantic import BaseModel, ConfigDict, Field -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_sandbox as logger from zhenxun.utils.pydantic_compat import model_dump if TYPE_CHECKING: - from zhenxun.services.ai.sandbox.drivers.base import BaseSandboxSession + from .drivers.base import BaseSandboxSession class LanguageProfile(BaseModel): @@ -337,7 +337,6 @@ class SandboxBlueprint(BaseModel): def with_local_dir(self, path: str, local_dir: str) -> "SandboxBlueprint": """声明预置本地宿主机完整目录""" - from zhenxun.services.log import logger for bm in self.bind_mounts: if path.startswith(bm.sandbox_path) or bm.sandbox_path.startswith(path): @@ -370,8 +369,6 @@ class SandboxBlueprint(BaseModel): """声明宿主机物理目录双向挂载 (Bind Mount)""" from pathlib import Path - from zhenxun.services.log import logger - if remote_path in self.entries: logger.warning( f"⚠️ [Sandbox] 目标沙箱路径 '{remote_path}' " diff --git a/zhenxun/services/ai/sandbox/protocols.py b/zhenxun/services/ai/sandbox/protocols.py index 08e104d5..6933f5b8 100644 --- a/zhenxun/services/ai/sandbox/protocols.py +++ b/zhenxun/services/ai/sandbox/protocols.py @@ -5,7 +5,7 @@ from dataclasses import dataclass from pathlib import Path from typing import Any, Protocol, runtime_checkable -from zhenxun.services.ai.sandbox.models import SandboxExecutionResult +from .models import SandboxExecutionResult class InteractiveTerminalSession(Protocol): diff --git a/zhenxun/services/ai/sandbox/registry.py b/zhenxun/services/ai/sandbox/registry.py index 209959e6..0f511535 100644 --- a/zhenxun/services/ai/sandbox/registry.py +++ b/zhenxun/services/ai/sandbox/registry.py @@ -1,10 +1,10 @@ from typing import TYPE_CHECKING, ClassVar -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_sandbox as logger if TYPE_CHECKING: - from zhenxun.services.ai.sandbox.addons.base import BaseSandboxExtension - from zhenxun.services.ai.sandbox.drivers.base import BaseSandboxClient + from .addons.base import BaseSandboxExtension + from .drivers.base import BaseSandboxClient class SandboxRegistry: diff --git a/zhenxun/services/ai/sandbox/runtimes.py b/zhenxun/services/ai/sandbox/runtimes.py index 13b70d14..2de245a5 100644 --- a/zhenxun/services/ai/sandbox/runtimes.py +++ b/zhenxun/services/ai/sandbox/runtimes.py @@ -7,13 +7,14 @@ import uuid import aiohttp -from zhenxun.services.ai.sandbox.drivers.base import BaseSandboxSession -from zhenxun.services.ai.sandbox.models import ( +from zhenxun.services.ai.utils.logger import log_sandbox as logger +from zhenxun.utils.utils import infer_plugin_namespace + +from .drivers.base import BaseSandboxSession +from .models import ( LanguageProfile, SandboxExecutionResult, ) -from zhenxun.services.log import logger -from zhenxun.utils.utils import infer_plugin_namespace def parse_shebang(script_path: str | Path) -> str | None: diff --git a/zhenxun/services/ai/tools/bridges/delegate.py b/zhenxun/services/ai/tools/bridges/delegate.py index 90e37aaa..9e66208d 100644 --- a/zhenxun/services/ai/tools/bridges/delegate.py +++ b/zhenxun/services/ai/tools/bridges/delegate.py @@ -1,19 +1,19 @@ +from __future__ import annotations + import json -from typing import TYPE_CHECKING, Any +from typing import Any from pydantic import BaseModel, Field -from zhenxun.services.ai.core.exceptions import ControlFlowExit +from zhenxun.services.ai.core.exceptions import AbortException, ControlFlowExit from zhenxun.services.ai.core.stream_events import ToolStreamChunkEvent +from zhenxun.services.ai.flow.base import BaseRunnable from zhenxun.services.ai.run import RunContext from zhenxun.services.ai.tools.core.tool import BaseTool from zhenxun.services.ai.tools.models import ToolResult -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_tool as logger from zhenxun.utils.pydantic_compat import model_dump -if TYPE_CHECKING: - from zhenxun.services.ai.flow.base import BaseRunnable - STRUCTURED_INPUT_PREAMBLE = ( "\n\n### 🛠️ [嵌套调用前置语境]\n" "你现在正在作为一个『工具/子节点』被外部主智能体调用。\n" @@ -30,24 +30,22 @@ class DelegateArgs(BaseModel): class DelegateTool(BaseTool): - """ - 将任意实现了 run() 方法的实体 (Agent/Team/Workflow 等) 包装为大模型可调用的工具。 - (SubRoutine 委派模式) - """ + """将可运行实体 (Agent/Team/Workflow) 包装为子例程委派工具。""" def __init__( self, - runnable: "BaseRunnable[Any]", + runnable: BaseRunnable[Any], name: str | None = None, description: str | None = None, + max_delegations: int = 3, ): - """ - 初始化委派工具,将可运行实体包装为子例程工具。 + """初始化委派工具。 参数: - runnable: 被包装的可运行实体,可以是 Agent、Team 或 Workflow 等。 - name: 自定义工具名称,若为 None 则默认推导为 runnable 实体名。 - description: 工具描述信息,用于指导大模型何时代用此工具。 + runnable: 被包装的可运行实体。 + name: 自定义工具名称。 + description: 工具描述信息。 + max_delegations: 允许向同一个实体连续委派且未获成功的最大次数限制。 """ resolved_name = name or getattr(runnable, "name", "SubRunnable") resolved_desc = description or getattr( @@ -61,6 +59,7 @@ class DelegateTool(BaseTool): super().__init__(name=final_name, description=resolved_desc) self.runnable = runnable self.args_schema = DelegateArgs + self.max_delegations = max_delegations async def execute( self, context: RunContext | None = None, **kwargs: Any @@ -70,12 +69,13 @@ class DelegateTool(BaseTool): counts = context.session.shared_state.setdefault("__delegate_counts__", {}) counts[self.name] = counts.get(self.name, 0) + 1 - if counts[self.name] > 3: + if counts[self.name] > self.max_delegations: return ToolResult( output=( - f"❌ 系统拦截:检测到无限委派死循环!\n" + f"❌ 系统拦截:委派重试次数已达上限。\n" f"你已经连续 {counts[self.name]} 次将子任务委派给下级实体 " - f"{self.name} 且未获最终成功。\n" + f"{self.name} 且未获最终成功" + f"(超出最大允许次数 {self.max_delegations})。\n" "请立即停止委派," "改变你的思考方向或直接向用户汇报失败结论!" ) @@ -86,8 +86,6 @@ class DelegateTool(BaseTool): logger.warning( f"⚠️ [DelegateTool] 委派深度超限 ({depth}),强制阻断: {self.name}" ) - from zhenxun.services.ai.core.exceptions import AbortException - raise AbortException( reason="嵌套层级过深,系统已强制拒绝执行委派", display=f"⚠️ {self.name} 嵌套层级过深", @@ -132,8 +130,6 @@ class DelegateTool(BaseTool): "DepthLimitExceeded" in final_output or "嵌套层级过深" in final_output ) if is_fatal: - from zhenxun.services.ai.core.exceptions import AbortException - raise AbortException( reason="下级实体遇到深度限制异常", display=f"⚠️ 实体 {self.name} 委派失败", @@ -141,8 +137,8 @@ class DelegateTool(BaseTool): usage = getattr(response, "usage", None) - if context and context.run.event_bus: - await context.run.event_bus.emit( + if context: + await context.run.emit( ToolStreamChunkEvent( tool_name=self.name, content=f"🧠 实体 {self.name} 执行完毕" ) @@ -156,8 +152,6 @@ class DelegateTool(BaseTool): raise e except Exception as e: logger.error(f"委派实体 {self.name} 执行失败: {e}", e=e) - from zhenxun.services.ai.core.exceptions import AbortException - raise AbortException( reason=f"Delegate Execution Error: {e}", display=f"❌ 实体 {self.name} 执行异常", diff --git a/zhenxun/services/ai/tools/bridges/handoff.py b/zhenxun/services/ai/tools/bridges/handoff.py index 0ac799a5..f7d46752 100644 --- a/zhenxun/services/ai/tools/bridges/handoff.py +++ b/zhenxun/services/ai/tools/bridges/handoff.py @@ -2,6 +2,8 @@ from typing import Any from pydantic import BaseModel, Field, create_model +from zhenxun.services.ai.core.messages import HandoffEvent +from zhenxun.services.ai.core.options import BaseOutputDefinition 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 @@ -18,6 +20,7 @@ class HandoffTool(BaseTool): target_name: str, target_description: str, input_schema: type[BaseModel] | Any | None = None, + max_handoffs: int = 3, ): """ 初始化移交工具,为模型赋予转移对话控制权到指定实体的能力。 @@ -26,7 +29,8 @@ class HandoffTool(BaseTool): target_name: 被转移的目标接收者(Agent 或负责人)的唯一标识名称。 target_description: 目标接收者的职责或专长说明,供模型决策是否移交。 input_schema: 自定义移交数据结构,指定转移时所需携带的结构化参数。 - """ + max_handoffs: 允许在同一个会话中向同一个实体发起移交的最大次数,防止无限踢皮球,默认 3。 + """ # noqa: E501 super().__init__( name=f"transfer_to_{target_name}", description=( @@ -34,11 +38,10 @@ class HandoffTool(BaseTool): ), ) self.target_name = target_name + self.max_handoffs = max_handoffs actual_schema = None if input_schema: - from zhenxun.services.ai.core.options import BaseOutputDefinition - if isinstance(input_schema, BaseOutputDefinition): actual_schema = input_schema.type_ else: @@ -80,19 +83,18 @@ class HandoffTool(BaseTool): if context: counts = context.session.shared_state.setdefault("__handoff_counts__", {}) counts[self.target_name] = counts.get(self.target_name, 0) + 1 - if counts[self.target_name] > 3: + if counts[self.target_name] > self.max_handoffs: return ToolResult( output=( - f"❌ 系统拦截:检测到严重的踢皮球现象!\n" - f"你所在的团队已经连续 {counts[self.target_name]} 次将任务" - f"移交给 {self.target_name},\n" - "但问题仍未解决。请立刻改变策略," - "由你亲自处理或得出最终结论,严禁再次移交!" + f"❌ 系统拦截:移交次数已达上限。\n" + f"你所在的团队已经向 {self.target_name} 尝试移交了 " + f"{counts[self.target_name]} 次," + f"超出了最大允许次数 ({self.max_handoffs})。\n" + "为防止任务陷入停滞,请立即改变策略," + "由你亲自处理当前任务或得出最终结论,严禁再次移交!" ) ).as_error() - from zhenxun.services.ai.core.messages import HandoffEvent - if context: context.run.add_event( HandoffEvent( diff --git a/zhenxun/services/ai/tools/bridges/matcher_bridge.py b/zhenxun/services/ai/tools/bridges/matcher_bridge.py index 83571e58..9b6df31c 100644 --- a/zhenxun/services/ai/tools/bridges/matcher_bridge.py +++ b/zhenxun/services/ai/tools/bridges/matcher_bridge.py @@ -24,8 +24,9 @@ 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 -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_tool as logger from zhenxun.utils.pydantic_compat import model_validate +from zhenxun.utils.utils import infer_plugin_namespace class MatcherAdapter(ABC): @@ -370,7 +371,8 @@ class MatcherTool(BaseTool): if not bot or not event: return ToolResult( - output="上下文缺少 bot 或 event,无法执行 Matcher" + output="❌ 当前处于无状态的自动化后台环境,缺乏真实的用户上下文(Event)," + "严禁调用该原生平台指令工具!请换用其他方法解决。" ).as_error() try: @@ -448,8 +450,6 @@ def bind_matcher( """ if require_prefix: - from zhenxun.utils.utils import infer_plugin_namespace - ns = infer_plugin_namespace(default="global") if ns and ns not in ("global", "unknown"): if not name.startswith(f"{ns}_"): diff --git a/zhenxun/services/ai/tools/core/capabilities.py b/zhenxun/services/ai/tools/core/capabilities.py index e2570a17..e837d013 100644 --- a/zhenxun/services/ai/tools/core/capabilities.py +++ b/zhenxun/services/ai/tools/core/capabilities.py @@ -21,7 +21,7 @@ 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 +from zhenxun.services.ai.utils.logger import log_tool as logger from zhenxun.utils.pydantic_compat import model_dump, parse_as _TOOL_RESULT_CACHE = SimpleMemoryCache(namespace="zhenxun_tool_cache") @@ -58,7 +58,7 @@ class CacheCapability(AbstractCapability): if not hasattr(tool, "_generate_cache_key"): return await handler(arguments) - cache_key = tool._generate_cache_key(arguments) + cache_key = getattr(tool, "_generate_cache_key")(arguments) cached_data = await _TOOL_RESULT_CACHE.get(cache_key) if cached_data is not None: cached_result = ( diff --git a/zhenxun/services/ai/tools/core/decorators.py b/zhenxun/services/ai/tools/core/decorators.py index 3ea6c4d3..67ca2ade 100644 --- a/zhenxun/services/ai/tools/core/decorators.py +++ b/zhenxun/services/ai/tools/core/decorators.py @@ -1,11 +1,23 @@ from __future__ import annotations from collections.abc import Callable +import inspect +import types from typing import Any from zhenxun.services.ai.capabilities import AbstractCapability from zhenxun.services.ai.run.context import RunContext -from zhenxun.services.ai.tools.core.capabilities import ( +from zhenxun.services.ai.tools.engine.registry import tool_provider_manager +from zhenxun.services.ai.tools.models import ( + EndRunResult, + ToolkitConfig, + ToolOptions, + ToolResult, +) +from zhenxun.services.ai.utils.logger import log_tool as logger +from zhenxun.utils.pydantic_compat import model_copy + +from .capabilities import ( AdminLevelCapability, ApprovalCapability, CacheCapability, @@ -15,14 +27,15 @@ from zhenxun.services.ai.tools.core.capabilities import ( LifecycleCapability, SuperuserCapability, ) -from zhenxun.services.ai.tools.models import ToolOptions -from zhenxun.utils.pydantic_compat import model_copy +from .tool import FunctionTool def toolkit( rules: list[ToolOptions] | ToolOptions | None = None, prefix: str = "", instructions: str | None = None, + auto_register: bool = False, + tags: list[str] | None = None, ): """ 类级别的工具箱装饰器,用于向内部所有 @tool 统一下发配置规则、名称前缀以及工具箱级系统提示词说明。 @@ -31,6 +44,8 @@ def toolkit( rules: 声明式规则集合或单个规则,应用于工具箱内所有工具。 prefix: 工具箱内所有工具的名称前缀,通常以下划线结尾。 instructions: 工具箱级别的系统提示词补充说明,大模型可见。 + auto_register: 是否在加载时自动实例化该类并注册到全局工具箱列表中。 + tags: 工具箱级别的路由标签,大模型路由系统将基于此发现整个工具箱。 返回: Callable: 装饰器函数,接收一个类并返回该类。 @@ -44,7 +59,8 @@ def toolkit( if isinstance(r, ToolOptions): merged_options = merged_options.merge(r) - from zhenxun.services.ai.tools.models import ToolkitConfig + if tags: + merged_options.tags = list(set(merged_options.tags + tags)) base_config = getattr(cls, "_default_config", None) or ToolkitConfig() new_config = model_copy(base_config, deep=True) @@ -62,6 +78,15 @@ def toolkit( if instructions is not None: cls.default_instructions = instructions + if auto_register: + try: + tool_provider_manager.register_toolkit(cls()) + except Exception as e: + logger.error( + f"自动注册工具箱 '{cls.__name__}' 失败" + f"(通常是因为自定义了必填参数的 __init__): {e}" + ) + return cls return decorator @@ -107,10 +132,6 @@ class ToolkitMethodDescriptor: def __get__(self, instance, owner): if instance is None: return self - import types - - from zhenxun.services.ai.tools.core.tool import FunctionTool - bound_func = types.MethodType(self.func, instance) return FunctionTool( func=bound_func, @@ -173,8 +194,6 @@ def tool( tool_name = name or func.__name__ tool_desc = description - import inspect - is_method = False if ( hasattr(func, "__qualname__") @@ -204,9 +223,9 @@ def tool( if not tool_name.startswith(f"{ns}_"): tool_name = f"{ns}_{tool_name}" - from zhenxun.services.ai.tools.core.tool import FunctionTool from zhenxun.services.ai.tools.engine.registry import tool_provider_manager - from zhenxun.services.log import logger + + from .tool import FunctionTool func_tool = FunctionTool( func=func, @@ -216,7 +235,6 @@ def tool( ) if auto_register: tool_provider_manager.register_tool(func_tool) - logger.debug(f"已按命名空间隔离注册了工具(Callable): '{tool_name}'") setattr(func_tool, "__tool_settings__", base_settings) return func_tool @@ -242,8 +260,6 @@ def direct_reply() -> ToolOptions: class DirectReplyCapability(AbstractCapability): async def wrap_tool_execute(self, context, tool_name, arguments, handler): result = await handler(arguments) - from zhenxun.services.ai.tools.models import EndRunResult, ToolResult - if isinstance(result, ToolResult): if getattr(result, "is_error", False): return result diff --git a/zhenxun/services/ai/tools/core/tool.py b/zhenxun/services/ai/tools/core/tool.py index 06b4ace7..f338cc68 100644 --- a/zhenxun/services/ai/tools/core/tool.py +++ b/zhenxun/services/ai/tools/core/tool.py @@ -16,16 +16,16 @@ from zhenxun.services.ai.core.exceptions import ( ) from zhenxun.services.ai.core.models import ToolDefinition 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.logger import log_tool as logger 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 +from .capabilities import InteractiveCapability from .schema import ( _parse_docstring, build_schema_hint, @@ -289,7 +289,6 @@ class BaseTool: async def _core_execution(self, context: RunContext, **kwargs: Any) -> ToolResult: """核心执行流水线 (Core Execution Pipeline)""" - _retries = context.run.tool_retries.get(self.name, 0) if ( self.settings.max_usage_count is not None diff --git a/zhenxun/services/ai/tools/core/toolkit.py b/zhenxun/services/ai/tools/core/toolkit.py index 9d53066f..c4ae69c3 100644 --- a/zhenxun/services/ai/tools/core/toolkit.py +++ b/zhenxun/services/ai/tools/core/toolkit.py @@ -29,6 +29,11 @@ class BaseToolkit: config: ToolkitConfig _default_config: ClassVar[ToolkitConfig] = ToolkitConfig() + @property + def name(self) -> str: + """工具箱的默认名称标识(取类名),主要用于全局注册表的 Hash 与展示""" + return self.__class__.__name__ + def __init__( self, prefix: str | None = None, diff --git a/zhenxun/services/ai/tools/engine/executor.py b/zhenxun/services/ai/tools/engine/executor.py index ec4f3efd..6f6abe9c 100644 --- a/zhenxun/services/ai/tools/engine/executor.py +++ b/zhenxun/services/ai/tools/engine/executor.py @@ -6,7 +6,7 @@ import asyncio from contextlib import asynccontextmanager import inspect import json -from typing import TYPE_CHECKING, Any, cast +from typing import Any, cast import json_repair from nonebot.adapters import Message as PlatformMessage @@ -26,8 +26,9 @@ from zhenxun.services.ai.core.stream_events import ( UserCustomEvent, ) from zhenxun.services.ai.message_builder import MessageBuilder -from zhenxun.services.ai.run.context import RunContext +from zhenxun.services.ai.run.context import RunContext, set_run_context from zhenxun.services.ai.run.di import DependencyInjector +from zhenxun.services.ai.tools.core.tool import BaseTool, register_tool_runner from zhenxun.services.ai.tools.models import ( StateSyncResult, ToolOptions, @@ -35,11 +36,9 @@ from zhenxun.services.ai.tools.models import ( ToolResultChunk, ValidatedToolCall, ) -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_tool as logger -if TYPE_CHECKING: - from zhenxun.services.ai.tools.core.tool import BaseTool - from zhenxun.services.ai.tools.engine.registry import ToolCollection +from .registry import ToolCollection class ToolExecutor: @@ -49,6 +48,7 @@ class ToolExecutor: """ def __init__(self): + """初始化工具执行器。""" pass def _get_combined_capability( @@ -126,8 +126,7 @@ class ToolExecutor: arguments, parsed_successfully = parsed, True logger.debug( "⚒️ 成功修复损坏的工具参数: " - f"{args_str} -> {repaired_str}", - "ToolExecutor", + f"{args_str} -> {repaired_str}" ) except Exception: pass @@ -145,7 +144,7 @@ class ToolExecutor: tool_name: str, executable: Any, event_bus: EventBus | None, - available_tools: "ToolCollection | dict[str, Any] | None" = None, + available_tools: ToolCollection | dict[str, Any] | None = None, ) -> RunContext: """准备/克隆工具调用所使用的隔离 RunContext""" safe_context = ( @@ -163,7 +162,7 @@ class ToolExecutor: async def validate_tool_call( self, tool_call: ToolCallPart, - available_tools: "ToolCollection | dict[str, Any] | None", + available_tools: ToolCollection | dict[str, Any] | None, context: RunContext | None = None, event_bus: EventBus | None = None, ) -> ValidatedToolCall: @@ -213,8 +212,6 @@ class ToolExecutor: async def inner_validate(args_inner): if isinstance(args_inner, dict) and hasattr(executable, "validate_args"): - import inspect - sig = inspect.signature(executable.validate_args) if "context" in sig.parameters: return await executable.validate_args( @@ -247,7 +244,7 @@ class ToolExecutor: async def execute_tool_call( self, validated: ValidatedToolCall, - available_tools: "ToolCollection | dict[str, Any] | None", + available_tools: ToolCollection | dict[str, Any] | None, context: RunContext | None = None, model_name: str | None = None, max_retries: int = 0, @@ -285,8 +282,6 @@ class ToolExecutor: available_tools, ) - from zhenxun.services.ai.run.context import set_run_context - combined_cap = self._get_combined_capability(executable, safe_context) async def inner_handler(args_inner: dict) -> Any: @@ -322,7 +317,7 @@ class ToolExecutor: async def execute_batch( self, tool_calls: list[ToolCallPart], - available_tools: "ToolCollection | dict[str, Any] | None", + available_tools: ToolCollection | dict[str, Any] | None, context: RunContext | None = None, model_name: str | None = None, max_retries: int = 0, @@ -415,11 +410,12 @@ class ToolExecutor: class ToolExecutionPolicy: """ - 工具执行策略 (Strategy Pattern)。 + 工具执行策略。 负责解析工具私有配置与系统全局配置,决定最大重试次数、Fallback 路由目标等流转行为。 """ def __init__(self, tool: BaseTool, global_max_retries: int = 0): + """初始化工具执行策略。""" self.tool = tool self.settings: ToolOptions = getattr(tool, "settings", ToolOptions()) self.metadata: dict[str, Any] = ( @@ -449,6 +445,7 @@ class ToolRunner(ABC): async def run( self, tool: BaseTool, context: RunContext, **kwargs: Any ) -> ToolResult: + """执行工具调用的抽象方法。""" pass @@ -461,6 +458,7 @@ class NativeToolRunner(ToolRunner): async def run( self, tool: BaseTool, context: RunContext, **kwargs: Any ) -> ToolResult: + """运行原生 Python 函数工具并返回结果。""" target_func = tool.get_execute_target() signature_target = tool.get_signature_target() @@ -498,8 +496,8 @@ class NativeToolRunner(ToolRunner): if tool and hasattr(tool, "settings") else False ) - if context.run.event_bus and not is_silent: - await context.run.event_bus.emit( + if not is_silent: + await context.run.emit( ToolStreamChunkEvent( tool_name=tool.name, content=chunk_obj.content, @@ -521,8 +519,7 @@ class NativeToolRunner(ToolRunner): else res ) parts = await MessageBuilder.unimsg_to_llm_parts(uni_msg) - if context and context.run.event_bus: - await context.run.event_bus.emit(UserCustomEvent(display=uni_msg)) + await context.run.emit(UserCustomEvent(display=uni_msg)) final_result = ToolResult(output=parts) else: final_result = ToolResult(output=res) @@ -530,6 +527,4 @@ class NativeToolRunner(ToolRunner): return final_result -from zhenxun.services.ai.tools.core.tool import register_tool_runner - register_tool_runner(NativeToolRunner) diff --git a/zhenxun/services/ai/tools/engine/registry.py b/zhenxun/services/ai/tools/engine/registry.py index 5521db79..e7103a8e 100644 --- a/zhenxun/services/ai/tools/engine/registry.py +++ b/zhenxun/services/ai/tools/engine/registry.py @@ -1,122 +1,76 @@ from __future__ import annotations +from abc import ABC, abstractmethod import asyncio from collections.abc import Callable, Iterable -import inspect from typing import Any, Generic, SupportsIndex, TypeVar, cast, overload from typing_extensions import Self -from nonebot.utils import is_coroutine_callable - from zhenxun.services.ai.core.exceptions import ConfigurationException from zhenxun.services.ai.core.protocols.tool import ( ToolExecutable, ToolProvider, + ToolResolvable, ) from zhenxun.services.ai.run.context import RunContext +from zhenxun.services.ai.tools.core.tool import FunctionTool from zhenxun.services.ai.tools.core.toolkit import BaseToolkit from zhenxun.services.ai.tools.models import Query, ResolvedToolPayload -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_tool as logger +from zhenxun.services.ai.utils.utils import parse_routing_string from zhenxun.utils.utils import infer_plugin_namespace T = TypeVar("T", bound=ToolExecutable) -class ToolCollection(list[T], Generic[T]): - """支持按索引和按名称获取的工具集合 (List + Dict)""" +class ToolCollection(Generic[T]): + """不可变的工具集合""" def __init__(self, iterable: Iterable[T] | None = None): """初始化工具集合。""" - super().__init__(iterable or []) + self._tuple: tuple[T, ...] = tuple(iterable) if iterable else () self._name_cache: dict[str, T] = {} self._build_name_cache() def _build_name_cache(self) -> None: """构建工具名称小写到工具实例的映射缓存。""" - self._name_cache = {} - for tool in self: + self._name_cache.clear() + for tool in self._tuple: name = getattr(tool, "name", None) if name: self._name_cache[name.lower()] = tool + def __iter__(self): + """获取工具元组的迭代器。""" + return iter(self._tuple) + + def __len__(self): + """获取工具集合中的工具数量。""" + return len(self._tuple) + + def __bool__(self): + """检查工具集合是否非空。""" + return bool(self._tuple) + @overload def __getitem__(self, key: SupportsIndex) -> T: ... @overload - def __getitem__(self, key: slice) -> list[T]: ... + def __getitem__(self, key: slice) -> tuple[T, ...]: ... @overload def __getitem__(self, key: str) -> T: ... def __getitem__(self, key: Any) -> Any: + """按索引、切片或名称获取工具。""" if isinstance(key, str): return self._name_cache[key.lower()] - return super().__getitem__(key) - - @overload - def __setitem__(self, key: SupportsIndex, value: T) -> None: ... - - @overload - def __setitem__(self, key: slice, value: Iterable[T]) -> None: ... - - @overload - def __setitem__(self, key: str, value: T) -> None: ... - - def __setitem__(self, key: Any, value: Any) -> None: - if isinstance(key, str): - name = key.lower() - if name in self._name_cache: - old_tool = self._name_cache[name] - try: - idx = super().index(old_tool) - super().__setitem__(idx, value) - except ValueError: - super().append(value) - else: - super().append(value) - self._name_cache[name] = value - else: - super().__setitem__(key, value) - self._name_cache[value.name.lower()] = value + return self._tuple[key] 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()] - try: - idx = super().index(old_tool) - super().__setitem__(idx, object) - except ValueError: - super().append(object) - else: - super().append(object) - 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: - del self._name_cache[name.lower()] - return tool - def filter_by_names(self, names: list[str] | None = None) -> "ToolCollection[T]": """根据名称列表筛选并返回新的工具子集合。""" if names is None: @@ -129,11 +83,6 @@ class ToolCollection(list[T], Generic[T]): ] ) - def clear(self) -> None: - """清空集合及所有名称缓存。""" - super().clear() - self._name_cache.clear() - def keys(self): """获取所有工具名称缓存的键。""" return self._name_cache.keys() @@ -147,155 +96,231 @@ class ToolCollection(list[T], Generic[T]): return self._name_cache.items() -class _StringResolver: - """字符串格式的工具路由解析器。 +class LocalToolProvider(ToolProvider): + """本地工具提供者。统一管理基于 @tool 注册的普通函数或工具箱。""" - 负责解析像 'ns.tool_name'、'ns.*' 或 'ns.#tag' 的语法路由。 - """ + def __init__(self): + """初始化本地工具提供者。""" + self._namespaced_tools: dict[str, list[Any]] = {} - def __init__( - self, name: str, manager: "ToolProviderManager", default_namespace: str - ): - """初始化字符串路由解析器。""" - self.name = name - self.manager = manager - self.default_namespace = default_namespace + def register_tool(self, tool: ToolExecutable, namespace: str): + """向指定命名空间注册单一工具。""" + if namespace not in self._namespaced_tools: + self._namespaced_tools[namespace] = [] + self._namespaced_tools[namespace].append(tool) - 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 = ( - await resolver() if is_coroutine_callable(resolver) else resolver() - ) - return await self.manager._normalize_to_resolver( - resolved, self.default_namespace - ).resolve(context) + def register_toolkit(self, toolkit: Any, namespace: str): + """向指定命名空间注册工具箱。""" + if namespace not in self._namespaced_tools: + self._namespaced_tools[namespace] = [] + self._namespaced_tools[namespace].append(toolkit) - s = self.name + async def initialize(self) -> None: + """初始化本地工具提供者。""" + pass - if "." in s: - ns, target = s.split(".", 1) - else: - ns = self.default_namespace - target = s + async def discover_tools(self) -> dict[str, ToolExecutable]: + """发现并获取所有注册的本地工具。""" + res = {} + for tools in self._namespaced_tools.values(): + for t in tools: + name = getattr(t, "name", getattr(t, "__class__", type).__name__) + res[name] = t + return res - from zhenxun.services.ai.tools.models import Query + async def get_tool_executable( + self, name: str, config: dict[str, Any] + ) -> ToolExecutable | None: + """根据名称获取本地工具实例。""" + tools = await self.discover_tools() + return tools.get(name) - if target == "*": - query = Query(namespace=ns) - elif target.startswith("#"): - tags = [t for t in target.split("#") if t] - query = Query(tags=tags, namespace=ns) - else: - query = Query(name=target, namespace=ns) - - logger.debug(f"🔍 [StringRouter] 语法解析: '{self.name}' -> {query}") - return await _QueryResolver( - query, self.manager, self.default_namespace - ).resolve(context) - - -class _QueryResolver: - """Query 查询对象格式的工具路由解析器。 - - 负责根据 namespace、标签或工具名称检索匹配的工具。 - """ - - 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() + async def query_tools(self, query: Query) -> list[ToolExecutable]: + """根据查询条件检索本地工具。""" + matched_tools = [] + target_namespace = query.namespace namespaces_to_search = [] - target_namespace = self.query.namespace or self.default_namespace - if target_namespace == "global": - namespaces_to_search = list(self.manager._namespaced_tools.keys()) + namespaces_to_search = self.get_all_namespaces() elif target_namespace: namespaces_to_search = [target_namespace] else: - raise ValueError(f"Query 对象必须显式指定 namespace 作用域: {self.query}") + namespaces_to_search = self.get_all_namespaces() for ns in namespaces_to_search: - if ns in self.manager._namespaced_tools: - for tool in self.manager._namespaced_tools[ns]: - if self.query.match(tool): - p = await tool.resolve(context) - if p: - payload.tools.extend(p.tools) - payload.injected_prompts.extend(p.injected_prompts) - payload.toolkits.extend(p.toolkits) + for tool in self.get_tools_by_namespace(ns): + if query.match(tool): + matched_tools.append(tool) - if self.query.name and not payload.tools and not self.query.tags: - specific = await self.manager.resolve_specific_tools([self.query.name]) - for t in specific: - if self.query.match(t): - payload.tools.append(t) + return matched_tools + + def get_tools_by_namespace(self, namespace: str) -> list[Any]: + """获取指定命名空间下的所有工具和工具箱。""" + return self._namespaced_tools.get(namespace, []) + + def get_all_namespaces(self) -> list[str]: + """获取所有已注册工具的命名空间列表。""" + return list(self._namespaced_tools.keys()) + + +class BaseToolResolver(ABC): + """工具载荷解析器抽象基类协议 (责任链模式)""" + + @abstractmethod + def match(self, item: Any) -> bool: + """判断解析器是否匹配当前工具对象。""" + pass + + @abstractmethod + async def resolve( + self, + item: Any, + manager: "ToolProviderManager", + default_namespace: str, + context: RunContext | None = None, + ) -> ResolvedToolPayload: + """解析工具对象并返回其载荷。""" + pass + + +class ProtocolResolver(BaseToolResolver): + """处理已经实现了 ToolResolvable 协议的对象""" + + def match(self, item: Any) -> bool: + """检查是否实现 ToolResolvable 协议。""" + return hasattr(item, "resolve") + + async def resolve( + self, + item: Any, + manager: "ToolProviderManager", + default_namespace: str, + context: RunContext | None = None, + ) -> ResolvedToolPayload: + """直接调用对象的 resolve 方法进行解析。""" + return await item.resolve(context) + + +class StringRoutingResolver(BaseToolResolver): + """处理带语法的路由字符串""" + + def match(self, item: Any) -> bool: + """检查是否为非宏的路由字符串。""" + return isinstance(item, str) + + async def resolve( + self, + item: Any, + manager: "ToolProviderManager", + default_namespace: str, + context: RunContext | None = None, + ) -> ResolvedToolPayload: + s = cast(str, item) + parsed_args = parse_routing_string(s, default_namespace) + query = Query(**parsed_args) + + logger.debug(f"🔍 [StringRouter] 语法解析: '{item}' -> {query}") + return await manager._resolve_single(query, default_namespace, context) + + +class QueryResolver(BaseToolResolver): + """处理 Query 查询对象""" + + def match(self, item: Any) -> bool: + """检查对象是否为 Query 实例。""" + return isinstance(item, Query) + + async def resolve( + self, + item: Any, + manager: "ToolProviderManager", + default_namespace: str, + context: RunContext | None = None, + ) -> ResolvedToolPayload: + query = cast(Query, item) + payload = ResolvedToolPayload() + + if query.namespace == "*": + query.namespace = "global" + elif not query.namespace: + query.namespace = default_namespace + + matched_tools = await manager.local_provider.query_tools(query) + for tool in matched_tools: + if hasattr(tool, "resolve"): + resolvable_tool = cast(ToolResolvable, tool) + p = await resolvable_tool.resolve(context) + if p: + for t in p.tools: + if query.match(t): + payload.tools.append(t) + payload.injected_prompts.extend(p.injected_prompts) + for tk in p.toolkits: + if tk not in payload.toolkits: + payload.toolkits.append(tk) + else: + payload.tools.append(tool) return payload -class _CallableResolver: - """普通 Python 函数/可调用对象格式的工具路由解析器。 +class CallableResolver(BaseToolResolver): + """处理普通 Python 函数""" - 负责将其包装为 FunctionTool 实例。 - """ + def match(self, item: Any) -> bool: + """检查对象是否为可调用对象。""" + return callable(item) - def __init__(self, func: Callable, manager: "ToolProviderManager"): - """初始化可调用对象路由解析器。""" - self.func = func - self.manager = manager - - async def resolve(self, context: RunContext | None = None) -> ResolvedToolPayload: - """将可调用对象转换为函数工具并返回其解析载荷。""" + async def resolve( + self, + item: Any, + manager: "ToolProviderManager", + default_namespace: str, + context: RunContext | None = None, + ) -> ResolvedToolPayload: + func = cast(Callable, item) for candidate in ( - getattr(self.func, "__tool_name__", None), - getattr(self.func, "__name__", None), + getattr(func, "__tool_name__", None), + getattr(func, "__name__", None), ): if candidate: - for ns_tools in self.manager._namespaced_tools.values(): - if t := ns_tools.get(candidate): - return await t.resolve(context) - from zhenxun.services.ai.tools.core.tool import FunctionTool + for ns in manager.local_provider.get_all_namespaces(): + for t in manager.local_provider.get_tools_by_namespace(ns): + if getattr(t, "name", None) == candidate: + return await t.resolve(context) - t = FunctionTool(func=self.func) + t = FunctionTool(func=func) return await t.resolve(context) -class _TypeAdapterResolver: - """基于自定义类型映射注册的工具路由解析器。负责调用对应类型的解析函数。""" +class DictToolResolver(BaseToolResolver): + """处理字典类型的动态/按需工具配置解析器""" - def __init__( + def match(self, item: Any) -> bool: + """检查对象是否为配置字典。""" + return isinstance(item, dict) + + async def resolve( self, item: Any, - resolver_func: Callable, manager: "ToolProviderManager", default_namespace: str, - ): - """初始化类型适配器解析器。""" - self.item = item - self.resolver_func = resolver_func - self.manager = manager - self.default_namespace = default_namespace + context: RunContext | None = None, + ) -> ResolvedToolPayload: + config = cast(dict, item) + name = config.get("name") + if not name: + raise ConfigurationException("工具配置字典必须包含 'name' 字段。") - async def resolve(self, context: RunContext | None = None) -> ResolvedToolPayload: - """执行类型适配器函数并解析返回对应的工具载荷。""" - resolved = ( - await self.resolver_func(self.item) - if is_coroutine_callable(self.resolver_func) - else self.resolver_func(self.item) - ) - return await self.manager._normalize_to_resolver( - resolved, self.default_namespace - ).resolve(context) + for provider in manager._providers: + executable = await provider.get_tool_executable(name, config) + if executable: + return await manager._resolve_single( + executable, default_namespace, context + ) + + raise ConfigurationException(f"没有为 ad-hoc 工具 '{name}' 找到合适的提供者。") class ToolProviderManager: @@ -314,25 +339,23 @@ class ToolProviderManager: if hasattr(self, "_initialized") and self._initialized: return - self._providers: list[ToolProvider] = [] - self._namespaced_tools: dict[str, ToolCollection] = {} - self._resolved_tools: ToolCollection | None = None + self.local_provider = LocalToolProvider() + self._providers: list[ToolProvider] = [self.local_provider] self._init_lock = asyncio.Lock() self._init_promise: asyncio.Task | None = None self._initialized = True - self._macro_resolvers: dict[str, Callable] = {} - self._type_resolvers: dict[type, Callable] = {} + self._resolvers: list[BaseToolResolver] = [ + ProtocolResolver(), + QueryResolver(), + StringRoutingResolver(), + DictToolResolver(), + CallableResolver(), + ] - 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_resolver(self, resolver: BaseToolResolver) -> None: + """注册自定义的工具载荷解析器(插在兜底的 CallableResolver 之前)。""" + self._resolvers.insert(-1, resolver) def register(self, provider: ToolProvider): """注册一个新的 ToolProvider。""" @@ -342,21 +365,26 @@ class ToolProviderManager: def register_tool(self, tool: ToolExecutable): """注册由 @tool 生成的单一工具""" - ns = infer_plugin_namespace() - if ns not in self._namespaced_tools: - self._namespaced_tools[ns] = ToolCollection() - self._namespaced_tools[ns].append(tool) - self._resolved_tools = None + ns = infer_plugin_namespace(default="global") + self.local_provider.register_tool(tool, ns) + tags = getattr(getattr(tool, "settings", None), "tags", []) + tag_str = f" | Tags: {tags}" if tags else "" + tool_name = getattr(tool, "name", "unknown") + logger.debug(f"已注册工具: '{tool_name}' -> Namespace: '{ns}'{tag_str}") def register_toolkit(self, toolkit: Any) -> None: """ 注册一个完整的 Toolkit 实例,使其可通过智能字符串路由(Tag或Name)被动态发现。 """ - ns = infer_plugin_namespace() - if ns not in self._namespaced_tools: - self._namespaced_tools[ns] = ToolCollection() - self._namespaced_tools[ns].append(toolkit) - self._resolved_tools = None + ns = infer_plugin_namespace(default="global") + self.local_provider.register_toolkit(toolkit, ns) + tk_name = getattr(toolkit, "__class__", type).__name__ + + config = getattr(toolkit, "config", None) + shared_options = getattr(config, "shared_options", None) if config else None + tags = getattr(shared_options, "tags", []) if shared_options else [] + tag_str = f" | Tags: {tags}" if tags else "" + logger.debug(f"已注册工具箱: '{tk_name}' -> Namespace: '{ns}'{tag_str}") async def initialize(self) -> None: """懒加载初始化所有已注册的 ToolProvider。""" @@ -375,28 +403,11 @@ class ToolProviderManager: await asyncio.gather(*init_tasks, return_exceptions=True) logger.info("所有工具提供者初始化完成。") - async def discover_tools( - self, - allowed_servers: list[str] | None = None, - excluded_servers: list[str] | None = None, + def _process_discovery_results( + self, results: list[Any], provider_indices: list[int] ) -> dict[str, ToolExecutable]: - """向所有已初始化的 ToolProvider 并发执行工具发现。""" - discover_tasks = [] - provider_indices = [] - for i, provider in enumerate(self._providers): - sig = inspect.signature(provider.discover_tools) - params_to_pass = {} - if "allowed_servers" in sig.parameters: - params_to_pass["allowed_servers"] = allowed_servers - if "excluded_servers" in sig.parameters: - params_to_pass["excluded_servers"] = excluded_servers - - discover_tasks.append(provider.discover_tools(**params_to_pass)) - provider_indices.append(i) - - results = await asyncio.gather(*discover_tasks, return_exceptions=True) - - provider_tools = {} + """处理并合并来自多个工具提供者的并发发现结果""" + provider_tools: dict[str, ToolExecutable] = {} for result_idx, provider_result in enumerate(results): provider = self._providers[provider_indices[result_idx]] provider_name = provider.__class__.__name__ @@ -406,7 +417,7 @@ class ToolProviderManager: f"提供者 '{provider_name}' 发现了 {len(provider_result)} 个工具。" ) for name, executable in provider_result.items(): - if provider_tools.get(name): + if name in provider_tools: logger.warning( f"发现重复的工具名称 '{name}',后发现的将覆盖前者。" ) @@ -417,106 +428,29 @@ class ToolProviderManager: ) return provider_tools - async def _query_engine( - self, - names: list[str] | None = None, - allowed_servers: list[str] | None = None, - excluded_servers: list[str] | None = None, - include_providers: bool = True, - ) -> ToolCollection: - """统一查询引擎:收敛所有本地与云端的工具检索逻辑""" - await self.initialize() - resolved = ToolCollection() + async def discover_tools(self) -> dict[str, ToolExecutable]: + """向所有已初始化的 ToolProvider 并发执行工具发现。""" + discover_tasks = [] + provider_indices = [] + for i, provider in enumerate(self._providers): + discover_tasks.append(provider.discover_tools()) + provider_indices.append(i) - for ns_tools in self._namespaced_tools.values(): - for t in ns_tools: - if names and t.name not in names: - continue - resolved.append(t) + results = await asyncio.gather(*discover_tasks, return_exceptions=True) + provider_tools = self._process_discovery_results(results, provider_indices) + return provider_tools - if not include_providers: - return resolved - - if names: - missing_names = [n for n in names if not resolved.get(n)] - for name in missing_names: - config = {"name": name} - for provider in self._providers: - try: - if executable := await provider.get_tool_executable( - name, config - ): - resolved.append(executable) - break - except Exception as exc: - logger.error( - f"provider '{provider.__class__.__name__}'" - f"解析工具 '{name}' 出错: {exc}" - ) - else: - provider_tools = await self.discover_tools( - allowed_servers, excluded_servers - ) - for t in provider_tools.values(): - resolved.append(t) - - return resolved - - async def get_resolved_tools( - self, - allowed_servers: list[str] | None = None, - 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 - or namespaces is not None + async def _resolve_single( + self, item: Any, default_namespace: str, context: RunContext | None + ) -> ResolvedToolPayload: + """通过责任链解析单个工具意图并返回载荷""" + for r in self._resolvers: + if r.match(item): + return await r.resolve(item, self, default_namespace, context) + raise TypeError( + f"严格协议校验失败: 工具对象 {type(item)} 必须实现 ToolResolvable 协议 " + "(包含 resolve 方法)。如果你想注册普通函数,请使用 @tool 装饰器。" ) - if not has_filters and self._resolved_tools is not None: - return self._resolved_tools - - tools = await self._query_engine( - allowed_servers=allowed_servers, excluded_servers=excluded_servers - ) - - if not has_filters: - self._resolved_tools = tools - 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: - """将任意工具配置或定义包装为标准的多态解析器对象。""" - if hasattr(item, "resolve"): - return item - - if isinstance(item, Query): - return _QueryResolver(item, self, default_ns) - if isinstance(item, str): - return _StringResolver(item, self, default_ns) - if type(item) in self._type_resolvers: - return _TypeAdapterResolver( - item, self._type_resolvers[type(item)], self, default_ns - ) - if callable(item): - return _CallableResolver(item, self) - - if not hasattr(item, "resolve"): - raise TypeError( - f"严格协议校验失败: 工具对象 {type(item)} 必须实现 ToolResolvable 协议 " - "(包含 resolve 方法)。如果你想注册普通函数,请使用 @tool 装饰器。" - ) - return item async def resolve_tools( self, @@ -547,14 +481,7 @@ class ToolProviderManager: _flatten(tool_definitions) - resolvers = [self._normalize_to_resolver(t, namespace) for t in defs] - - for i, r in enumerate(resolvers): - if asyncio.iscoroutine(r): - r = await r - resolvers[i] = self._normalize_to_resolver(r, namespace) - - tasks = [r.resolve(context) for r in resolvers] + tasks = [self._resolve_single(t, namespace, context) for t in defs] payloads = await asyncio.gather(*tasks, return_exceptions=False) final_payload = ResolvedToolPayload() @@ -568,34 +495,10 @@ class ToolProviderManager: for t in p.tools: if not getattr(t, "parent_toolkit", None): t.parent_toolkit = global_toolkit - final_payload.tools.append(t) - final_payload.injected_prompts.extend(p.injected_prompts) - final_payload.toolkits.extend(p.toolkits) + final_payload.merge(p) - final_payload.tools = ToolCollection(final_payload.tools) return final_payload tool_provider_manager = ToolProviderManager() - - -async def _dict_ad_hoc_resolver(config: dict): - """针对字典类型的 ad-hoc 工具配置的类型解析器。""" - name = config.get("name") - if not name: - raise ConfigurationException( - "工具配置字典必须包含 'name' 字段。", - ) - - for provider in tool_provider_manager._providers: - executable = await provider.get_tool_executable(name, config) - if executable: - return executable - - raise ConfigurationException( - f"没有为 ad-hoc 工具 '{name}' 找到合适的提供者。", - ) - - -tool_provider_manager.register_type_resolver(dict, _dict_ad_hoc_resolver) diff --git a/zhenxun/services/ai/tools/models.py b/zhenxun/services/ai/tools/models.py index 299b436b..37b576d8 100644 --- a/zhenxun/services/ai/tools/models.py +++ b/zhenxun/services/ai/tools/models.py @@ -3,12 +3,14 @@ """ from dataclasses import dataclass, field +import fnmatch from typing import TYPE_CHECKING, Any, Literal from pydantic import BaseModel, ConfigDict, Field from zhenxun.services.ai.core.messages import ToolCallPart, UsageInfo from zhenxun.services.ai.run.context import RunContext +from zhenxun.services.ai.utils.logger import log_tool as logger from zhenxun.utils.pydantic_compat import model_dump, model_validate if TYPE_CHECKING: @@ -210,12 +212,12 @@ class ToolOverride(BaseModel): async def resolve(self, context: RunContext | None = None) -> "ResolvedToolPayload": from zhenxun.services.ai.tools.engine.registry import tool_provider_manager - from zhenxun.services.ai.tools.models import ResolvedToolPayload - from zhenxun.services.log import logger - found_tools = await tool_provider_manager.resolve_specific_tools([self.name]) - if found_tools: - base_tool = found_tools[0] + payload = await tool_provider_manager.resolve_tools( + [self.name], context=context + ) + if payload and payload.tools: + base_tool = payload.tools[0] if hasattr(base_tool, "clone_with_options"): cloned_tool = base_tool.clone_with_options(self) if hasattr(cloned_tool, "resolve"): @@ -231,17 +233,8 @@ class ToolOverride(BaseModel): return ResolvedToolPayload() -class GlobalToolFilter(BaseModel): - """全局宏观工具过滤器""" - - allowed_servers: list[str] | None = None - """仅允许的服务端名称列表""" - excluded_servers: list[str] | None = None - """需要排除的服务端名称列表""" - - class ValidatedToolCall(BaseModel): - """工具调用验证结果载体(解耦验证与执行)""" + """工具调用验证结果载体""" model_config = ConfigDict(arbitrary_types_allowed=True) @@ -265,22 +258,68 @@ class Query(BaseModel): 用于在 Agent 中精确或批量筛选加载特定命名空间、特定标签的工具。 """ - name: str | None = Field(default=None) - """如果提供,则必须与工具的名称完全一致。""" + name: str | list[str] | None = Field(default=None) + """如果提供,则工具的最终解析名称必须等于该字符串或在列表中。""" + toolkit: str | list[str] | None = Field(default=None) + """如果提供,则工具必须属于指定的 Toolkit (支持字符串或列表)。""" tags: list[str] | None = Field(default=None) """如果提供,则工具必须包含这里列出的所有标签 (交集/AND匹配)。""" + exclude_tags: list[str] | None = Field(default=None) + """如果提供,则工具不能包含这里列出的任何标签 (排斥过滤)。""" namespace: str | None = Field(default=None) - """必填(由系统补充或显式声明)。限制搜索的插件命名空间,'global' 将跨全插件搜索。""" + """限制搜索的插件命名空间。如果不指定,将自动推导为调用者所在的插件; + 'global' 将跨全插件搜索。""" metadata_filter: dict[str, Any] | None = Field(default=None) """如果提供,则工具的 metadata 必须包含这里列出的所有键值对。""" def match(self, tool: "BaseTool") -> bool: """判断某个工具或工具箱是否符合当前 Query 的筛选条件""" + + def _match_pattern(val: str, pattern: str | list[str]) -> bool: + patterns = [pattern] if isinstance(pattern, str) else pattern + for p in patterns: + if "*" in p or "?" in p: + if fnmatch.fnmatch(val, p): + return True + elif val == p: + return True + return False + + is_toolkit_itself = hasattr(tool, "get_tools") and not hasattr(tool, "execute") + + if self.toolkit: + if is_toolkit_itself: + tk_name = getattr( + tool, "name", getattr(tool, "__class__", type).__name__ + ) + else: + parent_tk = getattr(tool, "parent_toolkit", None) + tk_name = ( + getattr( + parent_tk, + "name", + getattr(parent_tk, "__class__", type).__name__, + ) + if parent_tk + else None + ) + + if not tk_name: + return False + + if not _match_pattern(tk_name, self.toolkit): + return False + + if is_toolkit_itself and (self.name or self.tags or self.exclude_tags): + return True + if self.name: tool_name = getattr(tool, "name", getattr(tool, "__class__", type).__name__) - if tool_name != self.name: + if not _match_pattern(tool_name, self.name): return False - if self.tags: + + tool_tags = [] + if self.tags or self.exclude_tags: tool_config = getattr(tool, "config", None) if ( tool_config @@ -291,9 +330,15 @@ class Query(BaseModel): else: tool_settings = getattr(tool, "settings", None) tool_tags = getattr(tool_settings, "tags", []) if tool_settings else [] + + if self.tags: if not all(tag in tool_tags for tag in self.tags): return False + if self.exclude_tags: + if any(tag in tool_tags for tag in self.exclude_tags): + return False + if self.metadata_filter: tool_settings = getattr(tool, "settings", None) tool_metadata = ( @@ -315,9 +360,15 @@ class ResolvedToolPayload: injected_prompts: list[str] = field(default_factory=list) toolkits: list[Any] = field(default_factory=list) + def merge(self, other: "ResolvedToolPayload") -> "ResolvedToolPayload": + """将另一个解析载荷合并到当前载荷中""" + self.tools.extend(other.tools) + self.injected_prompts.extend(other.injected_prompts) + self.toolkits.extend(other.toolkits) + return self + __all__ = [ - "GlobalToolFilter", "Query", "ResolvedToolPayload", "ToolOptions", diff --git a/zhenxun/services/ai/tools/providers/builtin/hitl.py b/zhenxun/services/ai/tools/providers/builtin/hitl.py index 63c2b83c..0442c74a 100644 --- a/zhenxun/services/ai/tools/providers/builtin/hitl.py +++ b/zhenxun/services/ai/tools/providers/builtin/hitl.py @@ -4,8 +4,8 @@ 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 -from zhenxun.services.log import logger +from zhenxun.services.ai.tools.models import ResolvedToolPayload, ToolResult +from zhenxun.services.ai.utils.logger import log_tool as logger class HITLToolkit(BaseToolkit): @@ -23,6 +23,12 @@ class HITLToolkit(BaseToolkit): "用户回答后,你将收到答案并可以继续任务。" ) + async def resolve(self, context: RunContext | None = None) -> ResolvedToolPayload: + """智能环境感知:如果处于无真实用户交互的后台环境(如定时任务),自动隐身以节省Token""" + if context is None or context.get_event() is None: + return ResolvedToolPayload() + return await super().resolve(context) + @tool( name="ask_user_for_help", description="向当前对话的用户提出问题以获取信息或指导。当你无法独立完成任务时,请调用此工具。", @@ -47,8 +53,8 @@ class HITLToolkit(BaseToolkit): raise e logger.info(f"收到用户求助回复: {user_reply}") - if context and context.run.event_bus: - await context.run.event_bus.emit( + if context: + await context.run.emit( ToolStreamChunkEvent( tool_name=context.call.tool_name, content="🗣️ 已收到用户的回复" ) diff --git a/zhenxun/services/ai/tools/providers/builtin/memory.py b/zhenxun/services/ai/tools/providers/builtin/memory.py index 5d2702f3..c3118e65 100644 --- a/zhenxun/services/ai/tools/providers/builtin/memory.py +++ b/zhenxun/services/ai/tools/providers/builtin/memory.py @@ -10,7 +10,7 @@ from zhenxun.services.ai.tools.core.decorators import tool from zhenxun.services.ai.tools.core.tool import BaseTool from zhenxun.services.ai.tools.core.toolkit import BaseToolkit from zhenxun.services.ai.tools.models import ToolOptions, ToolResult -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_tool as logger class MemoryManagementToolkit(BaseToolkit): diff --git a/zhenxun/services/ai/tools/providers/builtin/rest_api.py b/zhenxun/services/ai/tools/providers/builtin/rest_api.py index b46da774..1b5f524b 100644 --- a/zhenxun/services/ai/tools/providers/builtin/rest_api.py +++ b/zhenxun/services/ai/tools/providers/builtin/rest_api.py @@ -8,7 +8,7 @@ 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 -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_tool as logger class RestApiToolkit(BaseToolkit): @@ -101,15 +101,15 @@ class RestApiToolkit(BaseToolkit): is_error = response.status_code >= 400 result = ToolResult(output=result_dict) - if context and context.run.event_bus: + if context: msg = ( f"🌐 已调用 API: {url}" if not is_error else f"❌ API 调用失败 (Status: {response.status_code})" ) - await context.run.event_bus.emit( + await context.run.emit( ToolStreamChunkEvent( - tool_name=context.call.tool_name if context else "make_request", + tool_name=context.call.tool_name, content=msg, ) ) @@ -120,10 +120,10 @@ class RestApiToolkit(BaseToolkit): except Exception as e: logger.error(f"RestApiToolkit 请求失败: {e}") - if context and context.run.event_bus: - await context.run.event_bus.emit( + if context: + await context.run.emit( ToolStreamChunkEvent( - tool_name=context.call.tool_name if context else "make_request", + tool_name=context.call.tool_name, content="❌ API 网络请求发生框架级错误", ) ) diff --git a/zhenxun/services/ai/tools/providers/builtin/sandbox.py b/zhenxun/services/ai/tools/providers/builtin/sandbox.py index c2169e07..8c4fc2c1 100644 --- a/zhenxun/services/ai/tools/providers/builtin/sandbox.py +++ b/zhenxun/services/ai/tools/providers/builtin/sandbox.py @@ -10,7 +10,7 @@ from zhenxun.services.ai.sandbox.models import ( 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.services.ai.utils.logger import log_tool as logger from zhenxun.utils.pydantic_compat import model_copy @@ -103,8 +103,8 @@ class SandboxToolkit(BaseToolkit): bp = model_copy(self.blueprint, deep=True) - if context.run.event_bus: - await context.run.event_bus.emit( + if context: + await context.run.emit( ToolStreamChunkEvent( tool_name="Sandbox", content="正在分析代码依赖并分配沙箱环境..." ) @@ -122,8 +122,8 @@ class SandboxToolkit(BaseToolkit): ) context.session.shared_state[state_key] = code_executor - if context.run.event_bus: - await context.run.event_bus.emit( + if context: + await context.run.emit( ToolStreamChunkEvent( tool_name="Sandbox", content=f"沙箱已就绪,正在后台执行 {language} 代码...", @@ -225,8 +225,8 @@ class SandboxToolkit(BaseToolkit): result = ToolResult( output=final_output if len(final_output) > 1 else final_output_text ) - if len(image_bytes_list) > 0 and context and context.run.event_bus: - await context.run.event_bus.emit(UserCustomEvent(display=final_output)) + if len(image_bytes_list) > 0 and context: + await context.run.emit(UserCustomEvent(display=final_output)) return result @tool( @@ -249,8 +249,8 @@ class SandboxToolkit(BaseToolkit): ) -> ToolResult: session_id = self.sandbox_session_id or context.session_id or "default" - if context.run.event_bus: - await context.run.event_bus.emit( + if context: + await context.run.emit( ToolStreamChunkEvent( tool_name="Sandbox", content=f"正在虚拟终端执行命令: {command} ..." ) @@ -319,8 +319,8 @@ class SandboxToolkit(BaseToolkit): await asyncio.sleep(1.5) output = await interactive_session.read_output(timeout=5) - if context and context.run.event_bus: - await context.run.event_bus.emit( + if context: + await context.run.emit( ToolStreamChunkEvent( tool_name=context.call.tool_name, content="⌨️ 已向后台进程发送输入" ) @@ -357,8 +357,8 @@ class SandboxToolkit(BaseToolkit): await interactive_session.interrupt() await asyncio.sleep(1) output = await interactive_session.read_output() - if context and context.run.event_bus: - await context.run.event_bus.emit( + if context: + await context.run.emit( ToolStreamChunkEvent( tool_name=context.call.tool_name, content="🛑 已强制中断后台进程" ) diff --git a/zhenxun/services/ai/tools/providers/mcp/provider.py b/zhenxun/services/ai/tools/providers/mcp/provider.py index 1a9a9f30..ec01ad83 100644 --- a/zhenxun/services/ai/tools/providers/mcp/provider.py +++ b/zhenxun/services/ai/tools/providers/mcp/provider.py @@ -16,10 +16,11 @@ from zhenxun.services.ai.core.protocols.tool import ( ) from zhenxun.services.ai.sandbox.models import SandboxBlueprint from zhenxun.services.ai.tools.models import ResolvedToolPayload -from zhenxun.services.ai.tools.providers.mcp.toolkit import MCPToolkit -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_tool as logger from zhenxun.utils.pydantic_compat import model_dump, model_validate, model_validator +from .toolkit import MCPToolkit + MCP_PATH = DATA_PATH / "ai" / "mcp.json" @@ -30,6 +31,8 @@ class MCPServerConfig(BaseModel): """传输协议类型:stdio / sse / streamable-http / sandbox_proxy""" url: str | None = Field(default=None) """远端地址(用于 sse 或 streamable-http)""" + headers: dict[str, str] | None = Field(default=None) + """HTTP 请求头(用于 sse 或 streamable-http 的接口鉴权)""" timeout: int = Field(default=30) """请求超时时间(秒)""" command: str | None = None @@ -51,6 +54,34 @@ class MCPServerConfig(BaseModel): sandbox_blueprint: SandboxBlueprint | None = Field(default=None) """沙箱环境装配配置(用于 sandbox_proxy 自动处理依赖)""" + @model_validator(mode="before") + @classmethod + def _normalize_config(cls, values: Any) -> Any: + if not isinstance(values, dict): + return values + + if "type" in values and "transport" not in values: + values["transport"] = values.pop("type") + + transport_val = values.get("transport") + if isinstance(transport_val, str) and transport_val.lower() == "streamablehttp": + values["transport"] = "streamable-http" + + headers = values.get("headers") + if isinstance(headers, dict): + new_headers = {} + for k, v in headers.items(): + if k.lower() == "authorization" and isinstance(v, str): + if not any( + v.startswith(prefix) + for prefix in ("Bearer ", "Basic ", "Digest ") + ): + v = f"Bearer {v}" + new_headers[k] = v + values["headers"] = new_headers + + return values + class MCPToolsConfig(BaseModel): mcpServers: dict[str, MCPServerConfig] = Field(default_factory=dict) @@ -193,6 +224,7 @@ class GlobalMCPProvider(ToolProvider): command=conf.command, args=conf.args, url=conf.url, + headers=conf.headers, env=conf.env, cwd=conf.cwd, install_command=conf.install_command, @@ -339,6 +371,7 @@ class MCPSource(BaseModel): command=self.config.command, args=self.config.args, url=self.config.url, + headers=self.config.headers, env=self.config.env, cwd=self.config.cwd, install_command=self.config.install_command, diff --git a/zhenxun/services/ai/tools/providers/mcp/toolkit.py b/zhenxun/services/ai/tools/providers/mcp/toolkit.py index 2d557673..36205a1c 100644 --- a/zhenxun/services/ai/tools/providers/mcp/toolkit.py +++ b/zhenxun/services/ai/tools/providers/mcp/toolkit.py @@ -23,7 +23,7 @@ from zhenxun.services.ai.sandbox.models import SandboxBlueprint from zhenxun.services.ai.tools.core.tool import BaseTool from zhenxun.services.ai.tools.core.toolkit import BaseToolkit from zhenxun.services.ai.tools.models import ToolkitConfig, ToolResult -from zhenxun.services.log import logger +from zhenxun.services.ai.utils.logger import log_tool as logger from zhenxun.utils.lifespan import LifespanManager from zhenxun.utils.pydantic_compat import model_dump @@ -270,10 +270,10 @@ class MCPToolkit(BaseToolkit): command: str | None = None, args: list[str] | None = None, url: str | None = None, + headers: dict[str, str] | None = None, env: dict | None = None, cwd: str | None = None, install_command: str | None = None, - isolation: Literal["shared", "per_session"] = "shared", timeout: int = 30, admin_level: int = 0, header_provider: Callable[[RunContext], dict[str, str]] | None = None, @@ -292,10 +292,10 @@ class MCPToolkit(BaseToolkit): command: 用于 stdio 或 sandbox_proxy 模式启动服务器的可执行命令。 args: 启动服务器时附加的命令行参数。 url: 用于 sse 或 streamable-http 模式的服务器连接 URL。 + headers: 静态配置的 HTTP 请求头。 env: 启动服务器时的进程环境变量。 cwd: 启动服务器时的进程工作目录。 install_command: 首次启动前执行的环境热装配/依赖安装命令。 - isolation: 环境隔离模式,可选 "shared" (共享模式) 或 "per_session" (每个会话独立)。 timeout: 初始化和请求网络接口时的超时秒数,默认 30。 admin_level: 调用工具所需的管理权限等级,默认 0 (无限制)。 header_provider: 动态生成 HTTP 头部信息的工厂函数。 @@ -312,6 +312,7 @@ class MCPToolkit(BaseToolkit): self.command = command self.args = args or [] self.url = url + self.headers = headers or {} self.env = env or {} self.cwd = cwd self.install_command = install_command @@ -587,8 +588,13 @@ class MCPToolkit(BaseToolkit): self._ready_event.clear() self._init_exception = None - dynamic_headers = {} + dynamic_headers = self.headers.copy() + if self.header_provider and context: + dynamic_headers.update(self.header_provider(context)) + dynamic_env = self.env.copy() + if self.env_provider and context: + dynamic_env.update(self.env_provider(context)) self._shared_task = asyncio.create_task( self._spawn_session_task(dynamic_headers, dynamic_env) diff --git a/zhenxun/services/ai/tools/providers/skills/capabilities.py b/zhenxun/services/ai/tools/providers/skills/capabilities.py index 4550477d..e8fc4a79 100644 --- a/zhenxun/services/ai/tools/providers/skills/capabilities.py +++ b/zhenxun/services/ai/tools/providers/skills/capabilities.py @@ -4,11 +4,12 @@ from typing import Any from zhenxun.services.ai.capabilities import AbstractCapability 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 from zhenxun.utils.utils import infer_plugin_namespace +from .manager import skill_manager +from .models import Skill, SkillSource +from .toolkit import SkillMetaToolkit + class SkillCapability(AbstractCapability): """技能库挂载能力组件""" diff --git a/zhenxun/services/ai/tools/providers/skills/manager.py b/zhenxun/services/ai/tools/providers/skills/manager.py index ce1a7a46..f1c84400 100644 --- a/zhenxun/services/ai/tools/providers/skills/manager.py +++ b/zhenxun/services/ai/tools/providers/skills/manager.py @@ -11,7 +11,11 @@ import aiofiles import yaml from zhenxun.configs.path_config import DATA_PATH -from zhenxun.services.ai.tools.providers.skills.models import ( +from zhenxun.services.ai.utils.logger import log_tool as logger +from zhenxun.utils.pydantic_compat import model_dump, model_validate +from zhenxun.utils.utils import infer_plugin_namespace + +from .models import ( INSTRUCTIONS, METADATA, RESOURCES, @@ -20,9 +24,6 @@ from zhenxun.services.ai.tools.providers.skills.models import ( SkillEnvConfig, SkillFrontmatter, ) -from zhenxun.services.log import logger -from zhenxun.utils.pydantic_compat import model_dump, model_validate -from zhenxun.utils.utils import infer_plugin_namespace class SkillConfigManager: diff --git a/zhenxun/services/ai/tools/providers/skills/toolkit.py b/zhenxun/services/ai/tools/providers/skills/toolkit.py index 877f7ee8..ec81b42d 100644 --- a/zhenxun/services/ai/tools/providers/skills/toolkit.py +++ b/zhenxun/services/ai/tools/providers/skills/toolkit.py @@ -10,14 +10,15 @@ from zhenxun.services.ai.sandbox.protocols import ( 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 ( +from zhenxun.services.ai.utils.logger import log_tool as logger +from zhenxun.utils.pydantic_compat import model_copy +from zhenxun.utils.utils import infer_plugin_namespace + +from .manager import ( skill_env_manager, skill_manager, ) -from zhenxun.services.ai.tools.providers.skills.models import Skill -from zhenxun.services.log import logger -from zhenxun.utils.pydantic_compat import model_copy -from zhenxun.utils.utils import infer_plugin_namespace +from .models import Skill class SkillSandboxExecutionMixin: diff --git a/zhenxun/services/ai/utils/logger.py b/zhenxun/services/ai/utils/logger.py new file mode 100644 index 00000000..376b28b8 --- /dev/null +++ b/zhenxun/services/ai/utils/logger.py @@ -0,0 +1,65 @@ +""" +AI 模块专属日志代理门面 +""" + +from typing import Any + +from zhenxun.services.log import logger as global_logger + + +class AILoggerProxy: + def __init__(self, module_name: str, emoji: str = ""): + self.module_name = module_name + self.emoji = emoji + self._cmd = f"AI|{self.module_name}" + + def _format(self, msg: str) -> str: + """自动在消息开头追加 Emoji(如果消息本身不包含的话)""" + if self.emoji and not str(msg).lstrip().startswith(self.emoji): + return f"{self.emoji} {msg}" + return msg + + def info(self, info: str, command: str | None = None, **kwargs: Any): + cmd = command or self._cmd + global_logger.info(self._format(info), command=cmd, **kwargs) + + def debug(self, info: str, command: str | None = None, **kwargs: Any): + cmd = command or self._cmd + global_logger.debug(self._format(info), command=cmd, **kwargs) + + def warning(self, info: str, command: str | None = None, **kwargs: Any): + cmd = command or self._cmd + global_logger.warning(self._format(info), command=cmd, **kwargs) + + def error(self, info: str, command: str | None = None, **kwargs: Any): + cmd = command or self._cmd + global_logger.error(self._format(info), command=cmd, **kwargs) + + def success( + self, + info: str, + command: str | None = None, + param: dict[str, Any] | None = None, + result: str = "", + ): + cmd = command or self._cmd + global_logger.success( + self._format(info), command=cmd, param=param, result=result + ) + + def trace(self, info: str, command: str | None = None, **kwargs: Any): + cmd = command or self._cmd + global_logger.trace(self._format(info), command=cmd, **kwargs) + + +log_llm = AILoggerProxy("LLM") +log_agent = AILoggerProxy("Agent") +log_team = AILoggerProxy("Team") +log_tool = AILoggerProxy("Tool") +log_sandbox = AILoggerProxy("Sandbox") +log_memory = AILoggerProxy("Memory") +log_rag = AILoggerProxy("RAG") +log_flow = AILoggerProxy("Flow") +log_core = AILoggerProxy("Core") +log_knowledge = AILoggerProxy("Knowledge") +log_capability = AILoggerProxy("Capability") diff --git a/zhenxun/services/ai/utils/runtime.py b/zhenxun/services/ai/utils/runtime.py index bd60b691..b573e22b 100644 --- a/zhenxun/services/ai/utils/runtime.py +++ b/zhenxun/services/ai/utils/runtime.py @@ -88,7 +88,8 @@ class ContextUtils: @staticmethod def generate_session_meta( bot: Bot, - event: Event, + event: Event | None = None, + deps: Any | None = None, scope_builder: ScopeBuilder | None = None, prefix: str = "", namespace: str | None = None, @@ -104,7 +105,8 @@ class ContextUtils: if scope_builder is None: scope_builder = Isolation.AGENT_USER() - deps = NoneBotDeps(bot=bot, event=event) + if deps is None: + deps = NoneBotDeps(bot=bot, event=event) if event else NoneBotDeps(bot=bot) selector = scope_builder.resolve( deps=deps, prefix=prefix, diff --git a/zhenxun/services/ai/utils/scope.py b/zhenxun/services/ai/utils/scope.py index 57e7b253..be564316 100644 --- a/zhenxun/services/ai/utils/scope.py +++ b/zhenxun/services/ai/utils/scope.py @@ -5,6 +5,8 @@ from typing_extensions import Self from pydantic import BaseModel, Field +from zhenxun.utils.utils import infer_plugin_namespace + class ScopeSelector(BaseModel): """领域驱动:统一的作用域与实体资源选择器""" @@ -196,7 +198,6 @@ class ScopeBuilder: ) -> ScopeSelector: """解析当前上下文依赖并生成作用域选择器实例。""" from zhenxun.services.ai.utils.runtime import ContextUtils - from zhenxun.utils.utils import infer_plugin_namespace selector = ScopeSelector(base_prefix=prefix) if "platform" in self._dims: diff --git a/zhenxun/services/ai/utils/utils.py b/zhenxun/services/ai/utils/utils.py index e33d860a..003e7573 100644 --- a/zhenxun/services/ai/utils/utils.py +++ b/zhenxun/services/ai/utils/utils.py @@ -3,6 +3,7 @@ from collections.abc import Callable import functools import inspect import threading +from typing import Any def wrap_to_async(func: Callable) -> Callable: @@ -55,3 +56,37 @@ def wrap_to_async(func: Callable) -> Callable: setattr(async_wrapper, "__is_async_wrapper__", True) setattr(async_wrapper, "_is_coroutine", True) return async_wrapper + + +def parse_routing_string(route_str: str, default_namespace: str) -> dict[str, Any]: + """ + 统一的声明式路由字符串解析函数。 + + 支持格式: + - "namespace.target" + - "*.target" (等价于 global.target) + - "target" (使用 default_namespace) + + target 支持: + - "*" (全选) + - "#tag1#tag2" (按标签筛选) + - "name" (按精确名称筛选) + """ + if "." in route_str: + ns, target = route_str.split(".", 1) + if ns == "*": + ns = "global" + else: + ns = default_namespace + target = route_str + + result: dict[str, Any] = {"namespace": ns} + + if target == "*": + pass + elif target.startswith("#"): + result["tags"] = [t for t in target.split("#") if t] + else: + result["name"] = target + + return result diff --git a/zhenxun/services/scheduler/engine.py b/zhenxun/services/scheduler/engine.py index a9073116..b43e3265 100644 --- a/zhenxun/services/scheduler/engine.py +++ b/zhenxun/services/scheduler/engine.py @@ -7,12 +7,12 @@ import asyncio from collections.abc import Callable -from datetime import datetime from functools import partial import random import time +from typing import cast -import nonebot +from apscheduler.jobstores.base import JobLookupError from nonebot.adapters import Bot from nonebot.dependencies import Dependent from nonebot.exception import FinishedException, PausedException, SkippedException @@ -30,45 +30,15 @@ from zhenxun.utils.decorator.retry import Retry from zhenxun.utils.platform import PlatformUtils from zhenxun.utils.pydantic_compat import parse_as +from .registry import scheduler_registry from .repository import ScheduleRepository -from .types import ExecutionPolicy, ScheduleContext +from .types import ExecutionPolicy, ScheduleContext, TargetType JOB_PREFIX = "zhenxun_schedule_" SCHEDULE_CONCURRENCY_KEY = "all_groups_concurrency_limit" _LAST_PRESSURE_SKIP = 0.0 -def _resolve_scheduler_bot(bot_id: str | None, log_target: str) -> Bot | None: - if bot_id: - try: - return nonebot.get_bot(bot_id) - except KeyError: - logger.warning(f"{log_target} 需要的 Bot {bot_id} 不在线,本次执行跳过。") - return None - - bots = list(nonebot.get_bots().values()) - if not bots: - logger.warning(f"{log_target} 当前没有可用 Bot,本次执行跳过。") - return None - if len(bots) == 1: - return bots[0] - - qq_client_bots = [ - bot for bot in bots if PlatformUtils.get_platform_scope(bot) == "qq_client" - ] - if len(qq_client_bots) == 1: - bot = qq_client_bots[0] - logger.warning( - f"{log_target} 未指定 Bot,多 Bot 在线,自动选择 OneBot {bot.self_id}。" - ) - return bot - - logger.warning( - f"{log_target} 未指定 Bot 且多 Bot 在线,无法安全选择," "本次执行跳过。" - ) - return None - - class APSchedulerAdapter: """封装对 APScheduler 的操作""" @@ -97,8 +67,10 @@ class APSchedulerAdapter: try: scheduler.remove_job(job_id) - except Exception: + except JobLookupError: pass + except Exception as e: + logger.debug(f"尝试覆盖 APScheduler 任务 {job_id} 时发生异常: {e}") if not isinstance(schedule.trigger_config, dict): logger.error( @@ -132,7 +104,7 @@ class APSchedulerAdapter: job_params["coalesce"] = False scheduler.add_job( - _execute_job, + _execute_persistent_job, trigger=schedule.trigger_type, **job_params, **trigger_params, @@ -154,8 +126,10 @@ class APSchedulerAdapter: try: scheduler.remove_job(job_id) logger.debug(f"已从APScheduler中移除任务: {job_id}") - except Exception: - pass + except JobLookupError: + logger.debug(f"APScheduler 中未找到任务: {job_id},无需移除。") + except Exception as e: + logger.debug(f"从 APScheduler 移除任务 {job_id} 时发生异常: {e}") @staticmethod def pause_job(schedule_id: int): @@ -168,8 +142,10 @@ class APSchedulerAdapter: job_id = APSchedulerAdapter._get_job_id(schedule_id) try: scheduler.pause_job(job_id) - except Exception: - pass + except JobLookupError: + logger.debug(f"APScheduler 中未找到任务: {job_id},无法暂停。") + except Exception as e: + logger.debug(f"暂停 APScheduler 任务 {job_id} 时发生异常: {e}") @staticmethod def resume_job(schedule_id: int): @@ -182,8 +158,8 @@ class APSchedulerAdapter: job_id = APSchedulerAdapter._get_job_id(schedule_id) try: scheduler.resume_job(job_id) - except Exception: - import asyncio + except JobLookupError: + logger.debug(f"APScheduler 中未找到任务: {job_id},尝试重新注册...") async def _re_add_job(): schedule = await ScheduleRepository.get_by_id(schedule_id) @@ -191,6 +167,8 @@ class APSchedulerAdapter: APSchedulerAdapter.add_or_reschedule_job(schedule) asyncio.create_task(_re_add_job()) # noqa: RUF006 + except Exception as e: + logger.debug(f"恢复 APScheduler 任务 {job_id} 时发生异常: {e}") @staticmethod def get_job_status(schedule_id: int) -> dict: @@ -236,38 +214,68 @@ class APSchedulerAdapter: return scheduler.add_job( - _execute_job, + _execute_ephemeral_job, trigger=trigger_type, id=job_id, misfire_grace_time=60, - args=[None], - kwargs={"context_override": context}, + kwargs={"context": context}, **trigger_config, ) logger.debug(f"已添加新的临时APScheduler任务: {job_id}") +async def _run_dependent_task( + func: Callable, injected_params: dict, bot: Bot, state: T_State +): + """将普通函数包装为 NoneBot Dependent 并执行的内部辅助方法""" + + async def wrapper(bot: Bot): + return await func(bot=bot, **injected_params) + + dependent = Dependent.parse(call=wrapper, allow_types=Matcher.HANDLER_PARAM_TYPES) + return await dependent(bot=bot, state=state) + + async def _execute_single_job_instance( - schedule: ScheduledJob, bot, group_id: str | None = None + schedule: ScheduledJob, bot, target_id: str | None = None ): """ 负责执行一个具体目标的任务实例。 """ - from .manager import scheduler_manager plugin_name = schedule.plugin_name - if group_id is None and schedule.target_type == "GROUP": - group_id = schedule.target_identifier - task_meta = scheduler_manager._registered_tasks.get(plugin_name) + actual_group_id = None + actual_user_id = None + + if schedule.target_type in ( + TargetType.GROUP, + TargetType.TAG, + TargetType.ALL_GROUPS, + ): + actual_group_id = target_id or ( + schedule.target_identifier + if schedule.target_type == TargetType.GROUP + else None + ) + elif schedule.target_type == TargetType.USER: + actual_user_id = target_id or schedule.target_identifier + + target_log = actual_group_id or actual_user_id + + task_meta = scheduler_registry.tasks.get(plugin_name) if not task_meta: logger.error(f"无法执行任务:插件 '{plugin_name}' 在执行期间变得不可用。") return - is_blocked = await CommonUtils.task_is_block(bot, plugin_name, group_id) + is_blocked = await CommonUtils.task_is_block(bot, plugin_name, actual_group_id) if is_blocked: - target_desc = f"群 {group_id}" if group_id else "全局" + target_desc = ( + f"群 {actual_group_id}" + if actual_group_id + else (f"用户 {actual_user_id}" if actual_user_id else "全局") + ) logger.info( f"插件 '{plugin_name}' 的定时任务在目标 [{target_desc}] " f"因功能被禁用而跳过执行。" @@ -279,7 +287,8 @@ async def _execute_single_job_instance( plugin_name=plugin_name, bot_id=bot.self_id, platform_scope=PlatformUtils.get_platform_scope(bot), - group_id=group_id, + group_id=actual_group_id, + user_id=actual_user_id, job_kwargs=schedule.job_kwargs if isinstance(schedule.job_kwargs, dict) else {}, ) state: T_State = {ScheduleContext: context} @@ -300,18 +309,13 @@ async def _execute_single_job_instance( injected_params["params"] = params_instance # type: ignore except Exception as e: logger.error( - f"任务 {schedule.id} (目标: {group_id}) 参数验证失败: {e}", e=e + f"任务 {schedule.id} (目标: {target_log}) 参数验证失败: {e}", + e=e, ) raise - async def wrapper(bot: Bot): - return await task_meta["func"](bot=bot, **injected_params) # type: ignore - - dependent = Dependent.parse( - call=wrapper, - allow_types=Matcher.HANDLER_PARAM_TYPES, - ) - return await dependent(bot=bot, state=state) + func = cast(Callable, task_meta["func"]) + return await _run_dependent_task(func, injected_params, bot, state) try: if policy.retries > 0: @@ -332,117 +336,37 @@ async def _execute_single_job_instance( exception=retry_exceptions, on_success=on_success_handler, on_failure=on_failure_handler, - log_name=f"ScheduledJob-{schedule.id}-{group_id or 'global'}", + log_name=f"ScheduledJob-{schedule.id}-{target_log or 'global'}", ) decorated_executor = retry_decorator(task_execution_coro) await decorated_executor() else: logger.info( - f"插件 '{plugin_name}' 开始为目标 [{group_id or '全局'}] " + f"插件 '{plugin_name}' 开始为目标 [{target_log or '全局'}] " f"执行定时任务 (ID: {schedule.id})。" ) await task_execution_coro() except (PausedException, FinishedException, SkippedException) as e: logger.warning( - f"定时任务 {schedule.id} (目标: {group_id}) 被中断: {type(e).__name__}" + f"定时任务 {schedule.id} (目标: {target_log}) 被中断: {type(e).__name__}" ) except Exception as e: logger.error( - f"执行定时任务 {schedule.id} (目标: {group_id}) " + f"执行定时任务 {schedule.id} (目标: {target_log}) " f"时发生未被策略处理的最终错误", e=e, ) -async def _execute_job( - schedule_id: int | None, - force: bool = False, - context_override: ScheduleContext | None = None, -): - """ - APScheduler 调度的入口函数,现在作为分发器。 - """ - from .manager import scheduler_manager - - schedule = None - - if context_override: - plugin_name = context_override.plugin_name - task_meta = scheduler_manager._registered_tasks.get(plugin_name) - if not task_meta or not task_meta["func"]: - logger.error(f"无法执行临时任务:函数 '{plugin_name}' 未注册。") - return - - try: - bot = _resolve_scheduler_bot( - context_override.bot_id, - f"临时任务 {plugin_name}", - ) - if bot is None: - return - context_override.platform_scope = PlatformUtils.get_platform_scope(bot) - logger.info(f"开始执行临时任务: {plugin_name}") - injected_params = {"context": context_override} - state: T_State = {ScheduleContext: context_override} - - async def wrapper(bot: Bot): - return await task_meta["func"](bot=bot, **injected_params) # type: ignore - - dependent = Dependent.parse( - call=wrapper, - allow_types=Matcher.HANDLER_PARAM_TYPES, - ) - await dependent(bot=bot, state=state) - logger.info(f"临时任务 '{plugin_name}' 执行完成。") - except Exception as e: - logger.error(f"执行临时任务 '{plugin_name}' 时发生错误", e=e) - return - - if schedule_id is None: - logger.error("执行持久化任务时 schedule_id 不能为空。") - return - - global _LAST_PRESSURE_SKIP - if should_pause_tasks(): - now = time.time() - if now - _LAST_PRESSURE_SKIP > 30: - _LAST_PRESSURE_SKIP = now - logger.info("scheduler paused due to message pressure") - return - - scheduler_manager._running_tasks.add(schedule_id) - try: - schedule = await ScheduleRepository.get_by_id(schedule_id) - if not schedule or (not schedule.is_enabled and not force): - logger.warning(f"定时任务 {schedule_id} 不存在或已禁用,跳过执行。") - return - - bot = _resolve_scheduler_bot(schedule.bot_id, f"任务 {schedule_id}") - if bot is None: - return - - resolver = scheduler_manager._target_resolvers.get(schedule.target_type) - if not resolver: - logger.error( - f"任务 {schedule.id} 的目标类型 '{schedule.target_type}' " - f"没有注册解析器,执行跳过。" - ) - raise ValueError(f"未知的目标类型: {schedule.target_type}") - - try: - resolved_targets = await resolver(schedule.target_identifier, bot) - except Exception as e: - logger.error(f"为任务 {schedule.id} 解析目标失败", e=e) - raise - - logger.info( - f"任务 {schedule.id} ({schedule.name or schedule.plugin_name}) 开始执行, " - f"目标类型: {schedule.target_type}, " - f"解析出 {len(resolved_targets)} 个目标" - ) +class ExecutionDispatcher: + """执行管线分发器:处理并发限制、串行间隔、随机散列打散逻辑""" + @staticmethod + async def dispatch( + schedule: ScheduledJob, bot: Bot, resolved_targets: list[str | None] + ): concurrency_limit = Config.get_config( "SchedulerManager", SCHEDULE_CONCURRENCY_KEY, 5 ) @@ -462,15 +386,10 @@ async def _execute_job( ) for i, target_id in enumerate(resolved_targets): if i > 0: - logger.debug( - f"任务 {schedule.id} 目标 [{target_id or '全局'}]: " - f"等待 {interval_seconds} 秒后执行。" - ) await asyncio.sleep(interval_seconds) - await _execute_single_job_instance(schedule, bot, group_id=target_id) + await _execute_single_job_instance(schedule, bot, target_id=target_id) else: spread_seconds = spread_config.get("spread", 1.0) - logger.debug( f"任务 {schedule.id}: 将在 {spread_seconds:.2f} 秒内分散执行 " f"{len(resolved_targets)} 个目标。" @@ -478,46 +397,112 @@ async def _execute_job( async def worker(target_id: str | None): delay = random.uniform(0.1, spread_seconds) - logger.debug( - f"任务 {schedule.id} 目标 [{target_id or '全局'}]: " - f"随机延迟 {delay:.2f} 秒后执行。" - ) await asyncio.sleep(delay) async with semaphore: await _execute_single_job_instance( - schedule, bot, group_id=target_id + schedule, bot, target_id=target_id ) tasks_to_run = [worker(target_id) for target_id in resolved_targets] if tasks_to_run: await asyncio.gather(*tasks_to_run, return_exceptions=True) - schedule.last_run_at = datetime.now() - schedule.last_run_status = "SUCCESS" - schedule.consecutive_failures = 0 - await schedule.save( - update_fields=["last_run_at", "last_run_status", "consecutive_failures"] + +async def _execute_ephemeral_job(context: ScheduleContext): + """ + 执行临时的、内存态的定时任务 (Ephemeral Job) + """ + plugin_name = context.plugin_name + task_meta = scheduler_registry.tasks.get(plugin_name) + if not task_meta or not task_meta["func"]: + logger.error(f"无法执行临时任务:函数 '{plugin_name}' 未注册。") + return + + try: + bot = PlatformUtils.resolve_bot( + bot_id=context.bot_id, + log_cmd=f"临时任务 {plugin_name}", ) + if bot is None: + return + + context.platform_scope = PlatformUtils.get_platform_scope(bot) + logger.info(f"开始执行临时任务: {plugin_name}") + injected_params = {"context": context} + state: T_State = {ScheduleContext: context} + + func = cast(Callable, task_meta["func"]) + await _run_dependent_task(func, injected_params, bot, state) + logger.info(f"临时任务 '{plugin_name}' 执行完成。") + except Exception as e: + logger.error(f"执行临时任务 '{plugin_name}' 时发生错误", e=e) + + +async def _execute_persistent_job(schedule_id: int, force: bool = False): + """ + 执行数据库中持久化的定时任务 (Persistent Job) + """ + schedule = None + + global _LAST_PRESSURE_SKIP + if should_pause_tasks(): + now = time.time() + if now - _LAST_PRESSURE_SKIP > 30: + _LAST_PRESSURE_SKIP = now + logger.info("scheduler paused due to message pressure") + return + + scheduler_registry.running_tasks.add(schedule_id) + try: + schedule = await ScheduleRepository.get_by_id(schedule_id) + if not schedule or (not schedule.is_enabled and not force): + logger.warning(f"定时任务 {schedule_id} 不存在或已禁用,跳过执行。") + return + + bot = PlatformUtils.resolve_bot( + bot_id=schedule.bot_id, log_cmd=f"任务 {schedule_id}" + ) + if bot is None: + return + + resolver = scheduler_registry.target_resolvers.get(schedule.target_type) + if not resolver: + logger.error( + f"任务 {schedule.id} 的目标类型 '{schedule.target_type}' " + f"没有注册解析器,执行跳过。" + ) + raise ValueError(f"未知的目标类型: {schedule.target_type}") + + try: + resolved_targets = await resolver(schedule.target_identifier, bot) + except Exception as e: + logger.error(f"为任务 {schedule.id} 解析目标失败", e=e) + raise + + logger.info( + f"任务 {schedule.id} ({schedule.name or schedule.plugin_name}) 开始执行, " + f"目标类型: {schedule.target_type}, " + f"解析出 {len(resolved_targets)} 个目标" + ) + + await ExecutionDispatcher.dispatch(schedule, bot, resolved_targets) + + await ScheduleRepository.update_run_status(schedule, is_success=True) if schedule.is_one_off: logger.info(f"一次性任务 {schedule.id} 执行成功,将被删除。") await ScheduledJob.filter(id=schedule.id).delete() APSchedulerAdapter.remove_job(schedule.id) if schedule.plugin_name.startswith("runtime_one_off__"): - scheduler_manager._registered_tasks.pop(schedule.plugin_name, None) + scheduler_registry.tasks.pop(schedule.plugin_name, None) logger.debug(f"已注销一次性运行时任务: {schedule.plugin_name}") except Exception as e: logger.error(f"执行任务 {schedule_id} 期间发生严重错误", e=e) if schedule: - schedule.last_run_at = datetime.now() - schedule.last_run_status = "FAILURE" - schedule.consecutive_failures = (schedule.consecutive_failures or 0) + 1 - await schedule.save( - update_fields=["last_run_at", "last_run_status", "consecutive_failures"] - ) + await ScheduleRepository.update_run_status(schedule, is_success=False) finally: if schedule_id is not None: - scheduler_manager._running_tasks.discard(schedule_id) + scheduler_registry.running_tasks.discard(schedule_id) diff --git a/zhenxun/services/scheduler/lifecycle.py b/zhenxun/services/scheduler/lifecycle.py index 6e41281b..67b79b40 100644 --- a/zhenxun/services/scheduler/lifecycle.py +++ b/zhenxun/services/scheduler/lifecycle.py @@ -10,8 +10,9 @@ from zhenxun.utils.pydantic_compat import model_dump from .engine import APSchedulerAdapter from .manager import scheduler_manager +from .registry import scheduler_registry from .repository import ScheduleRepository -from .types import ScheduleContext +from .types import JobConfig, ScheduleContext @PriorityLifecycle.on_startup(priority=90) @@ -21,7 +22,7 @@ async def _load_schedules_from_db(): schedules = await ScheduleRepository.get_all_enabled() count = 0 for schedule in schedules: - if schedule.plugin_name in scheduler_manager._registered_tasks: + if schedule.plugin_name in scheduler_registry.tasks: APSchedulerAdapter.add_or_reschedule_job(schedule) count += 1 else: @@ -30,7 +31,7 @@ async def _load_schedules_from_db(): logger.info("正在检查并注册声明式默认任务...") declared_count = 0 - for task_info in scheduler_manager._declared_tasks: + for task_info in scheduler_registry.persistent_declarations: plugin_name = task_info.plugin_name group_id = task_info.group_id bot_id = task_info.bot_id @@ -45,21 +46,21 @@ async def _load_schedules_from_db(): if not exists: logger.info(f"为插件 '{plugin_name}' 注册新的默认定时任务...") - trigger_config_dict = model_dump( - task_info.trigger, exclude={"trigger_type"} - ) - target_type = "GROUP" if group_id else "GLOBAL" target_identifier = group_id or "" + config = JobConfig( + trigger=task_info.trigger, + job_kwargs=task_info.job_kwargs, + bot_id=bot_id, + source="PLUGIN_DEFAULT", + ) + schedule = await scheduler_manager.add_schedule( plugin_name=plugin_name, target_type=target_type, target_identifier=target_identifier, - trigger_type=task_info.trigger.trigger_type, - trigger_config=trigger_config_dict, - job_kwargs=task_info.job_kwargs, - bot_id=bot_id, + config=config, ) if schedule: declared_count += 1 @@ -74,7 +75,7 @@ async def _load_schedules_from_db(): logger.info("正在调度声明式临时任务...") ephemeral_count = 0 - for declaration in scheduler_manager._ephemeral_declared_tasks: + for declaration in scheduler_registry.ephemeral_declarations: try: job_id = f"runtime::{declaration.plugin_name}::{declaration.func.__name__}" diff --git a/zhenxun/services/scheduler/manager.py b/zhenxun/services/scheduler/manager.py index 3ded5516..404367a8 100644 --- a/zhenxun/services/scheduler/manager.py +++ b/zhenxun/services/scheduler/manager.py @@ -10,69 +10,66 @@ from __future__ import annotations from collections.abc import Awaitable, Callable, Coroutine from datetime import datetime import inspect +import json from typing import Any, ClassVar import uuid -from arclet.alconna import Alconna, Option +from arclet.alconna import Alconna import nonebot from nonebot.adapters import Bot -from pydantic import BaseModel +from pydantic import BaseModel, ValidationError from zhenxun.configs.config import Config from zhenxun.models.scheduled_job import ScheduledJob from zhenxun.services.log import logger -from zhenxun.utils.pydantic_compat import model_dump, model_validate +from zhenxun.utils.pydantic_compat import dump_json_safely, model_dump, model_validate -from .engine import APSchedulerAdapter +from .engine import APSchedulerAdapter, _execute_persistent_job +from .registry import scheduler_registry from .repository import ScheduleRepository from .targeting import ( ScheduleTargeter, + _resolve_all_groups, + _resolve_global_or_user, + _resolve_group, + _resolve_tag, + _resolve_user, ) from .types import ( BaseTrigger, EphemeralJobDeclaration, - ExecutionOptions, ExecutionPolicy, + JobConfig, ScheduleContext, ScheduledJobDeclaration, + TargetType, + Trigger, ) class SchedulerManager: - ALL_GROUPS: ClassVar[str] = "__ALL_GROUPS__" - _registered_tasks: ClassVar[ - dict[ - str, - dict[str, Callable | type[BaseModel] | int | list[Option] | Alconna | None], - ] - ] = {} - _declared_tasks: ClassVar[list[ScheduledJobDeclaration]] = [] - _ephemeral_declared_tasks: ClassVar[list[EphemeralJobDeclaration]] = [] - _running_tasks: ClassVar[set] = set() - _target_resolvers: ClassVar[ - dict[str, Callable[[str, Bot], Awaitable[list[str | None]]]] - ] = {} + ALL_GROUPS: ClassVar[str] = scheduler_registry.ALL_GROUPS def __init__(self): self._register_builtin_resolvers() def _register_builtin_resolvers(self): """在管理器初始化时注册所有内置的目标解析器。""" - from .targeting import ( - _resolve_all_groups, - _resolve_global_or_user, - _resolve_group, - _resolve_tag, - _resolve_user, - ) - - if "GROUP" in self._target_resolvers: + if TargetType.GROUP.value in scheduler_registry.target_resolvers: return - self.register_target_resolver("GROUP", _resolve_group) - self.register_target_resolver("TAG", _resolve_tag) - self.register_target_resolver("ALL_GROUPS", _resolve_all_groups) - self.register_target_resolver("GLOBAL", _resolve_global_or_user) - self.register_target_resolver("USER", _resolve_user) + scheduler_registry.register_target_resolver( + TargetType.GROUP.value, _resolve_group + ) + scheduler_registry.register_target_resolver(TargetType.TAG.value, _resolve_tag) + scheduler_registry.register_target_resolver( + TargetType.ALL_GROUPS.value, _resolve_all_groups + ) + scheduler_registry.register_target_resolver( + TargetType.GLOBAL.value, _resolve_global_or_user + ) + scheduler_registry.register_target_resolver( + TargetType.USER.value, _resolve_user + ) logger.debug("已注册所有内置的定时任务目标解析器。") def register_target_resolver( @@ -83,20 +80,14 @@ class SchedulerManager: """ 注册一个新的目标类型解析器。 """ - if target_type in self._target_resolvers: + if target_type in scheduler_registry.target_resolvers: logger.warning(f"目标解析器 '{target_type}' 已存在,将被覆盖。") - self._target_resolvers[target_type.upper()] = resolver_func - logger.info(f"已注册新的定时任务目标解析器: '{target_type}'") + scheduler_registry.register_target_resolver(target_type, resolver_func) + logger.debug(f"已注册新的定时任务目标解析器: '{target_type}'") def target(self, **filters: Any) -> ScheduleTargeter: """ 创建目标选择器以执行批量操作 - - 参数: - **filters: 过滤条件,支持plugin_name、group_id、bot_id等字段。 - - 返回: - ScheduleTargeter: 目标选择器对象,可用于批量操作。 """ return ScheduleTargeter(self, **filters) @@ -112,7 +103,20 @@ class SchedulerManager: default_interval: int | None = None, ): """ - 声明式定时任务的统一装饰器。 + 声明式定时任务的统一装饰器 + + 参数: + trigger: 定时触发器配置。 + group_id: 默认目标群号,如果不指定则为全局任务。 + bot_id: 执行此任务时应优先匹配的 Bot 标识符。 + default_params: 任务函数的默认参数模型实例。 + policy: 执行和重试策略配置。 + default_jitter: 默认抖动延迟时间(秒)。 + default_spread: 默认并发散列随机延迟打散的最大延迟范围(秒)。 + default_interval: 默认串行派发任务时的固定等待间隔(秒)。 + + 返回: + Callable: 装饰器函数,用于包裹并注册声明式任务。 """ def decorator(func: Callable[..., Coroutine]) -> Callable[..., Coroutine]: @@ -123,7 +127,6 @@ class SchedulerManager: plugin_name = plugin.name params_model = None - from .types import ScheduleContext for param in inspect.signature(func).parameters.values(): if ( @@ -134,9 +137,9 @@ class SchedulerManager: params_model = param.annotation break - if plugin_name in self._registered_tasks: + if plugin_name in scheduler_registry.tasks: logger.warning(f"插件 '{plugin_name}' 的定时任务已被重复注册。") - self._registered_tasks[plugin_name] = { + scheduler_registry.tasks[plugin_name] = { "func": func, "model": params_model, "default_jitter": default_jitter, @@ -155,7 +158,7 @@ class SchedulerManager: trigger=trigger, job_kwargs=job_kwargs, ) - self._declared_tasks.append(task_declaration) + scheduler_registry.persistent_declarations.append(task_declaration) logger.debug( f"发现声明式定时任务 '{plugin_name}',将在启动时进行注册。" ) @@ -178,7 +181,9 @@ class SchedulerManager: raise ValueError(f"函数 {func.__name__} 不在任何已加载的插件中。") plugin_name = plugin.name - self._registered_tasks[f"ephemeral::{plugin_name}::{func.__name__}"] = { + scheduler_registry.tasks[ + f"ephemeral::{plugin_name}::{func.__name__}" + ] = { "func": func, "model": None, } @@ -188,7 +193,7 @@ class SchedulerManager: func=func, trigger=trigger, ) - self._ephemeral_declared_tasks.append(declaration) + scheduler_registry.ephemeral_declarations.append(declaration) logger.debug( f"发现临时定时任务 '{plugin_name}:{func.__name__}',将在启动时调度" ) @@ -211,12 +216,24 @@ class SchedulerManager: ) -> Callable: """ 注册可调度的任务函数 + + 参数: + plugin_name: 插件名称。 + params_model: 参数的模型定义,如为 None 则任务不接收额外参数。 + cli_parser: 用于解析命令行参数的 Alconna 解析器。 + default_permission: 运行此定时任务所需的默认权限等级。 + default_jitter: 默认抖动延迟时间(秒)。 + default_spread: 默认并发散列随机延迟打散的最大延迟范围(秒)。 + default_interval: 默认串行派发任务时的固定等待间隔(秒)。 + + 返回: + Callable: 装饰器函数,用于包裹并注册任务。 """ def decorator(func: Callable[..., Coroutine]) -> Callable[..., Coroutine]: - if plugin_name in self._registered_tasks: + if plugin_name in scheduler_registry.tasks: logger.warning(f"插件 '{plugin_name}' 的定时任务已被重复注册。") - self._registered_tasks[plugin_name] = { + scheduler_registry.tasks[plugin_name] = { "func": func, "model": params_model, "cli_parser": cli_parser, @@ -237,7 +254,7 @@ class SchedulerManager: """ 获取已注册插件列表 """ - return list(self._registered_tasks.keys()) + return list(scheduler_registry.tasks.keys()) async def run_at(self, func: Callable[..., Coroutine], trigger: BaseTrigger) -> str: """ @@ -252,14 +269,18 @@ class SchedulerManager: bot_id=None, platform_scope=None, group_id=None, + user_id=None, job_kwargs={}, ) + trigger_config_dict = json.loads(dump_json_safely(trigger)) + trigger_config_dict.pop("trigger_type", None) + APSchedulerAdapter.add_ephemeral_job( job_id=job_id, func=func, trigger_type=trigger.trigger_type, - trigger_config=model_dump(trigger, exclude={"trigger_type"}), + trigger_config=trigger_config_dict, context=context, ) logger.info(f"已动态调度一个临时任务 (ID: {job_id}),将在 {trigger} 触发。") @@ -279,26 +300,40 @@ class SchedulerManager: required_permission: int = 5, ) -> "ScheduledJob | None": """ - 编程式API,用于动态调度一个持久化的、一次性的任务。 + 编程式API,用于动态调度一个持久化的、一次性的任务 + + 参数: + func: 待执行的协程函数。 + trigger: 定时触发器配置。 + user_id: 目标用户 ID,与 group_id 互斥。 + group_id: 目标群号,与 user_id 互斥。 + bot_id: 指定运行此任务的 Bot 标识符。 + job_kwargs: 传递给任务函数的实际参数字典。 + name: 任务的可读别名。 + created_by: 任务创建者的标识。 + required_permission: 运行该任务需要的最低权限等级。 + + 返回: + ScheduledJob | None: 持久化定时任务的数据模型实例,如果失败则返回 None。 """ if user_id and group_id: raise ValueError("user_id 和 group_id 不能同时提供。") temp_plugin_name = f"runtime_one_off__{func.__module__}.{func.__name__}__{uuid.uuid4().hex[:8]}" # noqa: E501 - self._registered_tasks[temp_plugin_name] = {"func": func, "model": None} + scheduler_registry.tasks[temp_plugin_name] = {"func": func, "model": None} logger.debug(f"为一次性任务动态注册临时插件: '{temp_plugin_name}'") - target_type = "USER" if user_id else ("GROUP" if group_id else "GLOBAL") + target_type = ( + TargetType.USER.value + if user_id + else (TargetType.GROUP.value if group_id else TargetType.GLOBAL.value) + ) target_identifier = user_id or group_id or "" - return await self.add_schedule( - plugin_name=temp_plugin_name, - target_type=target_type, - target_identifier=target_identifier, - trigger_type=trigger.trigger_type, - trigger_config=model_dump(trigger, exclude={"trigger_type"}), - job_kwargs=job_kwargs, + config = JobConfig( + trigger=trigger, + job_kwargs=job_kwargs or {}, bot_id=bot_id, name=name, created_by=created_by, @@ -306,6 +341,13 @@ class SchedulerManager: is_one_off=True, ) + return await self.add_schedule( + plugin_name=temp_plugin_name, + target_type=target_type, + target_identifier=target_identifier, + config=config, + ) + async def add_daily_task( self, plugin_name: str, @@ -318,6 +360,18 @@ class SchedulerManager: ) -> "ScheduledJob | None": """ 添加每日定时任务 + + 参数: + plugin_name: 插件名称。 + group_id: 目标群号,如果不指定则为全局任务。 + hour: 触发的小时数(0-23)。 + minute: 触发的分钟数(0-59)。 + second: 触发的秒数(0-59)。 + job_kwargs: 传递给任务函数的参数字典。 + bot_id: 运行此任务的首选 Bot。 + + 返回: + ScheduledJob | None: 定时任务的数据模型对象,如果失败则返回 None。 """ trigger_config = { "hour": hour, @@ -325,14 +379,15 @@ class SchedulerManager: "second": second, "timezone": Config.get_config("SchedulerManager", "SCHEDULER_TIMEZONE"), } + + trigger = Trigger.cron(**trigger_config) return await self.add_schedule( plugin_name, - target_type="GROUP" if group_id else "GLOBAL", + target_type=TargetType.GROUP.value if group_id else TargetType.GLOBAL.value, target_identifier=group_id or "", - trigger_type="cron", - trigger_config=trigger_config, - job_kwargs=job_kwargs, - bot_id=bot_id, + config=JobConfig( + trigger=trigger, job_kwargs=job_kwargs or {}, bot_id=bot_id + ), ) async def add_interval_task( @@ -351,6 +406,21 @@ class SchedulerManager: ) -> "ScheduledJob | None": """ 添加间隔性定时任务 + + 参数: + plugin_name: 插件名称。 + group_id: 目标群号,如果不指定则为全局任务。 + weeks: 间隔的周数。 + days: 间隔的天数。 + hours: 间隔的小时数。 + minutes: 间隔的分钟数。 + seconds: 间隔的秒数。 + start_date: 间隔计算的起始时间。 + job_kwargs: 传递给任务函数的参数字典。 + bot_id: 运行此任务的首选 Bot。 + + 返回: + ScheduledJob | None: 定时任务的数据模型对象,如果失败则返回 None。 """ trigger_config = { "weeks": weeks, @@ -361,23 +431,22 @@ class SchedulerManager: "start_date": start_date, } trigger_config = {k: v for k, v in trigger_config.items() if v} + + trigger = Trigger.interval(**trigger_config) return await self.add_schedule( plugin_name, - target_type="GROUP" if group_id else "GLOBAL", + target_type=TargetType.GROUP.value if group_id else TargetType.GLOBAL.value, target_identifier=group_id or "", - trigger_type="interval", - trigger_config=trigger_config, - job_kwargs=job_kwargs, - bot_id=bot_id, + config=JobConfig( + trigger=trigger, job_kwargs=job_kwargs or {}, bot_id=bot_id + ), ) def _validate_and_prepare_kwargs( self, plugin_name: str, job_kwargs: dict | None ) -> tuple[bool, str | dict]: """验证并准备任务参数,应用默认值""" - from pydantic import ValidationError - - task_meta = self._registered_tasks.get(plugin_name) + task_meta = scheduler_registry.tasks.get(plugin_name) if not task_meta: return False, f"插件 '{plugin_name}' 未注册。" @@ -410,17 +479,7 @@ class SchedulerManager: plugin_name: str, target_type: str, target_identifier: str, - trigger_type: str, - trigger_config: dict, - job_kwargs: dict | None = None, - bot_id: str | None = None, - *, - name: str | None = None, - created_by: str | None = None, - required_permission: int = 5, - source: str = "USER", - is_one_off: bool = False, - execution_options: dict | None = None, + config: JobConfig, ) -> "ScheduledJob | None": """ 添加定时任务(通用方法) @@ -429,51 +488,43 @@ class SchedulerManager: plugin_name: 插件名称。 target_type: 目标类型 (GROUP, USER, TAG, ALL_GROUPS, GLOBAL)。 target_identifier: 目标标识符。 - trigger_type: 触发器类型 (cron, interval, date)。 - trigger_config: 触发器配置字典。 - job_kwargs: 传递给任务函数的额外参数。 - bot_id: Bot ID约束。 - name: 任务别名。 - created_by: 创建者ID。 - required_permission: 管理此任务所需的权限。 - source: 任务来源 (USER, PLUGIN_DEFAULT)。 - is_one_off: 是否为一次性任务。 - execution_options: 任务执行的额外选项 (例如: jitter, spread)。 - - 返回: - ScheduledJob | None: 创建的任务信息,失败时返回None。 + config: JobConfig 参数配置聚合对象。 """ - if plugin_name not in self._registered_tasks: + if plugin_name not in scheduler_registry.tasks: logger.error(f"插件 '{plugin_name}' 没有注册可用的定时任务。") return None - is_valid, result = self._validate_and_prepare_kwargs(plugin_name, job_kwargs) + is_valid, result = self._validate_and_prepare_kwargs( + plugin_name, config.job_kwargs + ) if not is_valid: logger.error(f"任务参数校验失败: {result}") return None - options_dict = execution_options or {} - validated_options = ExecutionOptions(**options_dict) - search_kwargs = { "plugin_name": plugin_name, "target_type": target_type, "target_identifier": target_identifier, } - if bot_id: - search_kwargs["bot_id"] = bot_id + if config.bot_id: + search_kwargs["bot_id"] = config.bot_id + + trigger_config_dict = json.loads(dump_json_safely(config.trigger)) + trigger_config_dict.pop("trigger_type", None) defaults = { - "name": name, - "trigger_type": trigger_type, - "trigger_config": trigger_config, + "name": config.name, + "trigger_type": config.trigger.trigger_type, + "trigger_config": trigger_config_dict, "job_kwargs": result, "is_enabled": True, - "created_by": created_by, - "required_permission": required_permission, - "source": source, - "is_one_off": is_one_off, - "execution_options": model_dump(validated_options, exclude_none=True), + "created_by": config.created_by, + "required_permission": config.required_permission, + "source": config.source, + "is_one_off": config.is_one_off, + "execution_options": model_dump( + config.execution_options, exclude_none=True + ), } defaults = {k: v for k, v in defaults.items() if v is not None} @@ -485,7 +536,7 @@ class SchedulerManager: action_str = "创建" if created else "更新" logger.info( - f"已成功{action_str}任务 '{name or plugin_name}' (ID: {schedule.id})" + f"已成功{action_str}任务 '{config.name or plugin_name}' (ID: {schedule.id})" ) return schedule @@ -494,7 +545,15 @@ class SchedulerManager: ) -> tuple[list[ScheduledJob], int]: """ 根据条件获取定时任务列表 - """ + + 参数: + page: 分页页码,从 1 开始。 + page_size: 每页的任务数量。 + **filters: 过滤条件,支持 plugin_name, target_type, is_enabled 等字段。 + + 返回: + tuple[list[ScheduledJob], int]: 包含任务对象列表和总符合条件的任务数量的元组。 + """ # noqa: E501 cleaned_filters = {k: v for k, v in filters.items() if v is not None} return await ScheduleRepository.query_schedules( page=page, page_size=page_size, **cleaned_filters @@ -523,7 +582,7 @@ class SchedulerManager: status_dict.update(status_from_scheduler) status_dict["is_enabled"] = ( "运行中" - if schedule_id in self._running_tasks + if schedule_id in scheduler_registry.running_tasks else ("启用" if schedule.is_enabled else "暂停") ) statuses.append(status_dict) @@ -539,6 +598,15 @@ class SchedulerManager: ) -> tuple[bool, str]: """ 更新定时任务配置 + + 参数: + schedule_id: 定时任务的ID。 + trigger_type: 触发器类型,例如 'cron' 或 'interval'。 + trigger_config: 触发器配置字典。 + job_kwargs: 更新后的任务参数字典。 + + 返回: + tuple[bool, str]: 包含执行结果(成功为 True)和对应状态消息的元组。 """ schedule = await ScheduleRepository.get_by_id(schedule_id) if not schedule: @@ -589,7 +657,7 @@ class SchedulerManager: status_text = ( "运行中" - if schedule_id in self._running_tasks + if schedule_id in scheduler_registry.running_tasks else ("启用" if schedule.is_enabled else "暂停") ) @@ -636,16 +704,14 @@ class SchedulerManager: """ 立即手动触发指定的定时任务 """ - from .engine import _execute_job - schedule = await ScheduleRepository.get_by_id(schedule_id) if not schedule: return False, f"未找到 ID 为 {schedule_id} 的定时任务。" - if schedule.plugin_name not in self._registered_tasks: + if schedule.plugin_name not in scheduler_registry.tasks: return False, f"插件 '{schedule.plugin_name}' 没有注册可用的定时任务。" try: - await _execute_job(schedule.id, force=True) + await _execute_persistent_job(schedule.id, force=True) return True, f"已手动触发任务 (ID: {schedule.id})。" except Exception as e: logger.error(f"手动触发任务失败: {e}") @@ -654,12 +720,6 @@ class SchedulerManager: async def get_schedule_by_id(self, schedule_id: int) -> "ScheduledJob | None": """ 通过ID获取任务对象的公共方法。 - - 参数: - schedule_id: 任务ID。 - - 返回: - ScheduledJob | None: 任务对象,不存在时返回None。 """ return await ScheduleRepository.get_by_id(schedule_id) diff --git a/zhenxun/services/scheduler/registry.py b/zhenxun/services/scheduler/registry.py new file mode 100644 index 00000000..540ff054 --- /dev/null +++ b/zhenxun/services/scheduler/registry.py @@ -0,0 +1,42 @@ +from collections.abc import Awaitable, Callable + +from arclet.alconna import Alconna, Option +from nonebot.adapters import Bot +from pydantic import BaseModel + +from zhenxun.services.log import logger + +from .types import EphemeralJobDeclaration, ScheduledJobDeclaration + + +class SchedulerRegistry: + """定时任务注册中心,统一管理所有任务的元数据和解析器""" + + ALL_GROUPS = "__ALL_GROUPS__" + + def __init__(self): + """初始化调度任务注册中心容器。""" + self.tasks: dict[ + str, + dict[str, Callable | type[BaseModel] | int | list[Option] | Alconna | None], + ] = {} + self.persistent_declarations: list[ScheduledJobDeclaration] = [] + self.ephemeral_declarations: list[EphemeralJobDeclaration] = [] + self.target_resolvers: dict[ + str, Callable[[str, Bot], Awaitable[list[str | None]]] + ] = {} + self.running_tasks: set[int] = set() + + def register_target_resolver( + self, + target_type: str, + resolver_func: Callable[[str, Bot], Awaitable[list[str | None]]], + ): + """注册指定执行目标的解析策略。""" + if target_type in self.target_resolvers: + logger.warning(f"目标解析器 '{target_type}' 已存在,将被覆盖。") + self.target_resolvers[target_type.upper()] = resolver_func + logger.info(f"已注册新的定时任务目标解析器: '{target_type}'") + + +scheduler_registry = SchedulerRegistry() diff --git a/zhenxun/services/scheduler/repository.py b/zhenxun/services/scheduler/repository.py index 7214503a..66886bfb 100644 --- a/zhenxun/services/scheduler/repository.py +++ b/zhenxun/services/scheduler/repository.py @@ -4,6 +4,7 @@ 封装所有对 ScheduledJob 模型的数据库操作,将数据访问逻辑与业务逻辑分离。 """ +from datetime import datetime from typing import Any from tortoise.queryset import QuerySet @@ -44,6 +45,23 @@ class ScheduleRepository: return await ScheduledJob.filter(plugin_name=plugin_name).all() return await ScheduledJob.all() + @staticmethod + async def update_run_status(schedule: ScheduledJob, is_success: bool): + """ + 统一更新任务的最后运行状态 + """ + schedule.last_run_at = datetime.now() + if is_success: + schedule.last_run_status = "SUCCESS" + schedule.consecutive_failures = 0 + else: + schedule.last_run_status = "FAILURE" + schedule.consecutive_failures = (schedule.consecutive_failures or 0) + 1 + + await schedule.save( + update_fields=["last_run_at", "last_run_status", "consecutive_failures"] + ) + @staticmethod async def save(schedule: ScheduledJob, update_fields: list[str] | None = None): """ diff --git a/zhenxun/services/scheduler/targeting.py b/zhenxun/services/scheduler/targeting.py index 2d8407fa..5aeed8c2 100644 --- a/zhenxun/services/scheduler/targeting.py +++ b/zhenxun/services/scheduler/targeting.py @@ -11,6 +11,9 @@ from nonebot.adapters import Bot from zhenxun.services.tags import tag_manager +from .engine import APSchedulerAdapter +from .repository import ScheduleRepository + __all__ = [ "ScheduleTargeter", "_resolve_all_groups", @@ -22,24 +25,29 @@ __all__ = [ async def _resolve_group(target_identifier: str, bot: Bot) -> list[str | None]: + """解析 GROUP 类型的执行目标。""" return [target_identifier] async def _resolve_tag(target_identifier: str, bot: Bot) -> list[str | None]: + """解析 TAG 类型的标签执行目标。""" result = await tag_manager.resolve_tag_to_group_ids(target_identifier) return result # type: ignore async def _resolve_user(target_identifier: str, bot: Bot) -> list[str | None]: + """解析 USER 类型的执行目标。""" return [target_identifier] async def _resolve_all_groups(target_identifier: str, bot: Bot) -> list[str | None]: + """解析 ALL_GROUPS 类型的全群执行目标。""" result = await tag_manager.resolve_tag_to_group_ids("@all", bot=bot) return result async def _resolve_global_or_user(target_identifier: str, bot: Bot) -> list[str | None]: + """解析 GLOBAL 类型的全局执行目标。""" return [None] @@ -66,8 +74,6 @@ class ScheduleTargeter: 返回: list[ScheduledJob]: 符合过滤条件的任务列表。 """ - from .repository import ScheduleRepository - query = ScheduleRepository.filter(**self._filters) return await query.all() @@ -147,9 +153,6 @@ class ScheduleTargeter: 返回: tuple[int, str]: (成功移除的任务数量, 操作结果消息)。 """ - from .engine import APSchedulerAdapter - from .repository import ScheduleRepository - schedules = await self._get_schedules() if not schedules: target_desc = self._generate_target_description() diff --git a/zhenxun/services/scheduler/triggers.py b/zhenxun/services/scheduler/triggers.py deleted file mode 100644 index 60523f18..00000000 --- a/zhenxun/services/scheduler/triggers.py +++ /dev/null @@ -1,80 +0,0 @@ -from datetime import datetime -from typing import Literal - -from pydantic import BaseModel, Field - - -class BaseTrigger(BaseModel): - """触发器配置的基类""" - - trigger_type: str = Field(..., exclude=True) - - -class CronTrigger(BaseTrigger): - """Cron 触发器配置""" - - trigger_type: Literal["cron"] = "cron" # type: ignore - year: int | str | None = None - month: int | str | None = None - day: int | str | None = None - week: int | str | None = None - day_of_week: int | str | None = None - hour: int | str | None = None - minute: int | str | None = None - second: int | str | None = None - start_date: datetime | str | None = None - end_date: datetime | str | None = None - timezone: str | None = None - jitter: int | None = None - - -class IntervalTrigger(BaseTrigger): - """Interval 触发器配置""" - - trigger_type: Literal["interval"] = "interval" # type: ignore - weeks: int = 0 - days: int = 0 - hours: int = 0 - minutes: int = 0 - seconds: int = 0 - start_date: datetime | str | None = None - end_date: datetime | str | None = None - timezone: str | None = None - jitter: int | None = None - - -class DateTrigger(BaseTrigger): - """Date 触发器配置""" - - trigger_type: Literal["date"] = "date" # type: ignore - run_date: datetime | str - timezone: str | None = None - - -class Trigger: - """ - 一个用于创建类型安全触发器配置的工厂类。 - 提供了流畅的、具备IDE自动补全功能的API。 - - 使用示例: - from zhenxun.services.scheduler import Trigger - - @scheduler.job(trigger=Trigger.cron(hour=8)) - async def my_task(): - ... - """ - - @staticmethod - def cron(**kwargs) -> CronTrigger: - """创建一个 Cron 触发器配置。""" - return CronTrigger(**kwargs) - - @staticmethod - def interval(**kwargs) -> IntervalTrigger: - """创建一个 Interval 触发器配置。""" - return IntervalTrigger(**kwargs) - - @staticmethod - def date(**kwargs) -> DateTrigger: - """创建一个 Date 触发器配置。""" - return DateTrigger(**kwargs) diff --git a/zhenxun/services/scheduler/types.py b/zhenxun/services/scheduler/types.py index d1561bda..e8373377 100644 --- a/zhenxun/services/scheduler/types.py +++ b/zhenxun/services/scheduler/types.py @@ -4,56 +4,101 @@ from collections.abc import Awaitable, Callable from datetime import datetime +from enum import Enum from typing import Any, Literal from pydantic import BaseModel, Field +from zhenxun.utils.pydantic_compat import model_validate + + +class TargetType(str, Enum): + """定时任务执行目标类型枚举""" + + GLOBAL = "GLOBAL" + """全局""" + ALL_GROUPS = "ALL_GROUPS" + """所有群组""" + GROUP = "GROUP" + """群组""" + USER = "USER" + """用户""" + TAG = "TAG" + """标签""" + class BaseTrigger(BaseModel): """触发器配置的基类""" trigger_type: str = Field(..., exclude=True) + """触发器类型""" class CronTrigger(BaseTrigger): """Cron 触发器配置""" trigger_type: Literal["cron"] = "cron" # type: ignore + """触发器类型""" year: int | str | None = None + """年""" month: int | str | None = None + """月""" day: int | str | None = None + """日""" week: int | str | None = None + """周""" day_of_week: int | str | None = None + """星期几""" hour: int | str | None = None + """小时""" minute: int | str | None = None + """分钟""" second: int | str | None = None + """秒""" start_date: datetime | str | None = None + """开始日期""" end_date: datetime | str | None = None + """结束日期""" timezone: str | None = None + """时区""" jitter: int | None = None + """运行抖动时间""" class IntervalTrigger(BaseTrigger): """Interval 触发器配置""" trigger_type: Literal["interval"] = "interval" # type: ignore + """触发器类型""" weeks: int = 0 + """周数""" days: int = 0 + """天数""" hours: int = 0 + """小时数""" minutes: int = 0 + """分钟数""" seconds: int = 0 + """秒数""" start_date: datetime | str | None = None + """开始日期""" end_date: datetime | str | None = None + """结束日期""" timezone: str | None = None + """时区""" jitter: int | None = None + """运行抖动时间""" class DateTrigger(BaseTrigger): """Date 触发器配置""" trigger_type: Literal["date"] = "date" # type: ignore + """触发器类型""" run_date: datetime | str + """运行日期""" timezone: str | None = None + """时区""" class Trigger: @@ -83,18 +128,18 @@ class ExecutionOptions(BaseModel): 封装定时任务的执行策略,包括重试和回调。 """ - jitter: int | None = Field(None, description="触发时间抖动(秒)") - spread: int | None = Field( - None, description="(并发模式)多目标执行的最大分散延迟(秒)" - ) - interval: int | None = Field( - None, description="多目标执行的固定间隔(秒),设置后将强制串行执行" - ) - concurrency_policy: Literal["ALLOW", "SKIP", "QUEUE"] = Field( - "ALLOW", description="并发策略" - ) + jitter: int | None = Field(None) + """触发时间抖动(秒)""" + spread: int | None = Field(None) + """(并发模式)多目标执行的最大分散延迟(秒)""" + interval: int | None = Field(None) + """多目标执行的固定间隔(秒),设置后将强制串行执行""" + concurrency_policy: Literal["ALLOW", "SKIP", "QUEUE"] = Field("ALLOW") + """并发策略""" retries: int = 0 + """重试次数""" retry_delay_seconds: int = 30 + """重试延迟时间(秒)""" class ScheduleContext(BaseModel): @@ -102,12 +147,20 @@ class ScheduleContext(BaseModel): 定时任务执行上下文,可通过依赖注入获取。 """ - schedule_id: int = Field(..., description="数据库中的任务ID") - plugin_name: str = Field(..., description="任务所属的插件名称") - bot_id: str | None = Field(None, description="执行任务的Bot ID") - platform_scope: str | None = Field(None, description="执行任务的细粒度平台作用域") - group_id: str | None = Field(None, description="当前执行实例的目标群组ID") - job_kwargs: dict = Field(default_factory=dict, description="任务配置的参数") + schedule_id: int = Field(...) + """数据库中的任务ID""" + plugin_name: str = Field(...) + """任务所属的插件名称""" + bot_id: str | None = Field(None) + """执行任务的Bot ID""" + platform_scope: str | None = Field(None) + """执行任务的细粒度平台作用域""" + group_id: str | None = Field(None) + """当前执行实例的目标群组ID""" + user_id: str | None = None + """当前执行实例的目标用户ID(私聊场景)""" + job_kwargs: dict = Field(default_factory=dict) + """任务配置的参数""" class ExecutionPolicy(BaseModel): @@ -116,13 +169,50 @@ class ExecutionPolicy(BaseModel): """ retries: int = 0 + """重试次数""" retry_delay_seconds: int = 30 + """重试延迟时间(秒)""" retry_backoff: bool = False + """是否使用退避算法延迟重试""" retry_on_exceptions: list[type[Exception]] | None = None + """触发重试的异常类型列表""" on_success_callback: Callable[[ScheduleContext, Any], Awaitable[None]] | None = None + """任务执行成功后的回调函数""" on_failure_callback: ( Callable[[ScheduleContext, Exception], Awaitable[None]] | None ) = None + """任务执行失败后的回调函数""" + + class Config: + arbitrary_types_allowed = True + + +class JobConfig(BaseModel): + """ + 定时任务参数配置聚合实体 (Parameter Object)。 + 封装了除核心身份标识(插件名、目标类型)之外的所有调度配置。 + """ + + trigger: BaseTrigger + """触发器配置""" + job_kwargs: dict[str, Any] = Field(default_factory=dict) + """任务执行参数""" + bot_id: str | None = None + """绑定的 Bot ID""" + name: str | None = None + """定时任务名称""" + created_by: str | None = None + """任务创建者标识""" + required_permission: int = 5 + """执行任务所需的权限等级""" + source: str = "USER" + """任务来源""" + is_one_off: bool = False + """是否为一次性任务""" + execution_options: ExecutionOptions = Field( + default_factory=lambda: model_validate(ExecutionOptions, {}) + ) + """执行控制选项配置""" class Config: arbitrary_types_allowed = True @@ -132,10 +222,15 @@ class ScheduledJobDeclaration(BaseModel): """用于在启动时声明默认定时任务的内部数据模型""" plugin_name: str + """插件名称""" group_id: str | None + """绑定的群组 ID""" bot_id: str | None + """绑定的 Bot ID""" trigger: BaseTrigger + """触发器配置""" job_kwargs: dict[str, Any] + """任务执行参数""" class Config: arbitrary_types_allowed = True @@ -145,8 +240,11 @@ class EphemeralJobDeclaration(BaseModel): """用于在启动时声明临时任务的内部数据模型""" plugin_name: str + """插件名称""" func: Callable[..., Awaitable[Any]] + """临时任务执行的异步函数""" trigger: BaseTrigger + """触发器配置""" class Config: arbitrary_types_allowed = True diff --git a/zhenxun/utils/pydantic_compat.py b/zhenxun/utils/pydantic_compat.py index 578b80ed..29f93b23 100644 --- a/zhenxun/utils/pydantic_compat.py +++ b/zhenxun/utils/pydantic_compat.py @@ -48,9 +48,14 @@ else: return type_validate_python(self.type_, obj) + from pydantic import root_validator + def model_validator(*args: Any, **kwargs: Any) -> Any: + mode = kwargs.get("mode", "after") + pre = mode == "before" + def decorator(func: Any) -> Any: - return func + return root_validator(pre=pre, allow_reuse=True)(func) return decorator @@ -68,6 +73,7 @@ __all__ = [ "model_dump_json", "model_fields", "model_json_schema", + "model_rebuild", "model_validate", "model_validator", "parse_as", @@ -180,3 +186,14 @@ def dump_json_safely(obj: Any, **kwargs) -> str: ) return json.dumps(obj, default=default_serializer, **kwargs) + + +def model_rebuild(model_class: type[BaseModel], **kwargs: Any) -> None: + """ + Pydantic V1/V2 兼容的前向引用重建函数。 + V2 调用 `model_rebuild()`,V1 调用 `update_forward_refs()`。 + """ + if PYDANTIC_V2: + model_class.model_rebuild(**kwargs) + else: + model_class.update_forward_refs(**kwargs)