mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-07 18:27:52 +02:00
Add output kwarg
This commit is contained in:
@@ -215,9 +215,13 @@ class Pregel(RunnableSerializable[dict[str, Any] | Any, dict[str, Any] | Any]):
|
||||
input: Iterator[dict[str, Any] | Any],
|
||||
run_manager: CallbackManagerForChainRun,
|
||||
config: RunnableConfig,
|
||||
*,
|
||||
output: str | Sequence[str] | None = None,
|
||||
) -> Iterator[tuple[dict[str, Any] | Any, CheckpointView]]:
|
||||
if config["recursion_limit"] < 1:
|
||||
raise ValueError("recursion_limit must be at least 1")
|
||||
# assign defaults
|
||||
output = output if output is not None else self.output
|
||||
# copy nodes to ignore mutations during execution
|
||||
processes = {**self.nodes}
|
||||
# get checkpoint from saver, or create an empty one
|
||||
@@ -256,24 +260,31 @@ class Pregel(RunnableSerializable[dict[str, Any] | Any, dict[str, Any] | Any]):
|
||||
# collect all writes to channels, without applying them yet
|
||||
pending_writes = deque[tuple[str, Any]]()
|
||||
|
||||
# prepare tasks with config
|
||||
tasks_w_config = [
|
||||
(
|
||||
proc,
|
||||
input,
|
||||
patch_config(
|
||||
config,
|
||||
run_name=name,
|
||||
callbacks=run_manager.get_child(f"graph:step:{step}"),
|
||||
configurable={
|
||||
# deque.extend is thread-safe
|
||||
CONFIG_KEY_SEND: pending_writes.extend,
|
||||
CONFIG_KEY_READ: read,
|
||||
},
|
||||
),
|
||||
)
|
||||
for proc, input, name in next_tasks
|
||||
]
|
||||
|
||||
# execute tasks, and wait for one to fail or all to finish.
|
||||
# each task is independent from all other concurrent tasks
|
||||
done, inflight = concurrent.futures.wait(
|
||||
[
|
||||
executor.submit(
|
||||
proc.invoke,
|
||||
input,
|
||||
patch_config(
|
||||
config,
|
||||
callbacks=run_manager.get_child(f"pregel:step:{step}"),
|
||||
configurable={
|
||||
# deque.extend is thread-safe
|
||||
CONFIG_KEY_SEND: pending_writes.extend,
|
||||
CONFIG_KEY_READ: read,
|
||||
},
|
||||
),
|
||||
)
|
||||
for proc, input, _ in next_tasks
|
||||
executor.submit(proc.invoke, input, config)
|
||||
for proc, input, config in tasks_w_config
|
||||
],
|
||||
return_when=concurrent.futures.FIRST_EXCEPTION,
|
||||
timeout=self.step_timeout,
|
||||
@@ -293,7 +304,7 @@ class Pregel(RunnableSerializable[dict[str, Any] | Any, dict[str, Any] | Any]):
|
||||
values=_updateable_channel_values(channels),
|
||||
step=step + 1,
|
||||
)
|
||||
yield map_output(self.output, pending_writes, channels), view
|
||||
yield map_output(output, pending_writes, channels), view
|
||||
# if view was updated, apply writes to channels
|
||||
_apply_writes_from_view(checkpoint, channels, view)
|
||||
|
||||
@@ -312,9 +323,13 @@ class Pregel(RunnableSerializable[dict[str, Any] | Any, dict[str, Any] | Any]):
|
||||
input: AsyncIterator[dict[str, Any] | Any],
|
||||
run_manager: AsyncCallbackManagerForChainRun,
|
||||
config: RunnableConfig,
|
||||
*,
|
||||
output: str | Sequence[str] | None = None,
|
||||
) -> AsyncIterator[tuple[dict[str, Any] | Any, CheckpointView]]:
|
||||
if config["recursion_limit"] < 1:
|
||||
raise ValueError("recursion_limit must be at least 1")
|
||||
# assign defaults
|
||||
output = output if output is not None else self.output
|
||||
# copy nodes to ignore mutations during execution
|
||||
processes = {**self.nodes}
|
||||
# get checkpoint from saver, or create an empty one
|
||||
@@ -351,27 +366,31 @@ class Pregel(RunnableSerializable[dict[str, Any] | Any, dict[str, Any] | Any]):
|
||||
# collect all writes to channels, without applying them yet
|
||||
pending_writes = deque[tuple[str, Any]]()
|
||||
|
||||
# prepare tasks with config
|
||||
tasks_w_config = [
|
||||
(
|
||||
proc,
|
||||
input,
|
||||
patch_config(
|
||||
config,
|
||||
run_name=name,
|
||||
callbacks=run_manager.get_child(f"graph:step:{step}"),
|
||||
configurable={
|
||||
# deque.extend is thread-safe
|
||||
CONFIG_KEY_SEND: pending_writes.extend,
|
||||
CONFIG_KEY_READ: read,
|
||||
},
|
||||
),
|
||||
)
|
||||
for proc, input, name in next_tasks
|
||||
]
|
||||
|
||||
# execute tasks, and wait for one to fail or all to finish.
|
||||
# each task is independent from all other concurrent tasks
|
||||
done, inflight = await asyncio.wait(
|
||||
[
|
||||
asyncio.create_task(
|
||||
proc.ainvoke(
|
||||
input,
|
||||
patch_config(
|
||||
config,
|
||||
callbacks=run_manager.get_child(
|
||||
f"pregel:step:{step}"
|
||||
),
|
||||
configurable={
|
||||
# deque.extend is thread-safe
|
||||
CONFIG_KEY_SEND: pending_writes.extend,
|
||||
CONFIG_KEY_READ: read,
|
||||
},
|
||||
),
|
||||
)
|
||||
)
|
||||
for proc, input, _ in next_tasks
|
||||
asyncio.create_task(proc.ainvoke(input, config))
|
||||
for proc, input, config in tasks_w_config
|
||||
],
|
||||
return_when=asyncio.FIRST_EXCEPTION,
|
||||
timeout=self.step_timeout,
|
||||
@@ -391,7 +410,7 @@ class Pregel(RunnableSerializable[dict[str, Any] | Any, dict[str, Any] | Any]):
|
||||
values=_updateable_channel_values(channels),
|
||||
step=step + 1,
|
||||
)
|
||||
yield map_output(self.output, pending_writes, channels), view
|
||||
yield map_output(output, pending_writes, channels), view
|
||||
# if view was updated, apply writes to channels
|
||||
_apply_writes_from_view(checkpoint, channels, view)
|
||||
|
||||
@@ -409,10 +428,12 @@ class Pregel(RunnableSerializable[dict[str, Any] | Any, dict[str, Any] | Any]):
|
||||
self,
|
||||
input: dict[str, Any] | Any,
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
output: str | Sequence[str] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> dict[str, Any] | Any:
|
||||
latest: dict[str, Any] | Any = None
|
||||
for chunk in self.stream(input, config, **kwargs):
|
||||
for chunk in self.stream(input, config, output=output, **kwargs):
|
||||
latest = chunk
|
||||
return latest
|
||||
|
||||
@@ -420,30 +441,36 @@ class Pregel(RunnableSerializable[dict[str, Any] | Any, dict[str, Any] | Any]):
|
||||
self,
|
||||
input: dict[str, Any] | Any,
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
output: str | Sequence[str] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterator[dict[str, Any] | Any]:
|
||||
return self.transform(iter([input]), config, **kwargs)
|
||||
return self.transform(iter([input]), config, output=output, **kwargs)
|
||||
|
||||
def transform(
|
||||
self,
|
||||
input: Iterator[dict[str, Any] | Any],
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
output: str | Sequence[str] | None = None,
|
||||
**kwargs: Any | None,
|
||||
) -> Iterator[dict[str, Any] | Any]:
|
||||
for output, _ in self._transform_stream_with_config(
|
||||
input, self._transform, config, **kwargs
|
||||
for out, _ in self._transform_stream_with_config(
|
||||
input, self._transform, config, output=output, **kwargs
|
||||
):
|
||||
if output is not None:
|
||||
yield output
|
||||
if out is not None:
|
||||
yield cast(dict[str, Any] | Any, out)
|
||||
|
||||
def step(
|
||||
self,
|
||||
input: dict[str, Any] | Any,
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
output: str | Sequence[str] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Iterator[tuple[dict[str, Any] | Any, CheckpointView]]:
|
||||
for tup in self._transform_stream_with_config(
|
||||
iter([input]), self._transform, config, **kwargs
|
||||
iter([input]), self._transform, config, output=output, **kwargs
|
||||
):
|
||||
yield cast(tuple[dict[str, Any] | Any, CheckpointView], tup)
|
||||
|
||||
@@ -451,10 +478,12 @@ class Pregel(RunnableSerializable[dict[str, Any] | Any, dict[str, Any] | Any]):
|
||||
self,
|
||||
input: dict[str, Any] | Any,
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
output: str | Sequence[str] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> dict[str, Any] | Any:
|
||||
latest: dict[str, Any] | Any = None
|
||||
async for chunk in self.astream(input, config, **kwargs):
|
||||
async for chunk in self.astream(input, config, output=output, **kwargs):
|
||||
latest = chunk
|
||||
return latest
|
||||
|
||||
@@ -462,37 +491,45 @@ class Pregel(RunnableSerializable[dict[str, Any] | Any, dict[str, Any] | Any]):
|
||||
self,
|
||||
input: dict[str, Any] | Any,
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
output: str | Sequence[str] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterator[dict[str, Any] | Any]:
|
||||
async def input_stream() -> AsyncIterator[dict[str, Any] | Any]:
|
||||
yield input
|
||||
|
||||
async for chunk in self.atransform(input_stream(), config, **kwargs):
|
||||
async for chunk in self.atransform(
|
||||
input_stream(), config, output=output, **kwargs
|
||||
):
|
||||
yield chunk
|
||||
|
||||
async def atransform(
|
||||
self,
|
||||
input: AsyncIterator[dict[str, Any] | Any],
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
output: str | Sequence[str] | None = None,
|
||||
**kwargs: Any | None,
|
||||
) -> AsyncIterator[dict[str, Any] | Any]:
|
||||
async for output, _ in self._atransform_stream_with_config(
|
||||
input, self._atransform, config, **kwargs
|
||||
async for out, _ in self._atransform_stream_with_config(
|
||||
input, self._atransform, config, output=output, **kwargs
|
||||
):
|
||||
if output is not None:
|
||||
yield output
|
||||
if out is not None:
|
||||
yield out
|
||||
|
||||
async def astep(
|
||||
self,
|
||||
input: dict[str, Any] | Any,
|
||||
config: RunnableConfig | None = None,
|
||||
*,
|
||||
output: str | Sequence[str] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterator[tuple[dict[str, Any] | Any, CheckpointView]]:
|
||||
async def input_stream() -> AsyncIterator[dict[str, Any] | Any]:
|
||||
yield input
|
||||
|
||||
async for tup in self._atransform_stream_with_config(
|
||||
input_stream(), self._atransform, config, **kwargs
|
||||
input_stream(), self._atransform, config, output=output, **kwargs
|
||||
):
|
||||
yield cast(tuple[dict[str, Any] | Any, CheckpointView], tup)
|
||||
|
||||
|
||||
@@ -42,6 +42,7 @@ def test_invoke_single_process_in_out(mocker: MockerFixture) -> None:
|
||||
assert app.input_schema.schema() == {"title": "PregelInput", "type": "integer"}
|
||||
assert app.output_schema.schema() == {"title": "PregelOutput", "type": "integer"}
|
||||
assert app.invoke(2) == 3
|
||||
assert app.invoke(2, output=["output"]) == {"output": 3}
|
||||
assert repr(app), "does not raise recursion error"
|
||||
|
||||
assert gapp.invoke(2) == 3
|
||||
@@ -243,6 +244,10 @@ def test_invoke_two_processes_in_dict_out(mocker: MockerFixture) -> None:
|
||||
)
|
||||
|
||||
assert [*app.stream({"input": 2, "inbox": 12})] == [13, 4] # [12 + 1, 2 + 1 + 1]
|
||||
assert [*app.stream({"input": 2, "inbox": 12}, output=["output"])] == [
|
||||
{"output": 13},
|
||||
{"output": 4},
|
||||
]
|
||||
|
||||
|
||||
def test_batch_two_processes_in_out() -> None:
|
||||
@@ -256,6 +261,13 @@ 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"]) == [
|
||||
{"output": 5},
|
||||
{"output": 4},
|
||||
{"output": 3},
|
||||
{"output": 5},
|
||||
{"output": 7},
|
||||
]
|
||||
|
||||
graph = Graph()
|
||||
graph.add_node("add_one", add_one_with_delay)
|
||||
|
||||
@@ -36,6 +36,7 @@ async def test_invoke_single_process_in_out(mocker: MockerFixture) -> None:
|
||||
assert app.input_schema.schema() == {"title": "PregelInput", "type": "integer"}
|
||||
assert app.output_schema.schema() == {"title": "PregelOutput", "type": "integer"}
|
||||
assert await app.ainvoke(2) == 3
|
||||
assert await app.ainvoke(2, output=["output"]) == {"output": 3}
|
||||
|
||||
|
||||
async def test_invoke_single_process_in_out_implicit_channels(
|
||||
@@ -195,6 +196,9 @@ 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})] == [13, 4]
|
||||
assert [
|
||||
c async for c in pubsub.astream({"input": 2, "inbox": 12}, output=["output"])
|
||||
] == [{"output": 13}, {"output": 4}]
|
||||
|
||||
|
||||
async def test_batch_two_processes_in_out() -> None:
|
||||
@@ -211,6 +215,13 @@ 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"]) == [
|
||||
{"output": 5},
|
||||
{"output": 4},
|
||||
{"output": 3},
|
||||
{"output": 5},
|
||||
{"output": 7},
|
||||
]
|
||||
|
||||
|
||||
async def test_invoke_many_processes_in_out(mocker: MockerFixture) -> None:
|
||||
|
||||
Reference in New Issue
Block a user