feat(langgraph): add type checking for matching node signatures vs input_schema for add_node (#5424)

This commit is contained in:
Sydney Runkle
2025-07-10 09:42:37 -04:00
committed by GitHub
parent e5ded1888b
commit 5f00938aa2
5 changed files with 181 additions and 56 deletions
+35 -31
View File
@@ -1,7 +1,9 @@
from __future__ import annotations
import sys
from collections.abc import Sequence
from typing import Any, NamedTuple, Protocol, Union
from dataclasses import dataclass
from typing import Any, Generic, Protocol, Union
from langchain_core.runnables import Runnable, RunnableConfig
from typing_extensions import TypeAlias
@@ -9,47 +11,49 @@ from typing_extensions import TypeAlias
from langgraph.constants import EMPTY_SEQ
from langgraph.store.base import BaseStore
from langgraph.types import CachePolicy, RetryPolicy, StreamWriter
from langgraph.typing import StateT_contra
from langgraph.typing import NodeInputT, NodeInputT_contra
_DC_SLOTS = {"slots": True} if sys.version_info >= (3, 10) else {}
class _Node(Protocol[StateT_contra]):
def __call__(self, state: StateT_contra) -> Any: ...
class _Node(Protocol[NodeInputT_contra]):
def __call__(self, state: NodeInputT_contra) -> Any: ...
class _NodeWithConfig(Protocol[StateT_contra]):
def __call__(self, state: StateT_contra, config: RunnableConfig) -> Any: ...
class _NodeWithConfig(Protocol[NodeInputT_contra]):
def __call__(self, state: NodeInputT_contra, config: RunnableConfig) -> Any: ...
class _NodeWithWriter(Protocol[StateT_contra]):
def __call__(self, state: StateT_contra, *, writer: StreamWriter) -> Any: ...
class _NodeWithWriter(Protocol[NodeInputT_contra]):
def __call__(self, state: NodeInputT_contra, *, writer: StreamWriter) -> Any: ...
class _NodeWithStore(Protocol[StateT_contra]):
def __call__(self, state: StateT_contra, *, store: BaseStore) -> Any: ...
class _NodeWithStore(Protocol[NodeInputT_contra]):
def __call__(self, state: NodeInputT_contra, *, store: BaseStore) -> Any: ...
class _NodeWithWriterStore(Protocol[StateT_contra]):
class _NodeWithWriterStore(Protocol[NodeInputT_contra]):
def __call__(
self, state: StateT_contra, *, writer: StreamWriter, store: BaseStore
self, state: NodeInputT_contra, *, writer: StreamWriter, store: BaseStore
) -> Any: ...
class _NodeWithConfigWriter(Protocol[StateT_contra]):
class _NodeWithConfigWriter(Protocol[NodeInputT_contra]):
def __call__(
self, state: StateT_contra, *, config: RunnableConfig, writer: StreamWriter
self, state: NodeInputT_contra, *, config: RunnableConfig, writer: StreamWriter
) -> Any: ...
class _NodeWithConfigStore(Protocol[StateT_contra]):
class _NodeWithConfigStore(Protocol[NodeInputT_contra]):
def __call__(
self, state: StateT_contra, *, config: RunnableConfig, store: BaseStore
self, state: NodeInputT_contra, *, config: RunnableConfig, store: BaseStore
) -> Any: ...
class _NodeWithConfigWriterStore(Protocol[StateT_contra]):
class _NodeWithConfigWriterStore(Protocol[NodeInputT_contra]):
def __call__(
self,
state: StateT_contra,
state: NodeInputT_contra,
*,
config: RunnableConfig,
writer: StreamWriter,
@@ -61,23 +65,23 @@ class _NodeWithConfigWriterStore(Protocol[StateT_contra]):
# we move to adding a context arg. Maybe what we do is we add support for kwargs with param spec
# this is purely for typing purposes though, so can easily change in the coming weeks.
StateNode: TypeAlias = Union[
_Node[StateT_contra],
_NodeWithConfig[StateT_contra],
_NodeWithWriter[StateT_contra],
_NodeWithStore[StateT_contra],
_NodeWithWriterStore[StateT_contra],
_NodeWithConfigWriter[StateT_contra],
_NodeWithConfigStore[StateT_contra],
_NodeWithConfigWriterStore[StateT_contra],
Runnable[StateT_contra, Any],
_Node[NodeInputT],
_NodeWithConfig[NodeInputT],
_NodeWithWriter[NodeInputT],
_NodeWithStore[NodeInputT],
_NodeWithWriterStore[NodeInputT],
_NodeWithConfigWriter[NodeInputT],
_NodeWithConfigStore[NodeInputT],
_NodeWithConfigWriterStore[NodeInputT],
Runnable[NodeInputT, Any],
]
# TODO: use a dataclass generic on NodeInputType
class StateNodeSpec(NamedTuple):
runnable: StateNode
@dataclass(**_DC_SLOTS)
class StateNodeSpec(Generic[NodeInputT]):
runnable: StateNode[NodeInputT]
metadata: dict[str, Any] | None
input_schema: type[Any]
input_schema: type[NodeInputT]
retry_policy: RetryPolicy | Sequence[RetryPolicy] | None
cache_policy: CachePolicy | None
ends: tuple[str, ...] | dict[str, str] | None = EMPTY_SEQ
+81 -20
View File
@@ -83,7 +83,7 @@ from langgraph.types import (
RetryPolicy,
Send,
)
from langgraph.typing import InputT, OutputT, StateT
from langgraph.typing import InputT, NodeInputT, OutputT, StateT
from langgraph.warnings import LangGraphDeprecatedSinceV05
__all__ = ("StateGraph", "CompiledStateGraph")
@@ -267,13 +267,31 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
*,
defer: bool = False,
metadata: dict[str, Any] | None = None,
input_schema: type[Any] | None = None,
input_schema: None = None,
retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None,
cache_policy: CachePolicy | None = None,
destinations: dict[str, str] | tuple[str, ...] | None = None,
**kwargs: Unpack[DeprecatedKwargs],
) -> Self:
"""Add a new node to the state graph.
"""Add a new node to the state graph, input schema is inferred as the state schema.
Will take the name of the function/runnable as the node name.
"""
...
@overload
def add_node(
self,
node: StateNode[NodeInputT],
*,
defer: bool = False,
metadata: dict[str, Any] | None = None,
input_schema: type[NodeInputT],
retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None,
cache_policy: CachePolicy | None = None,
destinations: dict[str, str] | tuple[str, ...] | None = None,
**kwargs: Unpack[DeprecatedKwargs],
) -> Self:
"""Add a new node to the state graph, input schema is specified.
Will take the name of the function/runnable as the node name.
"""
...
@@ -286,23 +304,40 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
*,
defer: bool = False,
metadata: dict[str, Any] | None = None,
input_schema: type[Any] | None = None,
input_schema: None = None,
retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None,
cache_policy: CachePolicy | None = None,
destinations: dict[str, str] | tuple[str, ...] | None = None,
**kwargs: Unpack[DeprecatedKwargs],
) -> Self:
"""Add a new node to the state graph."""
"""Add a new node to the state graph, input schema is inferred as the state schema."""
...
@overload
def add_node(
self,
node: str,
action: StateNode[NodeInputT],
*,
defer: bool = False,
metadata: dict[str, Any] | None = None,
input_schema: type[NodeInputT],
retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None,
cache_policy: CachePolicy | None = None,
destinations: dict[str, str] | tuple[str, ...] | None = None,
**kwargs: Unpack[DeprecatedKwargs],
) -> Self:
"""Add a new node to the state graph, input schema is specified."""
...
def add_node(
self,
node: str | StateNode[StateT],
action: StateNode[StateT] | None = None,
node: str | StateNode[StateT] | StateNode[NodeInputT],
action: StateNode[StateT] | StateNode[NodeInputT] | None = None,
*,
defer: bool = False,
metadata: dict[str, Any] | None = None,
input_schema: type[Any] | None = None,
input_schema: type[NodeInputT] | None = None,
retry_policy: RetryPolicy | Sequence[RetryPolicy] | None = None,
cache_policy: CachePolicy | None = None,
destinations: dict[str, str] | tuple[str, ...] | None = None,
@@ -375,7 +410,7 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
category=LangGraphDeprecatedSinceV05,
)
if input_schema is None:
input_schema = cast(Union[type[InputT], None], input_)
input_schema = cast(Union[type[NodeInputT], None], input_)
if not isinstance(node, str):
action = node
@@ -412,6 +447,8 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
f"'{character}' is a reserved character and is not allowed in the node names."
)
inferred_input_schema = None
ends: tuple[str, ...] | dict[str, str] = EMPTY_SEQ
try:
if (
@@ -432,7 +469,7 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
)
if input_hint := hints.get(first_parameter_name):
if isinstance(input_hint, type) and get_type_hints(input_hint):
input_schema = input_hint
inferred_input_schema = input_hint
if rtn := hints.get("return"):
# Handle Union types
rtn_origin = get_origin(rtn)
@@ -460,17 +497,41 @@ class StateGraph(Generic[StateT, InputT, OutputT]):
if destinations is not None:
ends = destinations
if input_schema is not None:
self.nodes[node] = StateNodeSpec[NodeInputT](
coerce_to_runnable(action, name=node, trace=False),
metadata,
input_schema=input_schema,
retry_policy=retry_policy,
cache_policy=cache_policy,
ends=ends,
defer=defer,
)
elif inferred_input_schema is not None:
self.nodes[node] = StateNodeSpec(
coerce_to_runnable(action, name=node, trace=False),
metadata,
input_schema=inferred_input_schema,
retry_policy=retry_policy,
cache_policy=cache_policy,
ends=ends,
defer=defer,
)
else:
self.nodes[node] = StateNodeSpec[StateT](
coerce_to_runnable(action, name=node, trace=False),
metadata,
input_schema=self.state_schema,
retry_policy=retry_policy,
cache_policy=cache_policy,
ends=ends,
defer=defer,
)
input_schema = input_schema or inferred_input_schema
if input_schema is not None:
self._add_schema(input_schema)
self.nodes[node] = StateNodeSpec(
coerce_to_runnable(action, name=node, trace=False),
metadata,
input_schema=input_schema or self.state_schema,
retry_policy=retry_policy,
cache_policy=cache_policy,
ends=ends,
defer=defer,
)
return self
def add_edge(self, start_key: str | list[str], end_key: str) -> Self:
@@ -923,7 +984,7 @@ class CompiledStateGraph(
writers=[ChannelWrite(write_entries)],
)
elif node is not None:
input_schema = node.input_schema if node else self.builder._state_schema
input_schema = node.input_schema if node else self.builder.state_schema
input_channels = list(self.builder.schemas[input_schema])
is_single_input = len(input_channels) == 1 and "__root__" in input_channels
if input_schema in self.schema_to_mapper:
+5 -4
View File
@@ -118,6 +118,7 @@ from langgraph.types import (
All,
CachePolicy,
Checkpointer,
Command,
Interrupt,
Send,
StateSnapshot,
@@ -2346,7 +2347,7 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou
def stream(
self,
input: InputT,
input: InputT | Command | None,
config: RunnableConfig | None = None,
*,
stream_mode: StreamMode | Sequence[StreamMode] | None = None,
@@ -2568,7 +2569,7 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou
async def astream(
self,
input: InputT,
input: InputT | Command | None,
config: RunnableConfig | None = None,
*,
stream_mode: StreamMode | Sequence[StreamMode] | None = None,
@@ -2812,7 +2813,7 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou
def invoke(
self,
input: InputT,
input: InputT | Command | None,
config: RunnableConfig | None = None,
*,
stream_mode: StreamMode = "values",
@@ -2887,7 +2888,7 @@ class Pregel(PregelProtocol[StateT, InputT, OutputT], Generic[StateT, InputT, Ou
async def ainvoke(
self,
input: InputT,
input: InputT | Command | None,
config: RunnableConfig | None = None,
*,
stream_mode: StreamMode = "values",
+4 -1
View File
@@ -27,6 +27,9 @@ InputT = TypeVar("InputT", bound=StateLike, default=StateT)
Defaults to `StateT`.
"""
OutputT = TypeVar("OutputT", bound=Union[StateLike, None], default=StateT)
"""Type variable used to represent the output of a state graph."""
NodeInputT = TypeVar("NodeInputT", bound=StateLike)
NodeInputT_contra = TypeVar("NodeInputT_contra", bound=StateLike, contravariant=True)
@@ -2,11 +2,13 @@ from dataclasses import dataclass
from operator import add
from typing import Annotated, Any
import pytest
from langchain_core.runnables import RunnableConfig
from pydantic import BaseModel
from typing_extensions import TypedDict
from langgraph.graph import StateGraph
from langgraph.types import Command
def test_typed_dict_state() -> None:
@@ -103,3 +105,57 @@ def test_input_state_specified() -> None:
new_graph.invoke({"something": 1})
new_graph.invoke({"something": 2, "info": ["hello", "world"]}) # type: ignore[arg-type]
@pytest.mark.skip("Purely for type checking")
def test_invoke_with_all_valid_types() -> None:
class State(TypedDict):
a: int
def a(state: State) -> Any: ...
graph = StateGraph(State).add_node("a", a).set_entry_point("a").compile()
graph.invoke({"a": 1})
graph.invoke(None)
graph.invoke(Command())
def test_add_node_with_explicit_input_schema() -> None:
class A(TypedDict):
a1: int
a2: str
class B(TypedDict):
b1: int
b2: str
class ANarrow(TypedDict):
a1: int
class BNarrow(TypedDict):
b1: int
class State(A, B): ...
def a(state: A) -> Any: ...
def b(state: B) -> Any: ...
workflow = StateGraph(State)
# input schema matches typed schemas
workflow.add_node("a", a, input_schema=A)
workflow.add_node("b", b, input_schema=B)
# input schema does not match typed schemas
workflow.add_node("a_wrong", a, input_schema=B) # type: ignore[arg-type]
workflow.add_node("b_wrong", b, input_schema=A) # type: ignore[arg-type]
# input schema is more broad than the typed schemas, which is allowed
# by the principles of contravariance
workflow.add_node("a_inclusive", a, input_schema=State)
workflow.add_node("b_inclusive", b, input_schema=State)
# input schema is more narrow than the typed schemas, which is not allowed
# because it violates the principles of contravariance
workflow.add_node("a_narrow", a, input_schema=ANarrow) # type: ignore[arg-type]
workflow.add_node("b_narrow", b, input_schema=BNarrow) # type: ignore[arg-type]