From 312f026e9c26e1085753e517d28802b3f2862571 Mon Sep 17 00:00:00 2001 From: William Fu-Hinthorn <13333726+hinthornw@users.noreply.github.com> Date: Wed, 12 Mar 2025 10:54:23 -0700 Subject: [PATCH 1/4] Add tests --- libs/langgraph/langgraph/graph/state.py | 155 +++++++++++++++++++--- libs/langgraph/tests/test_pregel.py | 118 +++++++++++++++- libs/langgraph/tests/test_pregel_async.py | 110 +++++++++++++++ 3 files changed, 364 insertions(+), 19 deletions(-) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 2dcd2e95e..2c6eca4f7 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -26,7 +26,7 @@ from typing import ( from langchain_core.runnables import Runnable, RunnableConfig from pydantic import BaseModel from pydantic.v1 import BaseModel as BaseModelV1 -from typing_extensions import Self +from typing_extensions import Annotated, Self from langgraph._api.deprecation import LangGraphDeprecationWarning from langgraph.channels.base import BaseChannel @@ -626,11 +626,13 @@ class StateGraph(Graph): compiled = CompiledStateGraph( builder=self, config_type=self.config_schema, - input_model=self.input - if len(self.channels) > 1 - and isclass(self.input) - and issubclass(self.input, (BaseModel, BaseModelV1)) - else None, + input_model=( + self.input + if len(self.channels) > 1 + and isclass(self.input) + and issubclass(self.input, (BaseModel, BaseModelV1)) + else None + ), nodes={}, channels={ **self.channels, @@ -940,23 +942,142 @@ def _pick_mapper( ) -> Optional[Callable[[Any], Any]]: if state_keys == ["__root__"]: return None - if issubclass(schema, dict): - return None - if issubclass(schema, BaseModel): - return partial(_coerce_state_pydantic, schema) - if issubclass(schema, BaseModelV1): - return partial(_coerce_state_pydantic_v1, schema) + if isclass(schema): + if issubclass(schema, dict): + return None + if issubclass(schema, BaseModel): + return partial(_coerce_state_pydantic, schema) + if issubclass(schema, BaseModelV1): + return partial(_coerce_state_pydantic_v1, schema) return partial(_coerce_state, schema) -def _coerce_state_pydantic(schema: Type[Any], input: dict[str, Any]) -> dict[str, Any]: - return schema.model_construct(**input) +def _coerce_state_pydantic( + schema: Type[Any], input_data: dict[str, Any], *, __depth__: int = 5 +) -> Any: + if not isinstance(input_data, dict) or __depth__ <= 0: + return input_data + + processed_input = {} + for field_name, field_value in input_data.items(): + if field_name not in schema.model_fields: + processed_input[field_name] = field_value + continue + + field_info = schema.model_fields[field_name] + field_type = field_info.annotation + processed_input[field_name] = _process_field_value( + field_type, field_value, __depth__ - 1 + ) + + return schema.model_construct(**processed_input) def _coerce_state_pydantic_v1( - schema: Type[Any], input: dict[str, Any] -) -> dict[str, Any]: - return schema.construct(**input) + schema: Type[Any], input_data: dict[str, Any], *, __depth__: int = 5 +) -> Any: + if not isinstance(input_data, dict) or __depth__ <= 0: + return input_data + + processed_input = {} + for field_name, field_value in input_data.items(): + if field_name not in schema.__fields__: + processed_input[field_name] = field_value + continue + + field_info = schema.__fields__[field_name] + field_type = field_info.annotation + processed_input[field_name] = _process_field_value( + field_type, field_value, __depth__ - 1 + ) + + return schema.construct(**processed_input) + + +def _process_field_value( + field_type: Type[Any], field_value: Any, __depth__: int +) -> Any: + if __depth__ <= 0 or field_value is None: + return field_value + origin = get_origin(field_type) + + if origin is Annotated: + real_type, *_ = get_args(field_type) + res = _process_field_value(real_type, field_value, __depth__) + return res + + if isclass(field_type): + is_class_ = True + try: + is_model = issubclass(field_type, BaseModel) + except TypeError: + is_class_ = False + is_model = False + if is_model: + if isinstance(field_value, dict): + return _coerce_state_pydantic( + field_type, field_value, __depth__=__depth__ + ) + return field_value + if is_class_ and issubclass(field_type, BaseModelV1): + if isinstance(field_value, dict): + return _coerce_state_pydantic_v1( + field_type, field_value, __depth__=__depth__ + ) + return field_value + + if origin is list or field_type is list: + if not isinstance(field_value, (list, tuple)): + raise TypeError( + f"Expected a list/tuple for {field_type}, got {type(field_value)}." + ) + (item_type,) = get_args(field_type) + return [ + _process_field_value(item_type, item, __depth__ - 1) for item in field_value + ] + + if origin is dict or field_type is dict: + if not isinstance(field_value, dict): + raise TypeError( + f"Expected a dict for {field_type}, got {type(field_value)}." + ) + key_type, val_type = get_args(field_type) + return { + _process_field_value(key_type, k, __depth__ - 1): _process_field_value( + val_type, v, __depth__ - 1 + ) + for k, v in field_value.items() + } + + if origin is tuple: + if not isinstance(field_value, (list, tuple)): + raise TypeError( + f"Expected a tuple/list for {field_type}, got {type(field_value)}." + ) + args = get_args(field_type) + # Handle tuple[type1, type2, ...] with fixed length and different types + result = [] + for i, arg in enumerate(args): + if i < len(field_value): + result.append(_process_field_value(arg, field_value[i], __depth__ - 1)) + else: + # If field_value is shorter than expected, use None for remaining positions + result.append(None) + # If field_value is longer than expected, truncate it + return tuple(result) + + if origin is Union: + for arg in get_args(field_type): + if arg is type(None): + # e.g. Optional + continue + try: + result = _process_field_value(arg, field_value, __depth__ - 1) + return result + except Exception: + pass # Fall back to the next union argument + + return field_value def _coerce_state(schema: Type[Any], input: dict[str, Any]) -> dict[str, Any]: diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 938ee63f3..279d2b892 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -2607,7 +2607,7 @@ def test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic1( arbitrary_types_allowed = True query: str - inner: InnerObject + inner: Annotated[InnerObject, lambda x, y: y] answer: Optional[str] = None docs: Annotated[list[str], sorted_add] client: Annotated[httpx.Client, Context(make_httpx_client)] @@ -2626,9 +2626,11 @@ def test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic1( docs: Optional[list[str]] = None def rewrite_query(data: State) -> State: + assert isinstance(data.inner, InnerObject) return {"query": f"query: {data.query}"} def analyzer_one(data: State) -> State: + assert isinstance(data.inner, InnerObject) return StateUpdate(query=f"analyzed: {data.query}") def retriever_one(data: State) -> State: @@ -2775,7 +2777,7 @@ def test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic2( model_config = ConfigDict(arbitrary_types_allowed=True) query: str - inner: InnerObject + inner: Annotated[InnerObject, lambda x, y: y] answer: Optional[str] = None docs: Annotated[list[str], sorted_add] client: Annotated[httpx.Client, Context(make_httpx_client)] @@ -2794,9 +2796,11 @@ def test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic2( docs: list[str] def rewrite_query(data: State) -> State: + assert isinstance(data.inner, InnerObject) return {"query": f"query: {data.query}"} def analyzer_one(data: State) -> State: + assert isinstance(data.inner, InnerObject) return StateUpdate(query=f"analyzed: {data.query}") def retriever_one(data: State) -> State: @@ -3027,6 +3031,116 @@ def test_in_one_fan_out_state_graph_waiting_edge_custom_state_class_pydantic_inp } +@pytest.mark.parametrize("version", ["v1", "v2"]) +def test_nested_pydantic_models(version: str) -> None: + """Test that nested Pydantic models are properly constructed from leaf nodes up.""" + + # Define nested Pydantic models + if version == "v1": + from pydantic.v1 import BaseModel, Field + else: + from pydantic import BaseModel, Field + + class NestedModel(BaseModel): + value: int + name: str + + # Forward reference model + class RecursiveModel(BaseModel): + value: str + child: Optional["RecursiveModel"] = None + + # Discriminated union models + class Cat(BaseModel): + pet_type: Literal["cat"] + meow: str + + class Dog(BaseModel): + pet_type: Literal["dog"] + bark: str + + # Cyclic reference model + class Person(BaseModel): + id: str + name: str + friends: list[str] = Field(default_factory=list) # IDs of friends + + class State(BaseModel): + # Basic nested model tests + top_level: str + nested: NestedModel + optional_nested: Optional[NestedModel] = None + dict_nested: dict[str, NestedModel] + list_nested: Annotated[ + Union[dict, list[dict[str, NestedModel]]], lambda x, y: (x or []) + [y] + ] + tuple_nested: tuple[str, NestedModel] + tuple_list_nested: list[tuple[int, NestedModel]] + complex_tuple: tuple[str, dict[str, tuple[int, NestedModel]]] + + # Forward reference test + recursive: RecursiveModel + + # Discriminated union test + pet: Union[Cat, Dog] + + # Cyclic reference test + people: dict[str, Person] # Map of ID -> Person + + inputs = { + # Basic nested models + "top_level": "initial", + "nested": {"value": 42, "name": "test"}, + "optional_nested": {"value": 10, "name": "optional"}, + "dict_nested": {"a": {"value": 5, "name": "a"}}, + "list_nested": [{"a": {"value": 6, "name": "b"}}], + "tuple_nested": ["tuple-key", {"value": 7, "name": "tuple-value"}], + "tuple_list_nested": [[1, {"value": 8, "name": "tuple-in-list"}]], + "complex_tuple": [ + "complex", + {"nested": [9, {"value": 10, "name": "deep"}]}, + ], + # Forward reference + "recursive": {"value": "parent", "child": {"value": "child", "child": None}}, + # Discriminated union (using a cat in this case) + "pet": {"pet_type": "cat", "meow": "meow!"}, + # Cyclic references + "people": { + "1": { + "id": "1", + "name": "Alice", + "friends": ["2", "3"], # Alice is friends with Bob and Charlie + }, + "2": { + "id": "2", + "name": "Bob", + "friends": ["1"], # Bob is friends with Alice + }, + "3": { + "id": "3", + "name": "Charlie", + "friends": ["1", "2"], # Charlie is friends with Alice and Bob + }, + }, + } + + update = {"top_level": "updated", "nested": {"value": 100, "name": "updated"}} + + def node_fn(state: State) -> dict: + assert state == State(**inputs) + return update + + builder = StateGraph(State) + builder.add_node("process", node_fn) + builder.set_entry_point("process") + builder.set_finish_point("process") + graph = builder.compile() + + result = graph.invoke(inputs.copy()) + + assert result == {**inputs, **update} + + @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) def test_in_one_fan_out_state_graph_waiting_edge_plus_regular( request: pytest.FixtureRequest, checkpointer_name: str diff --git a/libs/langgraph/tests/test_pregel_async.py b/libs/langgraph/tests/test_pregel_async.py index e131ffe44..105a0c9e6 100644 --- a/libs/langgraph/tests/test_pregel_async.py +++ b/libs/langgraph/tests/test_pregel_async.py @@ -4511,6 +4511,116 @@ async def test_in_one_fan_out_state_graph_waiting_edge_via_branch( ] +@pytest.mark.parametrize("version", ["v1", "v2"]) +async def test_nested_pydantic_models(version: str) -> None: + """Test that nested Pydantic models are properly constructed from leaf nodes up.""" + + # Define nested Pydantic models + if version == "v1": + from pydantic.v1 import BaseModel, Field + else: + from pydantic import BaseModel, Field + + class NestedModel(BaseModel): + value: int + name: str + + # Forward reference model + class RecursiveModel(BaseModel): + value: str + child: Optional["RecursiveModel"] = None + + # Discriminated union models + class Cat(BaseModel): + pet_type: Literal["cat"] + meow: str + + class Dog(BaseModel): + pet_type: Literal["dog"] + bark: str + + # Cyclic reference model + class Person(BaseModel): + id: str + name: str + friends: list[str] = Field(default_factory=list) # IDs of friends + + class State(BaseModel): + # Basic nested model tests + top_level: str + nested: NestedModel + optional_nested: Optional[NestedModel] = None + dict_nested: dict[str, NestedModel] + list_nested: Annotated[ + Union[dict, list[dict[str, NestedModel]]], lambda x, y: (x or []) + [y] + ] + tuple_nested: tuple[str, NestedModel] + tuple_list_nested: list[tuple[int, NestedModel]] + complex_tuple: tuple[str, dict[str, tuple[int, NestedModel]]] + + # Forward reference test + recursive: RecursiveModel + + # Discriminated union test + pet: Union[Cat, Dog] + + # Cyclic reference test + people: dict[str, Person] # Map of ID -> Person + + inputs = { + # Basic nested models + "top_level": "initial", + "nested": {"value": 42, "name": "test"}, + "optional_nested": {"value": 10, "name": "optional"}, + "dict_nested": {"a": {"value": 5, "name": "a"}}, + "list_nested": [{"a": {"value": 6, "name": "b"}}], + "tuple_nested": ["tuple-key", {"value": 7, "name": "tuple-value"}], + "tuple_list_nested": [[1, {"value": 8, "name": "tuple-in-list"}]], + "complex_tuple": [ + "complex", + {"nested": [9, {"value": 10, "name": "deep"}]}, + ], + # Forward reference + "recursive": {"value": "parent", "child": {"value": "child", "child": None}}, + # Discriminated union (using a cat in this case) + "pet": {"pet_type": "cat", "meow": "meow!"}, + # Cyclic references + "people": { + "1": { + "id": "1", + "name": "Alice", + "friends": ["2", "3"], # Alice is friends with Bob and Charlie + }, + "2": { + "id": "2", + "name": "Bob", + "friends": ["1"], # Bob is friends with Alice + }, + "3": { + "id": "3", + "name": "Charlie", + "friends": ["1", "2"], # Charlie is friends with Alice and Bob + }, + }, + } + + update = {"top_level": "updated", "nested": {"value": 100, "name": "updated"}} + + async def node_fn(state: State) -> dict: + assert state == State(**inputs) + return update + + builder = StateGraph(State) + builder.add_node("process", node_fn) + builder.set_entry_point("process") + builder.set_finish_point("process") + graph = builder.compile() + + result = await graph.ainvoke(inputs.copy()) + + assert result == {**inputs, **update} + + @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_ASYNC) async def test_in_one_fan_out_state_graph_waiting_edge_custom_state_class( snapshot: SnapshotAssertion, mocker: MockerFixture, checkpointer_name: str From d88f59eea4de46f5f29c01c4b1f692645119dc0c Mon Sep 17 00:00:00 2001 From: William Fu-Hinthorn <13333726+hinthornw@users.noreply.github.com> Date: Wed, 12 Mar 2025 17:51:54 -0700 Subject: [PATCH 2/4] Cache --- libs/langgraph/langgraph/graph/state.py | 256 ++++++++++++------------ libs/langgraph/tests/test_pregel.py | 11 +- 2 files changed, 142 insertions(+), 125 deletions(-) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 2c6eca4f7..a6f997d25 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -945,139 +945,149 @@ def _pick_mapper( if isclass(schema): if issubclass(schema, dict): return None - if issubclass(schema, BaseModel): - return partial(_coerce_state_pydantic, schema) - if issubclass(schema, BaseModelV1): - return partial(_coerce_state_pydantic_v1, schema) + if issubclass(schema, (BaseModel, BaseModelV1)): + return _SchemaCoercionMapper(schema) return partial(_coerce_state, schema) -def _coerce_state_pydantic( - schema: Type[Any], input_data: dict[str, Any], *, __depth__: int = 5 -) -> Any: - if not isinstance(input_data, dict) or __depth__ <= 0: - return input_data +class _SchemaCoercionMapper: + _cache: dict[tuple[Type[Any], int], "_SchemaCoercionMapper"] = {} - processed_input = {} - for field_name, field_value in input_data.items(): - if field_name not in schema.model_fields: - processed_input[field_name] = field_value - continue + def __new__(cls, schema: Type[Any], max_depth: int = 5) -> "_SchemaCoercionMapper": + key = (schema, max_depth) + if key in cls._cache: + return cls._cache[key] + inst = super().__new__(cls) + cls._cache[key] = inst + return inst - field_info = schema.model_fields[field_name] - field_type = field_info.annotation - processed_input[field_name] = _process_field_value( - field_type, field_value, __depth__ - 1 - ) + def __init__(self, schema: Type[Any], max_depth: int = 5): + if hasattr(self, "_inited"): + return + self._inited = True + self.schema = schema + self.max_depth = max_depth + if hasattr(schema, "model_fields") and hasattr(schema, "model_construct"): + self._fields = {n: f.annotation for n, f in schema.model_fields.items()} + self._construct = schema.model_construct + elif hasattr(schema, "__fields__") and callable( + getattr(schema, "construct", None) + ): + self._fields = {n: f.annotation for n, f in schema.__fields__.items()} + self._construct = schema.construct + else: + raise TypeError("Schema is neither valid Pydantic v1 nor v2 model.") + self._field_coercers: Optional[dict[str, Callable[[Any, Any], Any]]] = None - return schema.model_construct(**processed_input) + def __call__(self, input_data: Any, depth: Optional[int] = None) -> Any: + return self.coerce(input_data, depth) + def coerce(self, input_data: Any, depth: Optional[int] = None) -> Any: + if depth is None: + depth = self.max_depth + if not isinstance(input_data, dict) or depth <= 0: + return input_data + processed = {} + if self._field_coercers is None: + self._field_coercers = { + n: self._build_coercer(t) for n, t in self._fields.items() + } + for k, v in input_data.items(): + fn = self._field_coercers.get(k) + processed[k] = fn(v, depth - 1) if fn else v + return self._construct(**processed) -def _coerce_state_pydantic_v1( - schema: Type[Any], input_data: dict[str, Any], *, __depth__: int = 5 -) -> Any: - if not isinstance(input_data, dict) or __depth__ <= 0: - return input_data - - processed_input = {} - for field_name, field_value in input_data.items(): - if field_name not in schema.__fields__: - processed_input[field_name] = field_value - continue - - field_info = schema.__fields__[field_name] - field_type = field_info.annotation - processed_input[field_name] = _process_field_value( - field_type, field_value, __depth__ - 1 - ) - - return schema.construct(**processed_input) - - -def _process_field_value( - field_type: Type[Any], field_value: Any, __depth__: int -) -> Any: - if __depth__ <= 0 or field_value is None: - return field_value - origin = get_origin(field_type) - - if origin is Annotated: - real_type, *_ = get_args(field_type) - res = _process_field_value(real_type, field_value, __depth__) - return res - - if isclass(field_type): - is_class_ = True - try: - is_model = issubclass(field_type, BaseModel) - except TypeError: - is_class_ = False - is_model = False - if is_model: - if isinstance(field_value, dict): - return _coerce_state_pydantic( - field_type, field_value, __depth__=__depth__ - ) - return field_value - if is_class_ and issubclass(field_type, BaseModelV1): - if isinstance(field_value, dict): - return _coerce_state_pydantic_v1( - field_type, field_value, __depth__=__depth__ - ) - return field_value - - if origin is list or field_type is list: - if not isinstance(field_value, (list, tuple)): - raise TypeError( - f"Expected a list/tuple for {field_type}, got {type(field_value)}." - ) - (item_type,) = get_args(field_type) - return [ - _process_field_value(item_type, item, __depth__ - 1) for item in field_value - ] - - if origin is dict or field_type is dict: - if not isinstance(field_value, dict): - raise TypeError( - f"Expected a dict for {field_type}, got {type(field_value)}." - ) - key_type, val_type = get_args(field_type) - return { - _process_field_value(key_type, k, __depth__ - 1): _process_field_value( - val_type, v, __depth__ - 1 - ) - for k, v in field_value.items() - } - - if origin is tuple: - if not isinstance(field_value, (list, tuple)): - raise TypeError( - f"Expected a tuple/list for {field_type}, got {type(field_value)}." - ) - args = get_args(field_type) - # Handle tuple[type1, type2, ...] with fixed length and different types - result = [] - for i, arg in enumerate(args): - if i < len(field_value): - result.append(_process_field_value(arg, field_value[i], __depth__ - 1)) - else: - # If field_value is shorter than expected, use None for remaining positions - result.append(None) - # If field_value is longer than expected, truncate it - return tuple(result) - - if origin is Union: - for arg in get_args(field_type): - if arg is type(None): - # e.g. Optional - continue + def _build_coercer(self, field_type: Any) -> Callable[[Any, Any], Any]: + origin = get_origin(field_type) + if origin is Annotated: + real_type, *_ = get_args(field_type) + sub = self._build_coercer(real_type) + return lambda v, d: sub(v, d) + if isclass(field_type): + is_class_ = True try: - result = _process_field_value(arg, field_value, __depth__ - 1) - return result - except Exception: - pass # Fall back to the next union argument + is_base_model = issubclass(field_type, BaseModel) + except TypeError: + is_class_ = False + is_base_model = False - return field_value + if is_base_model: + mapper = _SchemaCoercionMapper(field_type, self.max_depth) + return lambda v, d: mapper.coerce(v, d) if isinstance(v, dict) else v + if is_class_ and issubclass(field_type, BaseModelV1): + mapper = _SchemaCoercionMapper(field_type, self.max_depth) + return lambda v, d: mapper.coerce(v, d) if isinstance(v, dict) else v + if origin is list or field_type is list: + args = get_args(field_type) + if len(args) != 1: + return lambda v, d: v + sub = self._build_coercer(args[0]) + + def list_coercer(v: Any, d: Any) -> Any: + if not isinstance(v, (list, tuple)): + raise TypeError(f"Expected list, got {type(v).__name__}") + return [sub(x, d - 1) for x in v] + + return list_coercer + if origin is dict or field_type is dict: + args = get_args(field_type) + if len(args) != 2: + + def plain_dict_coercer(v: Any, d: Any) -> Any: + if not isinstance(v, dict): + raise TypeError(f"Expected dict, got {type(v).__name__}") + return v + + return plain_dict_coercer + k_sub = self._build_coercer(args[0]) + v_sub = self._build_coercer(args[1]) + + def dict_coercer(v: Any, d: Any) -> Any: + if not isinstance(v, dict): + raise TypeError(f"Expected dict, got {type(v).__name__}") + return {k_sub(k, d - 1): v_sub(val, d - 1) for k, val in v.items()} + + return dict_coercer + + if origin is tuple: + targs = get_args(field_type) + if not targs: + return lambda v, d: v + subs = [self._build_coercer(a) for a in targs] + + def tuple_coercer(v: Any, d: Any) -> Any: + if not isinstance(v, (list, tuple)): + raise TypeError(f"Expected tuple-like, got {type(v).__name__}") + out = [] + for i, sp in enumerate(subs): + out.append(sp(v[i] if i < len(v) else None, d - 1)) + return tuple(out) + + return tuple_coercer + if origin is Union: + uargs = get_args(field_type) + subs, none_in_union = [], False + for arg in uargs: + if arg is type(None): + none_in_union = True + else: + subs.append(self._build_coercer(arg)) + + def union_coercer(v: Any, d: Any) -> Any: + if v is None and none_in_union: + return None + err = None + for sp in subs: + try: + return sp(v, d - 1) + except Exception as e: + err = e + if err: + raise err + return v + + return union_coercer + return lambda v, d: v def _coerce_state(schema: Type[Any], input: dict[str, Any]) -> dict[str, Any]: diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 279d2b892..5a4c9bfa8 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -3069,7 +3069,7 @@ def test_nested_pydantic_models(version: str) -> None: # Basic nested model tests top_level: str nested: NestedModel - optional_nested: Optional[NestedModel] = None + optional_nested: Annotated[Optional[NestedModel], lambda x, y: y, "Foo"] dict_nested: dict[str, NestedModel] list_nested: Annotated[ Union[dict, list[dict[str, NestedModel]]], lambda x, y: (x or []) + [y] @@ -3126,8 +3126,10 @@ def test_nested_pydantic_models(version: str) -> None: update = {"top_level": "updated", "nested": {"value": 100, "name": "updated"}} + expected = State(**inputs) + def node_fn(state: State) -> dict: - assert state == State(**inputs) + assert state == expected return update builder = StateGraph(State) @@ -3140,6 +3142,11 @@ def test_nested_pydantic_models(version: str) -> None: assert result == {**inputs, **update} + new_inputs = inputs.copy() + new_inputs["list_nested"] = {"foo": "bar"} + expected = State(**new_inputs) + assert {**new_inputs, **update} == graph.invoke(new_inputs.copy()) + @pytest.mark.parametrize("checkpointer_name", ALL_CHECKPOINTERS_SYNC) def test_in_one_fan_out_state_graph_waiting_edge_plus_regular( From ffbcdd1eccdf6dd48e14122fcebe76eda65a3da2 Mon Sep 17 00:00:00 2001 From: William Fu-Hinthorn <13333726+hinthornw@users.noreply.github.com> Date: Wed, 12 Mar 2025 18:29:55 -0700 Subject: [PATCH 3/4] weakref --- libs/langgraph/langgraph/graph/state.py | 15 ++++++++++----- 1 file changed, 10 insertions(+), 5 deletions(-) diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index a6f997d25..228aa87a3 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -2,6 +2,7 @@ import inspect import logging import typing import warnings +import weakref from functools import partial from inspect import isclass, isfunction, ismethod, signature from types import FunctionType @@ -951,14 +952,18 @@ def _pick_mapper( class _SchemaCoercionMapper: - _cache: dict[tuple[Type[Any], int], "_SchemaCoercionMapper"] = {} + _cache: weakref.WeakKeyDictionary[Type[Any], dict[int, "_SchemaCoercionMapper"]] = ( + weakref.WeakKeyDictionary() + ) def __new__(cls, schema: Type[Any], max_depth: int = 5) -> "_SchemaCoercionMapper": - key = (schema, max_depth) - if key in cls._cache: - return cls._cache[key] + if schema not in cls._cache: + cls._cache[schema] = {} + if max_depth in cls._cache[schema]: + return cls._cache[schema][max_depth] + inst = super().__new__(cls) - cls._cache[key] = inst + cls._cache[schema][max_depth] = inst return inst def __init__(self, schema: Type[Any], max_depth: int = 5): From 919282fead9f3908b31ca3f6a497ed20eb2af67c Mon Sep 17 00:00:00 2001 From: William Fu-Hinthorn <13333726+hinthornw@users.noreply.github.com> Date: Thu, 13 Mar 2025 09:35:26 -0700 Subject: [PATCH 4/4] Move files --- .../langgraph/langgraph/graph/schema_utils.py | 162 ++++++++++++++++++ libs/langgraph/langgraph/graph/state.py | 150 +--------------- 2 files changed, 165 insertions(+), 147 deletions(-) create mode 100644 libs/langgraph/langgraph/graph/schema_utils.py diff --git a/libs/langgraph/langgraph/graph/schema_utils.py b/libs/langgraph/langgraph/graph/schema_utils.py new file mode 100644 index 000000000..bf8d5c4a0 --- /dev/null +++ b/libs/langgraph/langgraph/graph/schema_utils.py @@ -0,0 +1,162 @@ +import logging +import weakref +from inspect import isclass +from typing import ( + Any, + Callable, + Optional, + Type, + Union, + get_args, + get_origin, +) + +from pydantic import BaseModel +from pydantic.v1 import BaseModel as BaseModelV1 +from typing_extensions import Annotated + +logger = logging.getLogger(__name__) + + +class SchemaCoercionMapper: + _cache: weakref.WeakKeyDictionary[Type[Any], dict[int, "SchemaCoercionMapper"]] = ( + weakref.WeakKeyDictionary() + ) + + def __new__(cls, schema: Type[Any], max_depth: int = 5) -> "SchemaCoercionMapper": + if schema not in cls._cache: + cls._cache[schema] = {} + if max_depth in cls._cache[schema]: + return cls._cache[schema][max_depth] + + inst = super().__new__(cls) + cls._cache[schema][max_depth] = inst + return inst + + def __init__(self, schema: Type[Any], max_depth: int = 5): + if hasattr(self, "_inited"): + return + self._inited = True + self.schema = schema + self.max_depth = max_depth + if hasattr(schema, "model_fields") and hasattr(schema, "model_construct"): + self._fields = {n: f.annotation for n, f in schema.model_fields.items()} + self._construct = schema.model_construct + elif hasattr(schema, "__fields__") and callable( + getattr(schema, "construct", None) + ): + self._fields = {n: f.annotation for n, f in schema.__fields__.items()} + self._construct = schema.construct + else: + raise TypeError("Schema is neither valid Pydantic v1 nor v2 model.") + self._field_coercers: Optional[dict[str, Callable[[Any, Any], Any]]] = None + + def __call__(self, input_data: Any, depth: Optional[int] = None) -> Any: + return self.coerce(input_data, depth) + + def coerce(self, input_data: Any, depth: Optional[int] = None) -> Any: + if depth is None: + depth = self.max_depth + if not isinstance(input_data, dict) or depth <= 0: + return input_data + processed = {} + if self._field_coercers is None: + self._field_coercers = { + n: self._build_coercer(t) for n, t in self._fields.items() + } + for k, v in input_data.items(): + fn = self._field_coercers.get(k) + processed[k] = fn(v, depth - 1) if fn else v + return self._construct(**processed) + + def _build_coercer(self, field_type: Any) -> Callable[[Any, Any], Any]: + origin = get_origin(field_type) + if origin is Annotated: + real_type, *_ = get_args(field_type) + sub = self._build_coercer(real_type) + return lambda v, d: sub(v, d) + if isclass(field_type): + is_class_ = True + try: + is_base_model = issubclass(field_type, BaseModel) + except TypeError: + is_class_ = False + is_base_model = False + + if is_base_model: + mapper = SchemaCoercionMapper(field_type, self.max_depth) + return lambda v, d: mapper.coerce(v, d) if isinstance(v, dict) else v + if is_class_ and issubclass(field_type, BaseModelV1): + mapper = SchemaCoercionMapper(field_type, self.max_depth) + return lambda v, d: mapper.coerce(v, d) if isinstance(v, dict) else v + if origin is list or field_type is list: + args = get_args(field_type) + if len(args) != 1: + return lambda v, d: v + sub = self._build_coercer(args[0]) + + def list_coercer(v: Any, d: Any) -> Any: + if not isinstance(v, (list, tuple)): + raise TypeError(f"Expected list, got {type(v).__name__}") + return [sub(x, d - 1) for x in v] + + return list_coercer + if origin is dict or field_type is dict: + args = get_args(field_type) + if len(args) != 2: + + def plain_dict_coercer(v: Any, d: Any) -> Any: + if not isinstance(v, dict): + raise TypeError(f"Expected dict, got {type(v).__name__}") + return v + + return plain_dict_coercer + k_sub = self._build_coercer(args[0]) + v_sub = self._build_coercer(args[1]) + + def dict_coercer(v: Any, d: Any) -> Any: + if not isinstance(v, dict): + raise TypeError(f"Expected dict, got {type(v).__name__}") + return {k_sub(k, d - 1): v_sub(val, d - 1) for k, val in v.items()} + + return dict_coercer + + if origin is tuple: + targs = get_args(field_type) + if not targs: + return lambda v, d: v + subs = [self._build_coercer(a) for a in targs] + + def tuple_coercer(v: Any, d: Any) -> Any: + if not isinstance(v, (list, tuple)): + raise TypeError(f"Expected tuple-like, got {type(v).__name__}") + out = [] + for i, sp in enumerate(subs): + out.append(sp(v[i] if i < len(v) else None, d - 1)) + return tuple(out) + + return tuple_coercer + if origin is Union: + uargs = get_args(field_type) + subs, none_in_union = [], False + for arg in uargs: + if arg is type(None): + none_in_union = True + else: + subs.append(self._build_coercer(arg)) + + def union_coercer(v: Any, d: Any) -> Any: + if v is None and none_in_union: + return None + err = None + for sp in subs: + try: + return sp(v, d - 1) + except Exception as e: + err = e + if err: + raise err + return v + + return union_coercer + return lambda v, d: v diff --git a/libs/langgraph/langgraph/graph/state.py b/libs/langgraph/langgraph/graph/state.py index 228aa87a3..f4143ca1f 100644 --- a/libs/langgraph/langgraph/graph/state.py +++ b/libs/langgraph/langgraph/graph/state.py @@ -2,7 +2,6 @@ import inspect import logging import typing import warnings -import weakref from functools import partial from inspect import isclass, isfunction, ismethod, signature from types import FunctionType @@ -27,7 +26,7 @@ from typing import ( from langchain_core.runnables import Runnable, RunnableConfig from pydantic import BaseModel from pydantic.v1 import BaseModel as BaseModelV1 -from typing_extensions import Annotated, Self +from typing_extensions import Self from langgraph._api.deprecation import LangGraphDeprecationWarning from langgraph.channels.base import BaseChannel @@ -51,6 +50,7 @@ from langgraph.graph.graph import ( Graph, Send, ) +from langgraph.graph.schema_utils import SchemaCoercionMapper from langgraph.managed.base import ( ChannelKeyPlaceholder, ChannelTypePlaceholder, @@ -947,154 +947,10 @@ def _pick_mapper( if issubclass(schema, dict): return None if issubclass(schema, (BaseModel, BaseModelV1)): - return _SchemaCoercionMapper(schema) + return SchemaCoercionMapper(schema) return partial(_coerce_state, schema) -class _SchemaCoercionMapper: - _cache: weakref.WeakKeyDictionary[Type[Any], dict[int, "_SchemaCoercionMapper"]] = ( - weakref.WeakKeyDictionary() - ) - - def __new__(cls, schema: Type[Any], max_depth: int = 5) -> "_SchemaCoercionMapper": - if schema not in cls._cache: - cls._cache[schema] = {} - if max_depth in cls._cache[schema]: - return cls._cache[schema][max_depth] - - inst = super().__new__(cls) - cls._cache[schema][max_depth] = inst - return inst - - def __init__(self, schema: Type[Any], max_depth: int = 5): - if hasattr(self, "_inited"): - return - self._inited = True - self.schema = schema - self.max_depth = max_depth - if hasattr(schema, "model_fields") and hasattr(schema, "model_construct"): - self._fields = {n: f.annotation for n, f in schema.model_fields.items()} - self._construct = schema.model_construct - elif hasattr(schema, "__fields__") and callable( - getattr(schema, "construct", None) - ): - self._fields = {n: f.annotation for n, f in schema.__fields__.items()} - self._construct = schema.construct - else: - raise TypeError("Schema is neither valid Pydantic v1 nor v2 model.") - self._field_coercers: Optional[dict[str, Callable[[Any, Any], Any]]] = None - - def __call__(self, input_data: Any, depth: Optional[int] = None) -> Any: - return self.coerce(input_data, depth) - - def coerce(self, input_data: Any, depth: Optional[int] = None) -> Any: - if depth is None: - depth = self.max_depth - if not isinstance(input_data, dict) or depth <= 0: - return input_data - processed = {} - if self._field_coercers is None: - self._field_coercers = { - n: self._build_coercer(t) for n, t in self._fields.items() - } - for k, v in input_data.items(): - fn = self._field_coercers.get(k) - processed[k] = fn(v, depth - 1) if fn else v - return self._construct(**processed) - - def _build_coercer(self, field_type: Any) -> Callable[[Any, Any], Any]: - origin = get_origin(field_type) - if origin is Annotated: - real_type, *_ = get_args(field_type) - sub = self._build_coercer(real_type) - return lambda v, d: sub(v, d) - if isclass(field_type): - is_class_ = True - try: - is_base_model = issubclass(field_type, BaseModel) - except TypeError: - is_class_ = False - is_base_model = False - - if is_base_model: - mapper = _SchemaCoercionMapper(field_type, self.max_depth) - return lambda v, d: mapper.coerce(v, d) if isinstance(v, dict) else v - if is_class_ and issubclass(field_type, BaseModelV1): - mapper = _SchemaCoercionMapper(field_type, self.max_depth) - return lambda v, d: mapper.coerce(v, d) if isinstance(v, dict) else v - if origin is list or field_type is list: - args = get_args(field_type) - if len(args) != 1: - return lambda v, d: v - sub = self._build_coercer(args[0]) - - def list_coercer(v: Any, d: Any) -> Any: - if not isinstance(v, (list, tuple)): - raise TypeError(f"Expected list, got {type(v).__name__}") - return [sub(x, d - 1) for x in v] - - return list_coercer - if origin is dict or field_type is dict: - args = get_args(field_type) - if len(args) != 2: - - def plain_dict_coercer(v: Any, d: Any) -> Any: - if not isinstance(v, dict): - raise TypeError(f"Expected dict, got {type(v).__name__}") - return v - - return plain_dict_coercer - k_sub = self._build_coercer(args[0]) - v_sub = self._build_coercer(args[1]) - - def dict_coercer(v: Any, d: Any) -> Any: - if not isinstance(v, dict): - raise TypeError(f"Expected dict, got {type(v).__name__}") - return {k_sub(k, d - 1): v_sub(val, d - 1) for k, val in v.items()} - - return dict_coercer - - if origin is tuple: - targs = get_args(field_type) - if not targs: - return lambda v, d: v - subs = [self._build_coercer(a) for a in targs] - - def tuple_coercer(v: Any, d: Any) -> Any: - if not isinstance(v, (list, tuple)): - raise TypeError(f"Expected tuple-like, got {type(v).__name__}") - out = [] - for i, sp in enumerate(subs): - out.append(sp(v[i] if i < len(v) else None, d - 1)) - return tuple(out) - - return tuple_coercer - if origin is Union: - uargs = get_args(field_type) - subs, none_in_union = [], False - for arg in uargs: - if arg is type(None): - none_in_union = True - else: - subs.append(self._build_coercer(arg)) - - def union_coercer(v: Any, d: Any) -> Any: - if v is None and none_in_union: - return None - err = None - for sp in subs: - try: - return sp(v, d - 1) - except Exception as e: - err = e - if err: - raise err - return v - - return union_coercer - return lambda v, d: v - - def _coerce_state(schema: Type[Any], input: dict[str, Any]) -> dict[str, Any]: return schema(**input)