chore: respect default values for schemas

This commit is contained in:
William Fu-Hinthorn
2026-02-19 07:17:02 -08:00
parent a2912cef67
commit 36d1e96398
3 changed files with 54 additions and 17 deletions
+6 -12
View File
@@ -3,10 +3,11 @@ from __future__ import annotations
import dataclasses
import types
import weakref
from collections.abc import Generator, Sequence
from typing import Annotated, Any, Optional, Union, get_origin, get_type_hints
from collections.abc import Callable, Generator, Sequence
from typing import Annotated, Any, Optional, Union, cast, get_origin, get_type_hints
from pydantic import BaseModel
from pydantic_core import PydanticUndefined
from typing_extensions import NotRequired, ReadOnly, Required
from langgraph._internal._typing import MISSING
@@ -104,17 +105,10 @@ def get_field_default(name: str, type_: Any, schema: type[Any]) -> Any:
if isinstance(schema, type) and issubclass(schema, BaseModel):
if name in schema.model_fields:
field = schema.model_fields[name]
# Check default_factory first (it takes precedence in Pydantic)
if field.default_factory is not None:
return field.default_factory() # type: ignore[call-arg]
# Check if default is set (not PydanticUndefined)
if (
hasattr(field.default, "__class__")
and getattr(field.default.__class__, "__name__", "")
== "PydanticUndefinedType"
):
pass # No default, fall through
else:
factory = cast(Callable[[], Any], field.default_factory)
return factory()
if field.default is not PydanticUndefined:
return field.default
if dataclasses.is_dataclass(schema):
field_info = next(
+9 -5
View File
@@ -1,4 +1,5 @@
import collections.abc
import copy
from collections.abc import Callable, Sequence
from typing import Any, Generic
@@ -69,7 +70,7 @@ class BinaryOperatorAggregate(Generic[Value], BaseChannel[Value, Value, Value]):
if typ in (collections.abc.Mapping, collections.abc.MutableMapping):
typ = dict
if default is not MISSING:
self.value = default
self.value = copy.deepcopy(default)
else:
try:
self.value = typ()
@@ -78,10 +79,13 @@ class BinaryOperatorAggregate(Generic[Value], BaseChannel[Value, Value, Value]):
def __eq__(self, value: object) -> bool:
return isinstance(value, BinaryOperatorAggregate) and (
value.operator is self.operator
if value.operator.__name__ != "<lambda>"
and self.operator.__name__ != "<lambda>"
else True
(
value.operator is self.operator
if value.operator.__name__ != "<lambda>"
and self.operator.__name__ != "<lambda>"
else True
)
and value.default == self.default
)
@property
+39
View File
@@ -122,6 +122,45 @@ def test_binop_with_default() -> None:
assert channel.get() == {"a": 1, "b": 2}
def test_binop_with_default_mutable_safety() -> None:
"""Mutable defaults should not be shared across channel instances."""
default = {"a": 1}
ch1 = BinaryOperatorAggregate(dict, operator.or_, default=default).from_checkpoint(
MISSING
)
ch2 = BinaryOperatorAggregate(dict, operator.or_, default=default).from_checkpoint(
MISSING
)
# Mutate ch1's value via a reducer that mutates in-place
def mutating_reducer(a: dict, b: dict) -> dict:
a.update(b)
return a
ch1.operator = mutating_reducer
ch1.update([{"b": 2}])
assert ch1.get() == {"a": 1, "b": 2}
# ch2 should be unaffected
assert ch2.get() == {"a": 1}
# Original default should be unaffected
assert default == {"a": 1}
def test_binop_with_default_multi_invoke() -> None:
"""Defaults should be fresh across multiple from_checkpoint calls."""
template = BinaryOperatorAggregate(dict, operator.or_, default={"a": 1})
# Simulate two separate runs
run1 = template.from_checkpoint(MISSING)
run1.update([{"b": 2}])
assert run1.get() == {"a": 1, "b": 2}
run2 = template.from_checkpoint(MISSING)
assert run2.get() == {"a": 1} # Should NOT see {"b": 2}
def test_untracked_value() -> None:
channel = UntrackedValue(dict).from_checkpoint(MISSING)
assert channel.ValueType is dict