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:
Nuno Campos
2024-09-04 11:01:38 -07:00
committed by GitHub
5 changed files with 312 additions and 278 deletions
+40 -228
View File
@@ -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}")
+71 -40
View File
@@ -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()
+193
View File
@@ -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")
+4 -4
View File
@@ -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",
+4 -6
View File
@@ -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",