mirror of
https://github.com/langchain-ai/langgraph.git
synced 2026-09-30 05:25:05 +02:00
chore: respect default values for schemas
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user