mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-10-10 14:20:04 +08:00
♻️ refactor(core): 重构 AI 编排框架与记忆及 RAG 子系统 (#2149)
* ♻️ refactor(core): 重构 AI 编排框架与记忆及 RAG 子系统 - 【重构】重构 `BaseRunnable` 并引入统一的 `RunIntent` 意图载体,规范 Agent、Team 和 Workflow 的执行流 - 【解耦】将中期记忆槽和长期向量记忆从 `MemoryConfig` 中解耦,转为独立的能力组件与工具箱进行管理 - 【记忆】移除 `MemoryReader` 和 `MemoryWriter`,统一封装为 `SessionMemoryContext` 会话记忆门面 - 【RAG】重构检索器与存储后端接口,统一采用 `QueryRequest` 进行多维度联合检索,并引入 `InMemoryScorer` 提升打分性能 - 【事件】优化 `EventBus` 异步事件分发机制,引入队列机制确保事件按序处理,避免并发竞态问题 - 【依赖注入】移除 `memory` 注入项,优化 `DependencyInjector` 的签名解析缓存以提升性能 * 🚨 auto fix by pre-commit hooks --------- Co-authored-by: webjoin111 <455457521@qq.com> Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
This commit is contained in:
co-authored by
webjoin111
pre-commit-ci[bot]
parent
922d092650
commit
52f7dbdedf
@@ -1,13 +1,14 @@
|
||||
from abc import ABC, abstractmethod
|
||||
import asyncio
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import Any
|
||||
from typing import cast
|
||||
|
||||
from zhenxun.services.ai.core.exceptions import (
|
||||
AbortException,
|
||||
ControlFlowExit,
|
||||
ToolFatalError,
|
||||
)
|
||||
from zhenxun.services.ai.core.stream_events import AgentStreamEvent
|
||||
from zhenxun.services.ai.run import RunContext
|
||||
from zhenxun.services.ai.utils.logger import log_flow as logger
|
||||
|
||||
@@ -23,6 +24,21 @@ from .types import (
|
||||
)
|
||||
|
||||
|
||||
class StreamCapturer:
|
||||
"""内部辅助类:透传工作流节点的内部流事件,并捕获最终的 StepOutput 产出值"""
|
||||
|
||||
def __init__(self, stream: AsyncIterator[AgentStreamEvent | StepOutput]):
|
||||
self._stream = stream
|
||||
self.output: StepOutput | None = None
|
||||
|
||||
async def __aiter__(self) -> AsyncIterator[AgentStreamEvent]:
|
||||
async for event in self._stream:
|
||||
if isinstance(event, StepOutput):
|
||||
self.output = event
|
||||
else:
|
||||
yield event
|
||||
|
||||
|
||||
class BaseNode(ABC):
|
||||
"""工作流节点统一抽象基类"""
|
||||
|
||||
@@ -50,7 +66,7 @@ class BaseNode(ABC):
|
||||
|
||||
async def _handle_execution_failure(
|
||||
self, e: BaseException, step_input: StepInput, context: RunContext, attempt: int
|
||||
) -> tuple[str, StepOutput | None, StepInput | None, Any]:
|
||||
) -> tuple[str, StepOutput | None, StepInput | None, "BaseNode | None"]:
|
||||
"""
|
||||
解析执行异常并应用容错策略
|
||||
"""
|
||||
@@ -128,19 +144,10 @@ class BaseNode(ABC):
|
||||
@abstractmethod
|
||||
async def run_stream(
|
||||
self, step_input: StepInput, context: RunContext
|
||||
) -> AsyncIterator[Any]:
|
||||
) -> AsyncIterator[AgentStreamEvent | StepOutput]:
|
||||
"""子类必须实现的核心流式执行逻辑"""
|
||||
yield None
|
||||
|
||||
async def _forward_stream(
|
||||
self, stream: AsyncIterator[Any], output_box: list[StepOutput]
|
||||
) -> AsyncIterator[Any]:
|
||||
"""辅助方法:转发内部流事件,并将最终的 StepOutput 拦截放入 output_box 列表中"""
|
||||
async for event in stream:
|
||||
if isinstance(event, StepOutput):
|
||||
output_box.append(event)
|
||||
else:
|
||||
yield event
|
||||
if False:
|
||||
yield cast(StepOutput, None)
|
||||
|
||||
async def aexecute(self, step_input: StepInput, context: RunContext) -> StepOutput:
|
||||
"""非流式执行(聚合流并返回最终结果),子类无需重写"""
|
||||
@@ -160,7 +167,7 @@ class BaseNode(ABC):
|
||||
|
||||
async def aexecute_stream(
|
||||
self, step_input: StepInput, context: RunContext
|
||||
) -> AsyncIterator[Any]:
|
||||
) -> AsyncIterator[AgentStreamEvent | StepOutput]:
|
||||
"""标准化模板方法:处理缓存快进、授权挂起、异常熔断与生命周期事件分发"""
|
||||
logger.debug(f" ⚙️ [节点] `{self.name}` 开始执行...")
|
||||
|
||||
@@ -222,4 +229,5 @@ class BaseNode(ABC):
|
||||
yield evt
|
||||
break
|
||||
|
||||
assert output is not None
|
||||
yield output
|
||||
|
||||
@@ -1,7 +1,5 @@
|
||||
import asyncio
|
||||
from collections.abc import AsyncIterator
|
||||
import contextlib
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
import uuid
|
||||
|
||||
from nonebot.params import Depends
|
||||
@@ -12,15 +10,16 @@ if TYPE_CHECKING:
|
||||
|
||||
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.core.models import CancellationToken
|
||||
from zhenxun.services.ai.core.stream_events import AgentStreamEvent, EventBus
|
||||
from zhenxun.services.ai.flow.core.base import BaseRunnable
|
||||
from zhenxun.services.ai.flow.core.models import BaseRuntimeConfig
|
||||
from zhenxun.services.ai.run.blackboard import BlackboardManager
|
||||
from zhenxun.services.ai.run.context import RunContext
|
||||
from zhenxun.services.ai.run.models import (
|
||||
AgentRunEnd,
|
||||
AgentRunError,
|
||||
AgentRunResult,
|
||||
StreamedRunResult,
|
||||
RunIntent,
|
||||
)
|
||||
from zhenxun.services.ai.tools.core.tool import FunctionTool
|
||||
from zhenxun.services.ai.utils.logger import log_flow as logger
|
||||
@@ -176,96 +175,37 @@ class Workflow(BaseRunnable[WorkflowRunResult]):
|
||||
返回:
|
||||
WorkflowRunResult: 包含执行状态、断点快照、各节点产出的全量工作流结果对象。
|
||||
"""
|
||||
session_id = (
|
||||
context.session_id if context and context.session_id else f"wf_{self.id}"
|
||||
)
|
||||
safe_context = context or RunContext(session_id=session_id)
|
||||
async with self.run_stream(
|
||||
prompt=prompt, context=context, **kwargs
|
||||
) as stream_result:
|
||||
res = await stream_result.get_run_result()
|
||||
return cast(WorkflowRunResult, res.structured_data)
|
||||
|
||||
if self.blackboard_schema and not safe_context.session.blackboard:
|
||||
safe_context.session.blackboard = BlackboardManager(
|
||||
async def _execute_stream(
|
||||
self,
|
||||
intent: RunIntent,
|
||||
context: RunContext,
|
||||
cancel_token: CancellationToken,
|
||||
event_bus: EventBus,
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterator[AgentStreamEvent]:
|
||||
"""统一核心流,不再自己维护 Task 和 EventBus"""
|
||||
|
||||
if self.blackboard_schema and not context.session.blackboard:
|
||||
context.session.blackboard = BlackboardManager(
|
||||
schema=self.blackboard_schema,
|
||||
initial_state=self.initial_blackboard_state,
|
||||
)
|
||||
|
||||
logger.debug(f"🏭 **工作流 [{self.name}] 启动**")
|
||||
|
||||
initial_input = StepInput(input=prompt)
|
||||
if kwargs:
|
||||
initial_input.additional_data.update(kwargs)
|
||||
|
||||
try:
|
||||
final_output = await self.root_steps.aexecute(initial_input, safe_context)
|
||||
|
||||
logger.debug(f"🏭 **工作流 [{self.name}] 运行结束**")
|
||||
|
||||
return self._build_result(initial_input, safe_context, final_output)
|
||||
|
||||
except BaseException as e:
|
||||
if isinstance(e, ControlFlowExit):
|
||||
logger.debug(f"⏭️ 工作流执行被业务控制流安全中止: {e}")
|
||||
dummy_output = StepOutput(content=str(e), success=False)
|
||||
return self._build_result(initial_input, safe_context, dummy_output)
|
||||
|
||||
raise e
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def run_stream(
|
||||
self,
|
||||
prompt: PromptInput | None = None,
|
||||
*,
|
||||
context: RunContext | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterator["StreamedRunResult[Any]"]:
|
||||
"""对齐 BaseRunnable 接口的流式上下文管理器"""
|
||||
event_bus = EventBus()
|
||||
if context:
|
||||
context.run.event_bus = event_bus
|
||||
|
||||
async def _execution_task():
|
||||
try:
|
||||
async for event in self._internal_stream(prompt, context, **kwargs):
|
||||
await event_bus.emit(event)
|
||||
except BaseException as e:
|
||||
await event_bus.emit(AgentRunError(error=e))
|
||||
finally:
|
||||
await event_bus.end()
|
||||
|
||||
task = asyncio.create_task(_execution_task())
|
||||
try:
|
||||
yield StreamedRunResult[Any](event_bus)
|
||||
finally:
|
||||
if not task.done():
|
||||
task.cancel()
|
||||
|
||||
async def _internal_stream(
|
||||
self,
|
||||
prompt: PromptInput | None = None,
|
||||
context: RunContext | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterator[Any]:
|
||||
"""流式执行工作流节点树的内部实现"""
|
||||
session_id = (
|
||||
context.session_id if context and context.session_id else f"wf_{self.id}"
|
||||
)
|
||||
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)
|
||||
initial_input = StepInput(input=intent.original_input, intent=intent)
|
||||
if kwargs:
|
||||
initial_input.additional_data.update(kwargs)
|
||||
|
||||
try:
|
||||
final_output = None
|
||||
async for event in self.root_steps.aexecute_stream(
|
||||
initial_input, safe_context
|
||||
):
|
||||
async for event in self.root_steps.aexecute_stream(initial_input, context):
|
||||
if isinstance(event, StepOutput):
|
||||
final_output = event
|
||||
else:
|
||||
@@ -274,17 +214,26 @@ class Workflow(BaseRunnable[WorkflowRunResult]):
|
||||
if final_output:
|
||||
logger.debug(f"🏭 **工作流 [{self.name}] 运行结束**")
|
||||
|
||||
wf_result = self._build_result(
|
||||
initial_input, safe_context, final_output
|
||||
)
|
||||
wf_result = self._build_result(initial_input, context, final_output)
|
||||
agent_res = AgentRunResult(
|
||||
output=wf_result.last_step_content,
|
||||
structured_data=wf_result,
|
||||
usage=UsageInfo(),
|
||||
)
|
||||
yield AgentRunEnd(result=agent_res)
|
||||
except Exception:
|
||||
pass
|
||||
except BaseException as e:
|
||||
if isinstance(e, ControlFlowExit):
|
||||
logger.debug(f"⏭️ 工作流执行被业务控制流安全中止: {e}")
|
||||
dummy_output = StepOutput(content=str(e), success=False)
|
||||
wf_result = self._build_result(initial_input, context, dummy_output)
|
||||
agent_res = AgentRunResult(
|
||||
output=wf_result.last_step_content,
|
||||
structured_data=wf_result,
|
||||
usage=UsageInfo(),
|
||||
)
|
||||
yield AgentRunEnd(result=agent_res)
|
||||
else:
|
||||
raise e
|
||||
|
||||
def as_tool(self, tool_name: str | None = None) -> FunctionTool:
|
||||
"""将工作流封装并导出为可供 Agent 直接调用的 FunctionTool 实例"""
|
||||
|
||||
@@ -1,16 +1,18 @@
|
||||
from abc import ABC, abstractmethod
|
||||
import asyncio
|
||||
from collections.abc import AsyncIterator, Callable, Sequence
|
||||
from collections.abc import AsyncIterator, Awaitable, Callable, Sequence
|
||||
import copy
|
||||
from typing import Any, cast
|
||||
from typing import cast
|
||||
|
||||
from zhenxun.services.ai.core.messages import PromptInput
|
||||
from zhenxun.services.ai.flow.base import BaseRunnable
|
||||
from zhenxun.services.ai.run import AgentTask, RunContext
|
||||
from zhenxun.services.ai.core.stream_events import AgentStreamEvent
|
||||
from zhenxun.services.ai.flow.core.base import BaseRunnable
|
||||
from zhenxun.services.ai.run import RunContext
|
||||
from zhenxun.services.ai.run.di import DependencyInjector
|
||||
from zhenxun.services.ai.run.models import AgentRunEnd
|
||||
from zhenxun.services.ai.run.models import AgentRunEnd, RunIntent
|
||||
from zhenxun.services.ai.utils.logger import log_flow as logger
|
||||
|
||||
from .base import BaseNode
|
||||
from .base import BaseNode, StreamCapturer
|
||||
from .policies import BaseFailurePolicy
|
||||
from .types import (
|
||||
StepInput,
|
||||
@@ -24,7 +26,7 @@ NodeSource = BaseNode | BaseRunnable | Callable
|
||||
|
||||
class Step(BaseNode):
|
||||
"""
|
||||
工作流中的最小执行单元门面 (Facade)。
|
||||
工作流中的最小执行单元门面
|
||||
对外部隐藏了 AgentNode 和 FunctionNode 的具体实现。
|
||||
当实例化 Step 时,底层会自动根据 executor 的类型返回专属的节点对象。
|
||||
"""
|
||||
@@ -35,7 +37,7 @@ class Step(BaseNode):
|
||||
if executor is None and len(args) > 1:
|
||||
executor = args[1]
|
||||
|
||||
from zhenxun.services.ai.flow.base import BaseRunnable
|
||||
from zhenxun.services.ai.flow.core.base import BaseRunnable
|
||||
|
||||
if isinstance(executor, BaseRunnable):
|
||||
return object.__new__(RunnableNode)
|
||||
@@ -75,7 +77,7 @@ class Step(BaseNode):
|
||||
|
||||
async def run_stream(
|
||||
self, step_input: StepInput, context: RunContext
|
||||
) -> AsyncIterator[Any]:
|
||||
) -> AsyncIterator[AgentStreamEvent | StepOutput]:
|
||||
if False:
|
||||
yield None
|
||||
raise NotImplementedError("这是一个外观门面,实际的执行发生在子类中。")
|
||||
@@ -86,32 +88,28 @@ class RunnableNode(Step):
|
||||
|
||||
async def run_stream(
|
||||
self, step_input: StepInput, context: RunContext
|
||||
) -> AsyncIterator[Any]:
|
||||
) -> AsyncIterator[AgentStreamEvent | StepOutput]:
|
||||
executor = cast(BaseRunnable, self.executor)
|
||||
|
||||
if isinstance(step_input.previous_step_content, AgentTask):
|
||||
prompt_data = step_input.previous_step_content
|
||||
prev_content_to_append = None
|
||||
else:
|
||||
prompt_data = self.prompt if self.prompt is not None else step_input.input
|
||||
prev_content_to_append = step_input.previous_step_content
|
||||
base_prompt = self.prompt if self.prompt is not None else step_input.input
|
||||
node_intent = RunIntent.from_input(base_prompt)
|
||||
|
||||
if isinstance(prompt_data, AgentTask):
|
||||
prompt_data = copy.copy(prompt_data)
|
||||
if prev_content_to_append:
|
||||
prev_content = str(prev_content_to_append)
|
||||
prompt_data.description = (
|
||||
prev_content = step_input.previous_step_content
|
||||
prompt_data = base_prompt
|
||||
|
||||
if prev_content:
|
||||
if node_intent.task_obj:
|
||||
task_clone = copy.copy(node_intent.task_obj)
|
||||
task_clone.description = (
|
||||
f"### 🔙 [上游节点执行输出]\n{prev_content}\n\n"
|
||||
f"### 🎯 [当前需执行的任务]\n{prompt_data.description}"
|
||||
f"### 🎯 [当前需执行的任务]\n{task_clone.description}"
|
||||
)
|
||||
context.run.user_input = prompt_data.description
|
||||
else:
|
||||
if prev_content_to_append:
|
||||
prompt_data = task_clone
|
||||
else:
|
||||
prompt_data = (
|
||||
f"[上游节点执行输出]:\n{prev_content_to_append}\n\n"
|
||||
f"[当前需执行的任务]:\n{prompt_data or ''}"
|
||||
f"[上游节点执行输出]:\n{prev_content}\n\n"
|
||||
f"[当前需执行的任务]:\n{node_intent.text}"
|
||||
)
|
||||
context.run.user_input = str(prompt_data) if prompt_data else ""
|
||||
|
||||
final_result = None
|
||||
sandbox_context = context.clone_for_member(self.name)
|
||||
@@ -136,7 +134,7 @@ class FunctionNode(Step):
|
||||
|
||||
async def run_stream(
|
||||
self, step_input: StepInput, context: RunContext
|
||||
) -> AsyncIterator[Any]:
|
||||
) -> AsyncIterator[AgentStreamEvent | StepOutput]:
|
||||
context.run.user_input = str(step_input.input) if step_input.input else ""
|
||||
|
||||
executor = cast(Callable, self.executor)
|
||||
@@ -170,21 +168,20 @@ class Steps(BaseNode):
|
||||
|
||||
async def run_stream(
|
||||
self, step_input: StepInput, context: RunContext
|
||||
) -> AsyncIterator[Any]:
|
||||
) -> AsyncIterator[AgentStreamEvent | StepOutput]:
|
||||
current_input = StepInput(
|
||||
input=step_input.input,
|
||||
intent=step_input.intent,
|
||||
previous_step_content=step_input.previous_step_content,
|
||||
additional_data=step_input.additional_data.copy(),
|
||||
)
|
||||
|
||||
all_outputs: list[StepOutput] = []
|
||||
for step_obj in self.steps:
|
||||
out_box: list[StepOutput] = []
|
||||
async for event in self._forward_stream(
|
||||
step_obj.aexecute_stream(current_input, context), out_box
|
||||
):
|
||||
capturer = StreamCapturer(step_obj.aexecute_stream(current_input, context))
|
||||
async for event in capturer:
|
||||
yield event
|
||||
step_out = out_box[0] if out_box else None
|
||||
step_out = capturer.output
|
||||
|
||||
if step_out:
|
||||
all_outputs.append(step_out)
|
||||
@@ -199,12 +196,47 @@ class Steps(BaseNode):
|
||||
)
|
||||
|
||||
|
||||
class Condition(BaseNode):
|
||||
class BranchingNode(BaseNode, ABC):
|
||||
"""
|
||||
处理基于条件或路由的单分支复合节点基类
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
async def select_branch(
|
||||
self, step_input: StepInput, context: RunContext
|
||||
) -> tuple[list[BaseNode], str, str]:
|
||||
"""
|
||||
子类实现此方法进行路由决策。
|
||||
返回: (目标节点列表, 分支标识名, 未命中有效分支时的兜底提示信息)
|
||||
"""
|
||||
pass
|
||||
|
||||
async def run_stream(
|
||||
self, step_input: StepInput, context: RunContext
|
||||
) -> AsyncIterator[AgentStreamEvent | StepOutput]:
|
||||
context.run.user_input = str(step_input.input) if step_input.input else ""
|
||||
target_steps, branch_name, fallback_msg = await self.select_branch(
|
||||
step_input, context
|
||||
)
|
||||
|
||||
if not target_steps:
|
||||
yield StepOutput(content=fallback_msg, success=True)
|
||||
return
|
||||
|
||||
steps_container = Steps(steps=target_steps, name=f"{self.name}_{branch_name}")
|
||||
capturer = StreamCapturer(steps_container.aexecute_stream(step_input, context))
|
||||
async for event in capturer:
|
||||
yield event
|
||||
if capturer.output:
|
||||
yield capturer.output
|
||||
|
||||
|
||||
class Condition(BranchingNode):
|
||||
"""根据条件函数的返回结果,决定走向 steps 还是 else_steps"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
evaluator: Any,
|
||||
evaluator: bool | Callable[..., bool | Awaitable[bool]],
|
||||
steps: Sequence[NodeSource],
|
||||
else_steps: Sequence[NodeSource] | None = None,
|
||||
name: str = "ConditionGroup",
|
||||
@@ -227,9 +259,9 @@ class Condition(BaseNode):
|
||||
def node_type(self) -> StepType:
|
||||
return StepType.CONDITION
|
||||
|
||||
async def run_stream(
|
||||
async def select_branch(
|
||||
self, step_input: StepInput, context: RunContext
|
||||
) -> AsyncIterator[Any]:
|
||||
) -> tuple[list[BaseNode], str, str]:
|
||||
if callable(self.evaluator):
|
||||
condition_result = await DependencyInjector.invoke(
|
||||
self.evaluator, {"step_input": step_input}, context
|
||||
@@ -238,32 +270,22 @@ class Condition(BaseNode):
|
||||
condition_result = bool(self.evaluator)
|
||||
|
||||
target_steps = self.steps if condition_result else self.else_steps
|
||||
branch_name = "if" if condition_result else "else"
|
||||
branch_name = "if_branch" if condition_result else "else_branch"
|
||||
fallback_msg = f"条件求值为 {condition_result},无对应步骤需执行。"
|
||||
|
||||
if not target_steps:
|
||||
yield StepOutput(
|
||||
content=f"条件求值为 {condition_result},无对应步骤需执行。",
|
||||
success=True,
|
||||
)
|
||||
return
|
||||
|
||||
steps_container = Steps(
|
||||
steps=target_steps, name=f"{self.name}_{branch_name}_branch"
|
||||
)
|
||||
out_box: list[StepOutput] = []
|
||||
async for event in self._forward_stream(
|
||||
steps_container.aexecute_stream(step_input, context), out_box
|
||||
):
|
||||
yield event
|
||||
if out_box:
|
||||
yield out_box[0]
|
||||
return target_steps, branch_name, fallback_msg
|
||||
|
||||
|
||||
class Router(BaseNode):
|
||||
class Router(BranchingNode):
|
||||
"""根据选择器函数的返回值(名称),从候选项中挑选步骤执行"""
|
||||
|
||||
def __init__(
|
||||
self, choices: Sequence[NodeSource], selector: Any, name: str = "RouterGroup"
|
||||
self,
|
||||
choices: Sequence[NodeSource],
|
||||
selector: str
|
||||
| list[str]
|
||||
| Callable[..., str | list[str] | Awaitable[str | list[str]]],
|
||||
name: str = "RouterGroup",
|
||||
):
|
||||
"""
|
||||
初始化选择路由器节点。
|
||||
@@ -285,9 +307,9 @@ class Router(BaseNode):
|
||||
def node_type(self) -> StepType:
|
||||
return StepType.ROUTER
|
||||
|
||||
async def run_stream(
|
||||
async def select_branch(
|
||||
self, step_input: StepInput, context: RunContext
|
||||
) -> AsyncIterator[Any]:
|
||||
) -> tuple[list[BaseNode], str, str]:
|
||||
if callable(self.selector):
|
||||
selected = await DependencyInjector.invoke(
|
||||
self.selector, {"step_input": step_input}, context
|
||||
@@ -308,18 +330,7 @@ class Router(BaseNode):
|
||||
else:
|
||||
target_steps.append(NodeFactory.build(s))
|
||||
|
||||
if not target_steps:
|
||||
yield StepOutput(content="没有命中任何有效路由分支。", success=True)
|
||||
return
|
||||
|
||||
steps_container = Steps(steps=target_steps, name=f"{self.name}_routed_steps")
|
||||
out_box: list[StepOutput] = []
|
||||
async for event in self._forward_stream(
|
||||
steps_container.aexecute_stream(step_input, context), out_box
|
||||
):
|
||||
yield event
|
||||
if out_box:
|
||||
yield out_box[0]
|
||||
return target_steps, "routed_steps", "没有命中任何有效路由分支。"
|
||||
|
||||
|
||||
class Loop(BaseNode):
|
||||
@@ -329,7 +340,7 @@ class Loop(BaseNode):
|
||||
self,
|
||||
steps: Sequence[NodeSource],
|
||||
max_iterations: int = 3,
|
||||
end_condition: Any = None,
|
||||
end_condition: bool | Callable[..., bool | Awaitable[bool]] | None = None,
|
||||
name: str = "LoopGroup",
|
||||
):
|
||||
"""
|
||||
@@ -353,7 +364,7 @@ class Loop(BaseNode):
|
||||
|
||||
async def run_stream(
|
||||
self, step_input: StepInput, context: RunContext
|
||||
) -> AsyncIterator[Any]:
|
||||
) -> AsyncIterator[AgentStreamEvent | StepOutput]:
|
||||
logger.debug(
|
||||
f" 🔁 开始循环: [Loop] `{self.name}` (最大 {self.max_iterations} 次)"
|
||||
)
|
||||
@@ -362,6 +373,7 @@ class Loop(BaseNode):
|
||||
all_results: list[StepOutput] = []
|
||||
current_input = StepInput(
|
||||
input=step_input.input,
|
||||
intent=step_input.intent,
|
||||
previous_step_content=step_input.previous_step_content,
|
||||
additional_data=step_input.additional_data.copy(),
|
||||
)
|
||||
@@ -372,12 +384,12 @@ class Loop(BaseNode):
|
||||
steps_container = Steps(
|
||||
steps=self.steps, name=f"{self.name}_iter_{iteration + 1}"
|
||||
)
|
||||
out_box: list[StepOutput] = []
|
||||
async for event in self._forward_stream(
|
||||
steps_container.aexecute_stream(current_input, context), out_box
|
||||
):
|
||||
capturer = StreamCapturer(
|
||||
steps_container.aexecute_stream(current_input, context)
|
||||
)
|
||||
async for event in capturer:
|
||||
yield event
|
||||
iter_output = out_box[0] if out_box else None
|
||||
iter_output = capturer.output
|
||||
|
||||
should_stop = False
|
||||
if iter_output:
|
||||
@@ -434,13 +446,13 @@ class Parallel(BaseNode):
|
||||
|
||||
async def run_stream(
|
||||
self, step_input: StepInput, context: RunContext
|
||||
) -> AsyncIterator[Any]:
|
||||
) -> AsyncIterator[AgentStreamEvent | StepOutput]:
|
||||
logger.debug(f" 🔀 [并发] `{self.name}` 开启了 {len(self.steps)} 个并发任务")
|
||||
|
||||
queue = asyncio.Queue()
|
||||
bg_tasks = []
|
||||
|
||||
async def worker(idx: int, s_obj: Any, c_ctx: RunContext):
|
||||
async def worker(idx: int, s_obj: BaseNode, c_ctx: RunContext):
|
||||
try:
|
||||
async for evt in s_obj.aexecute_stream(step_input, c_ctx):
|
||||
await queue.put(("event", evt))
|
||||
@@ -524,7 +536,7 @@ class NodeFactory:
|
||||
cls,
|
||||
executor: NodeSource,
|
||||
name: str | None = None,
|
||||
failure_policy: Any = None,
|
||||
failure_policy: BaseFailurePolicy | None = None,
|
||||
) -> BaseNode:
|
||||
"""底层物理实例化分发"""
|
||||
kwargs = {
|
||||
|
||||
@@ -1,10 +1,14 @@
|
||||
import copy
|
||||
from enum import Enum
|
||||
from typing import Any
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from .base import BaseNode
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from zhenxun.services.ai.llm.api import generate_structured
|
||||
from zhenxun.services.ai.run.context import RunContext
|
||||
from zhenxun.services.ai.utils.logger import log_flow as logger
|
||||
|
||||
from .types import StepInput
|
||||
@@ -26,7 +30,7 @@ class PolicyResult(BaseModel):
|
||||
"""执行延迟或重试前需要等待的缓冲秒数"""
|
||||
new_input: StepInput | None = None
|
||||
"""用于动态纠错自愈时替换传入的新参数结构"""
|
||||
fallback_node: Any | None = None
|
||||
fallback_node: "BaseNode | None" = None
|
||||
"""策略裁定降级时所指定的备用工作流节点"""
|
||||
healer_agent_name: str | None = None
|
||||
"""执行了高级自愈的大模型或修复者名称"""
|
||||
@@ -36,7 +40,11 @@ class BaseFailurePolicy:
|
||||
"""错误处理策略抽象基类"""
|
||||
|
||||
async def handle_failure(
|
||||
self, node: Any, exception: BaseException, step_input: StepInput, context: Any
|
||||
self,
|
||||
node: "BaseNode",
|
||||
exception: BaseException,
|
||||
step_input: StepInput,
|
||||
context: RunContext,
|
||||
) -> PolicyResult:
|
||||
"""
|
||||
处理节点执行失败的策略入口方法。
|
||||
@@ -57,7 +65,11 @@ class AbortPolicy(BaseFailurePolicy):
|
||||
"""直接中断策略"""
|
||||
|
||||
async def handle_failure(
|
||||
self, node: Any, exception: BaseException, step_input: StepInput, context: Any
|
||||
self,
|
||||
node: "BaseNode",
|
||||
exception: BaseException,
|
||||
step_input: StepInput,
|
||||
context: RunContext,
|
||||
) -> PolicyResult:
|
||||
return PolicyResult(action=PolicyAction.ABORT)
|
||||
|
||||
@@ -66,7 +78,11 @@ class SkipPolicy(BaseFailurePolicy):
|
||||
"""跳过并继续策略"""
|
||||
|
||||
async def handle_failure(
|
||||
self, node: Any, exception: BaseException, step_input: StepInput, context: Any
|
||||
self,
|
||||
node: "BaseNode",
|
||||
exception: BaseException,
|
||||
step_input: StepInput,
|
||||
context: RunContext,
|
||||
) -> PolicyResult:
|
||||
return PolicyResult(action=PolicyAction.CONTINUE)
|
||||
|
||||
@@ -86,7 +102,11 @@ class RetryPolicy(BaseFailurePolicy):
|
||||
self.delay = delay
|
||||
|
||||
async def handle_failure(
|
||||
self, node: Any, exception: BaseException, step_input: StepInput, context: Any
|
||||
self,
|
||||
node: "BaseNode",
|
||||
exception: BaseException,
|
||||
step_input: StepInput,
|
||||
context: RunContext,
|
||||
) -> PolicyResult:
|
||||
counts = context.state.setdefault("__retry_counts__", {})
|
||||
key = f"{node.name}_{id(self)}"
|
||||
@@ -100,7 +120,7 @@ class RetryPolicy(BaseFailurePolicy):
|
||||
class FallbackPolicy(BaseFailurePolicy):
|
||||
"""降级路由策略"""
|
||||
|
||||
def __init__(self, fallback_node: Any):
|
||||
def __init__(self, fallback_node: "BaseNode"):
|
||||
"""
|
||||
初始化降级路由策略。
|
||||
|
||||
@@ -110,7 +130,11 @@ class FallbackPolicy(BaseFailurePolicy):
|
||||
self.fallback_node = fallback_node
|
||||
|
||||
async def handle_failure(
|
||||
self, node: Any, exception: BaseException, step_input: StepInput, context: Any
|
||||
self,
|
||||
node: "BaseNode",
|
||||
exception: BaseException,
|
||||
step_input: StepInput,
|
||||
context: RunContext,
|
||||
) -> PolicyResult:
|
||||
return PolicyResult(
|
||||
action=PolicyAction.FALLBACK, fallback_node=self.fallback_node
|
||||
@@ -132,7 +156,11 @@ class SelfHealingPolicy(BaseFailurePolicy):
|
||||
self.max_retries = max_retries
|
||||
|
||||
async def handle_failure(
|
||||
self, node: Any, exception: BaseException, step_input: StepInput, context: Any
|
||||
self,
|
||||
node: "BaseNode",
|
||||
exception: BaseException,
|
||||
step_input: StepInput,
|
||||
context: RunContext,
|
||||
) -> PolicyResult:
|
||||
counts = context.state.setdefault("__heal_counts__", {})
|
||||
key = f"{node.name}_{id(self)}"
|
||||
|
||||
@@ -3,6 +3,8 @@ from typing import Any
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from zhenxun.services.ai.run.models import RunIntent
|
||||
|
||||
|
||||
class StepType(str, Enum):
|
||||
FUNCTION = "Function"
|
||||
@@ -20,6 +22,9 @@ class StepInput(BaseModel):
|
||||
input: Any = Field(default=None)
|
||||
"""继承自 Workflow 的初始输入"""
|
||||
|
||||
intent: RunIntent | None = Field(default=None)
|
||||
"""归一化后的意图载体,避免下游节点猜测解析 input 的原始类型"""
|
||||
|
||||
previous_step_content: Any = Field(default=None)
|
||||
"""上一个执行步骤产生的直接输出内容"""
|
||||
|
||||
|
||||
Reference in New Issue
Block a user