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 <eyurtsev@gmail.com>
This commit is contained in:
Nuno Campos
2025-03-04 16:14:54 -05:00
committed by GitHub
co-authored by Eugene Yurtsev
parent 3d4e5c0471
commit 256e92bfb3
5 changed files with 168 additions and 7 deletions
+34 -3
View File
@@ -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()
@@ -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
@@ -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"}'
+58 -3
View File
@@ -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]
+43
View File
@@ -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