From 5de9b354162bfee89b859ad72a404fb59488c35f Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Thu, 19 Sep 2024 13:01:54 -0700 Subject: [PATCH] Finish --- .../langgraph/checkpoint/base/__init__.py | 4 +- libs/langgraph/Makefile | 2 +- .../channels/dynamic_barrier_value.py | 2 +- .../langgraph/channels/ephemeral_value.py | 4 +- libs/langgraph/langgraph/graph/graph.py | 79 ++++++++++++------- libs/langgraph/langgraph/graph/message.py | 4 +- libs/langgraph/langgraph/graph/state.py | 38 ++++++--- .../langgraph/prebuilt/chat_agent_executor.py | 2 +- .../langgraph/langgraph/prebuilt/tool_node.py | 2 +- libs/langgraph/langgraph/pregel/__init__.py | 44 ++++++----- libs/langgraph/langgraph/pregel/algo.py | 10 ++- libs/langgraph/langgraph/pregel/debug.py | 4 +- libs/langgraph/langgraph/pregel/loop.py | 22 ++++-- libs/langgraph/langgraph/pregel/messages.py | 13 +-- libs/langgraph/langgraph/pregel/read.py | 12 +-- libs/langgraph/langgraph/pregel/runner.py | 25 +++--- libs/langgraph/langgraph/pregel/validate.py | 16 ++-- libs/langgraph/langgraph/pregel/write.py | 5 +- libs/langgraph/langgraph/utils/runnable.py | 2 +- libs/langgraph/pyproject.toml | 2 +- 20 files changed, 171 insertions(+), 121 deletions(-) diff --git a/libs/checkpoint/langgraph/checkpoint/base/__init__.py b/libs/checkpoint/langgraph/checkpoint/base/__init__.py index 822389cce..ae98e5df5 100644 --- a/libs/checkpoint/langgraph/checkpoint/base/__init__.py +++ b/libs/checkpoint/langgraph/checkpoint/base/__init__.py @@ -13,11 +13,11 @@ from typing import ( Sequence, Tuple, TypedDict, + TypeVar, Union, ) from langchain_core.runnables import ConfigurableFieldSpec, RunnableConfig -from typing_extensions import TypeVar from langgraph.checkpoint.base.id import uuid6 from langgraph.checkpoint.serde.base import SerializerProtocol, maybe_add_typed_methods @@ -29,7 +29,7 @@ from langgraph.checkpoint.serde.types import ( SendProtocol, ) -V = TypeVar("V", int, float, str, default=int) +V = TypeVar("V", int, float, str) PendingWrite = Tuple[str, str, Any] diff --git a/libs/langgraph/Makefile b/libs/langgraph/Makefile index 62ef303bb..7ec12f80a 100644 --- a/libs/langgraph/Makefile +++ b/libs/langgraph/Makefile @@ -75,7 +75,7 @@ lint lint_diff lint_package lint_tests: [ "$(PYTHON_FILES)" = "" ] || poetry run ruff format $(PYTHON_FILES) --diff [ "$(PYTHON_FILES)" = "" ] || poetry run ruff check --select I $(PYTHON_FILES) [ "$(PYTHON_FILES)" = "" ] || mkdir -p $(MYPY_CACHE) - [ "$(PYTHON_FILES)" = "" ] || poetry run mypy $(PYTHON_FILES) --cache-dir $(MYPY_CACHE) + [ "$(PYTHON_FILES)" != "langgraph" ] || poetry run mypy $(PYTHON_FILES) --cache-dir $(MYPY_CACHE) format format_diff: poetry run ruff format $(PYTHON_FILES) diff --git a/libs/langgraph/langgraph/channels/dynamic_barrier_value.py b/libs/langgraph/langgraph/channels/dynamic_barrier_value.py index bbecd3d8b..dfa77f350 100644 --- a/libs/langgraph/langgraph/channels/dynamic_barrier_value.py +++ b/libs/langgraph/langgraph/channels/dynamic_barrier_value.py @@ -6,7 +6,7 @@ from langgraph.channels.base import BaseChannel, Value from langgraph.errors import EmptyChannelError, InvalidUpdateError -class WaitForNames(NamedTuple): +class WaitForNames(NamedTuple, Generic[Value]): names: set[Value] diff --git a/libs/langgraph/langgraph/channels/ephemeral_value.py b/libs/langgraph/langgraph/channels/ephemeral_value.py index 5beba22eb..537a8763c 100644 --- a/libs/langgraph/langgraph/channels/ephemeral_value.py +++ b/libs/langgraph/langgraph/channels/ephemeral_value.py @@ -1,4 +1,4 @@ -from typing import Generic, Optional, Sequence, Type +from typing import Any, Generic, Optional, Sequence, Type from typing_extensions import Self @@ -11,7 +11,7 @@ class EphemeralValue(Generic[Value], BaseChannel[Value, Value, Value]): __slots__ = ("value", "guard") - def __init__(self, typ: Type[Value], guard: bool = True) -> None: + def __init__(self, typ: Any, guard: bool = True) -> None: super().__init__(typ) self.guard = guard diff --git a/libs/langgraph/langgraph/graph/graph.py b/libs/langgraph/langgraph/graph/graph.py index 706e5253c..93fe28b3d 100644 --- a/libs/langgraph/langgraph/graph/graph.py +++ b/libs/langgraph/langgraph/graph/graph.py @@ -57,7 +57,7 @@ class Branch(NamedTuple): def run( self, writer: Callable[ - [list[Union[str, Send]], RunnableConfig], Optional[ChannelWrite] + [Sequence[Union[str, Send]], RunnableConfig], Optional[ChannelWrite] ], reader: Optional[Callable[[RunnableConfig], Any]] = None, ) -> RunnableCallable: @@ -79,7 +79,7 @@ class Branch(NamedTuple): *, reader: Optional[Callable[[RunnableConfig], Any]], writer: Callable[ - [list[Union[str, Send]], RunnableConfig], Optional[ChannelWrite] + [Sequence[Union[str, Send]], RunnableConfig], Optional[ChannelWrite] ], ) -> Runnable: if reader: @@ -100,7 +100,7 @@ class Branch(NamedTuple): *, reader: Optional[Callable[[RunnableConfig], Any]], writer: Callable[ - [list[Union[str, Send]], RunnableConfig], Optional[ChannelWrite] + [Sequence[Union[str, Send]], RunnableConfig], Optional[ChannelWrite] ], ) -> Runnable: if reader: @@ -117,18 +117,20 @@ class Branch(NamedTuple): def _finish( self, writer: Callable[ - [list[Union[str, Send]], RunnableConfig], Optional[ChannelWrite] + [Sequence[Union[str, Send]], RunnableConfig], Optional[ChannelWrite] ], input: Any, result: Any, config: RunnableConfig, - ): + ) -> Union[Runnable, Any]: if not isinstance(result, list): result = [result] if self.ends: - destinations = [r if isinstance(r, Send) else self.ends[r] for r in result] + destinations: Sequence[Union[Send, str]] = [ + r if isinstance(r, Send) else self.ends[r] for r in result + ] else: - destinations = result + destinations = cast(Sequence[Union[Send, str]], result) if any(dest is None or dest == START for dest in destinations): raise ValueError("Branch did not return a valid destination") if any(p.node == END for p in destinations if isinstance(p, Send)): @@ -186,14 +188,20 @@ class Graph: ) if not isinstance(node, str): action = node - node = getattr(action, "name", action.__name__) + node = getattr(action, "name", getattr(action, "__name__")) + if node is None: + raise ValueError( + "Node name must be provided if action is not a function" + ) + if action is None: + raise RuntimeError if node in self.nodes: raise ValueError(f"Node `{node}` already present.") if node == END or node == START: raise ValueError(f"Node `{node}` is reserved.") - self.nodes[node] = NodeSpec( - coerce_to_runnable(action, name=node, trace=False), metadata + self.nodes[cast(str, node)] = NodeSpec( + coerce_to_runnable(action, name=cast(str, node), trace=False), metadata ) def add_edge(self, start_key: str, end_key: str) -> None: @@ -257,16 +265,20 @@ class Graph: # coerce path_map to a dictionary try: if isinstance(path_map, dict): - path_map = path_map.copy() + path_map_ = path_map.copy() elif isinstance(path_map, list): - path_map = {name: name for name in path_map} - elif rtn_type := get_type_hints(path.__call__).get( - "return" - ) or get_type_hints(path).get("return"): + path_map_ = {name: name for name in path_map} + elif callable(path) and ( + rtn_type := get_type_hints(path.__call__).get("return") + if hasattr(path, "__call__") + else get_type_hints(path).get("return") + ): if get_origin(rtn_type) is Literal: - path_map = {name: name for name in get_args(rtn_type)} + path_map_ = {name: name for name in get_args(rtn_type)} + else: + path_map_ = None except Exception: - pass + path_map_ = None # find a name for the condition path = coerce_to_runnable(path, name=None, trace=True) name = path.name or "condition" @@ -276,7 +288,7 @@ class Graph: f"Branch with name `{path.name}` already exists for node " f"`{source}`" ) # save it - self.branches[source][name] = Branch(path, path_map, then) + self.branches[source][name] = Branch(path, path_map_, then) def set_entry_point(self, key: str) -> None: """Specifies the first node to be called in the graph. @@ -405,7 +417,6 @@ class Graph: # create empty compiled graph compiled = CompiledGraph( - builder=self, nodes={}, channels={START: EphemeralValue(Any), END: EphemeralValue(Any)}, input_channels=START, @@ -418,6 +429,7 @@ class Graph: auto_validate=False, debug=debug, ) + compiled.builder = self # attach nodes, edges, and branches for key, node in self.nodes.items(): @@ -437,10 +449,6 @@ class Graph: class CompiledGraph(Pregel): builder: Graph - def __init__(self, *, builder: Graph, **kwargs): - super().__init__(**kwargs) - self.builder = builder - def attach_node(self, key: str, node: NodeSpec) -> None: self.channels[key] = EphemeralValue(Any) self.nodes[key] = ( @@ -463,7 +471,7 @@ class CompiledGraph(Pregel): def attach_branch(self, start: str, name: str, branch: Branch) -> None: def branch_writer( - packets: list[Union[str, Send]], config: RunnableConfig + packets: Sequence[Union[str, Send]], config: RunnableConfig ) -> Optional[ChannelWrite]: writes = [ ( @@ -473,7 +481,10 @@ class CompiledGraph(Pregel): ) for p in packets ] - return ChannelWrite(writes, tags=[TAG_HIDDEN]) + return ChannelWrite( + cast(Sequence[Union[ChannelWriteEntry, Send]], writes), + tags=[TAG_HIDDEN], + ) # add hidden start node if start == START and start not in self.nodes: @@ -489,7 +500,7 @@ class CompiledGraph(Pregel): channel_name = f"branch:{start}:{name}:{end}" self.channels[channel_name] = EphemeralValue(Any) self.nodes[end].triggers.append(channel_name) - self.nodes[end].channels.append(channel_name) + cast(list[str], self.nodes[end].channels).append(channel_name) def get_graph( self, @@ -504,17 +515,25 @@ class CompiledGraph(Pregel): } end_nodes: dict[str, DrawableNode] = {} if xray: - subgraphs = dict(self.get_subgraphs()) + subgraphs = { + k: v for k, v in self.get_subgraphs() if isinstance(v, CompiledGraph) + } else: subgraphs = {} def add_edge( - start: str, end: str, label: Optional[str] = None, conditional: bool = False + start: str, + end: str, + label: Optional[Hashable] = None, + conditional: bool = False, ) -> None: if end == END and END not in end_nodes: end_nodes[END] = graph.add_node(self.get_output_schema(config), END) return graph.add_edge( - start_nodes[start], end_nodes[end], label, conditional + start_nodes[start], + end_nodes[end], + str(label) if label is not None else None, + conditional, ) for key, n in self.builder.nodes.items(): @@ -563,7 +582,7 @@ class CompiledGraph(Pregel): elif branch.then is not None: ends = {k: k for k in default_ends if k not in (END, branch.then)} else: - ends = default_ends + ends = cast(dict[Hashable, str], default_ends) for label, end in ends.items(): add_edge( start, diff --git a/libs/langgraph/langgraph/graph/message.py b/libs/langgraph/langgraph/graph/message.py index 34bc9c090..6575bd10c 100644 --- a/libs/langgraph/langgraph/graph/message.py +++ b/libs/langgraph/langgraph/graph/message.py @@ -63,9 +63,9 @@ def add_messages(left: Messages, right: Messages) -> Messages: """ # coerce to list if not isinstance(left, list): - left = [left] + left = [left] # type: ignore[assignment] if not isinstance(right, list): - right = [right] + right = [right] # type: ignore[assignment] # coerce to message left = [ message_chunk_to_message(cast(BaseMessageChunk, m)) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index e66cff53d..803434316 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -7,11 +7,13 @@ from inspect import isclass, isfunction, signature from typing import ( Any, Callable, + Literal, NamedTuple, Optional, Sequence, Type, Union, + cast, get_origin, get_type_hints, overload, @@ -122,7 +124,7 @@ class StateGraph(Graph): >>> print(step1) {'x': [0.5, 0.75]}""" - nodes: dict[str, StateNodeSpec] + nodes: dict[str, StateNodeSpec] # type: ignore[assignment] channels: dict[str, BaseChannel] managed: dict[str, ManagedValueSpec] schemas: dict[Type[Any], dict[str, Union[BaseChannel, ManagedValueSpec]]] @@ -302,7 +304,7 @@ class StateGraph(Graph): if not isinstance(node, str): action = node if isinstance(action, Runnable): - node = action.name + node = action.get_name() else: node = getattr(action, "__name__", action.__class__.__name__) if node is None: @@ -323,13 +325,15 @@ class StateGraph(Graph): raise ValueError( "Node name must be provided if action is not a function" ) + if action is None: + raise RuntimeError if node in self.nodes: raise ValueError(f"Node `{node}` already present.") if node == END or node == START: raise ValueError(f"Node `{node}` is reserved.") for character in (NS_SEP, NS_END): - if character in node: + if character in cast(str, node): raise ValueError( f"'{character}' is a reserved character and is not allowed in the node names." ) @@ -349,8 +353,8 @@ class StateGraph(Graph): pass if input is not None: self._add_schema(input) - self.nodes[node] = StateNodeSpec( - coerce_to_runnable(action, name=node, trace=False), + self.nodes[cast(str, node)] = StateNodeSpec( + coerce_to_runnable(action, name=cast(str, node), trace=False), metadata, input=input or self.schema, retry_policy=retry, @@ -449,7 +453,6 @@ class StateGraph(Graph): ) compiled = CompiledStateGraph( - builder=self, config_type=self.config_schema, nodes={}, channels={ @@ -468,6 +471,7 @@ class StateGraph(Graph): debug=debug, store=store, ) + compiled.builder = self compiled.attach_node(START, None) for key, node in self.nodes.items(): @@ -618,7 +622,7 @@ class CompiledStateGraph(CompiledGraph): def attach_branch(self, start: str, name: str, branch: Branch) -> None: def branch_writer( - packets: list[Union[str, Send]], config: RunnableConfig + packets: Sequence[Union[str, Send]], config: RunnableConfig ) -> None: if filtered := [p for p in packets if p != END]: writes = [ @@ -638,7 +642,9 @@ class CompiledStateGraph(CompiledGraph): ), ) ) - ChannelWrite.do_write(config, writes) + ChannelWrite.do_write( + config, cast(Sequence[Union[Send, ChannelWriteEntry]], writes) + ) # attach branch publisher schema = ( @@ -708,11 +714,23 @@ def _get_channels( if name != "__slots__" } return ( - {k: v for k, v in all_keys.items() if not is_managed_value(v)}, + {k: v for k, v in all_keys.items() if isinstance(v, BaseChannel)}, {k: v for k, v in all_keys.items() if is_managed_value(v)}, ) +@overload +def _get_channel( + name: str, annotation: Any, *, allow_managed: Literal[False] +) -> BaseChannel: ... + + +@overload +def _get_channel( + name: str, annotation: Any, *, allow_managed: Literal[True] = True +) -> Union[BaseChannel, ManagedValueSpec]: ... + + def _get_channel( name: str, annotation: Any, *, allow_managed: bool = True ) -> Union[BaseChannel, ManagedValueSpec]: @@ -728,7 +746,7 @@ def _get_channel( channel.key = name return channel - fallback = LastValue(annotation) + fallback: LastValue = LastValue(annotation) fallback.key = name return fallback diff --git a/libs/langgraph/langgraph/prebuilt/chat_agent_executor.py b/libs/langgraph/langgraph/prebuilt/chat_agent_executor.py index d4a6157cb..dd2cddb5e 100644 --- a/libs/langgraph/langgraph/prebuilt/chat_agent_executor.py +++ b/libs/langgraph/langgraph/prebuilt/chat_agent_executor.py @@ -419,7 +419,7 @@ def create_react_agent( raise ValueError(f"Missing required key(s) {missing_keys} in state_schema") if isinstance(tools, ToolExecutor): - tool_classes = tools.tools + tool_classes: Sequence[BaseTool] = tools.tools tool_node = ToolNode(tool_classes) elif isinstance(tools, ToolNode): tool_classes = list(tools.tools_by_name.values()) diff --git a/libs/langgraph/langgraph/prebuilt/tool_node.py b/libs/langgraph/langgraph/prebuilt/tool_node.py index c0f7c83f6..52b80f75a 100644 --- a/libs/langgraph/langgraph/prebuilt/tool_node.py +++ b/libs/langgraph/langgraph/prebuilt/tool_node.py @@ -382,7 +382,7 @@ def _get_state_args(tool: BaseTool) -> Dict[str, Optional[str]]: full_schema = tool.get_input_schema() tool_args_to_state_fields: Dict = {} - def _is_injection(type_arg: Any): + def _is_injection(type_arg: Any) -> bool: if isinstance(type_arg, InjectedState) or ( isinstance(type_arg, type) and issubclass(type_arg, InjectedState) ): diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index 3cf791ef8..9f16ac698 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -138,7 +138,7 @@ class Channel: ) return PregelNode( channels=cast( - Union[Mapping[None, str], Mapping[str, str]], + Union[list[str], Mapping[str, str]], ( {key: channels} if isinstance(channels, str) and key is not None @@ -305,7 +305,9 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): @property def InputType(self) -> Any: if isinstance(self.input_channels, str): - return self.channels[self.input_channels].UpdateType + channel = self.channels[self.input_channels] + if isinstance(channel, BaseChannel): + return channel.UpdateType def get_input_schema( self, config: Optional[RunnableConfig] = None @@ -317,9 +319,9 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): return create_model( self.get_name("Input"), field_definitions={ - k: (self.channels[k].UpdateType, None) + k: (c.UpdateType, None) for k in self.input_channels or self.channels.keys() - if isinstance(self.channels[k], BaseChannel) + if (c := self.channels[k]) and isinstance(c, BaseChannel) }, ) @@ -335,7 +337,9 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): @property def OutputType(self) -> Any: if isinstance(self.output_channels, str): - return self.channels[self.output_channels].ValueType + channel = self.channels[self.output_channels] + if isinstance(channel, BaseChannel): + return channel.ValueType def get_output_schema( self, config: Optional[RunnableConfig] = None @@ -347,9 +351,9 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): return create_model( self.get_name("Output"), field_definitions={ - k: (self.channels[k].ValueType, None) + k: (c.ValueType, None) for k in self.output_channels - if isinstance(self.channels[k], BaseChannel) + if (c := self.channels[k]) and isinstance(c, BaseChannel) }, ) @@ -1050,8 +1054,8 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): bool, set[StreamMode], Union[str, Sequence[str]], - Optional[Sequence[str]], - Optional[Sequence[str]], + Union[All, Sequence[str]], + Union[All, Sequence[str]], Optional[BaseCheckpointSaver], ]: debug = debug if debug is not None else self.debug @@ -1199,8 +1203,8 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): debug, stream_modes, output_keys, - interrupt_before, - interrupt_after, + interrupt_before_, + interrupt_after_, checkpointer, ) = self._defaults( config, @@ -1253,7 +1257,7 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): else: return waiter else: - get_waiter = None + get_waiter = None # type: ignore[assignment] # 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 @@ -1261,8 +1265,8 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): # with channel updates applied only at the transition between steps while loop.tick( input_keys=self.input_channels, - interrupt_before=interrupt_before, - interrupt_after=interrupt_after, + interrupt_before=interrupt_before_, + interrupt_after=interrupt_after_, manager=run_manager, ): for _ in runner.tick( @@ -1397,7 +1401,7 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): # if running from astream_log() run each proc with streaming do_stream = next( ( - h + cast(_StreamingCallbackHandler, h) for h in run_manager.handlers if isinstance(h, _StreamingCallbackHandler) ), @@ -1415,8 +1419,8 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): debug, stream_modes, output_keys, - interrupt_before, - interrupt_after, + interrupt_before_, + interrupt_after_, checkpointer, ) = self._defaults( config, @@ -1457,7 +1461,7 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): def get_waiter() -> asyncio.Task[None]: return aioloop.create_task(stream.wait()) else: - get_waiter = None + get_waiter = None # type: ignore[assignment] # 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 @@ -1465,8 +1469,8 @@ class Pregel(Runnable[Union[dict[str, Any], Any], Union[dict[str, Any], Any]]): # with channel updates applied only at the transition between steps while loop.tick( input_keys=self.input_channels, - interrupt_before=interrupt_before, - interrupt_after=interrupt_after, + interrupt_before=interrupt_before_, + interrupt_after=interrupt_after_, manager=run_manager, ): async for _ in runner.atick( diff --git a/libs/langgraph/langgraph/pregel/algo.py b/libs/langgraph/langgraph/pregel/algo.py index bf4baeb0f..40a10a40f 100644 --- a/libs/langgraph/langgraph/pregel/algo.py +++ b/libs/langgraph/langgraph/pregel/algo.py @@ -52,6 +52,8 @@ from langgraph.pregel.read import PregelNode from langgraph.pregel.types import All, PregelExecutableTask, PregelTask from langgraph.utils.config import merge_configs, patch_config +GetNextVersion = Callable[[Optional[V], BaseChannel], V] + EMPTY_SEQ: tuple[str, ...] = tuple() @@ -173,7 +175,7 @@ def apply_writes( checkpoint: Checkpoint, channels: Mapping[str, BaseChannel], tasks: Iterable[WritesProtocol], - get_next_version: Optional[Callable[[Optional[V], BaseChannel], V]], + get_next_version: Optional[GetNextVersion], ) -> dict[str, list[Any]]: # update seen versions for task in tasks: @@ -200,7 +202,7 @@ def apply_writes( }: if channels[chan].consume() and get_next_version is not None: checkpoint["channel_versions"][chan] = get_next_version( - max_version, # type: ignore[arg-type] + max_version, channels[chan], ) @@ -234,7 +236,7 @@ def apply_writes( if chan in channels: if channels[chan].update(vals) and get_next_version is not None: checkpoint["channel_versions"][chan] = get_next_version( - max_version, # type: ignore[arg-type] + max_version, channels[chan], ) updated_channels.add(chan) @@ -244,7 +246,7 @@ def apply_writes( if chan not in updated_channels: if channels[chan].update([]) and get_next_version is not None: checkpoint["channel_versions"][chan] = get_next_version( - max_version, # type: ignore[arg-type] + max_version, channels[chan], ) diff --git a/libs/langgraph/langgraph/pregel/debug.py b/libs/langgraph/langgraph/pregel/debug.py index 782d5a13c..56f1eb9c6 100644 --- a/libs/langgraph/langgraph/pregel/debug.py +++ b/libs/langgraph/langgraph/pregel/debug.py @@ -82,7 +82,7 @@ TASK_NAMESPACE = UUID("6ba7b831-9dad-11d1-80b4-00c04fd430c8") def map_debug_tasks( - step: int, tasks: list[PregelExecutableTask] + step: int, tasks: Iterable[PregelExecutableTask] ) -> Iterator[DebugOutputTask]: ts = datetime.now(timezone.utc).isoformat() for task in tasks: @@ -132,7 +132,7 @@ def map_debug_checkpoint( stream_channels: Union[str, Sequence[str]], metadata: CheckpointMetadata, checkpoint: Checkpoint, - tasks: list[PregelExecutableTask], + tasks: Iterable[PregelExecutableTask], pending_writes: list[PendingWrite], ) -> Iterator[DebugOutputCheckpoint]: yield { diff --git a/libs/langgraph/langgraph/pregel/loop.py b/libs/langgraph/langgraph/pregel/loop.py index ba98b1605..49baa1846 100644 --- a/libs/langgraph/langgraph/pregel/loop.py +++ b/libs/langgraph/langgraph/pregel/loop.py @@ -64,6 +64,7 @@ from langgraph.managed.base import ( WritableManagedValue, ) from langgraph.pregel.algo import ( + GetNextVersion, PregelTaskWrites, apply_writes, increment, @@ -92,7 +93,7 @@ from langgraph.pregel.io import ( ) from langgraph.pregel.manager import AsyncChannelsManager, ChannelsManager from langgraph.pregel.read import PregelNode -from langgraph.pregel.types import PregelExecutableTask, StreamMode +from langgraph.pregel.types import All, PregelExecutableTask, StreamMode from langgraph.pregel.utils import get_new_channel_versions from langgraph.store.base import BaseStore from langgraph.store.batch import AsyncBatchedStore @@ -146,7 +147,7 @@ class PregelLoop: skip_done_tasks: bool is_nested: bool - checkpointer_get_next_version: Callable[[Optional[V]], V] + checkpointer_get_next_version: GetNextVersion checkpointer_put_writes: Optional[ Callable[[RunnableConfig, Sequence[tuple[str, Any]], str], Any] ] @@ -281,8 +282,8 @@ class PregelLoop: self, *, input_keys: Union[str, Sequence[str]], - interrupt_after: Sequence[str] = EMPTY_SEQ, - interrupt_before: Sequence[str] = EMPTY_SEQ, + interrupt_after: Union[All, Sequence[str]] = EMPTY_SEQ, + interrupt_before: Union[All, Sequence[str]] = EMPTY_SEQ, manager: Union[None, AsyncParentRunManager, ParentRunManager] = None, ) -> bool: """Execute a single iteration of the Pregel loop. @@ -681,6 +682,10 @@ class SyncPregelLoop(PregelLoop, ContextManager): if self.config.get("configurable", {}).get( CONFIG_KEY_ENSURE_LATEST ) and self.checkpoint_config["configurable"].get("checkpoint_id"): + if self.checkpointer is None: + raise RuntimeError( + "Cannot ensure latest checkpoint without checkpointer" + ) saved = self.checkpointer.get_tuple( patch_configurable(self.checkpoint_config, {"checkpoint_id": None}) ) @@ -771,7 +776,7 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager): self.checkpointer_put_writes = checkpointer.aput_writes else: self.checkpointer_get_next_version = increment - self._checkpointer_put_after_previous = None # type: ignore[method-assign] + self._checkpointer_put_after_previous = None # type: ignore[assignment] self.checkpointer_put_writes = None async def _checkpointer_put_after_previous( @@ -801,6 +806,10 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager): if self.config.get("configurable", {}).get( CONFIG_KEY_ENSURE_LATEST ) and self.checkpoint_config["configurable"].get("checkpoint_id"): + if self.checkpointer is None: + raise RuntimeError( + "Cannot ensure latest checkpoint without checkpointer" + ) saved = await self.checkpointer.aget_tuple( patch_configurable(self.checkpoint_config, {"checkpoint_id": None}) ) @@ -858,6 +867,3 @@ class AsyncPregelLoop(PregelLoop, AsyncContextManager): return await asyncio.shield( self.stack.__aexit__(exc_type, exc_value, traceback) ) - - -EMPTY_SEQ = tuple() diff --git a/libs/langgraph/langgraph/pregel/messages.py b/libs/langgraph/langgraph/pregel/messages.py index 0a96f0fee..7c3f90b10 100644 --- a/libs/langgraph/langgraph/pregel/messages.py +++ b/libs/langgraph/langgraph/pregel/messages.py @@ -7,7 +7,6 @@ from typing import ( List, Optional, Sequence, - Tuple, Union, cast, ) @@ -21,20 +20,16 @@ from langchain_core.tracers._streaming import T, _StreamingCallbackHandler from langgraph.constants import NS_SEP from langgraph.pregel.loop import StreamChunk +Meta = tuple[tuple[str, ...], dict[str, Any]] + class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler): def __init__(self, stream: Callable[[StreamChunk], None]): self.stream = stream - self.metadata: dict[UUID, tuple[tuple[str, ...], dict[str, Any]]] = {} + self.metadata: dict[UUID, Meta] = {} self.seen: set[Union[int, str]] = set() - def _emit( - self, - meta: Tuple[str, dict[str, Any]], - message: BaseMessage, - *, - dedupe: bool = False, - ) -> None: + def _emit(self, meta: Meta, message: BaseMessage, *, dedupe: bool = False) -> None: ident = id(message) if dedupe and message.id in self.seen: return diff --git a/libs/langgraph/langgraph/pregel/read.py b/libs/langgraph/langgraph/pregel/read.py index 79643a090..3ad988b89 100644 --- a/libs/langgraph/langgraph/pregel/read.py +++ b/libs/langgraph/langgraph/pregel/read.py @@ -18,7 +18,7 @@ from langchain_core.runnables import ( RunnablePassthrough, RunnableSerializable, ) -from langchain_core.runnables.base import Input, Other, Output, coerce_to_runnable +from langchain_core.runnables.base import Input, Other, coerce_to_runnable from langchain_core.runnables.utils import ConfigurableFieldSpec from langgraph.constants import CONFIG_KEY_READ @@ -206,7 +206,7 @@ class PregelNode(Runnable): Mapping[str, Runnable[Any, Other] | Callable[[Any], Other]], ], ) -> PregelNode: - if ChannelWrite.is_writer(other): + if isinstance(other, Runnable) and ChannelWrite.is_writer(other): return self.copy(update=dict(writers=[*self.writers, other])) elif self.bound is DEFAULT_BOUND: return self.copy(update=dict(bound=coerce_to_runnable(other))) @@ -237,7 +237,7 @@ class PregelNode(Runnable): input: Input, config: Optional[RunnableConfig] = None, **kwargs: Optional[Any], - ) -> Output: + ) -> Any: return self.bound.invoke( input, merge_configs({"metadata": self.metadata, "tags": self.tags}, config), @@ -249,7 +249,7 @@ class PregelNode(Runnable): input: Input, config: Optional[RunnableConfig] = None, **kwargs: Optional[Any], - ) -> Output: + ) -> Any: return await self.bound.ainvoke( input, merge_configs({"metadata": self.metadata, "tags": self.tags}, config), @@ -261,7 +261,7 @@ class PregelNode(Runnable): input: Input, config: Optional[RunnableConfig] = None, **kwargs: Optional[Any], - ) -> Iterator[Output]: + ) -> Iterator[Any]: yield from self.bound.stream( input, merge_configs({"metadata": self.metadata, "tags": self.tags}, config), @@ -273,7 +273,7 @@ class PregelNode(Runnable): input: Input, config: Optional[RunnableConfig] = None, **kwargs: Optional[Any], - ) -> AsyncIterator[Output]: + ) -> AsyncIterator[Any]: async for item in self.bound.astream( input, merge_configs({"metadata": self.metadata, "tags": self.tags}, config), diff --git a/libs/langgraph/langgraph/pregel/runner.py b/libs/langgraph/langgraph/pregel/runner.py index 7e1122485..14e84352f 100644 --- a/libs/langgraph/langgraph/pregel/runner.py +++ b/libs/langgraph/langgraph/pregel/runner.py @@ -5,11 +5,13 @@ from typing import ( Any, AsyncIterator, Callable, + Iterable, Iterator, Optional, Sequence, Type, Union, + cast, ) from langgraph.constants import ERROR, INTERRUPT, NO_WRITES @@ -33,7 +35,7 @@ class PregelRunner: def tick( self, - tasks: Sequence[PregelExecutableTask], + tasks: Iterable[PregelExecutableTask], *, reraise: bool = True, timeout: Optional[float] = None, @@ -106,7 +108,7 @@ class PregelRunner: async def atick( self, - tasks: Sequence[PregelExecutableTask], + tasks: Iterable[PregelExecutableTask], *, reraise: bool = True, timeout: Optional[float] = None, @@ -141,14 +143,17 @@ class PregelRunner: for t in tasks: if not t.writes: futures[ - self.submit( - arun_with_retry, - t, - retry_policy, - stream=self.use_astream, - __name__=t.name, - __cancel_on_exit__=True, - __reraise_on_exit__=reraise, + cast( + asyncio.Future, + self.submit( + arun_with_retry, + t, + retry_policy, + stream=self.use_astream, + __name__=t.name, + __cancel_on_exit__=True, + __reraise_on_exit__=reraise, + ), ) ] = t all_futures = futures.copy() diff --git a/libs/langgraph/langgraph/pregel/validate.py b/libs/langgraph/langgraph/pregel/validate.py index 8627642e9..232014240 100644 --- a/libs/langgraph/langgraph/pregel/validate.py +++ b/libs/langgraph/langgraph/pregel/validate.py @@ -1,4 +1,4 @@ -from typing import Mapping, Optional, Sequence, Union +from typing import Any, Mapping, Optional, Sequence, Union from langgraph.channels.base import BaseChannel from langgraph.constants import RESERVED @@ -65,18 +65,18 @@ def validate_graph( raise ValueError(f"Output channel '{chan}' not in 'channels'") if interrupt_after_nodes != "*": - for node in interrupt_after_nodes: - if node not in nodes: - raise ValueError(f"Node {node} not in nodes") + for n in interrupt_after_nodes: + if n not in nodes: + raise ValueError(f"Node {n} not in nodes") if interrupt_before_nodes != "*": - for node in interrupt_before_nodes: - if node not in nodes: - raise ValueError(f"Node {node} not in nodes") + for n in interrupt_before_nodes: + if n not in nodes: + raise ValueError(f"Node {n} not in nodes") def validate_keys( keys: Optional[Union[str, Sequence[str]]], - channels: Mapping[str, BaseChannel], + channels: Mapping[str, Any], ) -> None: if isinstance(keys, str): if keys not in channels: diff --git a/libs/langgraph/langgraph/pregel/write.py b/libs/langgraph/langgraph/pregel/write.py index 9c3b7782d..2adcab757 100644 --- a/libs/langgraph/langgraph/pregel/write.py +++ b/libs/langgraph/langgraph/pregel/write.py @@ -9,6 +9,7 @@ from typing import ( Sequence, TypeVar, Union, + cast, ) from langchain_core.runnables import Runnable, RunnableConfig @@ -34,7 +35,7 @@ class ChannelWriteEntry(NamedTuple): class ChannelWrite(RunnableCallable): - writes: Sequence[Union[ChannelWriteEntry, Send]] + writes: list[Union[ChannelWriteEntry, Send]] """ Sequence of write entries, each of which is a tuple of: - channel name @@ -54,7 +55,7 @@ class ChannelWrite(RunnableCallable): require_at_least_one_of: Optional[Sequence[str]] = None, ): super().__init__(func=self._write, afunc=self._awrite, name=None, tags=tags) - self.writes = writes + self.writes = cast(list[Union[ChannelWriteEntry, Send]], writes) self.require_at_least_one_of = require_at_least_one_of def get_name( diff --git a/libs/langgraph/langgraph/utils/runnable.py b/libs/langgraph/langgraph/utils/runnable.py index f0b16442d..56d5d5df4 100644 --- a/libs/langgraph/langgraph/utils/runnable.py +++ b/libs/langgraph/langgraph/utils/runnable.py @@ -269,7 +269,7 @@ class RunnableSeq(Runnable): if isinstance(step, RunnableSequence): steps_flat.extend(step.steps) elif isinstance(step, RunnableSeq): - steps_flat.extend(step.steps) # type: ignore[has-type] + steps_flat.extend(step.steps) else: steps_flat.append(coerce_to_runnable(step, name=None, trace=True)) if len(steps_flat) < 2: diff --git a/libs/langgraph/pyproject.toml b/libs/langgraph/pyproject.toml index 954ca2973..4975cabb0 100644 --- a/libs/langgraph/pyproject.toml +++ b/libs/langgraph/pyproject.toml @@ -57,7 +57,7 @@ warn_no_return = "False" warn_unused_ignores = "True" warn_redundant_casts = "True" allow_redefinition = "True" -disable_error_code = "typeddict-item, return-value, override" +disable_error_code = "typeddict-item, return-value, override, has-type" [tool.coverage.run] omit = ["tests/*"]