mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-30 21:45:08 +02:00
1stpass
This commit is contained in:
@@ -0,0 +1,13 @@
|
||||
from langgraph.graph_engine.state import (
|
||||
AdvancedStateGraph,
|
||||
CompiledGraphEngine,
|
||||
publish_to_channel,
|
||||
wait_for,
|
||||
)
|
||||
|
||||
__all__ = (
|
||||
"AdvancedStateGraph",
|
||||
"CompiledGraphEngine",
|
||||
"publish_to_channel",
|
||||
"wait_for",
|
||||
)
|
||||
@@ -0,0 +1,352 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import contextvars
|
||||
import inspect
|
||||
from collections.abc import Callable, Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Generic, TypeVar, cast
|
||||
|
||||
from langgraph.types import Command, Send
|
||||
|
||||
StateT = TypeVar("StateT")
|
||||
|
||||
_CURRENT_RUN: contextvars.ContextVar[_GraphEngineRun | None] = contextvars.ContextVar(
|
||||
"langgraph_graph_engine_run", default=None
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _ChannelSpec:
|
||||
typ: Any
|
||||
maxsize: int
|
||||
|
||||
|
||||
class AdvancedStateGraph(Generic[StateT]):
|
||||
"""Experimental in-memory graph engine with async channels."""
|
||||
|
||||
def __init__(self, state_schema: type[StateT]) -> None:
|
||||
self.state_schema = state_schema
|
||||
self._nodes: dict[str, Callable[..., Any]] = {}
|
||||
self._channels: dict[str, _ChannelSpec] = {}
|
||||
self._entry_point: str | None = None
|
||||
self._finish_point: str | None = None
|
||||
|
||||
def add_node(
|
||||
self,
|
||||
name_or_node: str | Callable[..., Any],
|
||||
node: Callable[..., Any] | None = None,
|
||||
) -> str:
|
||||
if node is None:
|
||||
if not callable(name_or_node):
|
||||
raise TypeError("add_node() expects a callable when name is omitted")
|
||||
node_name = _infer_node_name(name_or_node)
|
||||
node_fn = name_or_node
|
||||
else:
|
||||
if not isinstance(name_or_node, str):
|
||||
raise TypeError("add_node() expects a string node name")
|
||||
node_name = name_or_node
|
||||
node_fn = node
|
||||
|
||||
if node_name in self._nodes:
|
||||
raise ValueError(f"Node `{node_name}` already exists")
|
||||
self._nodes[node_name] = node_fn
|
||||
return node_name
|
||||
|
||||
def node(self, node: Callable[..., Any]) -> str:
|
||||
"""Register a node using the function name as node id."""
|
||||
return self.add_node(node)
|
||||
|
||||
def add_async_channel(
|
||||
self, name: str, typ: Any, maxsize: int | None = None
|
||||
) -> None:
|
||||
if name in self._channels:
|
||||
raise ValueError(f"Channel `{name}` already exists")
|
||||
self._channels[name] = _ChannelSpec(typ=typ, maxsize=maxsize or 0)
|
||||
|
||||
def set_entry_point(self, name_or_node: str | Callable[..., Any]) -> None:
|
||||
self._entry_point = self._resolve_node_name(name_or_node)
|
||||
|
||||
def set_finish_point(self, name_or_node: str | Callable[..., Any]) -> None:
|
||||
self._finish_point = self._resolve_node_name(name_or_node)
|
||||
|
||||
def add_entry_node(self, node: Callable[..., Any]) -> str:
|
||||
node_name = self.add_node(node)
|
||||
self.set_entry_point(node_name)
|
||||
return node_name
|
||||
|
||||
def add_finish_node(self, node: Callable[..., Any]) -> str:
|
||||
node_name = self.add_node(node)
|
||||
self.set_finish_point(node_name)
|
||||
return node_name
|
||||
|
||||
def _resolve_node_name(self, name_or_node: str | Callable[..., Any]) -> str:
|
||||
if isinstance(name_or_node, str):
|
||||
return name_or_node
|
||||
node_name = _infer_node_name(name_or_node)
|
||||
if node_name not in self._nodes:
|
||||
self._nodes[node_name] = name_or_node
|
||||
return node_name
|
||||
|
||||
def compile(self) -> CompiledGraphEngine[StateT]:
|
||||
if self._entry_point is None:
|
||||
raise ValueError("Entry point is not set")
|
||||
if self._finish_point is None:
|
||||
raise ValueError("Finish point is not set")
|
||||
if self._entry_point not in self._nodes:
|
||||
raise ValueError(f"Entry point node `{self._entry_point}` does not exist")
|
||||
if self._finish_point not in self._nodes:
|
||||
raise ValueError(f"Finish point node `{self._finish_point}` does not exist")
|
||||
return CompiledGraphEngine(
|
||||
nodes=dict(self._nodes),
|
||||
channels=dict(self._channels),
|
||||
entry_point=self._entry_point,
|
||||
finish_point=self._finish_point,
|
||||
)
|
||||
|
||||
|
||||
class CompiledGraphEngine(Generic[StateT]):
|
||||
"""Executable runtime for `AdvancedStateGraph`."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
nodes: dict[str, Callable[..., Any]],
|
||||
channels: dict[str, _ChannelSpec],
|
||||
entry_point: str,
|
||||
finish_point: str,
|
||||
) -> None:
|
||||
self._nodes = nodes
|
||||
self._channels = channels
|
||||
self._entry_point = entry_point
|
||||
self._finish_point = finish_point
|
||||
self._active_run: _GraphEngineRun | None = None
|
||||
self._run_lock = asyncio.Lock()
|
||||
|
||||
async def ainvoke(self, initial_state: StateT) -> StateT:
|
||||
async with self._run_lock:
|
||||
if self._active_run is not None:
|
||||
raise RuntimeError("Graph engine already has an active run")
|
||||
run = _GraphEngineRun(
|
||||
nodes=self._nodes,
|
||||
channel_specs=self._channels,
|
||||
entry_point=self._entry_point,
|
||||
finish_point=self._finish_point,
|
||||
)
|
||||
self._active_run = run
|
||||
try:
|
||||
return await run.run(initial_state)
|
||||
finally:
|
||||
async with self._run_lock:
|
||||
if self._active_run is run:
|
||||
self._active_run = None
|
||||
|
||||
async def apublish_to_channel(self, channel: str, value: Any) -> None:
|
||||
run = self._active_run
|
||||
if run is None:
|
||||
raise RuntimeError("No active graph run to publish to")
|
||||
await run.publish(channel, value)
|
||||
|
||||
|
||||
class _GraphEngineRun:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
nodes: dict[str, Callable[..., Any]],
|
||||
channel_specs: dict[str, _ChannelSpec],
|
||||
entry_point: str,
|
||||
finish_point: str,
|
||||
) -> None:
|
||||
self._nodes = nodes
|
||||
self._entry_point = entry_point
|
||||
self._finish_point = finish_point
|
||||
self._channels: dict[str, asyncio.Queue[Any]] = {
|
||||
name: asyncio.Queue(maxsize=spec.maxsize)
|
||||
for name, spec in channel_specs.items()
|
||||
}
|
||||
self._tasks: set[asyncio.Task[list[Send]]] = set()
|
||||
self._finished = False
|
||||
self._state: Any = None
|
||||
|
||||
async def run(self, initial_state: StateT) -> StateT:
|
||||
self._state = initial_state
|
||||
self._schedule(Send(self._entry_point, initial_state))
|
||||
try:
|
||||
while self._tasks and not self._finished:
|
||||
done, _ = await asyncio.wait(
|
||||
self._tasks, return_when=asyncio.FIRST_COMPLETED
|
||||
)
|
||||
for task in done:
|
||||
self._tasks.remove(task)
|
||||
exc = task.exception()
|
||||
if exc is not None:
|
||||
await self._cancel_all_tasks()
|
||||
raise exc
|
||||
sends = task.result()
|
||||
for send in sends:
|
||||
self._schedule(send)
|
||||
if self._finished:
|
||||
await self._cancel_all_tasks()
|
||||
return cast(StateT, self._state)
|
||||
finally:
|
||||
await self._cancel_all_tasks()
|
||||
|
||||
async def publish(self, channel: str, value: Any) -> None:
|
||||
queue = self._get_channel(channel)
|
||||
await queue.put(value)
|
||||
|
||||
def publish_nowait(self, channel: str, value: Any) -> None:
|
||||
queue = self._get_channel(channel)
|
||||
queue.put_nowait(value)
|
||||
|
||||
async def wait_for(self, channel: str, n: int = 1) -> Any:
|
||||
if n < 1:
|
||||
raise ValueError("wait_for count `n` must be >= 1")
|
||||
queue = self._get_channel(channel)
|
||||
if n == 1:
|
||||
return await queue.get()
|
||||
values: list[Any] = []
|
||||
for _ in range(n):
|
||||
values.append(await queue.get())
|
||||
return values
|
||||
|
||||
def _get_channel(self, channel: str) -> asyncio.Queue[Any]:
|
||||
if channel not in self._channels:
|
||||
raise ValueError(f"Unknown channel `{channel}`")
|
||||
return self._channels[channel]
|
||||
|
||||
def _schedule(self, send: Send) -> None:
|
||||
if self._finished:
|
||||
return
|
||||
task: asyncio.Task[list[Send]] = asyncio.create_task(self._execute_send(send))
|
||||
self._tasks.add(task)
|
||||
|
||||
async def _cancel_all_tasks(self) -> None:
|
||||
if not self._tasks:
|
||||
return
|
||||
to_cancel = list(self._tasks)
|
||||
for task in to_cancel:
|
||||
task.cancel()
|
||||
await asyncio.gather(*to_cancel, return_exceptions=True)
|
||||
self._tasks.clear()
|
||||
|
||||
async def _execute_send(self, send: Send) -> list[Send]:
|
||||
node_name = _resolve_target_name(send.node)
|
||||
if node_name not in self._nodes:
|
||||
raise ValueError(f"Unknown node `{node_name}`")
|
||||
node = self._nodes[node_name]
|
||||
|
||||
token = _CURRENT_RUN.set(self)
|
||||
try:
|
||||
result = node(send.arg)
|
||||
if inspect.isawaitable(result):
|
||||
result = await result
|
||||
finally:
|
||||
_CURRENT_RUN.reset(token)
|
||||
|
||||
if isinstance(result, Command):
|
||||
self._apply_update(result.update)
|
||||
next_sends = _normalize_goto(result.goto, default_arg=self._state)
|
||||
else:
|
||||
self._apply_update(result)
|
||||
next_sends = _normalize_result_to_sends(result, default_arg=self._state)
|
||||
|
||||
if node_name == self._finish_point:
|
||||
self._finished = True
|
||||
return []
|
||||
return next_sends
|
||||
|
||||
def _apply_update(self, update: Any) -> None:
|
||||
if update is None:
|
||||
return
|
||||
if isinstance(update, Mapping):
|
||||
if isinstance(self._state, Mapping):
|
||||
# Keep semantics simple: in-place update for mapping-like state.
|
||||
cast(dict[str, Any], self._state).update(update)
|
||||
return
|
||||
if isinstance(update, Sequence) and not isinstance(update, (str, bytes)):
|
||||
pairs = list(update)
|
||||
if all(
|
||||
isinstance(item, tuple) and len(item) == 2 and isinstance(item[0], str)
|
||||
for item in pairs
|
||||
):
|
||||
if isinstance(self._state, Mapping):
|
||||
cast(dict[str, Any], self._state).update(
|
||||
cast(dict[str, Any], pairs)
|
||||
)
|
||||
return
|
||||
|
||||
|
||||
def _normalize_result_to_sends(result: Any, *, default_arg: Any) -> list[Send]:
|
||||
if result is None:
|
||||
return []
|
||||
if isinstance(result, Send):
|
||||
return [result]
|
||||
if callable(result):
|
||||
return [Send(_infer_node_name(result), default_arg)]
|
||||
if isinstance(result, str):
|
||||
return [Send(result, default_arg)]
|
||||
if isinstance(result, Sequence) and not isinstance(result, (str, bytes)):
|
||||
sends: list[Send] = []
|
||||
for item in result:
|
||||
if isinstance(item, Send):
|
||||
sends.append(item)
|
||||
elif callable(item):
|
||||
sends.append(Send(_infer_node_name(item), default_arg))
|
||||
elif isinstance(item, str):
|
||||
sends.append(Send(item, default_arg))
|
||||
return sends
|
||||
return []
|
||||
|
||||
|
||||
def _normalize_goto(goto: Any, *, default_arg: Any) -> list[Send]:
|
||||
if not goto:
|
||||
return []
|
||||
if isinstance(goto, Send):
|
||||
return [goto]
|
||||
if callable(goto):
|
||||
return [Send(_infer_node_name(goto), default_arg)]
|
||||
if isinstance(goto, str):
|
||||
return [Send(goto, default_arg)]
|
||||
if isinstance(goto, Sequence):
|
||||
sends: list[Send] = []
|
||||
for item in goto:
|
||||
if isinstance(item, Send):
|
||||
sends.append(item)
|
||||
elif callable(item):
|
||||
sends.append(Send(_infer_node_name(item), default_arg))
|
||||
elif isinstance(item, str):
|
||||
sends.append(Send(item, default_arg))
|
||||
return sends
|
||||
return []
|
||||
|
||||
|
||||
async def wait_for(channel: str, n: int = 1) -> Any:
|
||||
run = _CURRENT_RUN.get()
|
||||
if run is None:
|
||||
raise RuntimeError("wait_for() can only be used inside graph_engine nodes")
|
||||
return await run.wait_for(channel, n=n)
|
||||
|
||||
|
||||
def publish_to_channel(channel: str, value: Any) -> None:
|
||||
run = _CURRENT_RUN.get()
|
||||
if run is None:
|
||||
raise RuntimeError(
|
||||
"publish_to_channel() can only be used inside graph_engine nodes"
|
||||
)
|
||||
run.publish_nowait(channel, value)
|
||||
|
||||
|
||||
def _infer_node_name(node: Callable[..., Any]) -> str:
|
||||
node_name = getattr(node, "__name__", "")
|
||||
if not node_name or node_name == "<lambda>":
|
||||
raise ValueError("Cannot infer node name from anonymous callable")
|
||||
return node_name
|
||||
|
||||
|
||||
def _resolve_target_name(target: Any) -> str:
|
||||
if isinstance(target, str):
|
||||
return target
|
||||
if callable(target):
|
||||
return _infer_node_name(target)
|
||||
raise ValueError(f"Unsupported node target type: {type(target)!r}")
|
||||
@@ -1,124 +1,174 @@
|
||||
from typing import Annotated
|
||||
import asyncio
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Literal
|
||||
|
||||
import pytest
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from langgraph.constants import END, START
|
||||
from langgraph.graph import StateGraph
|
||||
from langgraph.graph_engine import AdvancedStateGraph, publish_to_channel, wait_for
|
||||
from langgraph.types import Command, Send
|
||||
|
||||
pytestmark = pytest.mark.anyio
|
||||
|
||||
|
||||
class MainAgentState(TypedDict):
|
||||
input: str
|
||||
output: list[str]
|
||||
|
||||
advanced_flow = AdvancedStateGraph(MainAgentState)
|
||||
advanced_flow.add_async_channel("inbox", str) # default to infinite buffer like Rust channel
|
||||
|
||||
prompt = "based on current state, decide to use tool, or kick off subagent, or complete"
|
||||
main_llm = new_mock_llm(prompt, [email_tool, slack_tool])
|
||||
|
||||
async def llm_node(state: MainAgentState):
|
||||
decisions = await main_llm.ainvoke(state)
|
||||
sends = []
|
||||
for decision in decisions:
|
||||
if decision.type == "end":
|
||||
return Command(goto=Send("end", decision.complete)) # NOTE: we can simplify to just Complete(decision.complete)
|
||||
elif decision.type == "sub_agent":
|
||||
sends.append(Send("sub_agent", decision.sub_agent)) # NOTE: we can simplify to just sends.apppend(subagent, decision.sub_agent)
|
||||
elif decision.type == "tool":
|
||||
sends.append(Send("tool", decision.tool))
|
||||
sends.append(Send("wait_node"))
|
||||
return Command(goto=sends)
|
||||
|
||||
async def wait_node(state: MainAgentState):
|
||||
# wait for at least one message on the inbox channel
|
||||
# this is a "lightweight" interrupt, that does not block the entire graph
|
||||
msgs = wait_for("inbox")
|
||||
|
||||
output = state.output
|
||||
if msg.type == "tool":
|
||||
output.append("tool: " + msg.payload)
|
||||
elif msg.type == "sub_agent":
|
||||
output.append("sub_agent: " + msg.payload)
|
||||
elif msg.type == "user_input":
|
||||
output.append("user_input: " + msg.payload)
|
||||
|
||||
# loop back to llm node with new output
|
||||
return Command(goto=Send("llm_node", output))
|
||||
done: str | None
|
||||
|
||||
|
||||
async def tool_node(tool_input: str):
|
||||
await asyncio.sleep(5)
|
||||
output = "tool completed for: " + tool_input
|
||||
publish_to_channel("inbox", {"type": "tool", "payload": output})
|
||||
# just complete without going to next node
|
||||
class SubAgentState(TypedDict):
|
||||
input: str
|
||||
output: str
|
||||
|
||||
async def order_food_node(state: str):
|
||||
return "order_food_node completed for: " + state
|
||||
|
||||
# sub agent uses regular/simple state graph
|
||||
sub_agent = StateGraph(str)
|
||||
async def research_node (state: str):
|
||||
await asyncio.sleep(10)
|
||||
return "research sub agent completed for: " + state
|
||||
sub_agent.add_node("research_node", research_node)
|
||||
sub_agent.add_edge(START, "research_node")
|
||||
sub_agent.add_edge("research_node", END)
|
||||
@dataclass(frozen=True)
|
||||
class Decision:
|
||||
type: Literal["end", "sub_agent", "tool"]
|
||||
sub_agent: str | None = None
|
||||
tool: str | None = None
|
||||
complete: str | None = None
|
||||
|
||||
async def sub_agent_node(sub_agent_input: str):
|
||||
sub_agent_output = sub_agent.invoke({"input": sub_agent_input})
|
||||
publish_to_channel("inbox", {"type": "sub_agent", "payload": sub_agent_output})
|
||||
# just complete without going to next node
|
||||
|
||||
advanced_flow.add_node("llm_node", llm_node)
|
||||
advanced_flow.add_node("wait_node", wait_node)
|
||||
advanced_flow.add_node("tool_node", tool_node)
|
||||
advanced_flow.add_node("sub_agent_node", sub_agent_node)
|
||||
advanced_flow.add_node("order_food_node", order_food_node)
|
||||
advanced_flow.set_entry_point("llm_node")
|
||||
advanced_flow.set_finish_point("order_food_node")
|
||||
class MockPlanner:
|
||||
def __init__(self) -> None:
|
||||
self.responses: list[list[Decision]] = []
|
||||
self._idx = 0
|
||||
|
||||
## NOTE: above can be simplified to:
|
||||
# advanced_flow.add_entry_node(llm_node)
|
||||
# advanced_flow.node(wait_node)
|
||||
# advanced_flow.node(tool_node)
|
||||
# advanced_flow.node(sub_agent_node)
|
||||
# advanced_flow.add_finish_node(order_food_node)
|
||||
async def ainvoke(self, _: MainAgentState) -> list[Decision]:
|
||||
if self._idx >= len(self.responses):
|
||||
return []
|
||||
response = self.responses[self._idx]
|
||||
self._idx += 1
|
||||
return response
|
||||
|
||||
main_agent = advanced_flow.compile()
|
||||
|
||||
async def test_async_sub_graph():
|
||||
main_llm.mock_response = [
|
||||
# first llm invoke
|
||||
def build_sub_agent() -> Any:
|
||||
# Sub-agent uses the regular/simple StateGraph API.
|
||||
sub_agent = StateGraph(SubAgentState)
|
||||
|
||||
async def research_node(state: SubAgentState) -> dict[str, str]:
|
||||
# Make timing deterministic for the prototype flow assertions.
|
||||
if state["input"] == "research lunch options":
|
||||
await asyncio.sleep(0.05)
|
||||
else:
|
||||
await asyncio.sleep(0.09)
|
||||
return {"output": f"research sub agent completed for: {state['input']}"}
|
||||
|
||||
sub_agent.add_node("research_node", research_node)
|
||||
sub_agent.add_edge(START, "research_node")
|
||||
sub_agent.add_edge("research_node", END)
|
||||
return sub_agent.compile()
|
||||
|
||||
|
||||
async def test_async_sub_graph() -> None:
|
||||
planner = MockPlanner()
|
||||
sub_agent = build_sub_agent()
|
||||
|
||||
advanced_flow = AdvancedStateGraph(MainAgentState)
|
||||
# Default behavior is an unbounded async channel (maxsize=None).
|
||||
advanced_flow.add_async_channel("inbox", dict)
|
||||
|
||||
async def llm_node(state: MainAgentState) -> Command:
|
||||
# Planner decides whether to call a tool, spawn a sub-agent, or finish.
|
||||
decisions = await planner.ainvoke(state)
|
||||
sends: list[Send] = []
|
||||
for decision in decisions:
|
||||
if decision.type == "end":
|
||||
# NOTE: this can be simplified further in the future with a dedicated
|
||||
# complete primitive, instead of routing to a finish node manually.
|
||||
return Command(
|
||||
goto=Send(
|
||||
order_food_node,
|
||||
decision.complete or "order flow completed",
|
||||
)
|
||||
)
|
||||
if decision.type == "sub_agent" and decision.sub_agent:
|
||||
sends.append(Send("sub_agent_node", decision.sub_agent))
|
||||
if decision.type == "tool" and decision.tool:
|
||||
sends.append(Send("tool_node", decision.tool))
|
||||
# Keep the main loop responsive: wait for one inbound message and continue.
|
||||
sends.append(Send("wait_node", state))
|
||||
return Command(goto=sends)
|
||||
|
||||
async def wait_node(state: MainAgentState) -> Command:
|
||||
# Lightweight interrupt: only this node blocks on inbox.
|
||||
msg = await wait_for("inbox")
|
||||
state["output"].append(f"{msg['type']}: {msg['payload']}")
|
||||
# Loop back to planner with updated output.
|
||||
return Command(goto=Send("llm_node", state))
|
||||
|
||||
async def tool_node(tool_input: str) -> None:
|
||||
await asyncio.sleep(0.03)
|
||||
# Fire-and-forget style completion: publish result to inbox and exit.
|
||||
# (i.e., just complete without explicitly going to a next node)
|
||||
publish_to_channel(
|
||||
"inbox",
|
||||
{"type": "tool", "payload": f"tool completed for: {tool_input}"},
|
||||
)
|
||||
|
||||
async def sub_agent_node(sub_agent_input: str) -> None:
|
||||
# Sub-agent remains a regular StateGraph, compiled independently.
|
||||
sub_agent_output = await sub_agent.ainvoke(
|
||||
{"input": sub_agent_input, "output": ""}
|
||||
)
|
||||
# Same pattern as tool node: publish result and complete current node.
|
||||
publish_to_channel(
|
||||
"inbox",
|
||||
{"type": "sub_agent", "payload": sub_agent_output["output"]},
|
||||
)
|
||||
|
||||
async def order_food_node(complete_message: str) -> dict[str, str]:
|
||||
return {"done": complete_message}
|
||||
|
||||
# NOTE: This is the simplified API shape:
|
||||
# - add_entry_node(fn)
|
||||
# - node(fn)
|
||||
# - add_finish_node(fn)
|
||||
# where node names are inferred from function names via reflection.
|
||||
advanced_flow.add_entry_node(llm_node)
|
||||
advanced_flow.node(wait_node)
|
||||
advanced_flow.node(tool_node)
|
||||
advanced_flow.node(sub_agent_node)
|
||||
advanced_flow.add_finish_node(order_food_node)
|
||||
main_agent = advanced_flow.compile()
|
||||
|
||||
planner.responses = [
|
||||
[
|
||||
{
|
||||
"type": "sub_agent",
|
||||
"sub_agent": "research_node"
|
||||
},
|
||||
{
|
||||
"type": "tool",
|
||||
"tool": "slack_tool"
|
||||
}
|
||||
# First planner pass triggers one sub-agent + one tool.
|
||||
Decision(type="sub_agent", sub_agent="research lunch options"),
|
||||
Decision(type="tool", tool="slack_tool"),
|
||||
],
|
||||
# 2nd llm invoke, after additional user input
|
||||
[
|
||||
{
|
||||
"type": "sub_agent",
|
||||
}
|
||||
],
|
||||
# 3rd llm invoke, after tool node completes
|
||||
[],
|
||||
# 4th llm invoke, after 1st sub agent node completes
|
||||
# Second planner pass triggers another sub-agent.
|
||||
[Decision(type="sub_agent", sub_agent="find vegetarian fallback")],
|
||||
[],
|
||||
# 5th llm invoke, after 2nd sub agent node completes
|
||||
[
|
||||
{
|
||||
"type": "end"
|
||||
}
|
||||
],
|
||||
[],
|
||||
# Final pass decides to end.
|
||||
[Decision(type="end", complete="order submitted")],
|
||||
]
|
||||
|
||||
started = main_agent.ainvoke({"input": "help me get something for lunch"})
|
||||
# provide more info after 2 seconds
|
||||
await asyncio.sleep(2)
|
||||
main_agent.apublish_to_channel("inbox", {"type": "user_input", "payload": "No spicy food please"})
|
||||
started = asyncio.create_task(
|
||||
main_agent.ainvoke(
|
||||
{"input": "help me get something for lunch", "output": [], "done": None}
|
||||
)
|
||||
)
|
||||
|
||||
# External input can be injected while graph execution is in progress.
|
||||
await asyncio.sleep(0.01)
|
||||
await main_agent.apublish_to_channel(
|
||||
"inbox", {"type": "user_input", "payload": "No spicy food please"}
|
||||
)
|
||||
result = await started
|
||||
|
||||
assert result == {
|
||||
"input": "help me get something for lunch",
|
||||
"output":[
|
||||
|
||||
]
|
||||
}
|
||||
"output": [
|
||||
"user_input: No spicy food please",
|
||||
"tool: tool completed for: slack_tool",
|
||||
"sub_agent: research sub agent completed for: research lunch options",
|
||||
"sub_agent: research sub agent completed for: find vegetarian fallback",
|
||||
],
|
||||
"done": "order submitted",
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user