From ec4e92792dbbc78fc6b74d1b5927f87e3f737303 Mon Sep 17 00:00:00 2001 From: William FH <13333726+hinthornw@users.noreply.github.com> Date: Fri, 6 Sep 2024 11:50:45 -0700 Subject: [PATCH] Optional types for dataclass state defs (#1557) --- libs/langgraph/langgraph/utils/fields.py | 21 ++++++---- libs/langgraph/tests/test_state.py | 50 ++++++++++++++++++++++++ 2 files changed, 64 insertions(+), 7 deletions(-) diff --git a/libs/langgraph/langgraph/utils/fields.py b/libs/langgraph/langgraph/utils/fields.py index 55a8e81be..a4c29a9ec 100644 --- a/libs/langgraph/langgraph/utils/fields.py +++ b/libs/langgraph/langgraph/utils/fields.py @@ -1,12 +1,7 @@ +import dataclasses from typing import Any, Optional, Type, Union -from typing_extensions import ( - Annotated, - NotRequired, - ReadOnly, - Required, - get_origin, -) +from typing_extensions import Annotated, NotRequired, ReadOnly, Required, get_origin def _is_optional_type(type_: Any) -> bool: @@ -92,6 +87,18 @@ def get_field_default(name: str, type_: Any, schema: Type[Any]) -> Any: return ... # Handle NotRequired[] for earlier versions of python return None + if dataclasses.is_dataclass(schema): + field_info = next( + (f for f in dataclasses.fields(schema) if f.name == name), None + ) + if field_info: + if ( + field_info.default is not dataclasses.MISSING + and field_info.default is not ... + ): + return field_info.default + elif field_info.default_factory is not dataclasses.MISSING: + return field_info.default_factory() # Note, we ignore ReadOnly attributes, # as they don't make much sense. (we don't care if you mutate the state in your node) # and mutating state in your node has no effect on our graph state. diff --git a/libs/langgraph/tests/test_state.py b/libs/langgraph/tests/test_state.py index 12cd7f653..0234c5e2f 100644 --- a/libs/langgraph/tests/test_state.py +++ b/libs/langgraph/tests/test_state.py @@ -1,4 +1,6 @@ +import inspect import warnings +from dataclasses import dataclass, field from typing import Annotated as Annotated2 from typing import Any, Optional @@ -130,3 +132,51 @@ def test_state_schema_optional_values(total_: bool): assert ( set(json_schema["properties"].keys()) == expected_required | expected_optional ) + + +@pytest.mark.parametrize("kw_only_", [False, True]) +def test_state_schema_default_values(kw_only_: bool): + kwargs = {} + if "kw_only" in inspect.signature(dataclass).parameters: + kwargs = {"kw_only": kw_only_} + + @dataclass(**kwargs) + class InputState: + val1: str + val2: Optional[int] + val3: Annotated[Optional[float], "optional annotated"] + val4: Optional[str] = None + val5: list[int] = field(default_factory=lambda: [1, 2, 3]) + val6: dict[str, int] = field(default_factory=lambda: {"a": 1}) + val7: str = field(default=...) + val8: Annotated[int, "some metadata"] = 42 + val9: Annotated[str, "more metadata"] = field(default="some foo") + val10: str = "default" + val11: Annotated[list[str], "annotated list"] = field( + default_factory=lambda: ["a", "b"] + ) + + builder = StateGraph(InputState) + builder.add_node("n", lambda x: x) + builder.add_edge("__start__", "n") + graph = builder.compile() + model = graph.input_schema + json_schema = model.schema() + + expected_required = {"val1", "val7"} + expected_optional = { + "val2", + "val3", + "val4", + "val5", + "val6", + "val8", + "val9", + "val10", + "val11", + } + + assert set(json_schema.get("required", set())) == expected_required + assert ( + set(json_schema["properties"].keys()) == expected_required | expected_optional + )