mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-28 04:25:08 +02:00
Add input_keys, rename output_keys, add interrupt
This commit is contained in:
@@ -161,6 +161,8 @@ class Pregel(
|
||||
|
||||
hidden: Sequence[str] = Field(default_factory=list)
|
||||
|
||||
interrupt: Sequence[str] = Field(default_factory=list)
|
||||
|
||||
input: Union[str, Sequence[str]] = "input"
|
||||
|
||||
step_timeout: Optional[float] = None
|
||||
@@ -229,13 +231,16 @@ class Pregel(
|
||||
run_manager: CallbackManagerForChainRun,
|
||||
config: RunnableConfig,
|
||||
*,
|
||||
output: Optional[Union[str, Sequence[str]]] = None,
|
||||
input_keys: Optional[Union[str, Sequence[str]]] = None,
|
||||
output_keys: Optional[Union[str, Sequence[str]]] = None,
|
||||
) -> Iterator[Union[dict[str, Any], Any]]:
|
||||
if config["recursion_limit"] < 1:
|
||||
raise ValueError("recursion_limit must be at least 1")
|
||||
# assign defaults
|
||||
if output is None:
|
||||
output = [chan for chan in self.channels if chan not in self.hidden]
|
||||
if output_keys is None:
|
||||
output_keys = [chan for chan in self.channels if chan not in self.hidden]
|
||||
if input_keys is None:
|
||||
input_keys = self.input
|
||||
# copy nodes to ignore mutations during execution
|
||||
processes = {**self.nodes}
|
||||
# get checkpoint from saver, or create an empty one
|
||||
@@ -249,7 +254,7 @@ class Pregel(
|
||||
_apply_writes(
|
||||
checkpoint,
|
||||
channels,
|
||||
deque(w for c in input for w in map_input(self.input, c)),
|
||||
deque(w for c in input for w in map_input(input_keys, c)),
|
||||
config,
|
||||
0,
|
||||
)
|
||||
@@ -314,10 +319,10 @@ class Pregel(
|
||||
print_checkpoint(step, channels)
|
||||
|
||||
# yield current value and checkpoint view
|
||||
if step_output := map_output(output, pending_writes, channels):
|
||||
if step_output := map_output(output_keys, pending_writes, channels):
|
||||
yield step_output
|
||||
# we can detect updates when output is multiple channels (ie. dict)
|
||||
if not isinstance(output, str):
|
||||
if not isinstance(output_keys, str):
|
||||
# if view was updated, apply writes to channels
|
||||
_apply_writes_from_view(checkpoint, channels, step_output)
|
||||
|
||||
@@ -326,6 +331,10 @@ class Pregel(
|
||||
checkpoint = create_checkpoint(checkpoint, channels)
|
||||
self.saver.put(config, checkpoint)
|
||||
|
||||
# interrupt if any channel written to is in interrupt list
|
||||
if any(chan for chan, _ in pending_writes if chan in self.interrupt):
|
||||
break
|
||||
|
||||
# save end of run checkpoint
|
||||
if self.saver is not None and self.saver.at == CheckpointAt.END_OF_RUN:
|
||||
checkpoint = create_checkpoint(checkpoint, channels)
|
||||
@@ -337,7 +346,8 @@ class Pregel(
|
||||
run_manager: AsyncCallbackManagerForChainRun,
|
||||
config: RunnableConfig,
|
||||
*,
|
||||
output: Optional[Union[str, Sequence[str]]] = None,
|
||||
input_keys: Optional[Union[str, Sequence[str]]] = None,
|
||||
output_keys: Optional[Union[str, Sequence[str]]] = None,
|
||||
) -> AsyncIterator[Union[dict[str, Any], Any]]:
|
||||
if config["recursion_limit"] < 1:
|
||||
raise ValueError("recursion_limit must be at least 1")
|
||||
@@ -351,8 +361,10 @@ class Pregel(
|
||||
None,
|
||||
)
|
||||
# assign defaults
|
||||
if output is None:
|
||||
output = [chan for chan in self.channels if chan not in self.hidden]
|
||||
if output_keys is None:
|
||||
output_keys = [chan for chan in self.channels if chan not in self.hidden]
|
||||
if input_keys is None:
|
||||
input_keys = self.input
|
||||
# copy nodes to ignore mutations during execution
|
||||
processes = {**self.nodes}
|
||||
# get checkpoint from saver, or create an empty one
|
||||
@@ -364,7 +376,7 @@ class Pregel(
|
||||
_apply_writes(
|
||||
checkpoint,
|
||||
channels,
|
||||
deque([w async for c in input for w in map_input(self.input, c)]),
|
||||
deque([w async for c in input for w in map_input(input_keys, c)]),
|
||||
config,
|
||||
0,
|
||||
)
|
||||
@@ -434,10 +446,10 @@ class Pregel(
|
||||
print_checkpoint(step, channels)
|
||||
|
||||
# yield current value and checkpoint view
|
||||
if step_output := map_output(output, pending_writes, channels):
|
||||
if step_output := map_output(output_keys, pending_writes, channels):
|
||||
yield step_output
|
||||
# we can detect updates when output is multiple channels (ie. dict)
|
||||
if not isinstance(output, str):
|
||||
if not isinstance(output_keys, str):
|
||||
# if view was updated, apply writes to channels
|
||||
_apply_writes_from_view(checkpoint, channels, step_output)
|
||||
|
||||
@@ -446,6 +458,10 @@ class Pregel(
|
||||
checkpoint = create_checkpoint(checkpoint, channels)
|
||||
await self.saver.aput(config, checkpoint)
|
||||
|
||||
# interrupt if any channel written to is in interrupt list
|
||||
if any(chan for chan, _ in pending_writes if chan in self.interrupt):
|
||||
break
|
||||
|
||||
# save end of run checkpoint
|
||||
if self.saver is not None and self.saver.at == CheckpointAt.END_OF_RUN:
|
||||
checkpoint = create_checkpoint(checkpoint, channels)
|
||||
@@ -456,14 +472,16 @@ class Pregel(
|
||||
input: Union[dict[str, Any], Any],
|
||||
config: Optional[RunnableConfig] = None,
|
||||
*,
|
||||
output: Optional[Union[str, Sequence[str]]] = None,
|
||||
output_keys: Optional[Union[str, Sequence[str]]] = None,
|
||||
input_keys: Optional[Union[str, Sequence[str]]] = None,
|
||||
**kwargs: Any,
|
||||
) -> Union[dict[str, Any], Any]:
|
||||
latest: Union[dict[str, Any], Any] = None
|
||||
for chunk in self.stream(
|
||||
input,
|
||||
config,
|
||||
output=output if output is not None else self.output,
|
||||
output_keys=output_keys if output_keys is not None else self.output,
|
||||
input_keys=input_keys,
|
||||
**kwargs,
|
||||
):
|
||||
latest = chunk
|
||||
@@ -474,21 +492,34 @@ class Pregel(
|
||||
input: Union[dict[str, Any], Any],
|
||||
config: Optional[RunnableConfig] = None,
|
||||
*,
|
||||
output: Optional[Union[str, Sequence[str]]] = None,
|
||||
output_keys: Optional[Union[str, Sequence[str]]] = None,
|
||||
input_keys: Optional[Union[str, Sequence[str]]] = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterator[Union[dict[str, Any], Any]]:
|
||||
return self.transform(iter([input]), config, output=output, **kwargs)
|
||||
return self.transform(
|
||||
iter([input]),
|
||||
config,
|
||||
output_keys=output_keys,
|
||||
input_keys=input_keys,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def transform(
|
||||
self,
|
||||
input: Iterator[Union[dict[str, Any], Any]],
|
||||
config: Optional[RunnableConfig] = None,
|
||||
*,
|
||||
output: Optional[Union[str, Sequence[str]]] = None,
|
||||
output_keys: Optional[Union[str, Sequence[str]]] = None,
|
||||
input_keys: Optional[Union[str, Sequence[str]]] = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterator[Union[dict[str, Any], Any]]:
|
||||
for chunk in self._transform_stream_with_config(
|
||||
input, self._transform, config, output=output, **kwargs
|
||||
input,
|
||||
self._transform,
|
||||
config,
|
||||
output_keys=output_keys,
|
||||
input_keys=input_keys,
|
||||
**kwargs,
|
||||
):
|
||||
yield chunk
|
||||
|
||||
@@ -497,14 +528,16 @@ class Pregel(
|
||||
input: Union[dict[str, Any], Any],
|
||||
config: Optional[RunnableConfig] = None,
|
||||
*,
|
||||
output: Optional[Union[str, Sequence[str]]] = None,
|
||||
output_keys: Optional[Union[str, Sequence[str]]] = None,
|
||||
input_keys: Optional[Union[str, Sequence[str]]] = None,
|
||||
**kwargs: Any,
|
||||
) -> Union[dict[str, Any], Any]:
|
||||
latest: Union[dict[str, Any], Any] = None
|
||||
async for chunk in self.astream(
|
||||
input,
|
||||
config,
|
||||
output=output if output is not None else self.output,
|
||||
output_keys=output_keys if output_keys is not None else self.output,
|
||||
input_keys=input_keys,
|
||||
**kwargs,
|
||||
):
|
||||
latest = chunk
|
||||
@@ -515,14 +548,19 @@ class Pregel(
|
||||
input: Union[dict[str, Any], Any],
|
||||
config: Optional[RunnableConfig] = None,
|
||||
*,
|
||||
output: Optional[Union[str, Sequence[str]]] = None,
|
||||
output_keys: Optional[Union[str, Sequence[str]]] = None,
|
||||
input_keys: Optional[Union[str, Sequence[str]]] = None,
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterator[Union[dict[str, Any], Any]]:
|
||||
async def input_stream() -> AsyncIterator[Union[dict[str, Any], Any]]:
|
||||
yield input
|
||||
|
||||
async for chunk in self.atransform(
|
||||
input_stream(), config, output=output, **kwargs
|
||||
input_stream(),
|
||||
config,
|
||||
output_keys=output_keys,
|
||||
input_keys=input_keys,
|
||||
**kwargs,
|
||||
):
|
||||
yield chunk
|
||||
|
||||
@@ -531,11 +569,17 @@ class Pregel(
|
||||
input: AsyncIterator[Union[dict[str, Any], Any]],
|
||||
config: Optional[RunnableConfig] = None,
|
||||
*,
|
||||
output: Optional[Union[str, Sequence[str]]] = None,
|
||||
output_keys: Optional[Union[str, Sequence[str]]] = None,
|
||||
input_keys: Optional[Union[str, Sequence[str]]] = None,
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterator[Union[dict[str, Any], Any]]:
|
||||
async for chunk in self._atransform_stream_with_config(
|
||||
input, self._atransform, config, output=output, **kwargs
|
||||
input,
|
||||
self._atransform,
|
||||
config,
|
||||
output_keys=output_keys,
|
||||
input_keys=input_keys,
|
||||
**kwargs,
|
||||
):
|
||||
yield chunk
|
||||
|
||||
|
||||
@@ -44,7 +44,7 @@ def test_invoke_single_process_in_out(mocker: MockerFixture) -> None:
|
||||
assert app.input_schema.schema() == {"title": "LangGraphInput", "type": "integer"}
|
||||
assert app.output_schema.schema() == {"title": "LangGraphOutput", "type": "integer"}
|
||||
assert app.invoke(2) == 3
|
||||
assert app.invoke(2, output=["output"]) == {"output": 3}
|
||||
assert app.invoke(2, output_keys=["output"]) == {"output": 3}
|
||||
assert repr(app), "does not raise recursion error"
|
||||
|
||||
assert gapp.invoke(2) == 3
|
||||
@@ -157,6 +157,8 @@ def test_invoke_two_processes_in_out(mocker: MockerFixture) -> None:
|
||||
|
||||
assert app.invoke(2) == 4
|
||||
|
||||
assert app.invoke(2, input_keys="inbox") == 3
|
||||
|
||||
for step, values in enumerate(app.stream(2), start=1):
|
||||
if step == 1:
|
||||
assert values == {
|
||||
@@ -238,7 +240,7 @@ def test_invoke_two_processes_in_dict_out(mocker: MockerFixture) -> None:
|
||||
input=["input", "inbox"],
|
||||
)
|
||||
|
||||
assert [*app.stream({"input": 2, "inbox": 12}, output="output")] == [
|
||||
assert [*app.stream({"input": 2, "inbox": 12}, output_keys="output")] == [
|
||||
13,
|
||||
4,
|
||||
] # [12 + 1, 2 + 1 + 1]
|
||||
@@ -259,7 +261,7 @@ def test_batch_two_processes_in_out() -> None:
|
||||
app = Pregel(nodes={"one": one, "two": two})
|
||||
|
||||
assert app.batch([3, 2, 1, 3, 5]) == [5, 4, 3, 5, 7]
|
||||
assert app.batch([3, 2, 1, 3, 5], output=["output"]) == [
|
||||
assert app.batch([3, 2, 1, 3, 5], output_keys=["output"]) == [
|
||||
{"output": 5},
|
||||
{"output": 4},
|
||||
{"output": 3},
|
||||
|
||||
@@ -42,7 +42,7 @@ async def test_invoke_single_process_in_out(mocker: MockerFixture) -> None:
|
||||
assert app.input_schema.schema() == {"title": "LangGraphInput", "type": "integer"}
|
||||
assert app.output_schema.schema() == {"title": "LangGraphOutput", "type": "integer"}
|
||||
assert await app.ainvoke(2) == 3
|
||||
assert await app.ainvoke(2, output=["output"]) == {"output": 3}
|
||||
assert await app.ainvoke(2, output_keys=["output"]) == {"output": 3}
|
||||
|
||||
assert await gapp.ainvoke(2) == 3
|
||||
|
||||
@@ -157,6 +157,8 @@ async def test_invoke_two_processes_in_out(mocker: MockerFixture) -> None:
|
||||
|
||||
assert await app.ainvoke(2) == 4
|
||||
|
||||
assert await app.ainvoke(2, input_keys="inbox") == 3
|
||||
|
||||
step = 0
|
||||
async for values in app.astream(2):
|
||||
step += 1
|
||||
@@ -247,7 +249,7 @@ async def test_invoke_two_processes_in_dict_out(mocker: MockerFixture) -> None:
|
||||
|
||||
# [12 + 1, 2 + 1 + 1]
|
||||
assert [
|
||||
c async for c in pubsub.astream({"input": 2, "inbox": 12}, output="output")
|
||||
c async for c in pubsub.astream({"input": 2, "inbox": 12}, output_keys="output")
|
||||
] == [13, 4]
|
||||
assert [c async for c in pubsub.astream({"input": 2, "inbox": 12})] == [
|
||||
{"inbox": [3], "output": 13},
|
||||
@@ -269,7 +271,7 @@ async def test_batch_two_processes_in_out() -> None:
|
||||
)
|
||||
|
||||
assert await app.abatch([3, 2, 1, 3, 5]) == [5, 4, 3, 5, 7]
|
||||
assert await app.abatch([3, 2, 1, 3, 5], output=["output"]) == [
|
||||
assert await app.abatch([3, 2, 1, 3, 5], output_keys=["output"]) == [
|
||||
{"output": 5},
|
||||
{"output": 4},
|
||||
{"output": 3},
|
||||
|
||||
Reference in New Issue
Block a user