mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-10-01 22:15:11 +02:00
Merge pull request #1598 from langchain-ai/nc/3sep/runner
Move runner code to standalone class, use in Pregel.stream/astream
This commit is contained in:
@@ -1,8 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import concurrent.futures
|
||||
import time
|
||||
from collections import deque
|
||||
from functools import partial
|
||||
from typing import (
|
||||
@@ -66,12 +64,11 @@ from langgraph.constants import (
|
||||
CONFIG_KEY_SEND,
|
||||
CONFIG_KEY_STREAM,
|
||||
CONFIG_KEY_TASK_ID,
|
||||
ERROR,
|
||||
INTERRUPT,
|
||||
NS_END,
|
||||
NS_SEP,
|
||||
)
|
||||
from langgraph.errors import GraphInterrupt, GraphRecursionError, InvalidUpdateError
|
||||
from langgraph.errors import GraphRecursionError, InvalidUpdateError
|
||||
from langgraph.managed.base import ManagedValueSpec
|
||||
from langgraph.pregel.algo import (
|
||||
apply_writes,
|
||||
@@ -80,17 +77,13 @@ from langgraph.pregel.algo import (
|
||||
prepare_next_tasks,
|
||||
)
|
||||
from langgraph.pregel.config import patch_checkpoint_map, patch_configurable
|
||||
from langgraph.pregel.debug import (
|
||||
print_step_checkpoint,
|
||||
print_step_tasks,
|
||||
print_step_writes,
|
||||
tasks_w_writes,
|
||||
)
|
||||
from langgraph.pregel.debug import tasks_w_writes
|
||||
from langgraph.pregel.io import read_channels
|
||||
from langgraph.pregel.loop import AsyncPregelLoop, SyncPregelLoop
|
||||
from langgraph.pregel.manager import AsyncChannelsManager, ChannelsManager
|
||||
from langgraph.pregel.read import PregelNode
|
||||
from langgraph.pregel.retry import RetryPolicy, arun_with_retry, run_with_retry
|
||||
from langgraph.pregel.retry import RetryPolicy
|
||||
from langgraph.pregel.runner import PregelRunner
|
||||
from langgraph.pregel.types import (
|
||||
All,
|
||||
PregelExecutableTask,
|
||||
@@ -1170,9 +1163,11 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]):
|
||||
```
|
||||
"""
|
||||
|
||||
stream = deque()
|
||||
|
||||
def output() -> Iterator:
|
||||
while loop.stream:
|
||||
ns, mode, payload = loop.stream.popleft()
|
||||
while stream:
|
||||
ns, mode, payload = stream.popleft()
|
||||
if mode in stream_modes:
|
||||
if subgraphs and isinstance(stream_mode, list):
|
||||
yield (tuple(ns.split(NS_SEP)) if ns else (), mode, payload)
|
||||
@@ -1217,6 +1212,7 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]):
|
||||
|
||||
with SyncPregelLoop(
|
||||
input,
|
||||
stream=stream.append,
|
||||
config=config,
|
||||
store=self.store,
|
||||
checkpointer=checkpointer,
|
||||
@@ -1224,7 +1220,14 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]):
|
||||
specs=self.channels,
|
||||
output_keys=output_keys,
|
||||
stream_keys=self.stream_channels_asis,
|
||||
debug=debug,
|
||||
) as loop:
|
||||
# create runner
|
||||
runner = PregelRunner(
|
||||
submit=loop.submit,
|
||||
put_writes=loop.put_writes,
|
||||
)
|
||||
# enable subgraph streaming
|
||||
if subgraphs:
|
||||
loop.config["configurable"][CONFIG_KEY_STREAM] = loop.stream
|
||||
# Similarly to Bulk Synchronous Parallel / Pregel model
|
||||
@@ -1238,85 +1241,14 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]):
|
||||
interrupt_after=interrupt_after,
|
||||
manager=run_manager,
|
||||
):
|
||||
# debug flag
|
||||
if debug:
|
||||
print_step_checkpoint(
|
||||
loop.checkpoint_metadata,
|
||||
loop.channels,
|
||||
self.stream_channels_list,
|
||||
)
|
||||
# emit output
|
||||
yield from output()
|
||||
# debug flag
|
||||
if debug:
|
||||
print_step_tasks(loop.step, loop.tasks)
|
||||
|
||||
# execute tasks, and wait for one to fail or all to finish.
|
||||
# each task is independent from all other concurrent tasks
|
||||
# yield updates/debug output as each task finishes
|
||||
futures = {
|
||||
loop.submit(
|
||||
run_with_retry,
|
||||
task,
|
||||
self.retry_policy,
|
||||
): task
|
||||
for task in loop.tasks
|
||||
if not task.writes
|
||||
}
|
||||
all_futures = futures.copy()
|
||||
end_time = (
|
||||
self.step_timeout + time.monotonic()
|
||||
if self.step_timeout
|
||||
else None
|
||||
)
|
||||
if not futures:
|
||||
done, inflight = set(), set()
|
||||
while futures:
|
||||
done, inflight = concurrent.futures.wait(
|
||||
futures,
|
||||
return_when=concurrent.futures.FIRST_COMPLETED,
|
||||
timeout=(
|
||||
max(0, end_time - time.monotonic())
|
||||
if end_time
|
||||
else None
|
||||
),
|
||||
)
|
||||
if not done:
|
||||
break # timed out
|
||||
for fut in done:
|
||||
task = futures.pop(fut)
|
||||
if exc := _exception(fut):
|
||||
# save error to checkpointer
|
||||
if isinstance(exc, GraphInterrupt):
|
||||
loop.put_writes(
|
||||
task.id, [(INTERRUPT, i) for i in exc.args[0]]
|
||||
)
|
||||
else:
|
||||
loop.put_writes(task.id, [(ERROR, exc)])
|
||||
|
||||
else:
|
||||
# save task writes to checkpointer
|
||||
loop.put_writes(task.id, task.writes)
|
||||
else:
|
||||
# remove references to loop vars
|
||||
del fut, task
|
||||
for _ in runner.tick(
|
||||
loop.tasks,
|
||||
timeout=self.step_timeout,
|
||||
retry_policy=self.retry_policy,
|
||||
):
|
||||
# emit output
|
||||
yield from output()
|
||||
# maybe stop other tasks
|
||||
if _should_stop_others(done):
|
||||
break
|
||||
|
||||
# panic on failure or timeout
|
||||
_panic_or_proceed(all_futures, loop.step)
|
||||
# don't keep futures around in memory longer than needed
|
||||
del done, inflight, futures
|
||||
# debug flag
|
||||
if debug:
|
||||
print_step_writes(
|
||||
loop.step,
|
||||
[w for t in loop.tasks for w in t.writes],
|
||||
self.stream_channels_list,
|
||||
)
|
||||
for o in output():
|
||||
yield o
|
||||
# emit output
|
||||
yield from output()
|
||||
# handle exit
|
||||
@@ -1412,9 +1344,11 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]):
|
||||
```
|
||||
"""
|
||||
|
||||
stream = deque()
|
||||
|
||||
def output() -> Iterator:
|
||||
while loop.stream:
|
||||
ns, mode, payload = loop.stream.popleft()
|
||||
while stream:
|
||||
ns, mode, payload = stream.popleft()
|
||||
if mode in stream_modes:
|
||||
if subgraphs and isinstance(stream_mode, list):
|
||||
yield (tuple(ns.split(NS_SEP)) if ns else (), mode, payload)
|
||||
@@ -1467,6 +1401,7 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]):
|
||||
)
|
||||
async with AsyncPregelLoop(
|
||||
input,
|
||||
stream=stream.append,
|
||||
config=config,
|
||||
store=self.store,
|
||||
checkpointer=checkpointer,
|
||||
@@ -1475,9 +1410,15 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]):
|
||||
output_keys=output_keys,
|
||||
stream_keys=self.stream_channels_asis,
|
||||
) as loop:
|
||||
# create runner
|
||||
runner = PregelRunner(
|
||||
submit=loop.submit,
|
||||
put_writes=loop.put_writes,
|
||||
use_astream=do_stream is not None,
|
||||
)
|
||||
# enable subgraph streaming
|
||||
if subgraphs:
|
||||
loop.config["configurable"][CONFIG_KEY_STREAM] = loop.stream
|
||||
aioloop = asyncio.get_event_loop()
|
||||
# Similarly to Bulk Synchronous Parallel / Pregel model
|
||||
# computation proceeds in steps, while there are channel updates
|
||||
# channel updates from step N are only visible in step N+1
|
||||
@@ -1489,89 +1430,14 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]):
|
||||
interrupt_after=interrupt_after,
|
||||
manager=run_manager,
|
||||
):
|
||||
# debug flag
|
||||
if debug:
|
||||
print_step_checkpoint(
|
||||
loop.checkpoint_metadata,
|
||||
loop.channels,
|
||||
self.stream_channels_list,
|
||||
)
|
||||
# emit output
|
||||
for o in output():
|
||||
yield o
|
||||
# debug flag
|
||||
if debug:
|
||||
print_step_tasks(loop.step, loop.tasks)
|
||||
|
||||
# execute tasks, and wait for one to fail or all to finish.
|
||||
# each task is independent from all other concurrent tasks
|
||||
# yield updates/debug output as each task finishes
|
||||
futures = {
|
||||
loop.submit(
|
||||
arun_with_retry,
|
||||
task,
|
||||
self.retry_policy,
|
||||
stream=do_stream,
|
||||
__name__=task.name,
|
||||
__cancel_on_exit__=True,
|
||||
): task
|
||||
for task in loop.tasks
|
||||
if not task.writes
|
||||
}
|
||||
all_futures = futures.copy()
|
||||
end_time = (
|
||||
self.step_timeout + aioloop.time()
|
||||
if self.step_timeout
|
||||
else None
|
||||
)
|
||||
if not futures:
|
||||
done, inflight = set(), set()
|
||||
while futures:
|
||||
done, inflight = await asyncio.wait(
|
||||
futures,
|
||||
return_when=asyncio.FIRST_COMPLETED,
|
||||
timeout=(
|
||||
max(0, end_time - aioloop.time()) if end_time else None
|
||||
),
|
||||
)
|
||||
if not done:
|
||||
break # timed out
|
||||
|
||||
for fut in done:
|
||||
task = futures.pop(fut)
|
||||
if exc := _exception(fut):
|
||||
# save error to checkpointer
|
||||
if isinstance(exc, GraphInterrupt):
|
||||
loop.put_writes(
|
||||
task.id, [(INTERRUPT, i) for i in exc.args[0]]
|
||||
)
|
||||
else:
|
||||
loop.put_writes(task.id, [(ERROR, exc)])
|
||||
|
||||
else:
|
||||
# save task writes to checkpointer
|
||||
loop.put_writes(task.id, task.writes)
|
||||
else:
|
||||
# remove references to loop vars
|
||||
del fut, task
|
||||
async for _ in runner.atick(
|
||||
loop.tasks,
|
||||
timeout=self.step_timeout,
|
||||
retry_policy=self.retry_policy,
|
||||
):
|
||||
# emit output
|
||||
for o in output():
|
||||
yield o
|
||||
# maybe stop other tasks
|
||||
if _should_stop_others(done):
|
||||
break
|
||||
|
||||
# panic on failure or timeout
|
||||
_panic_or_proceed(all_futures, loop.step, asyncio.TimeoutError)
|
||||
# don't keep futures around in memory longer than needed
|
||||
del done, inflight, futures
|
||||
# debug flag
|
||||
if debug:
|
||||
print_step_writes(
|
||||
loop.step,
|
||||
[w for t in loop.tasks for w in t.writes],
|
||||
self.stream_channels_list,
|
||||
)
|
||||
# emit output
|
||||
for o in output():
|
||||
yield o
|
||||
@@ -1692,57 +1558,3 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]):
|
||||
return latest
|
||||
else:
|
||||
return chunks
|
||||
|
||||
|
||||
def _should_stop_others(
|
||||
done: Union[set[concurrent.futures.Future[Any]], set[asyncio.Task[Any]]],
|
||||
) -> bool:
|
||||
for fut in done:
|
||||
if fut.cancelled():
|
||||
return True
|
||||
if exc := fut.exception():
|
||||
return not isinstance(exc, GraphInterrupt)
|
||||
else:
|
||||
return False
|
||||
|
||||
|
||||
def _exception(
|
||||
fut: Union[concurrent.futures.Future[Any], asyncio.Task[Any]],
|
||||
) -> Optional[BaseException]:
|
||||
if fut.cancelled():
|
||||
if isinstance(fut, asyncio.Task):
|
||||
return asyncio.CancelledError()
|
||||
else:
|
||||
return concurrent.futures.CancelledError()
|
||||
else:
|
||||
return fut.exception()
|
||||
|
||||
|
||||
def _panic_or_proceed(
|
||||
futs: Union[set[concurrent.futures.Future[Any]], set[asyncio.Task[Any]]],
|
||||
step: int,
|
||||
timeout_exc_cls: Type[Exception] = TimeoutError,
|
||||
) -> None:
|
||||
done: set[Union[concurrent.futures.Future[Any], asyncio.Task[Any]]] = set()
|
||||
inflight: set[Union[concurrent.futures.Future[Any], asyncio.Task[Any]]] = set()
|
||||
for fut in futs:
|
||||
if fut.done():
|
||||
done.add(fut)
|
||||
else:
|
||||
inflight.add(fut)
|
||||
while done:
|
||||
# if any task failed
|
||||
if exc := _exception(done.pop()):
|
||||
# cancel all pending tasks
|
||||
while inflight:
|
||||
inflight.pop().cancel()
|
||||
# raise the exception
|
||||
raise exc
|
||||
|
||||
if inflight:
|
||||
# if we got here means we timed out
|
||||
while inflight:
|
||||
# cancel all pending tasks
|
||||
inflight.pop().cancel()
|
||||
# raise timeout error
|
||||
raise timeout_exc_cls(f"Timed out at step {step}")
|
||||
|
||||
@@ -2,7 +2,6 @@ import asyncio
|
||||
import concurrent.futures
|
||||
from collections import deque
|
||||
from contextlib import AsyncExitStack, ExitStack
|
||||
from itertools import tee
|
||||
from types import TracebackType
|
||||
from typing import (
|
||||
Any,
|
||||
@@ -65,6 +64,9 @@ from langgraph.pregel.debug import (
|
||||
map_debug_checkpoint,
|
||||
map_debug_task_results,
|
||||
map_debug_tasks,
|
||||
print_step_checkpoint,
|
||||
print_step_tasks,
|
||||
print_step_writes,
|
||||
)
|
||||
from langgraph.pregel.executor import (
|
||||
AsyncBackgroundExecutor,
|
||||
@@ -92,24 +94,16 @@ EMPTY_SEQ = ()
|
||||
|
||||
|
||||
class StreamProtocol(Protocol):
|
||||
def extend(self, values: Iterable[Tuple[str, str, Any]]) -> None: ...
|
||||
def popleft(self) -> Tuple[str, str, Any]: ...
|
||||
def __bool__(self) -> bool: ...
|
||||
def __call__(self, values: Iterable[Tuple[str, str, Any]]) -> None: ...
|
||||
|
||||
|
||||
class DuplexStream(StreamProtocol):
|
||||
def __init__(self, *streams: StreamProtocol) -> None:
|
||||
self.streams = streams
|
||||
def __init__(self, *queues: StreamProtocol) -> None:
|
||||
self.queues = queues
|
||||
|
||||
def extend(self, values: Iterable[Tuple[str, str, Any]]) -> None:
|
||||
for stream, vv in zip(self.streams, tee(values, len(self.streams))):
|
||||
stream.extend(vv)
|
||||
|
||||
def popleft(self) -> Tuple[str, str, Any]:
|
||||
return self.streams[0].popleft()
|
||||
|
||||
def __bool__(self) -> bool:
|
||||
return bool(self.streams[0])
|
||||
def __call__(self, value: Tuple[str, str, Any]) -> None:
|
||||
for queue in self.queues:
|
||||
queue(value)
|
||||
|
||||
|
||||
class PregelLoop:
|
||||
@@ -121,8 +115,9 @@ class PregelLoop:
|
||||
specs: Mapping[str, Union[BaseChannel, ManagedValueSpec]]
|
||||
output_keys: Union[str, Sequence[str]]
|
||||
stream_keys: Union[str, Sequence[str]]
|
||||
is_nested: bool
|
||||
stream: Optional[StreamProtocol]
|
||||
skip_done_tasks: bool
|
||||
is_nested: bool
|
||||
|
||||
checkpointer_get_next_version: Callable[[Optional[V]], V]
|
||||
checkpointer_put_writes: Optional[
|
||||
@@ -154,7 +149,6 @@ class PregelLoop:
|
||||
"pending", "done", "interrupt_before", "interrupt_after", "out_of_steps"
|
||||
]
|
||||
tasks: Sequence[PregelExecutableTask]
|
||||
stream: StreamProtocol
|
||||
output: Union[None, dict[str, Any], Any] = None
|
||||
|
||||
# public
|
||||
@@ -163,6 +157,7 @@ class PregelLoop:
|
||||
self,
|
||||
input: Optional[Any],
|
||||
*,
|
||||
stream: Optional[StreamProtocol],
|
||||
config: RunnableConfig,
|
||||
store: Optional[BaseStore],
|
||||
checkpointer: Optional[BaseCheckpointSaver],
|
||||
@@ -170,8 +165,9 @@ class PregelLoop:
|
||||
specs: Mapping[str, Union[BaseChannel, ManagedValueSpec]],
|
||||
output_keys: Union[str, Sequence[str]],
|
||||
stream_keys: Union[str, Sequence[str]],
|
||||
debug: bool = False,
|
||||
) -> None:
|
||||
self.stream = deque()
|
||||
self.stream = stream
|
||||
self.input = input
|
||||
self.config = config
|
||||
self.store = store
|
||||
@@ -182,6 +178,7 @@ class PregelLoop:
|
||||
self.stream_keys = stream_keys
|
||||
self.is_nested = CONFIG_KEY_TASK_ID in self.config.get("configurable", {})
|
||||
self.skip_done_tasks = "checkpoint_id" not in config["configurable"]
|
||||
self.debug = debug
|
||||
if CONFIG_KEY_STREAM in config["configurable"]:
|
||||
self.stream = DuplexStream(
|
||||
self.stream, config["configurable"][CONFIG_KEY_STREAM]
|
||||
@@ -229,24 +226,6 @@ class PregelLoop:
|
||||
)
|
||||
self._output_writes(task_id, writes)
|
||||
|
||||
def _output_writes(
|
||||
self, task_id: str, writes: Sequence[tuple[str, Any]], *, cached: bool = False
|
||||
) -> None:
|
||||
if task := next((t for t in self.tasks if t.id == task_id), None):
|
||||
self.stream.extend(
|
||||
(self.config["configurable"].get("checkpoint_ns", ""), "updates", v)
|
||||
for v in map_output_updates(self.output_keys, [(task, writes)], cached)
|
||||
)
|
||||
if not cached:
|
||||
self.stream.extend(
|
||||
(self.config["configurable"].get("checkpoint_ns", ""), "debug", v)
|
||||
for v in map_debug_task_results(
|
||||
self.step,
|
||||
[(task, writes)],
|
||||
self.stream_keys,
|
||||
)
|
||||
)
|
||||
|
||||
def tick(
|
||||
self,
|
||||
*,
|
||||
@@ -265,6 +244,15 @@ class PregelLoop:
|
||||
self._first(input_keys=input_keys)
|
||||
elif all(task.writes for task in self.tasks):
|
||||
writes = [w for t in self.tasks for w in t.writes]
|
||||
# debug flag
|
||||
if self.debug:
|
||||
print_step_writes(
|
||||
self.step,
|
||||
writes,
|
||||
[self.stream_keys]
|
||||
if isinstance(self.stream_keys, str)
|
||||
else self.stream_keys,
|
||||
)
|
||||
# all tasks have finished
|
||||
mv_writes = apply_writes(
|
||||
self.checkpoint,
|
||||
@@ -276,7 +264,7 @@ class PregelLoop:
|
||||
for key, values in mv_writes.items():
|
||||
self._update_mv(key, values)
|
||||
# produce values output
|
||||
self.stream.extend(
|
||||
self._emit(
|
||||
(self.config["configurable"].get("checkpoint_ns", ""), "values", v)
|
||||
for v in map_output_values(self.output_keys, writes, self.channels)
|
||||
)
|
||||
@@ -324,7 +312,7 @@ class PregelLoop:
|
||||
|
||||
# produce debug output
|
||||
if self._checkpointer_put_after_previous is not None:
|
||||
self.stream.extend(
|
||||
self._emit(
|
||||
(self.config["configurable"].get("checkpoint_ns", ""), "debug", v)
|
||||
for v in map_debug_checkpoint(
|
||||
self.step - 1, # printing checkpoint for previous step
|
||||
@@ -373,11 +361,15 @@ class PregelLoop:
|
||||
return False
|
||||
|
||||
# produce debug output
|
||||
self.stream.extend(
|
||||
self._emit(
|
||||
(self.config["configurable"].get("checkpoint_ns", ""), "debug", v)
|
||||
for v in map_debug_tasks(self.step, self.tasks)
|
||||
)
|
||||
|
||||
# debug flag
|
||||
if self.debug:
|
||||
print_step_tasks(self.step, self.tasks)
|
||||
|
||||
return True
|
||||
|
||||
# private
|
||||
@@ -399,7 +391,7 @@ class PregelLoop:
|
||||
version = self.checkpoint["channel_versions"][k]
|
||||
self.checkpoint["versions_seen"][INTERRUPT][k] = version
|
||||
# produce values output
|
||||
self.stream.extend(
|
||||
self._emit(
|
||||
(self.config["configurable"].get("checkpoint_ns", ""), "values", v)
|
||||
for v in map_output_values(self.output_keys, True, self.channels)
|
||||
)
|
||||
@@ -436,6 +428,15 @@ class PregelLoop:
|
||||
metadata["parents"] = self.config["configurable"].get(
|
||||
CONFIG_KEY_CHECKPOINT_MAP, {}
|
||||
)
|
||||
# debug flag
|
||||
if self.debug:
|
||||
print_step_checkpoint(
|
||||
metadata,
|
||||
self.channels,
|
||||
[self.stream_keys]
|
||||
if isinstance(self.stream_keys, str)
|
||||
else self.stream_keys,
|
||||
)
|
||||
# bail if no checkpointer
|
||||
if self._checkpointer_put_after_previous is not None:
|
||||
# create new checkpoint
|
||||
@@ -499,12 +500,35 @@ class PregelLoop:
|
||||
# suppress interrupt
|
||||
return True
|
||||
|
||||
def _emit(self, values: Sequence[tuple[str, str, Any]]) -> None:
|
||||
if self.stream is None:
|
||||
return
|
||||
for v in values:
|
||||
self.stream(v)
|
||||
|
||||
def _output_writes(
|
||||
self, task_id: str, writes: Sequence[tuple[str, Any]], *, cached: bool = False
|
||||
) -> None:
|
||||
if task := next((t for t in self.tasks if t.id == task_id), None):
|
||||
self._emit(
|
||||
(self.config["configurable"].get("checkpoint_ns", ""), "updates", v)
|
||||
for v in map_output_updates(self.output_keys, [(task, writes)], cached)
|
||||
)
|
||||
if not cached:
|
||||
self._emit(
|
||||
(self.config["configurable"].get("checkpoint_ns", ""), "debug", v)
|
||||
for v in map_debug_task_results(
|
||||
self.step, [(task, writes)], self.stream_keys
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
class SyncPregelLoop(PregelLoop, ContextManager):
|
||||
def __init__(
|
||||
self,
|
||||
input: Optional[Any],
|
||||
*,
|
||||
stream: Optional[StreamProtocol],
|
||||
config: RunnableConfig,
|
||||
store: Optional[BaseStore],
|
||||
checkpointer: Optional[BaseCheckpointSaver],
|
||||
@@ -512,9 +536,11 @@ class SyncPregelLoop(PregelLoop, ContextManager):
|
||||
specs: Mapping[str, Union[BaseChannel, ManagedValueSpec]],
|
||||
output_keys: Union[str, Sequence[str]] = EMPTY_SEQ,
|
||||
stream_keys: Union[str, Sequence[str]] = EMPTY_SEQ,
|
||||
debug: bool = False,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
input,
|
||||
stream=stream,
|
||||
config=config,
|
||||
checkpointer=checkpointer,
|
||||
store=store,
|
||||
@@ -522,6 +548,7 @@ class SyncPregelLoop(PregelLoop, ContextManager):
|
||||
specs=specs,
|
||||
output_keys=output_keys,
|
||||
stream_keys=stream_keys,
|
||||
debug=debug,
|
||||
)
|
||||
self.stack = ExitStack()
|
||||
if checkpointer:
|
||||
@@ -596,6 +623,7 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager):
|
||||
self,
|
||||
input: Optional[Any],
|
||||
*,
|
||||
stream: Optional[StreamProtocol],
|
||||
config: RunnableConfig,
|
||||
store: Optional[BaseStore],
|
||||
checkpointer: Optional[BaseCheckpointSaver],
|
||||
@@ -603,9 +631,11 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager):
|
||||
specs: Mapping[str, Union[BaseChannel, ManagedValueSpec]],
|
||||
output_keys: Union[str, Sequence[str]] = EMPTY_SEQ,
|
||||
stream_keys: Union[str, Sequence[str]] = EMPTY_SEQ,
|
||||
debug: bool = False,
|
||||
) -> None:
|
||||
super().__init__(
|
||||
input,
|
||||
stream=stream,
|
||||
config=config,
|
||||
checkpointer=checkpointer,
|
||||
store=store,
|
||||
@@ -613,6 +643,7 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager):
|
||||
specs=specs,
|
||||
output_keys=output_keys,
|
||||
stream_keys=stream_keys,
|
||||
debug=debug,
|
||||
)
|
||||
self.store = AsyncBatchedStore(self.store) if self.store else None
|
||||
self.stack = AsyncExitStack()
|
||||
|
||||
@@ -0,0 +1,193 @@
|
||||
import asyncio
|
||||
import concurrent.futures
|
||||
import time
|
||||
from typing import (
|
||||
Any,
|
||||
AsyncIterator,
|
||||
Callable,
|
||||
Iterator,
|
||||
Optional,
|
||||
Sequence,
|
||||
Type,
|
||||
Union,
|
||||
)
|
||||
|
||||
from langgraph.constants import ERROR, INTERRUPT
|
||||
from langgraph.errors import GraphInterrupt
|
||||
from langgraph.pregel.executor import Submit
|
||||
from langgraph.pregel.retry import arun_with_retry, run_with_retry
|
||||
from langgraph.pregel.types import PregelExecutableTask, RetryPolicy
|
||||
|
||||
|
||||
class PregelRunner:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
submit: Submit,
|
||||
put_writes: Callable[[str, Sequence[tuple[str, Any]]], None],
|
||||
use_astream: bool = False,
|
||||
) -> None:
|
||||
self.submit = submit
|
||||
self.put_writes = put_writes
|
||||
self.use_astream = use_astream
|
||||
|
||||
def tick(
|
||||
self,
|
||||
tasks: list[PregelExecutableTask],
|
||||
*,
|
||||
timeout: Optional[float] = None,
|
||||
retry_policy: Optional[RetryPolicy] = None,
|
||||
) -> Iterator[None]:
|
||||
# execute tasks, and wait for one to fail or all to finish.
|
||||
# each task is independent from all other concurrent tasks
|
||||
# yield updates/debug output as each task finishes
|
||||
futures = {
|
||||
self.submit(
|
||||
run_with_retry,
|
||||
task,
|
||||
retry_policy,
|
||||
): task
|
||||
for task in tasks
|
||||
if not task.writes
|
||||
}
|
||||
all_futures = futures.copy()
|
||||
end_time = timeout + time.monotonic() if timeout else None
|
||||
while futures:
|
||||
done, _ = concurrent.futures.wait(
|
||||
futures,
|
||||
return_when=concurrent.futures.FIRST_COMPLETED,
|
||||
timeout=(max(0, end_time - time.monotonic()) if end_time else None),
|
||||
)
|
||||
if not done:
|
||||
break # timed out
|
||||
for fut in done:
|
||||
task = futures.pop(fut)
|
||||
if exc := _exception(fut):
|
||||
if isinstance(exc, GraphInterrupt):
|
||||
# save interrupt to checkpointer
|
||||
self.put_writes(task.id, [(INTERRUPT, i) for i in exc.args[0]])
|
||||
else:
|
||||
# save error to checkpointer
|
||||
self.put_writes(task.id, [(ERROR, exc)])
|
||||
|
||||
else:
|
||||
# save task writes to checkpointer
|
||||
self.put_writes(task.id, task.writes)
|
||||
else:
|
||||
# remove references to loop vars
|
||||
del fut, task
|
||||
# maybe stop other tasks
|
||||
if _should_stop_others(done):
|
||||
break
|
||||
# give control back to the caller
|
||||
yield
|
||||
# panic on failure or timeout
|
||||
_panic_or_proceed(all_futures)
|
||||
|
||||
async def atick(
|
||||
self,
|
||||
tasks: list[PregelExecutableTask],
|
||||
*,
|
||||
timeout: Optional[float] = None,
|
||||
retry_policy: Optional[RetryPolicy] = None,
|
||||
) -> AsyncIterator[None]:
|
||||
loop = asyncio.get_event_loop()
|
||||
# execute tasks, and wait for one to fail or all to finish.
|
||||
# each task is independent from all other concurrent tasks
|
||||
# yield updates/debug output as each task finishes
|
||||
futures = {
|
||||
self.submit(
|
||||
arun_with_retry,
|
||||
task,
|
||||
retry_policy,
|
||||
stream=self.use_astream,
|
||||
__name__=task.name,
|
||||
__cancel_on_exit__=True,
|
||||
): task
|
||||
for task in tasks
|
||||
if not task.writes
|
||||
}
|
||||
all_futures = futures.copy()
|
||||
end_time = timeout + loop.time() if timeout else None
|
||||
while futures:
|
||||
done, _ = await asyncio.wait(
|
||||
futures,
|
||||
return_when=asyncio.FIRST_COMPLETED,
|
||||
timeout=(max(0, end_time - loop.time()) if end_time else None),
|
||||
)
|
||||
if not done:
|
||||
break # timed out
|
||||
for fut in done:
|
||||
task = futures.pop(fut)
|
||||
if exc := _exception(fut):
|
||||
if isinstance(exc, GraphInterrupt):
|
||||
# save interrupt to checkpointer
|
||||
self.put_writes(task.id, [(INTERRUPT, i) for i in exc.args[0]])
|
||||
else:
|
||||
# save error to checkpointer
|
||||
self.put_writes(task.id, [(ERROR, exc)])
|
||||
else:
|
||||
# save task writes to checkpointer
|
||||
self.put_writes(task.id, task.writes)
|
||||
else:
|
||||
# remove references to loop vars
|
||||
del fut, task
|
||||
# maybe stop other tasks
|
||||
if _should_stop_others(done):
|
||||
break
|
||||
# give control back to the caller
|
||||
yield
|
||||
# panic on failure or timeout
|
||||
_panic_or_proceed(all_futures, asyncio.TimeoutError)
|
||||
|
||||
|
||||
def _should_stop_others(
|
||||
done: Union[set[concurrent.futures.Future[Any]], set[asyncio.Task[Any]]],
|
||||
) -> bool:
|
||||
for fut in done:
|
||||
if fut.cancelled():
|
||||
return True
|
||||
if exc := fut.exception():
|
||||
return not isinstance(exc, GraphInterrupt)
|
||||
else:
|
||||
return False
|
||||
|
||||
|
||||
def _exception(
|
||||
fut: Union[concurrent.futures.Future[Any], asyncio.Task[Any]],
|
||||
) -> Optional[BaseException]:
|
||||
if fut.cancelled():
|
||||
if isinstance(fut, asyncio.Task):
|
||||
return asyncio.CancelledError()
|
||||
else:
|
||||
return concurrent.futures.CancelledError()
|
||||
else:
|
||||
return fut.exception()
|
||||
|
||||
|
||||
def _panic_or_proceed(
|
||||
futs: Union[set[concurrent.futures.Future[Any]], set[asyncio.Task[Any]]],
|
||||
timeout_exc_cls: Type[Exception] = TimeoutError,
|
||||
) -> None:
|
||||
done: set[Union[concurrent.futures.Future[Any], asyncio.Task[Any]]] = set()
|
||||
inflight: set[Union[concurrent.futures.Future[Any], asyncio.Task[Any]]] = set()
|
||||
for fut in futs:
|
||||
if fut.done():
|
||||
done.add(fut)
|
||||
else:
|
||||
inflight.add(fut)
|
||||
while done:
|
||||
# if any task failed
|
||||
if exc := _exception(done.pop()):
|
||||
# cancel all pending tasks
|
||||
while inflight:
|
||||
inflight.pop().cancel()
|
||||
# raise the exception
|
||||
raise exc
|
||||
if inflight:
|
||||
# if we got here means we timed out
|
||||
while inflight:
|
||||
# cancel all pending tasks
|
||||
inflight.pop().cancel()
|
||||
# raise timeout error
|
||||
raise timeout_exc_cls("Timed out")
|
||||
@@ -1957,8 +1957,6 @@ def test_channel_enter_exit_timing(mocker: MockerFixture) -> None:
|
||||
def test_conditional_graph(
|
||||
snapshot: SnapshotAssertion, request: pytest.FixtureRequest, checkpointer_name: str
|
||||
) -> None:
|
||||
from copy import deepcopy
|
||||
|
||||
from langchain_core.agents import AgentAction, AgentFinish
|
||||
from langchain_core.language_models.fake import FakeStreamingListLLM
|
||||
from langchain_core.prompts import PromptTemplate
|
||||
@@ -2000,12 +1998,15 @@ def test_conditional_graph(
|
||||
|
||||
# Define tool execution logic
|
||||
def execute_tools(data: dict) -> dict:
|
||||
data = data.copy()
|
||||
agent_action: AgentAction = data.pop("agent_outcome")
|
||||
observation = {t.name: t for t in tools}[agent_action.tool].invoke(
|
||||
agent_action.tool_input
|
||||
)
|
||||
if data.get("intermediate_steps") is None:
|
||||
data["intermediate_steps"] = []
|
||||
else:
|
||||
data["intermediate_steps"] = data["intermediate_steps"].copy()
|
||||
data["intermediate_steps"].append([agent_action, observation])
|
||||
return data
|
||||
|
||||
@@ -2066,8 +2067,7 @@ def test_conditional_graph(
|
||||
),
|
||||
}
|
||||
|
||||
# deepcopy because the nodes mutate the data
|
||||
assert [deepcopy(c) for c in app.stream({"input": "what is weather in sf"})] == [
|
||||
assert [c for c in app.stream({"input": "what is weather in sf"})] == [
|
||||
{
|
||||
"agent": {
|
||||
"input": "what is weather in sf",
|
||||
|
||||
@@ -2204,8 +2204,6 @@ async def test_channel_enter_exit_timing(mocker: MockerFixture) -> None:
|
||||
|
||||
@pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC)
|
||||
async def test_conditional_graph(checkpointer_name: str) -> None:
|
||||
from copy import deepcopy
|
||||
|
||||
from langchain_core.agents import AgentAction, AgentFinish
|
||||
from langchain_core.language_models.fake import FakeStreamingListLLM
|
||||
from langchain_core.prompts import PromptTemplate
|
||||
@@ -2243,12 +2241,15 @@ async def test_conditional_graph(checkpointer_name: str) -> None:
|
||||
|
||||
# Define tool execution logic
|
||||
async def execute_tools(data: dict) -> dict:
|
||||
data = data.copy()
|
||||
agent_action: AgentAction = data.pop("agent_outcome")
|
||||
observation = await {t.name: t for t in tools}[agent_action.tool].ainvoke(
|
||||
agent_action.tool_input
|
||||
)
|
||||
if data.get("intermediate_steps") is None:
|
||||
data["intermediate_steps"] = []
|
||||
else:
|
||||
data["intermediate_steps"] = data["intermediate_steps"].copy()
|
||||
data["intermediate_steps"].append([agent_action, observation])
|
||||
return data
|
||||
|
||||
@@ -2301,10 +2302,7 @@ async def test_conditional_graph(checkpointer_name: str) -> None:
|
||||
),
|
||||
}
|
||||
|
||||
# deepcopy because the nodes mutate the data
|
||||
assert [
|
||||
deepcopy(c) async for c in app.astream({"input": "what is weather in sf"})
|
||||
] == [
|
||||
assert [c async for c in app.astream({"input": "what is weather in sf"})] == [
|
||||
{
|
||||
"agent": {
|
||||
"input": "what is weather in sf",
|
||||
|
||||
Reference in New Issue
Block a user