diff --git a/libs/langgraph/langgraph/graph/schema_utils.py b/libs/langgraph/langgraph/graph/schema_utils.py index de769a660..4ce4ce83a 100644 --- a/libs/langgraph/langgraph/graph/schema_utils.py +++ b/libs/langgraph/langgraph/graph/schema_utils.py @@ -13,7 +13,7 @@ from typing import ( get_type_hints, ) -from pydantic import BaseModel +from pydantic import BaseModel, TypeAdapter __all__ = ["SchemaCoercionMapper"] @@ -243,67 +243,10 @@ _IDENTITY_TYPES: tuple[type[Any], ...] = ( type(None), ) -try: - # Pydantic v2. - from pydantic import TypeAdapter - try: - import pydantic.v1.types as v1_types_ - from pydantic.v1 import parse_obj_as - - v1_types = tuple( - v for k, v in vars(v1_types_).items() if k in v1_types_.__all__ - ) - except ImportError: - v1_types = () - - def parse_obj_as(tp: Any, v: Any) -> Any: # type: ignore - return v - - try: - from pydantic.v1 import parse_obj_as - from pydantic.v1.main import create_model - except ImportError: - create_model = None # type: ignore - - def _get_v1_parser(tp: Any) -> Any: - if create_model is not None: - try: - parser = create_model( - f"ParsingModel[{tp}]", - __root__=(tp, ...), - ) - return lambda v: parser(__root__=v).__root__ # type: ignore - except RuntimeError: - return lambda v: v - return lambda v: parse_obj_as(tp, v) - - @functools.lru_cache(maxsize=2048) - def _adapter_for(tp: Any) -> Callable[[Any], Any]: # noqa: D401 - if tp in v1_types: - return _get_v1_parser(tp) - try: - return TypeAdapter( - tp, config={"arbitrary_types_allowed": True} - ).validate_python - except TypeError: - # Delayed classes like ConstrainedList - return _get_v1_parser(tp) - -except ImportError: - # Pydantic V1 - from pydantic.v1.main import create_model - - @functools.lru_cache(maxsize=2048) - def _adapter_for(tp: Any) -> Callable[[Any], Any]: # noqa: D401 - try: - parser = create_model( - f"ParsingModel[{tp}]", - __root__=(tp, ...), - ) - return lambda v: parser(__root__=v).__root__ # type: ignore - except RuntimeError: - return lambda v: v +@functools.lru_cache(maxsize=2048) +def _adapter_for(tp: Any) -> Callable[[Any], Any]: # noqa: D401 + return TypeAdapter(tp, config={"arbitrary_types_allowed": True}).validate_python def _get_adapter(tp: Any) -> Callable[[Any], Any]: diff --git a/libs/langgraph/tests/test_pydantic.py b/libs/langgraph/tests/test_pydantic.py index 61da471b6..625998d01 100644 --- a/libs/langgraph/tests/test_pydantic.py +++ b/libs/langgraph/tests/test_pydantic.py @@ -8,8 +8,18 @@ import uuid from enum import Enum from typing import Annotated, Literal, Optional, Union -import pytest -from pydantic import BaseModel, field_validator, model_validator +from pydantic import ( + BaseModel, + ByteSize, + Field, + SecretStr, + confloat, + conint, + conlist, + constr, + field_validator, + model_validator, +) from langgraph.constants import END, START from langgraph.graph.state import StateGraph @@ -47,36 +57,9 @@ def test_is_supported_by_pydantic() -> None: assert is_supported_by_pydantic(PydanticModel) is True -@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 - # Import necessary modules - - if version == "v1": - from pydantic.v1 import ( # type: ignore - BaseModel, - ByteSize, - Field, - SecretStr, - confloat, - conint, - conlist, - constr, - ) - else: - from pydantic import ( # type: ignore - BaseModel, - ByteSize, - Field, - SecretStr, - confloat, - conint, - conlist, - constr, - ) - class NestedModel(BaseModel): value: int name: str diff --git a/libs/langgraph/tests/test_state.py b/libs/langgraph/tests/test_state.py index 1e8c2c623..011dab41e 100644 --- a/libs/langgraph/tests/test_state.py +++ b/libs/langgraph/tests/test_state.py @@ -7,7 +7,7 @@ from typing import Annotated as Annotated2 import pytest from langchain_core.runnables import RunnableConfig, RunnableLambda -from pydantic.v1 import BaseModel +from pydantic import BaseModel from typing_extensions import NotRequired, Required, TypedDict from langgraph.graph.state import StateGraph, _get_node_name, _warn_invalid_state_schema