♻️ 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:
Rumio
2026-07-14 16:48:33 +08:00
committed by GitHub
co-authored by webjoin111 pre-commit-ci[bot]
parent 922d092650
commit 52f7dbdedf
66 changed files with 2131 additions and 2353 deletions
+23 -15
View File
@@ -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
+39 -90
View File
@@ -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 实例"""
+95 -83
View File
@@ -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 = {
+37 -9
View File
@@ -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)
"""上一个执行步骤产生的直接输出内容"""