mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-20 06:35:46 +02:00
878 lines
30 KiB
Python
878 lines
30 KiB
Python
from __future__ import annotations
|
|
|
|
import atexit
|
|
import asyncio
|
|
import inspect
|
|
import os
|
|
import threading
|
|
from collections.abc import Callable, Coroutine, Sequence
|
|
from concurrent.futures import ThreadPoolExecutor
|
|
from dataclasses import dataclass
|
|
from datetime import timedelta
|
|
from typing import Any, Generic, TypeVar, cast
|
|
|
|
from saf_python_sdk.langgraph_rust_core import PyRustEngine # type: ignore[import-untyped]
|
|
|
|
from saf_python_sdk.types import Command, Send
|
|
|
|
StateT = TypeVar("StateT")
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class _ChannelSpec:
|
|
typ: Any
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class NodeStateOption:
|
|
locked_fields: tuple[str, ...] = ()
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ChannelCondition:
|
|
channel: str
|
|
min: int = 1
|
|
max: int = 0
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class TimerCondition:
|
|
seconds: float
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class AnyOfCondition:
|
|
conditions: tuple[WaitCondition, ...]
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class AllOfCondition:
|
|
conditions: tuple[WaitCondition, ...]
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ConditionResult:
|
|
met: bool
|
|
channel_name: str | None = None
|
|
values: list[Any] | None = None
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class WaitForResult:
|
|
conditions: list[ConditionResult]
|
|
|
|
|
|
WaitCondition = ChannelCondition | TimerCondition
|
|
|
|
_EXECUTOR_LOCK = threading.Lock()
|
|
_EXECUTOR: ThreadPoolExecutor | None = None
|
|
|
|
|
|
def _advanced_graph_executor() -> ThreadPoolExecutor:
|
|
global _EXECUTOR
|
|
with _EXECUTOR_LOCK:
|
|
if _EXECUTOR is None:
|
|
worker_count = int(os.getenv("LANGGRAPH_ADVANCED_GRAPH_PY_THREADS", "256"))
|
|
worker_count = max(worker_count, 1)
|
|
_EXECUTOR = ThreadPoolExecutor(
|
|
max_workers=worker_count,
|
|
thread_name_prefix="saf-advanced-py",
|
|
)
|
|
atexit.register(_shutdown_advanced_graph_executor)
|
|
return _EXECUTOR
|
|
|
|
|
|
def _shutdown_advanced_graph_executor() -> None:
|
|
global _EXECUTOR
|
|
with _EXECUTOR_LOCK:
|
|
if _EXECUTOR is not None:
|
|
_EXECUTOR.shutdown(wait=False, cancel_futures=False)
|
|
_EXECUTOR = None
|
|
|
|
|
|
class WaitRequested(Exception):
|
|
def __init__(self, payload: dict[str, Any]) -> None:
|
|
super().__init__("wait requested")
|
|
self.payload = payload
|
|
|
|
|
|
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._node_options: dict[str, NodeStateOption] = {}
|
|
self._async_channels: dict[str, _ChannelSpec] = {}
|
|
self._custom_output_streams: 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,
|
|
*,
|
|
state_option: dict[str, Any] | NodeStateOption | 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
|
|
self._node_options[node_name] = _normalize_node_state_option(state_option)
|
|
return node_name
|
|
|
|
def add_async_channel(self, name: str, typ: Any) -> None:
|
|
if name in self._async_channels:
|
|
raise ValueError(f"Channel `{name}` already exists")
|
|
self._async_channels[name] = _ChannelSpec(typ=typ)
|
|
|
|
def add_custom_outout_stream(self, name: str, typ: Any) -> None:
|
|
if name in self._custom_output_streams:
|
|
raise ValueError(f"Custom output stream `{name}` already exists")
|
|
self._custom_output_streams[name] = _ChannelSpec(typ=typ)
|
|
|
|
# Alias with corrected spelling.
|
|
def add_custom_output_stream(self, name: str, typ: Any) -> None:
|
|
self.add_custom_outout_stream(name, typ)
|
|
|
|
def add_entry_node(
|
|
self,
|
|
node: Callable[..., Any],
|
|
*,
|
|
state_option: dict[str, Any] | NodeStateOption | None = None,
|
|
) -> str:
|
|
node_name = self.add_node(node, state_option=state_option)
|
|
self._entry_point = self._resolve_node_name(node_name)
|
|
return node_name
|
|
|
|
def add_finish_node(
|
|
self,
|
|
node: Callable[..., Any],
|
|
*,
|
|
state_option: dict[str, Any] | NodeStateOption | None = None,
|
|
) -> str:
|
|
node_name = self.add_node(node, state_option=state_option)
|
|
self._finish_point = self._resolve_node_name(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._entry_point not in self._nodes:
|
|
raise ValueError(f"Entry point node `{self._entry_point}` does not exist")
|
|
if self._finish_point is not None and 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),
|
|
node_options=dict(self._node_options),
|
|
async_channels=dict(self._async_channels),
|
|
custom_output_streams=dict(self._custom_output_streams),
|
|
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]],
|
|
node_options: dict[str, NodeStateOption],
|
|
async_channels: dict[str, _ChannelSpec],
|
|
custom_output_streams: dict[str, _ChannelSpec],
|
|
entry_point: str,
|
|
finish_point: str | None,
|
|
) -> None:
|
|
self._nodes = nodes
|
|
self._node_options = node_options
|
|
self._async_channels = async_channels
|
|
self._custom_output_streams = custom_output_streams
|
|
self._entry_point = entry_point
|
|
self._finish_point = finish_point
|
|
|
|
async def ainvoke(self, initial_state: StateT) -> StateT:
|
|
handler = await self.astart(initial_state)
|
|
return await handler
|
|
|
|
async def astart(
|
|
self, initial_state: StateT, *, stream_mode: str | None = None
|
|
) -> GraphRunHandler[StateT]:
|
|
run = _GraphEngineRun(
|
|
nodes=self._nodes,
|
|
async_channel_specs=self._async_channels,
|
|
custom_output_stream_specs=self._custom_output_streams,
|
|
node_options=self._node_options,
|
|
entry_point=self._entry_point,
|
|
finish_point=self._finish_point,
|
|
stream_mode=stream_mode,
|
|
)
|
|
task = asyncio.create_task(run.run(initial_state))
|
|
return GraphRunHandler(run=run, task=task)
|
|
|
|
|
|
class Context:
|
|
"""Per-run context injected into advanced graph nodes."""
|
|
|
|
def __init__(self, run: _GraphEngineRun) -> None:
|
|
self._run = run
|
|
|
|
async def wait_for(
|
|
self, target: WaitCondition | AnyOfCondition | AllOfCondition
|
|
) -> WaitForResult:
|
|
resumed = self._run._consume_resume_event(target)
|
|
if resumed is not None:
|
|
return resumed
|
|
raise WaitRequested(_target_to_suspend_payload(target))
|
|
|
|
def is_resume(self) -> bool:
|
|
return self._run._is_resume_execution()
|
|
|
|
# Go-style alias.
|
|
def IsResume(self) -> bool:
|
|
return self.is_resume()
|
|
|
|
def publish_to_channel(self, channel: str, value: Any) -> None:
|
|
self._run.publish_nowait(channel, value)
|
|
|
|
async def apublish_to_channel(self, channel: str, value: Any) -> None:
|
|
await self._run.publish(channel, value)
|
|
|
|
def send_custom_stream_event(self, stream_name: str, value: Any) -> None:
|
|
self._run.send_custom_stream_event(stream_name, value)
|
|
|
|
|
|
class GraphRunHandler(Generic[StateT]):
|
|
"""Handle for an active in-memory run."""
|
|
|
|
def __init__(self, *, run: _GraphEngineRun, task: asyncio.Task[StateT]) -> None:
|
|
self._run = run
|
|
self._task = task
|
|
|
|
async def apublish_to_channel(self, channel: str, value: Any) -> None:
|
|
if self._task.done():
|
|
raise RuntimeError("Run has already completed")
|
|
await self._run.publish(channel, value)
|
|
|
|
async def receive_stream(self, stream_name: str) -> Any | None:
|
|
# Use a separate thread pool from graph execution to avoid deadlock
|
|
# when LANGGRAPH_ADVANCED_GRAPH_PY_THREADS is configured to 1.
|
|
return await asyncio.to_thread(self._run.receive_stream_sync, stream_name)
|
|
|
|
def close_all_streams(self) -> None:
|
|
self._run.close_all_streams_sync()
|
|
|
|
async def aresult(self) -> StateT:
|
|
return await self._task
|
|
|
|
def __await__(self) -> Any:
|
|
return self._task.__await__()
|
|
|
|
|
|
class _GraphEngineRun:
|
|
def __init__(
|
|
self,
|
|
*,
|
|
nodes: dict[str, Callable[..., Any]],
|
|
async_channel_specs: dict[str, _ChannelSpec],
|
|
custom_output_stream_specs: dict[str, _ChannelSpec],
|
|
node_options: dict[str, NodeStateOption],
|
|
entry_point: str,
|
|
finish_point: str | None,
|
|
stream_mode: str | None,
|
|
) -> None:
|
|
self._nodes = nodes
|
|
self._entry_point = entry_point
|
|
self._node_options = node_options
|
|
self._finish_point = finish_point
|
|
self._stream_mode = stream_mode
|
|
self._rust_engine = PyRustEngine()
|
|
for name in async_channel_specs:
|
|
self._rust_engine.add_async_channel(name)
|
|
for stream_name in custom_output_stream_specs:
|
|
self._rust_engine.add_custom_output_stream(stream_name)
|
|
self._stream_ready = threading.Event()
|
|
if self._stream_mode is None:
|
|
self._stream_ready.set()
|
|
self._tasks: set[asyncio.Task[list[Send]]] = set()
|
|
self._finished = False
|
|
self._state: Any = None
|
|
self._local = threading.local()
|
|
self.context = Context(self)
|
|
|
|
async def run(self, initial_state: StateT) -> StateT:
|
|
finish_point = self._finish_point or ""
|
|
if self._stream_mode is not None:
|
|
try:
|
|
self._rust_engine.start_stream(self._stream_mode)
|
|
finally:
|
|
self._stream_ready.set()
|
|
loop = asyncio.get_running_loop()
|
|
try:
|
|
result_obj = await loop.run_in_executor(
|
|
_advanced_graph_executor(),
|
|
self._rust_engine.run_graph_py,
|
|
self._entry_point,
|
|
finish_point,
|
|
initial_state,
|
|
self._execute_node_for_rust,
|
|
None,
|
|
_node_locked_fields_payload(self._node_options),
|
|
)
|
|
finally:
|
|
self._stream_ready.set()
|
|
self._state = result_obj
|
|
return cast(StateT, self._state)
|
|
|
|
async def publish(self, channel: str, value: Any) -> None:
|
|
loop = asyncio.get_running_loop()
|
|
await loop.run_in_executor(
|
|
_advanced_graph_executor(),
|
|
self._publish_sync,
|
|
channel,
|
|
value,
|
|
)
|
|
|
|
def publish_nowait(self, channel: str, value: Any) -> None:
|
|
self._publish_sync(channel, value)
|
|
|
|
async def wait_for(
|
|
self, target: WaitCondition | AnyOfCondition | AllOfCondition
|
|
) -> WaitForResult:
|
|
if isinstance(target, ChannelCondition):
|
|
value = await self._wait_for_channel_values(
|
|
target.channel, min=target.min, max=target.max
|
|
)
|
|
return WaitForResult(
|
|
conditions=[
|
|
ConditionResult(
|
|
met=True,
|
|
channel_name=target.channel,
|
|
values=_normalize_channel_values(value),
|
|
)
|
|
]
|
|
)
|
|
if isinstance(target, TimerCondition):
|
|
loop = asyncio.get_running_loop()
|
|
await loop.run_in_executor(
|
|
_advanced_graph_executor(),
|
|
self._rust_engine.wait_timer,
|
|
target.seconds,
|
|
)
|
|
return WaitForResult(conditions=[ConditionResult(met=True)])
|
|
if isinstance(target, AnyOfCondition):
|
|
raw_event = await self._wait_for_any_of(target)
|
|
return _wait_for_result_from_any_of_event(target, raw_event)
|
|
if isinstance(target, AllOfCondition):
|
|
raw_event = await self._wait_for_all_of(target)
|
|
return _wait_for_result_from_all_of_event(target, raw_event)
|
|
raise ValueError(f"Unsupported wait condition type: {type(target)!r}")
|
|
|
|
async def _wait_for_channel_values(
|
|
self, channel: str, min: int, max: int = 0
|
|
) -> Any:
|
|
if min < 1:
|
|
raise ValueError("wait_for count `min` must be >= 1")
|
|
if max < 0:
|
|
raise ValueError("wait_for max count `max` must be >= 0")
|
|
if max != 0 and max < min:
|
|
raise ValueError("wait_for max count `max` must be 0 or >= min")
|
|
loop = asyncio.get_running_loop()
|
|
event = await loop.run_in_executor(
|
|
_advanced_graph_executor(),
|
|
self._rust_engine.wait_channel,
|
|
channel,
|
|
min,
|
|
max,
|
|
)
|
|
return event["value"]
|
|
|
|
async def _wait_for_any_of(self, condition: AnyOfCondition) -> dict[str, Any]:
|
|
if not condition.conditions:
|
|
raise ValueError("any_of() requires at least one condition")
|
|
payload = {
|
|
"conditions": [_condition_to_rust(cond) for cond in condition.conditions]
|
|
}
|
|
loop = asyncio.get_running_loop()
|
|
return await loop.run_in_executor(
|
|
_advanced_graph_executor(),
|
|
self._rust_engine.wait_any_of_obj,
|
|
payload,
|
|
)
|
|
|
|
async def _wait_for_all_of(self, condition: AllOfCondition) -> dict[str, Any]:
|
|
if not condition.conditions:
|
|
raise ValueError("all_of() requires at least one condition")
|
|
payload = {
|
|
"conditions": [_condition_to_rust(cond) for cond in condition.conditions]
|
|
}
|
|
loop = asyncio.get_running_loop()
|
|
return await loop.run_in_executor(
|
|
_advanced_graph_executor(),
|
|
self._rust_engine.wait_all_of_obj,
|
|
payload,
|
|
)
|
|
|
|
def _publish_sync(self, channel: str, value: Any) -> None:
|
|
self._rust_engine.publish_obj(channel, value)
|
|
|
|
def send_custom_stream_event(self, stream_name: str, value: Any) -> None:
|
|
self._rust_engine.send_custom_stream_event_obj(stream_name, value)
|
|
|
|
def receive_stream_sync(self, stream_name: str) -> Any | None:
|
|
self._stream_ready.wait()
|
|
return self._rust_engine.receive_stream_obj(stream_name)
|
|
|
|
def close_all_streams_sync(self) -> None:
|
|
self._rust_engine.close_all_streams()
|
|
|
|
def _execute_node_for_rust(
|
|
self, node_name: str, node_input: Any, state: Any
|
|
) -> dict[str, Any]:
|
|
node_input, resume_event = _unwrap_resume_input(node_input)
|
|
self._set_resume_event(resume_event)
|
|
self._set_is_resume(resume_event is not None)
|
|
if node_name not in self._nodes:
|
|
raise ValueError(f"Unknown node `{node_name}`")
|
|
node = self._nodes[node_name]
|
|
try:
|
|
result = _invoke_node(node, self.context, node_input, state)
|
|
if inspect.isawaitable(result):
|
|
result = self._run_awaitable_in_worker(
|
|
cast(Coroutine[Any, Any, Any], result)
|
|
)
|
|
except WaitRequested as suspend:
|
|
return {"suspend": suspend.payload}
|
|
finally:
|
|
self._set_resume_event(None)
|
|
self._set_is_resume(False)
|
|
|
|
if isinstance(result, Command):
|
|
update = result.update
|
|
sends = _normalize_goto(result.goto, default_input=node_input)
|
|
else:
|
|
update = result
|
|
sends = _normalize_result_to_sends(result, default_input=node_input)
|
|
|
|
return {
|
|
"update": update,
|
|
"sends": [
|
|
{"node": _resolve_target_name(send.node), "arg": send.arg}
|
|
for send in sends
|
|
],
|
|
}
|
|
|
|
def _set_resume_event(self, event: dict[str, Any] | None) -> None:
|
|
self._local.resume_event = event
|
|
|
|
def _set_is_resume(self, is_resume: bool) -> None:
|
|
self._local.is_resume = is_resume
|
|
|
|
def _is_resume_execution(self) -> bool:
|
|
return bool(getattr(self._local, "is_resume", False))
|
|
|
|
def _consume_resume_event(
|
|
self, target: WaitCondition | AnyOfCondition | AllOfCondition
|
|
) -> WaitForResult | None:
|
|
event = cast(dict[str, Any] | None, getattr(self._local, "resume_event", None))
|
|
if event is None:
|
|
return None
|
|
self._local.resume_event = None
|
|
return _wait_for_result_from_resume_event(target, event)
|
|
|
|
def _run_awaitable_in_worker(self, awaitable: Coroutine[Any, Any, Any]) -> Any:
|
|
# Create and close a dedicated loop per execution to avoid
|
|
# interpreter-shutdown warnings from lingering thread-local loops.
|
|
return asyncio.run(awaitable)
|
|
|
|
|
|
def _normalize_result_to_sends(result: Any, *, default_input: Any) -> list[Send]:
|
|
if result is None:
|
|
return []
|
|
if isinstance(result, Send):
|
|
return [result]
|
|
if callable(result):
|
|
return [Send(_infer_node_name(result), default_input)]
|
|
if isinstance(result, str):
|
|
return [Send(result, default_input)]
|
|
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_input))
|
|
elif isinstance(item, str):
|
|
sends.append(Send(item, default_input))
|
|
return sends
|
|
return []
|
|
|
|
|
|
def _normalize_goto(goto: Any, *, default_input: Any) -> list[Send]:
|
|
if not goto:
|
|
return []
|
|
if isinstance(goto, Send):
|
|
return [goto]
|
|
if callable(goto):
|
|
return [Send(_infer_node_name(goto), default_input)]
|
|
if isinstance(goto, str):
|
|
return [Send(goto, default_input)]
|
|
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_input))
|
|
elif isinstance(item, str):
|
|
sends.append(Send(item, default_input))
|
|
return sends
|
|
return []
|
|
|
|
|
|
def _normalize_node_state_option(
|
|
state_option: dict[str, Any] | NodeStateOption | None,
|
|
) -> NodeStateOption:
|
|
if state_option is None:
|
|
return NodeStateOption()
|
|
if isinstance(state_option, NodeStateOption):
|
|
return state_option
|
|
if not isinstance(state_option, dict):
|
|
raise TypeError("state_option must be a dict or NodeStateOption")
|
|
locked_fields_raw = state_option.get("locked_fields", ())
|
|
if locked_fields_raw is None:
|
|
return NodeStateOption()
|
|
if not isinstance(locked_fields_raw, Sequence) or isinstance(
|
|
locked_fields_raw, (str, bytes)
|
|
):
|
|
raise TypeError("state_option['locked_fields'] must be a sequence of strings")
|
|
locked_fields: list[str] = []
|
|
for field in locked_fields_raw:
|
|
if not isinstance(field, str):
|
|
raise TypeError("locked field names must be strings")
|
|
if not field:
|
|
continue
|
|
locked_fields.append(field)
|
|
return NodeStateOption(locked_fields=tuple(locked_fields))
|
|
|
|
|
|
def _node_locked_fields_payload(
|
|
node_options: dict[str, NodeStateOption],
|
|
) -> dict[str, list[str]]:
|
|
payload: dict[str, list[str]] = {}
|
|
for node_name, option in node_options.items():
|
|
if not option.locked_fields:
|
|
continue
|
|
payload[node_name] = list(option.locked_fields)
|
|
return payload
|
|
|
|
|
|
def channel_condition(channel: str, min: int = 1, max: int = 0) -> ChannelCondition:
|
|
if min < 1:
|
|
raise ValueError("channel_condition `min` must be >= 1")
|
|
if max < 0:
|
|
raise ValueError("channel_condition `max` must be >= 0")
|
|
if max != 0 and max < min:
|
|
raise ValueError("channel_condition `max` must be 0 or >= min")
|
|
return ChannelCondition(channel=channel, min=min, max=max)
|
|
|
|
|
|
def timer_condition(
|
|
timeout: float | timedelta | None = None,
|
|
*,
|
|
seconds: float | None = None,
|
|
minutes: float | None = None,
|
|
) -> TimerCondition:
|
|
if timeout is not None and (seconds is not None or minutes is not None):
|
|
raise ValueError(
|
|
"Provide either `timeout` or named `seconds`/`minutes`, not both"
|
|
)
|
|
|
|
if isinstance(timeout, timedelta):
|
|
resolved_seconds = timeout.total_seconds()
|
|
elif isinstance(timeout, (int, float)):
|
|
resolved_seconds = float(timeout)
|
|
else:
|
|
resolved_seconds = 0.0
|
|
if seconds is not None:
|
|
resolved_seconds += float(seconds)
|
|
if minutes is not None:
|
|
resolved_seconds += float(minutes) * 60.0
|
|
|
|
if resolved_seconds <= 0:
|
|
raise ValueError("timer_condition must be greater than 0 seconds")
|
|
return TimerCondition(seconds=resolved_seconds)
|
|
|
|
|
|
def any_of(*conditions: WaitCondition) -> AnyOfCondition:
|
|
if not conditions:
|
|
raise ValueError("any_of() requires at least one condition")
|
|
return AnyOfCondition(conditions=tuple(conditions))
|
|
|
|
|
|
def all_of(*conditions: WaitCondition) -> AllOfCondition:
|
|
if not conditions:
|
|
raise ValueError("all_of() requires at least one condition")
|
|
return AllOfCondition(conditions=tuple(conditions))
|
|
|
|
|
|
def _normalize_channel_values(value: Any) -> list[Any]:
|
|
if isinstance(value, list):
|
|
return value
|
|
return [value]
|
|
|
|
|
|
def _wait_for_result_from_resume_event(
|
|
target: WaitCondition | AnyOfCondition | AllOfCondition, event: dict[str, Any]
|
|
) -> WaitForResult:
|
|
if isinstance(target, ChannelCondition):
|
|
return WaitForResult(
|
|
conditions=[
|
|
ConditionResult(
|
|
met=True,
|
|
channel_name=target.channel,
|
|
values=_normalize_channel_values(event.get("value")),
|
|
)
|
|
]
|
|
)
|
|
if isinstance(target, TimerCondition):
|
|
return WaitForResult(conditions=[ConditionResult(met=True)])
|
|
if isinstance(target, AllOfCondition):
|
|
return _wait_for_result_from_all_of_event(target, event)
|
|
return _wait_for_result_from_any_of_event(target, event)
|
|
|
|
|
|
def _wait_for_result_from_any_of_event(
|
|
target: AnyOfCondition, event: dict[str, Any]
|
|
) -> WaitForResult:
|
|
results = [ConditionResult(met=False) for _ in target.conditions]
|
|
condition = event.get("condition")
|
|
|
|
if condition == "timer":
|
|
for idx, cond in enumerate(target.conditions):
|
|
if isinstance(cond, TimerCondition):
|
|
results[idx] = ConditionResult(met=True)
|
|
break
|
|
return WaitForResult(conditions=results)
|
|
|
|
if condition != "channel":
|
|
return WaitForResult(conditions=results)
|
|
|
|
channel = cast(str | None, event.get("channel"))
|
|
value = event.get("value")
|
|
|
|
if channel == "__any_of__" and isinstance(value, list):
|
|
matched = list(value)
|
|
cursor = 0
|
|
for idx, cond in enumerate(target.conditions):
|
|
if not isinstance(cond, ChannelCondition):
|
|
continue
|
|
if cursor >= len(matched):
|
|
continue
|
|
item = matched[cursor]
|
|
if (
|
|
isinstance(item, dict)
|
|
and item.get("channel") == cond.channel
|
|
and "value" in item
|
|
):
|
|
results[idx] = ConditionResult(
|
|
met=True,
|
|
channel_name=cond.channel,
|
|
values=_normalize_channel_values(item.get("value")),
|
|
)
|
|
cursor += 1
|
|
return WaitForResult(conditions=results)
|
|
|
|
for idx, cond in enumerate(target.conditions):
|
|
if isinstance(cond, ChannelCondition) and cond.channel == channel:
|
|
results[idx] = ConditionResult(
|
|
met=True,
|
|
channel_name=cond.channel,
|
|
values=_normalize_channel_values(value),
|
|
)
|
|
break
|
|
return WaitForResult(conditions=results)
|
|
|
|
|
|
def _wait_for_result_from_all_of_event(
|
|
target: AllOfCondition, event: dict[str, Any]
|
|
) -> WaitForResult:
|
|
results = [ConditionResult(met=False) for _ in target.conditions]
|
|
condition = event.get("condition")
|
|
|
|
# all_of completion implies all timer conditions are satisfied.
|
|
for idx, cond in enumerate(target.conditions):
|
|
if isinstance(cond, TimerCondition):
|
|
results[idx] = ConditionResult(met=True)
|
|
|
|
if condition != "channel":
|
|
return WaitForResult(conditions=results)
|
|
|
|
channel = cast(str | None, event.get("channel"))
|
|
value = event.get("value")
|
|
|
|
if channel == "__all_of__" and isinstance(value, list):
|
|
matched_by_channel: dict[str, Any] = {}
|
|
for item in value:
|
|
if isinstance(item, dict) and isinstance(item.get("channel"), str):
|
|
matched_by_channel[cast(str, item["channel"])] = item.get("value")
|
|
|
|
for idx, cond in enumerate(target.conditions):
|
|
if not isinstance(cond, ChannelCondition):
|
|
continue
|
|
if cond.channel not in matched_by_channel:
|
|
continue
|
|
results[idx] = ConditionResult(
|
|
met=True,
|
|
channel_name=cond.channel,
|
|
values=_normalize_channel_values(matched_by_channel[cond.channel]),
|
|
)
|
|
return WaitForResult(conditions=results)
|
|
|
|
for idx, cond in enumerate(target.conditions):
|
|
if isinstance(cond, ChannelCondition) and cond.channel == channel:
|
|
results[idx] = ConditionResult(
|
|
met=True,
|
|
channel_name=cond.channel,
|
|
values=_normalize_channel_values(value),
|
|
)
|
|
break
|
|
return WaitForResult(conditions=results)
|
|
|
|
|
|
def _condition_to_rust(condition: WaitCondition) -> dict[str, Any]:
|
|
if isinstance(condition, ChannelCondition):
|
|
return {
|
|
"kind": "channel",
|
|
"channel": condition.channel,
|
|
"min": condition.min,
|
|
"max": condition.max,
|
|
}
|
|
if isinstance(condition, TimerCondition):
|
|
return {"kind": "timer", "seconds": condition.seconds}
|
|
raise TypeError(f"Unsupported condition type: {type(condition)!r}")
|
|
|
|
|
|
def _target_to_suspend_payload(
|
|
target: WaitCondition | AnyOfCondition | AllOfCondition,
|
|
) -> dict[str, Any]:
|
|
if isinstance(target, AnyOfCondition):
|
|
return {
|
|
"kind": "any_of",
|
|
"any_of": {
|
|
"conditions": [_condition_to_rust(cond) for cond in target.conditions]
|
|
},
|
|
}
|
|
if isinstance(target, AllOfCondition):
|
|
return {
|
|
"kind": "all_of",
|
|
"all_of": {
|
|
"conditions": [_condition_to_rust(cond) for cond in target.conditions]
|
|
},
|
|
}
|
|
return {"kind": "condition", "condition": _condition_to_rust(target)}
|
|
|
|
|
|
def _unwrap_resume_input(node_input: Any) -> tuple[Any, dict[str, Any] | None]:
|
|
if not isinstance(node_input, dict):
|
|
return node_input, None
|
|
if "__lg_resume_arg__" not in node_input or "__lg_resume_event__" not in node_input:
|
|
return node_input, None
|
|
resume_arg = node_input["__lg_resume_arg__"]
|
|
resume_event = node_input["__lg_resume_event__"]
|
|
if isinstance(resume_event, dict):
|
|
return resume_arg, resume_event
|
|
return resume_arg, None
|
|
|
|
|
|
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}")
|
|
|
|
|
|
def _invoke_node(node: Callable[..., Any], ctx: Context, node_input: Any, state: Any) -> Any:
|
|
try:
|
|
params = list(inspect.signature(node).parameters.values())
|
|
except (TypeError, ValueError):
|
|
params = []
|
|
|
|
if not params:
|
|
return node()
|
|
|
|
names = [param.name.lower() for param in params]
|
|
has_ctx = [("ctx" in name or "context" in name) for name in names]
|
|
has_state = [("state" in name) for name in names]
|
|
has_input = [("input" in name) for name in names]
|
|
|
|
kwargs: dict[str, Any] = {}
|
|
unresolved = False
|
|
for idx, param in enumerate(params):
|
|
if has_ctx[idx]:
|
|
kwargs[param.name] = ctx
|
|
elif has_state[idx]:
|
|
kwargs[param.name] = state
|
|
elif has_input[idx]:
|
|
kwargs[param.name] = node_input
|
|
else:
|
|
unresolved = True
|
|
|
|
if kwargs and not unresolved:
|
|
return node(**kwargs)
|
|
|
|
if len(params) == 1:
|
|
if has_ctx[0]:
|
|
return node(ctx)
|
|
if has_state[0]:
|
|
return node(state)
|
|
return node(node_input)
|
|
|
|
if len(params) == 2:
|
|
if has_ctx[0] and has_state[1]:
|
|
return node(ctx, state)
|
|
if has_ctx[0] and has_input[1]:
|
|
return node(ctx, node_input)
|
|
if has_input[0] and has_state[1]:
|
|
return node(node_input, state)
|
|
if has_state[0] and has_input[1]:
|
|
return node(state, node_input)
|
|
if has_state[0]:
|
|
return node(state, node_input)
|
|
if has_state[1]:
|
|
return node(node_input, state)
|
|
if has_ctx[0]:
|
|
return node(ctx, node_input)
|
|
return node(node_input, state)
|
|
|
|
return node(ctx, node_input, state)
|
|
|