mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-08-26 01:22:24 +02:00
Merge branch 'main' into wfh/output_schema
This commit is contained in:
@@ -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[<type>] 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.
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user