Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions docs/source/utils.rst
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,13 @@ Example:
Links
-----

After a link writes the other trait, ``link`` uses the protected
``_should_update`` method to check that the original trait still has the value
from the change event. The default accepts identical objects and otherwise
uses inequality. Subclasses can override this post-propagation consistency
check for values with custom comparison behavior; it does not suppress the
write to the other trait.

.. autoclass:: link

.. autoclass:: directional_link
79 changes: 79 additions & 0 deletions tests/test_traitlets.py
Original file line number Diff line number Diff line change
Expand Up @@ -2054,6 +2054,83 @@ def _value_validate(self, proposal):


class TestLink(TestCase):
def test_non_scalar_identity_comparison_preserves_observer_chain(self):
class AmbiguousComparison:
def __bool__(self):
raise ValueError("ambiguous comparison")

class Value:
__hash__ = None

def __eq__(self, other):
return AmbiguousComparison()

__ne__ = __eq__

class A(HasTraits):
value = Instance(Value)

a = A(value=Value())
b = A(value=Value())
c = link((a, "value"), (b, "value"))
observed = []
a.observe(lambda change: observed.append("a"), names="value")
b.observe(lambda change: observed.append("b"), names="value")

a.value = Value()
b.value = Value()

self.assertEqual(observed, ["b", "a", "a", "b"])
self.assertIs(a.value, b.value)
self.assertFalse(c.updating)

def test_should_update_can_be_overridden(self):
class Value:
__hash__ = None

def __init__(self, value):
self.value = value

def __eq__(self, other):
return isinstance(other, Value) and self.value == other.value

def __ne__(self, other):
return not self == other

class A(HasTraits):
value = Instance(Value)

class IdentityLink(link):
def _should_update(self, old, new):
return old is not new

def replace_with_equal_distinct(owner, change):
if not replaced:
replaced.append(True)
owner.value = Value(change.new.value)

for link_type, raises in ((link, False), (IdentityLink, True)):
for owner_name in ("source", "target"):
source = A(value=Value(1))
target = A(value=Value(1))
replaced = []
owner = source if owner_name == "source" else target
original = source if owner_name == "source" else target
opposite = target if owner_name == "source" else source
opposite.observe(
lambda change, original=original: replace_with_equal_distinct(original, change),
names="value",
)
connection = link_type((source, "value"), (target, "value"))
replaced.clear()

if raises:
with self.assertRaises(TraitError):
owner.value = Value(2)
else:
owner.value = Value(2)
self.assertFalse(connection.updating)

def test_connect_same(self):
"""Verify two traitlets of the same type can be linked together using link."""

Expand Down Expand Up @@ -2203,6 +2280,7 @@ def another_update(self, change):
mc = MyClass()
l = link((mc, "i"), (mc, "j"))
self.assertRaises(TraitError, setattr, mc, "i", 2)
self.assertFalse(l.updating)

def test_link_broken_at_target(self):
class MyClass(HasTraits):
Expand All @@ -2216,6 +2294,7 @@ def another_update(self, change):
mc = MyClass()
l = link((mc, "i"), (mc, "j"))
self.assertRaises(TraitError, setattr, mc, "j", 2)
self.assertFalse(l.updating)


class TestDirectionalLink(TestCase):
Expand Down
14 changes: 12 additions & 2 deletions traitlets/traitlets.py
Original file line number Diff line number Diff line change
Expand Up @@ -357,12 +357,22 @@ def _busy_updating(self) -> t.Any:
finally:
self.updating = False

def _should_update(self, old: t.Any, new: t.Any) -> bool:
"""Return whether two values should be considered different.

Subclasses can override this method when a trait value has comparison
semantics that cannot be reduced to a boolean with ``!=``.
"""
if old is new:
return False
return bool(old != new)

def _update_target(self, change: t.Any) -> None:
if self.updating:
return
with self._busy_updating():
setattr(self.target[0], self.target[1], self._transform(change.new))
if getattr(self.source[0], self.source[1]) != change.new:
if self._should_update(getattr(self.source[0], self.source[1]), change.new):
raise TraitError(
f"Broken link {self}: the source value changed while updating the target."
)
Expand All @@ -372,7 +382,7 @@ def _update_source(self, change: t.Any) -> None:
return
with self._busy_updating():
setattr(self.source[0], self.source[1], self._transform_inv(change.new))
if getattr(self.target[0], self.target[1]) != change.new:
if self._should_update(getattr(self.target[0], self.target[1]), change.new):
raise TraitError(
f"Broken link {self}: the target value changed while updating the source."
)
Expand Down
Loading