mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-25 17:12:26 +02:00
Merge pull request #1415 from langchain-ai/nc/21aug/loop-rm-pregel-arg
Remove graph arg to PregelLoop
This commit is contained in:
@@ -340,7 +340,10 @@ class Pregel(
|
||||
@property
|
||||
def stream_channels_asis(self) -> Union[str, Sequence[str]]:
|
||||
return self.stream_channels or [
|
||||
k for k in self.channels if not isinstance(self.channels[k], Context)
|
||||
k
|
||||
for k in self.channels
|
||||
if isinstance(self.channels[k], BaseChannel)
|
||||
and not isinstance(self.channels[k], Context)
|
||||
]
|
||||
|
||||
def get_state(self, config: RunnableConfig) -> StateSnapshot:
|
||||
@@ -956,7 +959,6 @@ class Pregel(
|
||||
config=config,
|
||||
store=self.store,
|
||||
checkpointer=checkpointer,
|
||||
graph=self,
|
||||
nodes=self.nodes,
|
||||
specs=self.channels,
|
||||
) as loop:
|
||||
@@ -966,7 +968,9 @@ class Pregel(
|
||||
# channels are guaranteed to be immutable for the duration of the step,
|
||||
# with channel updates applied only at the transition between steps
|
||||
while loop.tick(
|
||||
input_keys=self.input_channels,
|
||||
output_keys=output_keys,
|
||||
stream_keys=self.stream_channels_asis,
|
||||
interrupt_before=interrupt_before,
|
||||
interrupt_after=interrupt_after,
|
||||
manager=run_manager,
|
||||
@@ -1212,7 +1216,6 @@ class Pregel(
|
||||
config=config,
|
||||
store=self.store,
|
||||
checkpointer=checkpointer,
|
||||
graph=self,
|
||||
nodes=self.nodes,
|
||||
specs=self.channels,
|
||||
) as loop:
|
||||
@@ -1223,7 +1226,9 @@ class Pregel(
|
||||
# channels are guaranteed to be immutable for the duration of the step,
|
||||
# with channel updates applied only at the transition between steps
|
||||
while loop.tick(
|
||||
input_keys=self.input_channels,
|
||||
output_keys=output_keys,
|
||||
stream_keys=self.stream_channels_asis,
|
||||
interrupt_before=interrupt_before,
|
||||
interrupt_after=interrupt_after,
|
||||
manager=run_manager,
|
||||
|
||||
@@ -4,7 +4,6 @@ from collections import deque
|
||||
from contextlib import AsyncExitStack, ExitStack
|
||||
from types import TracebackType
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
AsyncContextManager,
|
||||
Callable,
|
||||
@@ -71,10 +70,6 @@ from langgraph.pregel.utils import get_new_channel_versions
|
||||
from langgraph.store.base import BaseStore
|
||||
from langgraph.store.batch import AsyncBatchedStore
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from langgraph.pregel import Pregel
|
||||
|
||||
|
||||
V = TypeVar("V")
|
||||
INPUT_DONE = object()
|
||||
INPUT_RESUMING = object()
|
||||
@@ -105,8 +100,6 @@ class PregelLoop:
|
||||
Any,
|
||||
]
|
||||
]
|
||||
graph: "Pregel"
|
||||
store: Optional[BaseStore]
|
||||
submit: Submit
|
||||
channels: Mapping[str, BaseChannel]
|
||||
managed: ManagedValueMapping
|
||||
@@ -136,14 +129,12 @@ class PregelLoop:
|
||||
checkpointer: Optional[BaseCheckpointSaver],
|
||||
nodes: Mapping[str, PregelNode],
|
||||
specs: Mapping[str, Union[BaseChannel, ManagedValueSpec]],
|
||||
graph: "Pregel",
|
||||
) -> None:
|
||||
self.stream = deque()
|
||||
self.input = input
|
||||
self.config = config
|
||||
self.store = store
|
||||
self.checkpointer = checkpointer
|
||||
self.graph = graph
|
||||
self.nodes = nodes
|
||||
self.specs = specs
|
||||
self.is_nested = CONFIG_KEY_READ in self.config.get("configurable", {})
|
||||
@@ -175,7 +166,9 @@ class PregelLoop:
|
||||
def tick(
|
||||
self,
|
||||
*,
|
||||
input_keys: Union[str, Sequence[str]],
|
||||
output_keys: Union[str, Sequence[str]] = EMPTY_SEQ,
|
||||
stream_keys: Union[str, Sequence[str]] = EMPTY_SEQ,
|
||||
interrupt_after: Sequence[str] = EMPTY_SEQ,
|
||||
interrupt_before: Sequence[str] = EMPTY_SEQ,
|
||||
manager: Union[None, AsyncParentRunManager, ParentRunManager] = None,
|
||||
@@ -187,7 +180,7 @@ class PregelLoop:
|
||||
raise RuntimeError("Cannot tick when status is no longer 'pending'")
|
||||
|
||||
if self.input not in (INPUT_DONE, INPUT_RESUMING):
|
||||
self._first()
|
||||
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]
|
||||
# all tasks have finished
|
||||
@@ -211,11 +204,7 @@ class PregelLoop:
|
||||
self._put_checkpoint(
|
||||
{
|
||||
"source": "loop",
|
||||
"writes": single(
|
||||
map_output_updates(output_keys, self.tasks)
|
||||
if self.graph.stream_mode == "updates"
|
||||
else map_output_values(output_keys, writes, self.channels)
|
||||
),
|
||||
"writes": single(map_output_updates(output_keys, self.tasks)),
|
||||
}
|
||||
)
|
||||
# after execution, check if we should interrupt
|
||||
@@ -258,7 +247,7 @@ class PregelLoop:
|
||||
self.step - 1, # printing checkpoint for previous step
|
||||
self.checkpoint_config,
|
||||
self.channels,
|
||||
self.graph.stream_channels_asis,
|
||||
stream_keys,
|
||||
self.checkpoint_metadata,
|
||||
self.checkpoint,
|
||||
self.tasks,
|
||||
@@ -282,6 +271,7 @@ class PregelLoop:
|
||||
# if all tasks have finished, re-tick
|
||||
if all(task.writes for task in self.tasks):
|
||||
return self.tick(
|
||||
input_keys=input_keys,
|
||||
output_keys=output_keys,
|
||||
interrupt_after=interrupt_after,
|
||||
interrupt_before=interrupt_before,
|
||||
@@ -306,7 +296,7 @@ class PregelLoop:
|
||||
|
||||
# private
|
||||
|
||||
def _first(self) -> None:
|
||||
def _first(self, *, input_keys: Union[str, Sequence[str]]) -> None:
|
||||
# resuming from previous checkpoint requires
|
||||
# - finding a previous checkpoint
|
||||
# - receiving None input (outer graph) or RESUMING flag (subgraph)
|
||||
@@ -323,7 +313,7 @@ class PregelLoop:
|
||||
version = self.checkpoint["channel_versions"][k]
|
||||
self.checkpoint["versions_seen"][INTERRUPT][k] = version
|
||||
# map inputs to channel updates
|
||||
elif input_writes := deque(map_input(self.graph.input_channels, self.input)):
|
||||
elif input_writes := deque(map_input(input_keys, self.input)):
|
||||
# discard any unfinished tasks from previous checkpoint
|
||||
discard_tasks = prepare_next_tasks(
|
||||
self.checkpoint,
|
||||
@@ -345,7 +335,7 @@ class PregelLoop:
|
||||
# save input checkpoint
|
||||
self._put_checkpoint({"source": "input", "writes": self.input})
|
||||
else:
|
||||
raise EmptyInputError(f"Received no input for {self.graph.input_channels}")
|
||||
raise EmptyInputError(f"Received no input for {input_keys}")
|
||||
# done with input
|
||||
self.input = INPUT_RESUMING if is_resuming else INPUT_DONE
|
||||
|
||||
@@ -430,13 +420,11 @@ class SyncPregelLoop(PregelLoop, ContextManager):
|
||||
checkpointer: Optional[BaseCheckpointSaver],
|
||||
nodes: Mapping[str, PregelNode],
|
||||
specs: Mapping[str, Union[BaseChannel, ManagedValueSpec]],
|
||||
graph: "Pregel",
|
||||
) -> None:
|
||||
super().__init__(
|
||||
input,
|
||||
config=config,
|
||||
checkpointer=checkpointer,
|
||||
graph=graph,
|
||||
store=store,
|
||||
nodes=nodes,
|
||||
specs=specs,
|
||||
@@ -504,7 +492,6 @@ class SyncPregelLoop(PregelLoop, ContextManager):
|
||||
traceback: Optional[TracebackType],
|
||||
) -> Optional[bool]:
|
||||
# unwind stack
|
||||
del self.graph
|
||||
return self.stack.__exit__(exc_type, exc_value, traceback)
|
||||
|
||||
|
||||
@@ -518,13 +505,11 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager):
|
||||
checkpointer: Optional[BaseCheckpointSaver],
|
||||
nodes: Mapping[str, PregelNode],
|
||||
specs: Mapping[str, Union[BaseChannel, ManagedValueSpec]],
|
||||
graph: "Pregel",
|
||||
) -> None:
|
||||
super().__init__(
|
||||
input,
|
||||
config=config,
|
||||
checkpointer=checkpointer,
|
||||
graph=graph,
|
||||
store=store,
|
||||
nodes=nodes,
|
||||
specs=specs,
|
||||
@@ -598,7 +583,6 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager):
|
||||
traceback: Optional[TracebackType],
|
||||
) -> Optional[bool]:
|
||||
# unwind stack
|
||||
del self.graph
|
||||
return await asyncio.shield(
|
||||
self.stack.__aexit__(exc_type, exc_value, traceback)
|
||||
)
|
||||
|
||||
@@ -657,7 +657,7 @@ def test_invoke_two_processes_in_out_interrupt(
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 6, "writes": 5},
|
||||
metadata={"source": "loop", "step": 6, "writes": {"two": 5}},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[1].config,
|
||||
),
|
||||
@@ -672,7 +672,7 @@ def test_invoke_two_processes_in_out_interrupt(
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 5, "writes": None},
|
||||
metadata={"source": "loop", "step": 5, "writes": {"one": None}},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[2].config,
|
||||
),
|
||||
@@ -702,7 +702,7 @@ def test_invoke_two_processes_in_out_interrupt(
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 3, "writes": None},
|
||||
metadata={"source": "loop", "step": 3, "writes": {"one": None}},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[4].config,
|
||||
),
|
||||
@@ -732,7 +732,7 @@ def test_invoke_two_processes_in_out_interrupt(
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 1, "writes": 4},
|
||||
metadata={"source": "loop", "step": 1, "writes": {"two": 4}},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[6].config,
|
||||
),
|
||||
@@ -747,7 +747,7 @@ def test_invoke_two_processes_in_out_interrupt(
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 0, "writes": None},
|
||||
metadata={"source": "loop", "step": 0, "writes": {"one": None}},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[7].config,
|
||||
),
|
||||
@@ -1995,12 +1995,14 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None:
|
||||
"step": 0,
|
||||
"writes": {
|
||||
"agent": {
|
||||
"input": "what is weather in sf",
|
||||
"agent_outcome": AgentAction(
|
||||
tool="search_api",
|
||||
tool_input="query",
|
||||
log="tool:search_api:query",
|
||||
),
|
||||
"agent": {
|
||||
"input": "what is weather in sf",
|
||||
"agent_outcome": AgentAction(
|
||||
tool="search_api",
|
||||
tool_input="query",
|
||||
log="tool:search_api:query",
|
||||
),
|
||||
}
|
||||
},
|
||||
},
|
||||
},
|
||||
@@ -2206,12 +2208,14 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None:
|
||||
"step": 0,
|
||||
"writes": {
|
||||
"agent": {
|
||||
"input": "what is weather in sf",
|
||||
"agent_outcome": AgentAction(
|
||||
tool="search_api",
|
||||
tool_input="query",
|
||||
log="tool:search_api:query",
|
||||
),
|
||||
"agent": {
|
||||
"input": "what is weather in sf",
|
||||
"agent_outcome": AgentAction(
|
||||
tool="search_api",
|
||||
tool_input="query",
|
||||
log="tool:search_api:query",
|
||||
),
|
||||
}
|
||||
}
|
||||
},
|
||||
},
|
||||
@@ -2411,12 +2415,14 @@ def test_conditional_graph(snapshot: SnapshotAssertion) -> None:
|
||||
"step": 0,
|
||||
"writes": {
|
||||
"agent": {
|
||||
"input": "what is weather in sf",
|
||||
"agent_outcome": AgentAction(
|
||||
tool="search_api",
|
||||
tool_input="query",
|
||||
log="tool:search_api:query",
|
||||
),
|
||||
"agent": {
|
||||
"input": "what is weather in sf",
|
||||
"agent_outcome": AgentAction(
|
||||
tool="search_api",
|
||||
tool_input="query",
|
||||
log="tool:search_api:query",
|
||||
),
|
||||
}
|
||||
}
|
||||
},
|
||||
},
|
||||
|
||||
@@ -876,7 +876,7 @@ async def test_invoke_two_processes_in_out_interrupt(
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 6, "writes": 5},
|
||||
metadata={"source": "loop", "step": 6, "writes": {"two": 5}},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[1].config,
|
||||
),
|
||||
@@ -891,7 +891,7 @@ async def test_invoke_two_processes_in_out_interrupt(
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 5, "writes": None},
|
||||
metadata={"source": "loop", "step": 5, "writes": {"one": None}},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[2].config,
|
||||
),
|
||||
@@ -921,7 +921,7 @@ async def test_invoke_two_processes_in_out_interrupt(
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 3, "writes": None},
|
||||
metadata={"source": "loop", "step": 3, "writes": {"one": None}},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[4].config,
|
||||
),
|
||||
@@ -951,7 +951,7 @@ async def test_invoke_two_processes_in_out_interrupt(
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 1, "writes": 4},
|
||||
metadata={"source": "loop", "step": 1, "writes": {"two": 4}},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[6].config,
|
||||
),
|
||||
@@ -966,7 +966,7 @@ async def test_invoke_two_processes_in_out_interrupt(
|
||||
"checkpoint_id": AnyStr(),
|
||||
}
|
||||
},
|
||||
metadata={"source": "loop", "step": 0, "writes": None},
|
||||
metadata={"source": "loop", "step": 0, "writes": {"one": None}},
|
||||
created_at=AnyStr(),
|
||||
parent_config=history[7].config,
|
||||
),
|
||||
@@ -2290,12 +2290,14 @@ async def test_conditional_graph() -> None:
|
||||
"step": 0,
|
||||
"writes": {
|
||||
"agent": {
|
||||
"input": "what is weather in sf",
|
||||
"agent_outcome": AgentAction(
|
||||
tool="search_api",
|
||||
tool_input="query",
|
||||
log="tool:search_api:query",
|
||||
),
|
||||
"agent": {
|
||||
"input": "what is weather in sf",
|
||||
"agent_outcome": AgentAction(
|
||||
tool="search_api",
|
||||
tool_input="query",
|
||||
log="tool:search_api:query",
|
||||
),
|
||||
}
|
||||
}
|
||||
},
|
||||
},
|
||||
@@ -2516,12 +2518,14 @@ async def test_conditional_graph() -> None:
|
||||
"step": 0,
|
||||
"writes": {
|
||||
"agent": {
|
||||
"input": "what is weather in sf",
|
||||
"agent_outcome": AgentAction(
|
||||
tool="search_api",
|
||||
tool_input="query",
|
||||
log="tool:search_api:query",
|
||||
),
|
||||
"agent": {
|
||||
"input": "what is weather in sf",
|
||||
"agent_outcome": AgentAction(
|
||||
tool="search_api",
|
||||
tool_input="query",
|
||||
log="tool:search_api:query",
|
||||
),
|
||||
}
|
||||
}
|
||||
},
|
||||
},
|
||||
@@ -2748,12 +2752,14 @@ async def test_conditional_graph() -> None:
|
||||
"step": 0,
|
||||
"writes": {
|
||||
"agent": {
|
||||
"input": "what is weather in sf",
|
||||
"agent_outcome": AgentAction(
|
||||
tool="search_api",
|
||||
tool_input="query",
|
||||
log="tool:search_api:query",
|
||||
),
|
||||
"agent": {
|
||||
"input": "what is weather in sf",
|
||||
"agent_outcome": AgentAction(
|
||||
tool="search_api",
|
||||
tool_input="query",
|
||||
log="tool:search_api:query",
|
||||
),
|
||||
}
|
||||
}
|
||||
},
|
||||
},
|
||||
|
||||
Reference in New Issue
Block a user