diff --git a/libs/langgraph/langgraph/_internal/_fields.py b/libs/langgraph/langgraph/_internal/_fields.py index 6ff3e04d8..4a81de142 100644 --- a/libs/langgraph/langgraph/_internal/_fields.py +++ b/libs/langgraph/langgraph/_internal/_fields.py @@ -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( diff --git a/libs/langgraph/langgraph/channels/binop.py b/libs/langgraph/langgraph/channels/binop.py index b1c1d1666..619ba906c 100644 --- a/libs/langgraph/langgraph/channels/binop.py +++ b/libs/langgraph/langgraph/channels/binop.py @@ -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__ != "" - and self.operator.__name__ != "" - else True + ( + value.operator is self.operator + if value.operator.__name__ != "" + and self.operator.__name__ != "" + else True + ) + and value.default == self.default ) @property diff --git a/libs/langgraph/tests/test_channels.py b/libs/langgraph/tests/test_channels.py index 4957c7cdb..060e5de77 100644 --- a/libs/langgraph/tests/test_channels.py +++ b/libs/langgraph/tests/test_channels.py @@ -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