From 256e92bfb339c3a5ae6c7b53805a83c3486879a7 Mon Sep 17 00:00:00 2001 From: Nuno Campos Date: Tue, 4 Mar 2025 13:14:54 -0800 Subject: [PATCH] Pregel.config_schema should use config_type directly when present (#3641) - the previous behavior of re-creating model through config_specs would lose custom annotations on config_type --------- Co-authored-by: Eugene Yurtsev --- libs/langgraph/langgraph/pregel/__init__.py | 37 ++++++++++- libs/langgraph/langgraph/utils/pydantic.py | 32 ++++++++++ .../tests/__snapshots__/test_pregel.ambr | 2 +- libs/langgraph/tests/test_pregel.py | 61 ++++++++++++++++++- libs/langgraph/tests/test_pydantic.py | 43 +++++++++++++ 5 files changed, 168 insertions(+), 7 deletions(-) create mode 100644 libs/langgraph/tests/test_pydantic.py diff --git a/libs/langgraph/langgraph/pregel/__init__.py b/libs/langgraph/langgraph/pregel/__init__.py index 68f4b8dc1..6948847d2 100644 --- a/libs/langgraph/langgraph/pregel/__init__.py +++ b/libs/langgraph/langgraph/pregel/__init__.py @@ -18,6 +18,7 @@ from typing import ( Type, Union, cast, + get_type_hints, overload, ) from uuid import UUID, uuid5 @@ -119,7 +120,7 @@ from langgraph.utils.config import ( recast_checkpoint_ns, ) from langgraph.utils.fields import get_enhanced_type_hints -from langgraph.utils.pydantic import create_model +from langgraph.utils.pydantic import create_model, is_supported_by_pydantic from langgraph.utils.queue import AsyncQueue, SyncQueue # type: ignore[attr-defined] WriteValue = Union[Callable[[Input], Output], Any] @@ -609,6 +610,36 @@ class Pregel(PregelProtocol): ] ] + def config_schema( + self, *, include: Optional[Sequence[str]] = None + ) -> Type[BaseModel]: + # If the config type is not set explicitly, we will try to infer it. + # If the config type is provided, but isn't directly supported by pydantic + # (e.g., vanilla python class), we will also delegate to the parent class, + # which handles cases where Pydantic doesn't support the type. + if self.config_type is None or not is_supported_by_pydantic(self.config_type): + return super().config_schema(include=include) + + include = include or [] + fields = { + "configurable": (self.config_type, None), + **{ + field_name: (field_type, None) + for field_name, field_type in get_type_hints(RunnableConfig).items() + if field_name in [i for i in include if i != "configurable"] + }, + } + return create_model(self.get_name("Config"), field_definitions=fields) + + def get_config_jsonschema( + self, *, include: Optional[Sequence[str]] = None + ) -> Dict[str, Any]: + schema = self.config_schema(include=include) + if hasattr(schema, "model_json_schema"): + return schema.model_json_schema() + else: + return schema.schema() + @property def InputType(self) -> Any: if isinstance(self.input_channels, str): @@ -634,7 +665,7 @@ class Pregel(PregelProtocol): def get_input_jsonschema( self, config: Optional[RunnableConfig] = None - ) -> Dict[All, Any]: + ) -> Dict[str, Any]: schema = self.get_input_schema(config) if hasattr(schema, "model_json_schema"): return schema.model_json_schema() @@ -666,7 +697,7 @@ class Pregel(PregelProtocol): def get_output_jsonschema( self, config: Optional[RunnableConfig] = None - ) -> Dict[All, Any]: + ) -> Dict[str, Any]: schema = self.get_output_schema(config) if hasattr(schema, "model_json_schema"): return schema.model_json_schema() diff --git a/libs/langgraph/langgraph/utils/pydantic.py b/libs/langgraph/langgraph/utils/pydantic.py index cd0984202..56cef30e6 100644 --- a/libs/langgraph/langgraph/utils/pydantic.py +++ b/libs/langgraph/langgraph/utils/pydantic.py @@ -1,5 +1,9 @@ +import sys +import typing +from dataclasses import is_dataclass from typing import Any, Dict, Optional, Union +import typing_extensions from pydantic import BaseModel from pydantic.v1 import BaseModel as BaseModelV1 @@ -35,3 +39,31 @@ def create_model( v1_kwargs["__root__"] = root return create_model(model_name, **v1_kwargs, **(field_definitions or {})) + + +def is_supported_by_pydantic(type_: Any) -> bool: + """Check if a given "complex" type is supported by pydantic. + + This will return False for primitive types like int, str, etc. + + The check is meant for container types like dataclasses, TypedDicts, etc. + """ + if is_dataclass(type_): + return True + + # Pydantic does not support mixing .v1 and root namespaces, so + # we only check for BaseModel (not pydantic.v1.BaseModel). + if isinstance(type_, type) and issubclass(type_, BaseModel): + return True + + if hasattr(type_, "__orig_bases__"): + for base in type_.__orig_bases__: + if base is typing_extensions.TypedDict: + return True + elif base is typing.TypedDict: # noqa: TID251 + # ignoring TID251 since it's OK to use typing.TypedDict in this case. + # Pydantic supports typing.TypedDict from Python 3.12 + # For older versions, only typing_extensions.TypedDict is supported. + if sys.version_info >= (3, 12): + return True + return False diff --git a/libs/langgraph/tests/__snapshots__/test_pregel.ambr b/libs/langgraph/tests/__snapshots__/test_pregel.ambr index 4d9622955..dc24e73fb 100644 --- a/libs/langgraph/tests/__snapshots__/test_pregel.ambr +++ b/libs/langgraph/tests/__snapshots__/test_pregel.ambr @@ -1528,7 +1528,7 @@ ''' # --- # name: test_state_graph_w_config_inherited_state_keys - '{"$defs": {"Configurable": {"properties": {"tools": {"default": null, "items": {"type": "string"}, "title": "Tools", "type": "array"}}, "title": "Configurable", "type": "object"}}, "properties": {"configurable": {"$ref": "#/$defs/Configurable", "default": null}}, "title": "LangGraphConfig", "type": "object"}' + '{"$defs": {"Config": {"properties": {"tools": {"items": {"type": "string"}, "title": "Tools", "type": "array"}}, "title": "Config", "type": "object"}}, "properties": {"configurable": {"$ref": "#/$defs/Config", "default": null}}, "title": "LangGraphConfig", "type": "object"}' # --- # name: test_state_graph_w_config_inherited_state_keys.1 '{"$defs": {"AgentAction": {"description": "Represents a request to execute an action by an agent.\\n\\nThe action consists of the name of the tool to execute and the input to pass\\nto the tool. The log is used to pass along extra information about the action.", "properties": {"tool": {"title": "Tool", "type": "string"}, "tool_input": {"anyOf": [{"type": "string"}, {"type": "object"}], "title": "Tool Input"}, "log": {"title": "Log", "type": "string"}, "type": {"const": "AgentAction", "default": "AgentAction", "enum": ["AgentAction"], "title": "Type", "type": "string"}}, "required": ["tool", "tool_input", "log"], "title": "AgentAction", "type": "object"}, "AgentFinish": {"description": "Final return value of an ActionAgent.\\n\\nAgents return an AgentFinish when they have reached a stopping condition.", "properties": {"return_values": {"title": "Return Values", "type": "object"}, "log": {"title": "Log", "type": "string"}, "type": {"const": "AgentFinish", "default": "AgentFinish", "enum": ["AgentFinish"], "title": "Type", "type": "string"}}, "required": ["return_values", "log"], "title": "AgentFinish", "type": "object"}}, "properties": {"input": {"title": "Input", "type": "string"}, "agent_outcome": {"anyOf": [{"$ref": "#/$defs/AgentAction"}, {"$ref": "#/$defs/AgentFinish"}, {"type": "null"}], "default": null, "title": "Agent Outcome"}, "intermediate_steps": {"default": null, "items": {"maxItems": 2, "minItems": 2, "prefixItems": [{"$ref": "#/$defs/AgentAction"}, {"type": "string"}], "type": "array"}, "title": "Intermediate Steps", "type": "array"}}, "required": ["input"], "title": "LangGraphInput", "type": "object"}' diff --git a/libs/langgraph/tests/test_pregel.py b/libs/langgraph/tests/test_pregel.py index 862c859b5..229f8e051 100644 --- a/libs/langgraph/tests/test_pregel.py +++ b/libs/langgraph/tests/test_pregel.py @@ -10,7 +10,7 @@ import warnings from collections import Counter, deque from concurrent.futures import ThreadPoolExecutor from contextlib import contextmanager -from dataclasses import dataclass +from dataclasses import dataclass, field from random import randrange from typing import ( Annotated, @@ -275,6 +275,61 @@ def test_checkpoint_errors() -> None: graph.invoke("", {"configurable": {"thread_id": "thread-1"}}) +def test_config_json_schema() -> None: + """Test that config json schema is generated properly.""" + chain = Channel.subscribe_to("input") | Channel.write_to("output") + + @dataclass + class Foo: + x: int + y: str = field(default="foo") + + app = Pregel( + nodes={ + "one": chain, + }, + channels={ + "ephemeral": EphemeralValue(Any), + "input": LastValue(int), + "output": LastValue(int), + }, + input_channels=["input", "ephemeral"], + output_channels="output", + config_type=Foo, + ) + + assert app.get_config_jsonschema() == { + "$defs": { + "Foo": { + "properties": { + "x": { + "title": "X", + "type": "integer", + }, + "y": { + "default": "foo", + "title": "Y", + "type": "string", + }, + }, + "required": [ + "x", + ], + "title": "Foo", + "type": "object", + }, + }, + "properties": { + "configurable": { + "$ref": "#/$defs/Foo", + "default": None, + }, + }, + "title": "LangGraphConfig", + "type": "object", + } + + def test_node_schemas_custom_output() -> None: class State(TypedDict): hello: str @@ -1444,7 +1499,7 @@ def test_imp_task(request: pytest.FixtureRequest, checkpointer_name: str) -> Non checkpointer = request.getfixturevalue(f"checkpointer_{checkpointer_name}") mapper_calls = 0 - class Config: + class Configurable: model: str @task() @@ -1454,7 +1509,7 @@ def test_imp_task(request: pytest.FixtureRequest, checkpointer_name: str) -> Non time.sleep(input / 100) return str(input) * 2 - @entrypoint(checkpointer=checkpointer, config_schema=Config) + @entrypoint(checkpointer=checkpointer, config_schema=Configurable) def graph(input: list[int]) -> list[str]: futures = [mapper(i) for i in input] mapped = [f.result() for f in futures] diff --git a/libs/langgraph/tests/test_pydantic.py b/libs/langgraph/tests/test_pydantic.py new file mode 100644 index 000000000..f1a350033 --- /dev/null +++ b/libs/langgraph/tests/test_pydantic.py @@ -0,0 +1,43 @@ +import sys +import typing + +import pydantic +import typing_extensions + +from langgraph.utils.pydantic import is_supported_by_pydantic + + +def test_is_supported_by_pydantic() -> None: + """Test if types are supported by pydantic.""" + + class TypedDictExtensions(typing_extensions.TypedDict): + x: int + + assert is_supported_by_pydantic(TypedDictExtensions) is True + + class VanillaClass: + x: int + + assert is_supported_by_pydantic(VanillaClass) is False + + class BuiltinTypedDict(typing.TypedDict): # noqa: TID251 + x: int + + if sys.version_info >= (3, 12): + assert is_supported_by_pydantic(BuiltinTypedDict) is True + else: + assert is_supported_by_pydantic(BuiltinTypedDict) is False + + class PydanticModel(pydantic.BaseModel): + x: int + + assert is_supported_by_pydantic(PydanticModel) is True + + if hasattr(pydantic, "v1"): + + class PydanticModelV1(pydantic.v1.BaseModel): + x: int + + assert is_supported_by_pydantic(PydanticModelV1) is False + + assert is_supported_by_pydantic(int) is False