Merge branch 'nc/state' of github.com:langchain-ai/permchain into nc/state

This commit is contained in:
Harrison Chase
2024-01-15 13:16:18 -08:00
6 changed files with 99 additions and 37 deletions
+5 -3
View File
@@ -29,6 +29,7 @@ class Graph:
self.nodes: dict[str, Runnable] = {}
self.edges = set[tuple[str, str]]()
self.branches: defaultdict[str, list[Branch]] = defaultdict(list)
self.support_multiple_edges = False
def add_node(self, key: str, action: RunnableLike) -> None:
if key in self.nodes:
@@ -46,8 +47,9 @@ class Graph:
if end_key not in self.nodes and end_key != END:
raise ValueError(f"Need to add_node `{end_key}` first")
# TODO: support multiple message passing
if start_key in set(start for start, _ in self.edges):
if not self.support_multiple_edges and start_key in set(
start for start, _ in self.edges
):
raise ValueError(f"Already found path for {start_key}")
self.edges.add((start_key, end_key))
@@ -111,7 +113,7 @@ class Graph:
outgoing = outgoing_edges[key]
edges_key = f"{key}:edges"
if outgoing or key in self.branches:
nodes[edges_key] = Channel.subscribe_to(key)
nodes[edges_key] = Channel.subscribe_to(key, tags=["langsmith:hidden"])
if outgoing:
nodes[edges_key] |= Channel.write_to(*[dest for dest in outgoing])
if key in self.branches:
+7 -3
View File
@@ -21,6 +21,8 @@ class StateGraph(Graph):
super().__init__()
self.schema = schema
self.channels = _get_channels(schema)
if any(isinstance(c, BinaryOperatorAggregate) for c in self.channels.values()):
self.support_multiple_edges = True
def compile(self) -> Pregel:
self.validate()
@@ -49,7 +51,9 @@ class StateGraph(Graph):
outgoing = outgoing_edges[key]
edges_key = f"{key}:edges"
if outgoing or key in self.branches:
nodes[edges_key] = Channel.subscribe_to(key) | ChannelRead(state_keys)
nodes[edges_key] = Channel.subscribe_to(
key, tags=["langsmith:hidden"]
) | ChannelRead(state_keys)
if outgoing:
nodes[edges_key] |= Channel.write_to(*[dest for dest in outgoing])
if key in self.branches:
@@ -59,12 +63,12 @@ class StateGraph(Graph):
)
nodes[START] = (
Channel.subscribe_to(f"{START}:inbox")
Channel.subscribe_to(f"{START}:inbox", tags=["langsmith:hidden"])
| _update_state
| Channel.write_to(START)
)
nodes[f"{START}:edges"] = (
Channel.subscribe_to(START)
Channel.subscribe_to(START, tags=["langsmith:hidden"])
| ChannelRead(state_keys)
| Channel.write_to(f"{self.entry_point}:inbox")
)
+75 -24
View File
@@ -84,8 +84,10 @@ class Channel:
def subscribe_to(
cls,
channels: str,
*,
key: Optional[str] = None,
when: Optional[Callable[[Any], bool]] = None,
tags: Optional[Sequence[str]] = None,
) -> ChannelInvoke:
...
@@ -94,8 +96,10 @@ class Channel:
def subscribe_to(
cls,
channels: Sequence[str],
*,
key: None = None,
when: Optional[Callable[[Any], bool]] = None,
tags: Optional[Sequence[str]] = None,
) -> ChannelInvoke:
...
@@ -103,8 +107,10 @@ class Channel:
def subscribe_to(
cls,
channels: Union[str, Sequence[str]],
*,
key: Optional[str] = None,
when: Optional[Callable[[Any], bool]] = None,
tags: Optional[Sequence[str]] = None,
) -> ChannelInvoke:
"""Runs process.invoke() each time channels are updated,
with a dict of the channel values as input."""
@@ -121,6 +127,7 @@ class Channel:
),
triggers=[channels] if isinstance(channels, str) else channels,
when=when,
tags=tags,
)
@classmethod
@@ -154,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
@@ -222,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
@@ -242,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,
)
@@ -307,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)
@@ -319,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)
@@ -330,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")
@@ -344,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
@@ -357,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,
)
@@ -427,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)
@@ -439,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)
@@ -449,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
@@ -467,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
@@ -490,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
@@ -508,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
@@ -524,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
+2 -1
View File
@@ -90,6 +90,7 @@ class ChannelInvoke(RunnableBindingBase):
channels: Mapping[None, str] | Mapping[str, str],
triggers: Sequence[str],
when: Optional[Callable[[Any], bool]] = None,
tags: Optional[Sequence[str]] = None,
*,
bound: Optional[Runnable[Any, Any]] = None,
kwargs: Optional[Mapping[str, Any]] = None,
@@ -102,7 +103,7 @@ class ChannelInvoke(RunnableBindingBase):
when=when,
bound=bound or default_bound,
kwargs=kwargs or {},
config=config,
config={**(config or {}), "tags": tags or []},
**other_kwargs,
)
+5 -3
View File
@@ -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},
+5 -3
View File
@@ -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},