diff --git a/CHANGELOG.md b/CHANGELOG.md index 55c979f..b52c5fd 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,32 @@ # Changelog +## Unreleased + +- Keep `UniqueList` membership in sync when replacing an indexed item, and + leave membership unchanged when the index is out of range, contributed by + @shkyyy18 in PR #51. +- Keep `UniqueList` membership in sync for slice assignment, `extend`, `pop`, + `remove`, `clear`, `+=` and `*=`, so `extend` and `+=` no longer add + duplicates and removed values can be added again. +- Accept one-shot iterables in `UniqueList` slice assignment, allow a slice to + reuse the values it replaces, and reject a slice that repeats a value. +- Make `copy.copy` and `copy.deepcopy` of a `UniqueList` return a working copy + with its own membership. +- Leave `UniqueList` membership unchanged when `insert` fails, and answer `in` + for an unhashable value the way a `list` does. +- Make `CastedDict` and `LazyCastedDict` load from `pickle` on every protocol, + including pickles written by 4.0.1, and stop `copy.copy` and `copy.deepcopy` + from casting the stored values a second time. +- Cast what `setdefault` and `|=` store in a `CastedDict` or `LazyCastedDict`. + A `None` default is stored as it is. +- Cast the key of a `LazyCastedDict` once when it is stored. It was cast twice. +- Let keyword arguments win over the mapping in `update` and the constructor + of the casted dicts, as `dict` does. +- Make `!=` the opposite of `==` for `SliceableDeque`, and compare unequal to a + set when an item is unhashable. +- Remove the iterable of mappings from the `DictUpdateArgs` type alias. The + code never accepted that shape. + ## 4.0.1 - 2026-08-30 - Allow `uv_build` 0.12.x, contributed by @felixonmars in PR #49. diff --git a/_python_utils_tests/clock.py b/_python_utils_tests/clock.py new file mode 100644 index 0000000..d0c5cbe --- /dev/null +++ b/_python_utils_tests/clock.py @@ -0,0 +1,56 @@ +"""A fake clock for the tests and doctests of ``python_utils.time``.""" + +import typing + +import pytest + +import python_utils.time + +#: The doctest that sleeps on the clock, by the name pytest gives it. +TIMEOUT_GENERATOR_DOCTEST: str = 'python_utils.time.timeout_generator' + + +class FakeClock: + """ + A clock that only moves when something sleeps on it. + + ``time.sleep`` promises to sleep at least as long as requested. A busy + machine sleeps tens of milliseconds longer, and that changes how many + items ``timeout_generator`` yields before its timeout. On this clock the + number of items depends on the arguments alone. + + Attributes: + now (float): The current time in seconds. + sleeps (list[float]): Every requested sleep, in order. + """ + + def __init__(self) -> None: + """Start at zero without any recorded sleeps.""" + self.now: float = 0.0 + self.sleeps: list[float] = [] + + def perf_counter(self) -> float: + """Return the current time, like ``time.perf_counter``.""" + return self.now + + def sleep(self, seconds: float) -> None: + """Record the sleep and move the clock forward, without waiting.""" + self.sleeps.append(seconds) + self.now += seconds + + +@pytest.fixture +def fake_clock(monkeypatch: pytest.MonkeyPatch) -> FakeClock: + """Replace the ``time`` module inside ``python_utils.time``.""" + clock: FakeClock = FakeClock() + monkeypatch.setattr(python_utils.time, 'time', clock) + return clock + + +@pytest.fixture(autouse=True) +def fake_clock_in_doctest(request: pytest.FixtureRequest) -> None: + """Run the ``timeout_generator`` doctest on the fake clock.""" + # pytest leaves `FixtureRequest.node` without a type. + node: pytest.Item = typing.cast(pytest.Item, request.node) + if node.name == TIMEOUT_GENERATOR_DOCTEST: + request.getfixturevalue('fake_clock') diff --git a/_python_utils_tests/test_containers.py b/_python_utils_tests/test_containers.py index ddee07e..f1c6527 100644 --- a/_python_utils_tests/test_containers.py +++ b/_python_utils_tests/test_containers.py @@ -1,9 +1,278 @@ """Tests for the container types in ``python_utils.containers``.""" +import collections +import collections.abc +import copy +import dataclasses +import pickle +import typing +import unittest.mock + import pytest from python_utils import containers +#: A callable that copies a ``UniqueList``, such as ``copy.copy``. +Copier = collections.abc.Callable[ + [containers.UniqueList[int]], containers.UniqueList[int] +] + +#: The protocols that can pickle a class with slots. The standard library +#: refuses such a class on protocol 0 and 1. +SLOTS_PROTOCOLS: range = range(2, pickle.HIGHEST_PROTOCOL + 1) + +#: A casted dict class, to run a test against both of them. +CastedDictType = type[containers.CastedDictBase[typing.Any, typing.Any]] + +#: ``CastedDict(int, int, {'1': '2'})`` and its lazy counterpart as pickled by +#: python-utils 4.0.1. That release could write protocol 2 and up but not read +#: them back. +LEGACY_DICT_PICKLES: tuple[tuple[CastedDictType, bytes], ...] = ( + ( + containers.CastedDict, + ( + b'ccopy_reg\n_reconstructor\np0\n(cpython_utils.containers\nCastedDict' + b'\np1\nc__builtin__\ndict\np2\n(dp3\nI1\nI2\nstp4\nRp5\n(dp6\n' + b'V_value_cast\np7\nc__builtin__\nlong\np8\nsV_key_cast\np9\ng8\nsb.' + ), + ), + ( + containers.CastedDict, + ( + b'\x80\x02cpython_utils.containers\nCastedDict\nq\x00)\x81q\x01K\x01K' + b'\x02s}q\x02(X\x0b\x00\x00\x00_value_castq\x03c__builtin__\nlong\nq' + b'\x04X\t\x00\x00\x00_key_castq\x05h\x04ub.' + ), + ), + ( + containers.CastedDict, + ( + b'\x80\x04\x95f\x00\x00\x00\x00\x00\x00\x00\x8c\x17python_utils.' + b'containers\x94\x8c\nCastedDict\x94\x93\x94)\x81\x94K\x01K\x02s}\x94' + b'(\x8c\x0b_value_cast\x94\x8c\x08builtins\x94\x8c\x03int\x94\x93\x94' + b'\x8c\t_key_cast\x94h\x08ub.' + ), + ), + ( + containers.LazyCastedDict, + ( + b'ccopy_reg\n_reconstructor\np0\n(cpython_utils.containers\n' + b'LazyCastedDict\np1\nc__builtin__\ndict\np2\n(dp3\nI1\nV2\np4\nstp5' + b'\nRp6\n(dp7\nV_value_cast\np8\nc__builtin__\nlong\np9\nsV_key_cast' + b'\np10\ng9\nsb.' + ), + ), + ( + containers.LazyCastedDict, + ( + b'\x80\x04\x95j\x00\x00\x00\x00\x00\x00\x00\x8c\x17python_utils.' + b'containers\x94\x8c\x0eLazyCastedDict\x94\x93\x94)\x81\x94K\x01K\x02s' + b'}\x94(\x8c\x0b_value_cast\x94\x8c\x08builtins\x94\x8c\x03int\x94\x93' + b'\x94\x8c\t_key_cast\x94h\x08ub.' + ), + ), +) + + +def raw_items( + values: dict[typing.Any, typing.Any], +) -> dict[typing.Any, typing.Any]: + """Return the items as they are stored, without any cast.""" + return dict[typing.Any, typing.Any].copy(values) + + +def double(value: int) -> int: + """Double a value, as a cast that must not be applied twice.""" + return value * 2 + + +class StrictUniqueList(containers.UniqueList[int]): + """A list that sets its duplicate policy on the class.""" + + on_duplicate: containers.OnDuplicate = 'raise' + + def __init__(self, *values: int) -> None: + """Fill the list without calling ``UniqueList.__init__``.""" + super(containers.UniqueList, self).__init__() + self._set = set() + for value in values: + self.append(value) + + +class SlottedUniqueList(containers.UniqueList[int]): + """A list that keeps extra attributes in slots.""" + + __slots__ = ('other', 'tag') + + other: str + tag: str + + +class SlottedCastedDict(containers.CastedDict[int, int]): + """A casted dict that keeps an extra attribute in a slot.""" + + __slots__ = ('tag',) + + tag: str + + +class ReducedCastedDict(containers.CastedDict[int, int]): + """A casted dict that brings its own pickle support.""" + + def __reduce__(self) -> tuple[typing.Any, ...]: + """Rebuild through the constructor.""" + return ReducedCastedDict, (int, int, raw_items(self)) + + +class StampedCastedDict(containers.CastedDict[int, int]): + """A casted dict that restores its own state, the textbook way.""" + + restored: bool = False + + def __setstate__(self, state: typing.Any) -> None: + """Restore the instance dictionary and leave a mark.""" + vars(self).update(state) + self.restored = True + + +class AuditedCastedDict(containers.CastedDict[int, int]): + """A casted dict that passes the default reduce value through.""" + + def __reduce_ex__(self, protocol: typing.SupportsIndex) -> typing.Any: + """Unpack the five default items and hand them on.""" + function: typing.Any + arguments: typing.Any + state: typing.Any + list_items: typing.Any + dict_items: typing.Any + function, arguments, state, list_items, dict_items = ( + super().__reduce_ex__(protocol) + ) + return function, arguments, state, list_items, dict_items + + +class BrokenHash: + """A value that can no longer be hashed once it is broken.""" + + def __init__(self) -> None: + """Start out hashable.""" + self.broken: bool = False + + def __hash__(self) -> int: + """Raise an error that is not a ``TypeError`` when broken.""" + if self.broken: + raise ValueError('broken hash') + + return id(self) + + +class Reconnecting: + """A mixin that drops a live handle from the state and reopens it.""" + + handle: str + + def __getstate__(self) -> dict[str, typing.Any]: + """Leave the handle out of the state.""" + state: dict[str, typing.Any] = dict(vars(self)) + state.pop('handle', None) + return state + + def __setstate__(self, state: dict[str, typing.Any]) -> None: + """Restore the attributes and reopen the handle.""" + vars(self).update(state) + self.handle = 'reopened' + + +class ReconnectingCastedDict(containers.CastedDict[int, int], Reconnecting): + """A casted dict whose ``__setstate__`` comes from a mixin after it.""" + + +class ReconnectingUniqueList(containers.UniqueList[int], Reconnecting): + """A unique list whose ``__setstate__`` comes from a mixin after it.""" + + +class VersionedCastedDict(containers.CastedDict[int, int]): + """A casted dict that adds a version to the default reduce state.""" + + version: int = 0 + + def __reduce_ex__(self, protocol: typing.SupportsIndex) -> typing.Any: + """Edit the default state, which is the instance dictionary.""" + function: typing.Any + arguments: typing.Any + state: typing.Any + list_items: typing.Any + dict_items: typing.Any + function, arguments, state, list_items, dict_items = ( + super().__reduce_ex__(protocol) + ) + state = {**state, 'version': 2} + return function, arguments, state, list_items, dict_items + + +class ShiftingIndex: + """An index that is valid the first time it is used and not after.""" + + def __init__(self) -> None: + """Start without any use.""" + self.uses: int = 0 + + def __index__(self) -> int: + """Return 0 once and an index out of range from then on.""" + self.uses += 1 + return 0 if self.uses == 1 else 99 + + +class CountingCast: + """A key cast that adds one and counts how often it is called.""" + + def __init__(self) -> None: + """Start without any calls.""" + self.calls: int = 0 + + def __call__(self, key: int) -> int: + """Return the key plus one.""" + self.calls += 1 + return key + 1 + + +class BareCastedDict(containers.CastedDict[str, int]): + """A casted dict without casts that never calls the base constructor.""" + + def __init__(self) -> None: + """Leave the casts at their class defaults.""" + + +@dataclasses.dataclass(unsafe_hash=True) +class Point: + """A hashable value whose hash changes when it is mutated.""" + + x: int + + +#: ``UniqueList(1, 2, on_duplicate='raise')`` as pickled by python-utils 4.0.1, +#: which stored the membership set in the instance state. +LEGACY_PICKLES: tuple[bytes, ...] = ( + ( + b'ccopy_reg\n_reconstructor\np0\n(cpython_utils.containers\nUniqueList' + b'\np1\nc__builtin__\nlist\np2\n(lp3\nI1\naI2\natp4\nRp5\n(dp6\n' + b'Von_duplicate\np7\nVraise\np8\nsV_set\np9\nc__builtin__\nset\np10\n' + b'((lp11\nI1\naI2\natp12\nRp13\nsb.' + ), + ( + b'\x80\x02cpython_utils.containers\nUniqueList\nq\x00)\x81q\x01(K\x01' + b'K\x02e}q\x02(X\x0c\x00\x00\x00on_duplicateq\x03X\x05\x00\x00\x00' + b'raiseq\x04X\x04\x00\x00\x00_setq\x05c__builtin__\nset\nq\x06]q\x07' + b'(K\x01K\x02e\x85q\x08Rq\tub.' + ), + ( + b'\x80\x04\x95^\x00\x00\x00\x00\x00\x00\x00\x8c\x17python_utils.' + b'containers\x94\x8c\nUniqueList\x94\x93\x94)\x81\x94(K\x01K\x02e}' + b'\x94(\x8c\x0con_duplicate\x94\x8c\x05raise\x94\x8c\x04_set\x94\x8f' + b'\x94(K\x01K\x02\x90ub.' + ), +) + def test_unique_list_ignore() -> None: """Ignore duplicate appends and block duplicate slice sets.""" @@ -78,3 +347,865 @@ def test_sliceable_deque_eq() -> None: assert d == {1, 2, 3} assert d == d assert d == containers.SliceableDeque([1, 2, 3]) + + +@pytest.mark.parametrize('on_duplicate', ['ignore', 'raise']) +@pytest.mark.parametrize('index', [0, -2]) +def test_unique_list_replace_membership( + on_duplicate: containers.OnDuplicate, index: int +) -> None: + """Release the old value after replacing an indexed item.""" + values = containers.UniqueList(1, 2, on_duplicate=on_duplicate) + values[index] = 3 + + assert values == [3, 2] + assert 1 not in values + assert 3 in values + values.append(1) + assert values == [3, 2, 1] + + +@pytest.mark.parametrize('on_duplicate', ['ignore', 'raise']) +@pytest.mark.parametrize('index', [2, -3]) +def test_unique_list_failed_replace_preserves_membership( + on_duplicate: containers.OnDuplicate, index: int +) -> None: + """Do not reserve a value when indexed assignment fails.""" + values = containers.UniqueList(1, 2, on_duplicate=on_duplicate) + with pytest.raises(IndexError): + values[index] = 3 + + assert values == [1, 2] + assert 3 not in values + values.append(3) + assert values == [1, 2, 3] + + +@pytest.mark.parametrize('on_duplicate', ['ignore', 'raise']) +def test_unique_list_replace_same_value( + on_duplicate: containers.OnDuplicate, +) -> None: + """Replacing an item with itself keeps membership intact.""" + values = containers.UniqueList(1, 2, on_duplicate=on_duplicate) + values[0] = 1 + assert values == [1, 2] + assert 1 in values + assert 2 in values + + +def test_unique_list_slice_replace_membership() -> None: + """Release the old values after replacing a slice.""" + values: containers.UniqueList[int] = containers.UniqueList( + 1, 2, 3, on_duplicate='raise' + ) + values[0:2] = [8, 9] + + assert values == [8, 9, 3] + assert 1 not in values + assert 8 in values + values.append(1) + assert values == [8, 9, 3, 1] + + +def test_unique_list_failed_slice_preserves_membership() -> None: + """Do not reserve values when slice assignment fails.""" + values: containers.UniqueList[int] = containers.UniqueList( + 1, 2, 3, 4, on_duplicate='raise' + ) + with pytest.raises(ValueError, match='extended slice'): + values[0:4:2] = [8, 9, 10] + + assert values == [1, 2, 3, 4] + assert 8 not in values + values.append(8) + assert values == [1, 2, 3, 4, 8] + + +def test_unique_list_slice_from_iterator() -> None: + """Assign a one-shot iterable to a slice without losing its values.""" + values: containers.UniqueList[int] = containers.UniqueList( + 1, 2, 3, on_duplicate='raise' + ) + values[0:2] = iter([8, 9]) + + assert values == [8, 9, 3] + assert 8 in values + assert 1 not in values + + +def test_unique_list_slice_reuses_replaced_values() -> None: + """Allow new slice values that only duplicate the values they replace.""" + values: containers.UniqueList[int] = containers.UniqueList( + 1, 2, 3, on_duplicate='raise' + ) + values[0:2] = (1, 2) + assert values == [1, 2, 3] + + values[0:2] = [2, 1] + assert values == [2, 1, 3] + + values[0:2] = [1, 9] + assert values == [1, 9, 3] + assert 2 not in values + + +def test_unique_list_slice_rejects_duplicates() -> None: + """Reject slice values that repeat themselves or a kept item.""" + values: containers.UniqueList[int] = containers.UniqueList( + 1, 2, 3, on_duplicate='raise' + ) + with pytest.raises(ValueError, match='Duplicate values'): + values[0:2] = [9, 9] + with pytest.raises(ValueError, match='Duplicate values'): + values[0:2] = [3, 9] + + assert values == [1, 2, 3] + assert 9 not in values + + +def test_unique_list_extend_ignore() -> None: + """Skip duplicates while extending and track the new members.""" + values: containers.UniqueList[int] = containers.UniqueList(1, 2) + values.extend(iter([2, 3, 3, 4])) + + assert values == [1, 2, 3, 4] + assert 3 in values + values[2] = 5 + assert values == [1, 2, 5, 4] + assert 3 not in values + + +@pytest.mark.parametrize('extra', [[3, 1], [3, 3]]) +def test_unique_list_extend_raise(extra: list[int]) -> None: + """Raise on a duplicate before extending with any of the values.""" + values: containers.UniqueList[int] = containers.UniqueList( + 1, 2, on_duplicate='raise' + ) + with pytest.raises(ValueError, match='Duplicate value'): + values.extend(extra) + + assert values == [1, 2] + assert 3 not in values + values.extend([3, 4]) + assert values == [1, 2, 3, 4] + + +def test_unique_list_iadd() -> None: + """Keep ``+=`` unique and return the same list.""" + values: containers.UniqueList[int] = containers.UniqueList(1, 2) + original: containers.UniqueList[int] = values + values += [2, 3] + + assert values is original + assert values == [1, 2, 3] + assert 3 in values + + +@pytest.mark.parametrize('on_duplicate', ['ignore', 'raise']) +def test_unique_list_imul_clear(on_duplicate: containers.OnDuplicate) -> None: + """Empty the list and its membership when multiplying by zero.""" + values: containers.UniqueList[int] = containers.UniqueList( + 1, 2, on_duplicate=on_duplicate + ) + values *= 1 + assert values == [1, 2] + + values *= 0 + assert values == [] + assert 1 not in values + values *= 2 + assert values == [] + values.append(1) + assert values == [1] + + +def test_unique_list_imul_repeat() -> None: + """Never repeat the items when multiplying in place.""" + values: containers.UniqueList[int] = containers.UniqueList(1, 2) + values *= 3 + assert values == [1, 2] + + strict: containers.UniqueList[int] = containers.UniqueList( + 1, 2, on_duplicate='raise' + ) + with pytest.raises(ValueError, match='Duplicate values'): + strict *= 2 + assert strict == [1, 2] + + +def test_unique_list_pop() -> None: + """Release a popped value so that it can be added again.""" + values: containers.UniqueList[int] = containers.UniqueList( + 1, 2, 3, on_duplicate='raise' + ) + last: int = values.pop() + first: int = values.pop(0) + assert (last, first) == (3, 1) + assert values == [2] + assert 1 not in values + assert 3 not in values + values.append(3) + assert values == [2, 3] + + with pytest.raises(IndexError): + values.pop(5) + assert values == [2, 3] + assert 2 in values + + +def test_unique_list_remove() -> None: + """Release a removed value so that it can be added again.""" + values: containers.UniqueList[int] = containers.UniqueList( + 1, 2, 3, on_duplicate='raise' + ) + values.remove(2) + assert values == [1, 3] + assert 2 not in values + values.append(2) + assert values == [1, 3, 2] + + with pytest.raises(ValueError): + values.remove(9) + assert values == [1, 3, 2] + + +def test_unique_list_clear() -> None: + """Release every value when the list is cleared.""" + values: containers.UniqueList[int] = containers.UniqueList( + 1, 2, on_duplicate='raise' + ) + values.clear() + assert values == [] + assert 1 not in values + values.extend([2, 1]) + assert values == [2, 1] + + +@pytest.mark.parametrize('on_duplicate', ['ignore', 'raise']) +@pytest.mark.parametrize('protocol', range(pickle.HIGHEST_PROTOCOL + 1)) +def test_unique_list_pickle( + on_duplicate: containers.OnDuplicate, protocol: int +) -> None: + """Round-trip through pickle with working membership.""" + values: containers.UniqueList[int] = containers.UniqueList( + 1, 2, on_duplicate=on_duplicate + ) + restored: containers.UniqueList[int] = pickle.loads( + pickle.dumps(values, protocol=protocol) + ) + + assert type(restored) is containers.UniqueList + assert restored == [1, 2] + assert restored.on_duplicate == on_duplicate + assert 1 in restored + restored.append(3) + assert restored == [1, 2, 3] + assert 3 not in values + + +@pytest.mark.parametrize( + 'data', LEGACY_PICKLES, ids=['protocol-0', 'protocol-2', 'protocol-4'] +) +def test_unique_list_legacy_pickle(data: bytes) -> None: + """Load a pickle written by python-utils 4.0.1.""" + restored: containers.UniqueList[int] = pickle.loads(data) + + assert type(restored) is containers.UniqueList + assert restored == [1, 2] + assert restored.on_duplicate == 'raise' + with pytest.raises(ValueError): + restored.append(1) + restored.append(3) + assert restored == [1, 2, 3] + + +@pytest.mark.parametrize('on_duplicate', ['ignore', 'raise']) +@pytest.mark.parametrize('copier', [copy.copy, copy.deepcopy]) +def test_unique_list_copy( + on_duplicate: containers.OnDuplicate, + copier: Copier, +) -> None: + """Copy the items and give the copy its own membership.""" + values: containers.UniqueList[int] = containers.UniqueList( + 1, 2, on_duplicate=on_duplicate + ) + copied: containers.UniqueList[int] = copier(values) + + assert type(copied) is containers.UniqueList + assert copied == [1, 2] + assert copied.on_duplicate == on_duplicate + copied.append(3) + assert copied == [1, 2, 3] + assert values == [1, 2] + assert 3 not in values + + +def test_unique_list_class_level_on_duplicate() -> None: + """Honour a duplicate policy that a subclass sets on the class.""" + values: StrictUniqueList = StrictUniqueList(1, 2) + + assert values.on_duplicate == 'raise' + with pytest.raises(ValueError, match='Duplicate value'): + values.append(1) + assert values == [1, 2] + + +@pytest.mark.parametrize('protocol', SLOTS_PROTOCOLS) +def test_unique_list_slots_pickle(protocol: int) -> None: + """Keep the slots of a subclass through pickle.""" + values: SlottedUniqueList = SlottedUniqueList(1, 2, on_duplicate='raise') + values.tag = 'spam' + values.other = 'eggs' + restored: SlottedUniqueList = pickle.loads( + pickle.dumps(values, protocol=protocol) + ) + + assert restored == [1, 2] + assert restored.on_duplicate == 'raise' + assert (restored.tag, restored.other) == ('spam', 'eggs') + assert set(vars(restored)) == {'on_duplicate', '_set'} + restored.append(3) + assert 3 in restored + + +@pytest.mark.parametrize('copier', [copy.copy, copy.deepcopy]) +def test_unique_list_slots_copy( + copier: collections.abc.Callable[[SlottedUniqueList], SlottedUniqueList], +) -> None: + """Keep the slots of a subclass through a copy.""" + values: SlottedUniqueList = SlottedUniqueList(1, 2, on_duplicate='raise') + values.tag = 'spam' + values.other = 'eggs' + copied: SlottedUniqueList = copier(values) + + assert copied == [1, 2] + assert copied.on_duplicate == 'raise' + assert (copied.tag, copied.other) == ('spam', 'eggs') + copied.append(3) + assert 3 not in values + + +def test_unique_list_pop_after_hash_change() -> None: + """Return a popped item even when its hash changed in the list.""" + point: Point = Point(1) + other: Point = Point(5) + values: containers.UniqueList[Point] = containers.UniqueList(other, point) + point.x = 2 + popped: Point = values.pop() + + assert popped is point + assert values == [other] + assert point not in values + assert other in values + values.append(point) + assert values == [other, point] + + +def test_unique_list_replace_after_hash_change() -> None: + """Replace an item by index even when its hash changed in the list.""" + point: Point = Point(1) + other: Point = Point(5) + values: containers.UniqueList[Point] = containers.UniqueList(point) + point.x = 2 + values[0] = other + + assert values == [other] + assert point not in values + assert other in values + + +def test_unique_list_remove_equal_unhashable() -> None: + """Release the stored item when an equal, unhashable value removes it.""" + values: containers.UniqueList[int] = containers.UniqueList(1, 2, 3) + values.remove(unittest.mock.ANY) + + assert values == [2, 3] + assert 1 not in values + values.append(1) + assert values == [2, 3, 1] + + +def test_unique_list_failed_insert_preserves_membership() -> None: + """Do not reserve a value when the insert itself fails.""" + values: containers.UniqueList[int] = containers.UniqueList(1, 2) + bad_index: typing.Any = 'spam' + with pytest.raises(TypeError): + values.insert(bad_index, 3) + + assert values == [1, 2] + assert 3 not in values + values.append(3) + assert values == [1, 2, 3] + + +def test_unique_list_contains_unhashable() -> None: + """Answer a membership test for an unhashable value like a list does.""" + values: containers.UniqueList[int] = containers.UniqueList(1, 2) + unhashable: typing.Any = [1] + + assert unhashable not in values + assert unittest.mock.ANY in values + + +@pytest.mark.parametrize( + 'dict_type', [containers.CastedDict, containers.LazyCastedDict] +) +@pytest.mark.parametrize('protocol', range(pickle.HIGHEST_PROTOCOL + 1)) +def test_casted_dict_pickle(dict_type: CastedDictType, protocol: int) -> None: + """Round-trip a casted dict through pickle without casting again.""" + values: containers.CastedDictBase[int, int] = dict_type(int, double) + values['1'] = 2 + restored: containers.CastedDictBase[int, int] = pickle.loads( + pickle.dumps(values, protocol=protocol) + ) + + assert type(restored) is dict_type + assert raw_items(restored) == raw_items(values) + assert restored[1] == values[1] + restored['3'] = 4 + assert restored[3] == 8 + + +@pytest.mark.parametrize( + 'dict_type,data', + LEGACY_DICT_PICKLES, + ids=['strict-0', 'strict-2', 'strict-4', 'lazy-0', 'lazy-4'], +) +def test_casted_dict_legacy_pickle( + dict_type: CastedDictType, data: bytes +) -> None: + """Load a pickle written by python-utils 4.0.1.""" + restored: containers.CastedDictBase[int, int] = pickle.loads(data) + + assert type(restored) is dict_type + assert restored[1] == 2 + restored['3'] = '4' + assert restored[3] == 4 + + +@pytest.mark.parametrize( + 'dict_type', [containers.CastedDict, containers.LazyCastedDict] +) +@pytest.mark.parametrize('copier', [copy.copy, copy.deepcopy]) +def test_casted_dict_copy( + dict_type: CastedDictType, + copier: collections.abc.Callable[ + [containers.CastedDictBase[int, int]], + containers.CastedDictBase[int, int], + ], +) -> None: + """Copy the stored items as they are, without casting them again.""" + values: containers.CastedDictBase[int, int] = dict_type(int, double) + values['1'] = 2 + copied: containers.CastedDictBase[int, int] = copier(values) + + assert type(copied) is dict_type + assert raw_items(copied) == raw_items(values) + assert copied[1] == values[1] + copied['3'] = 4 + assert copied[3] == 8 + assert 3 not in values + + +def test_casted_dict_deepcopy_cycle() -> None: + """Deep-copy a casted dict that contains itself.""" + values: containers.CastedDict[str, typing.Any] = containers.CastedDict( + str, None + ) + values['self'] = values + copied: containers.CastedDict[str, typing.Any] = copy.deepcopy(values) + + assert copied['self'] is copied + assert copied is not values + + +@pytest.mark.parametrize('protocol', SLOTS_PROTOCOLS) +def test_casted_dict_slots_pickle(protocol: int) -> None: + """Keep the slots of a subclass through pickle.""" + values: SlottedCastedDict = SlottedCastedDict(int, int) + values['1'] = '2' + values.tag = 'spam' + restored: SlottedCastedDict = pickle.loads( + pickle.dumps(values, protocol=protocol) + ) + + assert restored == {1: 2} + assert restored.tag == 'spam' + restored['3'] = '4' + assert restored[3] == 4 + + +def test_casted_dict_default_slots_state() -> None: + """Accept the default state of a class with slots, as 4.0.1 wrote it.""" + values: SlottedCastedDict = SlottedCastedDict.__new__(SlottedCastedDict) + values.__setstate__( + ({'_key_cast': int, '_value_cast': int}, {'tag': 'spam'}) + ) + values['1'] = '2' + + assert values == {1: 2} + assert values.tag == 'spam' + + +def test_casted_dict_slots_copy() -> None: + """Keep the slots of a subclass through a copy.""" + values: SlottedCastedDict = SlottedCastedDict(int, int) + values['1'] = '2' + values.tag = 'spam' + copied: SlottedCastedDict = copy.copy(values) + + assert copied == {1: 2} + assert copied.tag == 'spam' + + +def test_casted_dict_custom_reduce() -> None: + """Leave a subclass with its own ``__reduce__`` alone.""" + values: ReducedCastedDict = ReducedCastedDict(int, int) + values['1'] = '2' + restored: ReducedCastedDict = pickle.loads(pickle.dumps(values)) + copied: ReducedCastedDict = copy.copy(values) + + assert type(restored) is ReducedCastedDict + assert restored == copied == {1: 2} + + +@pytest.mark.parametrize('copier', [copy.copy, copy.deepcopy]) +def test_casted_dict_copy_without_attributes( + copier: collections.abc.Callable[[BareCastedDict], BareCastedDict], +) -> None: + """Copy a dict whose instance holds no attributes at all.""" + values: BareCastedDict = BareCastedDict() + values['spam'] = 1 + copied: BareCastedDict = copier(values) + + assert type(copied) is BareCastedDict + assert copied == {'spam': 1} + assert vars(copied) == {} + + +def test_casted_dict_setdefault() -> None: + """Cast the key and the value that ``setdefault`` stores.""" + values: containers.CastedDict[int, int] = containers.CastedDict(int, int) + first: int = values.setdefault('1', '2') + second: int = values.setdefault('1', '9') + third: int = values.setdefault(1, '9') + + assert (first, second, third) == (2, 2, 2) + assert raw_items(values) == {1: 2} + + +def test_casted_dict_setdefault_without_casts() -> None: + """Behave like ``dict.setdefault`` when there are no casts.""" + values: containers.CastedDict[str, typing.Any] = containers.CastedDict() + missing: typing.Any = values.setdefault('spam') + present: typing.Any = values.setdefault('spam', 'eggs') + + assert missing is None + assert present is None + assert values == {'spam': None} + + +def test_lazy_casted_dict_setdefault() -> None: + """Cast the key that ``setdefault`` stores and keep the value raw.""" + values: containers.LazyCastedDict[int, int] = containers.LazyCastedDict( + int, int + ) + values.setdefault('1', '2') + values.setdefault(1, '9') + + assert raw_items(values) == {1: '2'} + assert values[1] == 2 + + +@pytest.mark.parametrize( + 'dict_type', [containers.CastedDict, containers.LazyCastedDict] +) +def test_casted_dict_ior(dict_type: CastedDictType) -> None: + """Cast what ``|=`` merges in and keep the same dict.""" + values: containers.CastedDictBase[int, int] = dict_type(int, int) + original: containers.CastedDictBase[int, int] = values + values |= {'1': '2'} + values |= [('3', '4')] + + assert values is original + assert list(values) == [1, 3] + assert (values[1], values[3]) == (2, 4) + + +def test_casted_dict_update_keyword_precedence() -> None: + """Let keyword arguments win over the mapping, like ``dict`` does.""" + plain: dict[str, int] = {'a': 0, 'b': 0} + plain.update({'a': 1, 'c': 3}, a=2, b=1) + values: containers.CastedDict[str, int] = containers.CastedDict( + None, None, {'a': 0, 'b': 0} + ) + values.update({'a': 1, 'c': 3}, a=2, b=1) + constructed: containers.CastedDict[str, int] = containers.CastedDict( + None, None, {'a': 1}, a=2 + ) + + assert values == plain + assert list(values) == list(plain) + assert constructed == {'a': 2} + + +def test_casted_dict_update_single_positional() -> None: + """Reject a second positional argument like ``dict.update`` does.""" + values: containers.CastedDict[int, int] = containers.CastedDict(int, int) + first: typing.Any = {'1': '2'} + second: typing.Any = {'3': '4'} + + # The message is the interpreter's own and differs on PyPy. + with pytest.raises(TypeError, match='update'): + values.update(first, second) + assert values == {} + + +def test_sliceable_deque_ne() -> None: + """Keep ``!=`` the opposite of ``==`` for every supported type.""" + values: containers.SliceableDeque[int] = containers.SliceableDeque( + [1, 2, 3] + ) + equal: list[typing.Any] = [ + [1, 2, 3], + (1, 2, 3), + {1, 2, 3}, + collections.deque([1, 2, 3]), + containers.SliceableDeque([1, 2, 3]), + ] + different: list[typing.Any] = [[1, 2], (3, 2, 1), {1, 2}, 'spam', None] + + for other in equal: + assert (values == other) is True + assert (values != other) is False + for other in different: + assert (values != other) is True + assert (values == other) is False + + +def test_sliceable_deque_eq_set_unhashable() -> None: + """Compare unequal to a set when an item cannot be hashed.""" + values: containers.SliceableDeque[list[int]] = containers.SliceableDeque( + [[1], [2]] + ) + other: typing.Any = {1, 2} + + assert (values == other) is False + assert (values != other) is True + + +@pytest.mark.parametrize('copier', [copy.copy, copy.deepcopy]) +def test_casted_dict_subclass_setstate_copy( + copier: collections.abc.Callable[[StampedCastedDict], StampedCastedDict], +) -> None: + """Hand a subclass with its own ``__setstate__`` the default state.""" + values: StampedCastedDict = StampedCastedDict(int, int) + values['1'] = '2' + copied: StampedCastedDict = copier(values) + + assert copied == {1: 2} + assert copied.restored + copied['3'] = '4' + assert copied[3] == 4 + + +@pytest.mark.parametrize('protocol', range(pickle.HIGHEST_PROTOCOL + 1)) +def test_casted_dict_subclass_setstate_pickle(protocol: int) -> None: + """Pickle a subclass with its own ``__setstate__`` on every protocol.""" + values: StampedCastedDict = StampedCastedDict(int, int) + values['1'] = '2' + restored: StampedCastedDict = pickle.loads( + pickle.dumps(values, protocol=protocol) + ) + + assert restored == {1: 2} + assert restored.restored + + +def test_casted_dict_reduce_shape() -> None: + """Keep the five items of the default reduce value.""" + values: AuditedCastedDict = AuditedCastedDict(int, int) + values['1'] = '2' + plain: containers.CastedDict[int, int] = containers.CastedDict(int, int) + plain['1'] = '2' + + assert len(plain.__reduce_ex__(pickle.HIGHEST_PROTOCOL)) == 5 + assert copy.copy(values) == {1: 2} + assert pickle.loads(pickle.dumps(values)) == {1: 2} + + +@pytest.mark.parametrize( + 'dict_type', [containers.CastedDict, containers.LazyCastedDict] +) +@pytest.mark.parametrize('protocol', range(2, pickle.HIGHEST_PROTOCOL + 1)) +def test_casted_dict_empty_pickle_is_default( + dict_type: CastedDictType, protocol: int +) -> None: + """Pickle an empty dict in the form that python-utils 4.0.1 can load.""" + values: containers.CastedDictBase[int, int] = dict_type(int, int) + state: typing.Any = values.__reduce_ex__(protocol)[2] + + assert state == {'_key_cast': int, '_value_cast': int} + assert pickle.loads(pickle.dumps(values, protocol=protocol)) == {} + + +def test_unique_list_pop_with_broken_hash() -> None: + """Return a popped item when no item in the list can be hashed.""" + first: BrokenHash = BrokenHash() + second: BrokenHash = BrokenHash() + values: containers.UniqueList[BrokenHash] = containers.UniqueList( + first, second + ) + first.broken = True + second.broken = True + popped: BrokenHash = values.pop() + values.remove(first) + + assert popped is second + assert values == [] + + +def test_unique_list_remove_missing() -> None: + """Raise the error of ``list.remove`` for a missing value.""" + values: containers.UniqueList[int] = containers.UniqueList(1, 2) + plain: list[int] = [1, 2] + with pytest.raises(ValueError) as expected: + plain.remove(9) + with pytest.raises(ValueError) as raised: + values.remove(9) + + assert str(raised.value) == str(expected.value) + assert values == [1, 2] + + +@pytest.mark.parametrize( + 'dict_type', [containers.CastedDict, containers.LazyCastedDict] +) +def test_casted_dict_key_cast_once(dict_type: CastedDictType) -> None: + """Cast a key exactly once when it is stored.""" + key_cast: CountingCast = CountingCast() + values: containers.CastedDictBase[int, str] = dict_type(key_cast, None) + values[1] = 'spam' + + assert raw_items(values) == {2: 'spam'} + assert key_cast.calls == 1 + + +@pytest.mark.parametrize( + 'dict_type', [containers.CastedDict, containers.LazyCastedDict] +) +def test_casted_dict_setdefault_key_cast_once( + dict_type: CastedDictType, +) -> None: + """Cast the key of ``setdefault`` exactly once, found or not.""" + key_cast: CountingCast = CountingCast() + values: containers.CastedDictBase[int, str] = dict_type(key_cast, None) + stored: str = values.setdefault(1, 'spam') + found: str = values.setdefault(1, 'eggs') + + assert (stored, found) == ('spam', 'spam') + assert raw_items(values) == {2: 'spam'} + assert key_cast.calls == 2 + + +@pytest.mark.parametrize('value_cast', [int, str]) +def test_casted_dict_setdefault_none( + value_cast: collections.abc.Callable[[typing.Any], typing.Any], +) -> None: + """Store a missing default as ``None`` without casting it.""" + values: containers.CastedDict[int, typing.Any] = containers.CastedDict( + int, value_cast + ) + implicit: typing.Any = values.setdefault('1') + explicit: typing.Any = values.setdefault('2', None) + + assert implicit is None + assert explicit is None + assert raw_items(values) == {1: None, 2: None} + + +@pytest.mark.parametrize('protocol', [0, pickle.HIGHEST_PROTOCOL]) +def test_casted_dict_mixin_setstate(protocol: int) -> None: + """Run a ``__setstate__`` from a mixin that comes after the dict.""" + values: ReconnectingCastedDict = ReconnectingCastedDict(int, int) + values['1'] = '2' + values.handle = 'open' + restored: ReconnectingCastedDict = pickle.loads( + pickle.dumps(values, protocol=protocol) + ) + copied: ReconnectingCastedDict = copy.copy(values) + + assert restored == copied == {1: 2} + assert (restored.handle, copied.handle) == ('reopened', 'reopened') + copied['3'] = '4' + assert copied[3] == 4 + + +@pytest.mark.parametrize('protocol', [0, pickle.HIGHEST_PROTOCOL]) +def test_unique_list_mixin_setstate(protocol: int) -> None: + """Run a ``__setstate__`` from a mixin that comes after the list.""" + values: ReconnectingUniqueList = ReconnectingUniqueList(1, 2) + values.handle = 'open' + restored: ReconnectingUniqueList = pickle.loads( + pickle.dumps(values, protocol=protocol) + ) + copied: ReconnectingUniqueList = copy.copy(values) + + assert restored == copied == [1, 2] + assert (restored.handle, copied.handle) == ('reopened', 'reopened') + copied.append(1) + copied.append(3) + assert copied == [1, 2, 3] + + +@pytest.mark.parametrize('copier', [copy.copy, copy.deepcopy]) +def test_casted_dict_subclass_reduce_ex_state( + copier: collections.abc.Callable[ + [VersionedCastedDict], VersionedCastedDict + ], +) -> None: + """Give a subclass ``__reduce_ex__`` the default state to work on.""" + values: VersionedCastedDict = VersionedCastedDict(int, int) + values['1'] = '2' + copied: VersionedCastedDict = copier(values) + restored: VersionedCastedDict = pickle.loads(pickle.dumps(values)) + + assert copied == restored == {1: 2} + assert (copied.version, restored.version) == (2, 2) + + +def test_unique_list_delete_after_hash_change() -> None: + """Delete an item or a slice even when a hash changed in the list.""" + first: Point = Point(1) + second: Point = Point(2) + third: Point = Point(3) + values: containers.UniqueList[Point] = containers.UniqueList( + first, second, third + ) + second.x = 99 + del values[1] + + assert values == [first, third] + assert second not in values + + third.x = 98 + del values[0:2] + + assert values == [] + values.append(first) + assert values == [first] + + +def test_unique_list_failed_assignment_takes_membership_back() -> None: + """Undo the membership when the list refuses an index assignment.""" + values: containers.UniqueList[int] = containers.UniqueList(1, 2) + with pytest.raises(IndexError): + values[ShiftingIndex()] = 3 + with pytest.raises(IndexError): + values[ShiftingIndex()] = 1 + + assert values == [1, 2] + assert 3 not in values + assert 1 in values + values.append(3) + assert values == [1, 2, 3] diff --git a/_python_utils_tests/test_lazy_imports.py b/_python_utils_tests/test_lazy_imports.py index da8143e..54ec64c 100644 --- a/_python_utils_tests/test_lazy_imports.py +++ b/_python_utils_tests/test_lazy_imports.py @@ -11,6 +11,7 @@ import pytest import python_utils +from _python_utils_tests import clock def _run_clean(code: str) -> subprocess.CompletedProcess[str]: @@ -94,10 +95,13 @@ def test_star_import_resolves_all_names() -> None: @pytest.mark.asyncio -async def test_aio_timeout_generator_default_iterable() -> None: +async def test_aio_timeout_generator_default_iterable( + fake_clock: clock.FakeClock, +) -> None: """Default the iterable to ``aio.acount`` when omitted.""" # With no iterable the generator defaults to ``aio.acount`` -- exercising # the lazy ``aio``/``asyncio`` import and the None-resolution branch. + # The fake clock stands still, so the timeout cannot end the loop early. count = 0 generator: collections.abc.AsyncGenerator[object, None] = ( python_utils.aio_timeout_generator(timeout=0.05, interval=0.0) diff --git a/_python_utils_tests/test_time.py b/_python_utils_tests/test_time.py index da26c9e..c9cca14 100644 --- a/_python_utils_tests/test_time.py +++ b/_python_utils_tests/test_time.py @@ -3,12 +3,18 @@ import asyncio import datetime import itertools +import time import pytest import python_utils +from _python_utils_tests import clock from python_utils import types +#: Far longer than every timeout in this module, so a generator that sleeps +#: this long is always interrupted first. +STALL: float = 10.0 + @pytest.mark.parametrize( 'timeout,interval,interval_multiplier,maximum_interval,iterable,result', @@ -28,6 +34,8 @@ ) @pytest.mark.asyncio async def test_aio_timeout_generator( + fake_clock: clock.FakeClock, + monkeypatch: pytest.MonkeyPatch, timeout: float, interval: float, interval_multiplier: float, @@ -36,6 +44,13 @@ async def test_aio_timeout_generator( result: int, ) -> None: """Stop the async generator near the configured timeout.""" + + async def sleep(delay: float) -> None: + """Let the fake clock pass the delay without waiting for it.""" + fake_clock.sleep(delay) + + monkeypatch.setattr(asyncio, 'sleep', sleep) + i = None async for i in python_utils.aio_timeout_generator( timeout, interval, iterable, maximum_interval=maximum_interval @@ -46,12 +61,13 @@ async def test_aio_timeout_generator( @pytest.mark.parametrize( - 'timeout,interval,interval_multiplier,maximum_interval,iterable,result', + 'timeout,interval,interval_multiplier,maximum_interval,iterable,result,' + 'sleeps', [ - (0.1, 0.06, 0.5, 0.1, 'abc', 'c'), - (0.1, 0.07, 0.5, 0.1, itertools.count, 2), - (0.1, 0.07, 0.5, 0.1, itertools.count(), 2), - (0.1, 0.06, 1.0, None, 'abc', 'c'), + (0.1, 0.06, 0.5, 0.1, 'abc', 'c', [0.06, 0.03, 0.015]), + (0.1, 0.07, 0.5, 0.1, itertools.count, 2, [0.07, 0.035]), + (0.1, 0.07, 0.5, 0.1, itertools.count(), 2, [0.07, 0.035]), + (0.1, 0.06, 1.0, None, 'abc', 'c', [0.06, 0.06]), ( datetime.timedelta(seconds=0.1), datetime.timedelta(seconds=0.06), @@ -59,10 +75,12 @@ async def test_aio_timeout_generator( datetime.timedelta(seconds=0.1), itertools.count, 2, + [0.06, 0.1], ), ], ) def test_timeout_generator( + fake_clock: clock.FakeClock, timeout: float, interval: float, interval_multiplier: float, @@ -73,8 +91,9 @@ def test_timeout_generator( types.Callable[..., types.Iterable[types.Any]], ], result: int, + sleeps: types.List[float], ) -> None: - """Stop the sync generator near the configured timeout.""" + """Stop the sync generator at the timeout and scale the interval.""" i = None for i in python_utils.timeout_generator( timeout=timeout, @@ -86,59 +105,97 @@ def test_timeout_generator( assert i is not None assert i == result + assert fake_clock.sleeps == pytest.approx(sleeps) -@pytest.mark.asyncio -async def test_aio_generator_timeout_detector() -> None: - """Raise or exit on per-item and total timeouts.""" - # Make pyright happy - i = None +def test_timeout_generator_real_clock() -> None: + """Keep yielding on the real clock until the timeout has passed.""" + timeout: float = 0.05 + interval: float = 0.01 + start: float = time.perf_counter() + items: types.List[int] = list( + python_utils.timeout_generator(timeout, interval, itertools.count()) + ) + elapsed: float = time.perf_counter() - start + + # A sleep can take longer than requested but never shorter, so these + # hold on any machine. The exact number of items does not. + assert items == list(range(len(items))) + assert len(items) <= timeout / interval + 2 + assert elapsed >= timeout + + +async def stalling_generator() -> types.AsyncGenerator[int, None]: + """Yield 0-4 without waiting, then stall before the next item.""" + for i in range(10): + if i == 5: + await asyncio.sleep(STALL) + yield i + + +def ticking_generator( + fake_clock: clock.FakeClock, +) -> types.AsyncGenerator[int, None]: + """Yield 0-9 and let 0.1 seconds pass on the fake clock for each item.""" async def generator() -> types.AsyncGenerator[int, None]: - """Yield 0-9 with increasing sleeps between items.""" + """Advance the fake clock before every item.""" for i in range(10): - await asyncio.sleep(i / 20.0) + fake_clock.sleep(0.1) yield i + return generator() + + +@pytest.mark.asyncio +async def test_aio_generator_timeout_detector( + fake_clock: clock.FakeClock, +) -> None: + """Raise or exit on per-item and total timeouts.""" + # Make pyright happy + i = None + detector = python_utils.aio_generator_timeout_detector # Test regular timeout with reraise with pytest.raises(asyncio.TimeoutError): - async for i in detector(generator(), 0.25): + async for i in detector(stalling_generator(), 0.05): pass # Test regular timeout with clean exit - async for i in detector(generator(), 0.25, on_timeout=None): + async for i in detector(stalling_generator(), 0.05, on_timeout=None): pass assert i == 4 # Test total timeout with reraise with pytest.raises(asyncio.TimeoutError): - async for i in detector(generator(), total_timeout=0.5): + async for i in detector( + ticking_generator(fake_clock), total_timeout=0.45 + ): pass # Test total timeout with clean exit - async for i in detector(generator(), total_timeout=0.5, on_timeout=None): + async for i in detector( + ticking_generator(fake_clock), total_timeout=0.45, on_timeout=None + ): pass assert i == 4 # Test stop iteration - async for i in detector(generator(), on_timeout=None): + async for i in detector(ticking_generator(fake_clock), on_timeout=None): pass + assert i == 9 + @pytest.mark.asyncio async def test_aio_generator_timeout_detector_decorator_reraise() -> None: """Reraise ``TimeoutError`` on a per-item timeout.""" - # Test regular timeout with reraise - @python_utils.aio_generator_timeout_detector_decorator(timeout=0.05) - async def generator_timeout() -> types.AsyncGenerator[int, None]: - """Yield with increasing delays to trip the timeout.""" - for i in range(10): - await asyncio.sleep(i / 100.0) - yield i + generator_timeout = python_utils.aio_generator_timeout_detector_decorator( + timeout=0.05 + )(stalling_generator) with pytest.raises(asyncio.TimeoutError): async for _ in generator_timeout(): @@ -152,14 +209,9 @@ async def test_aio_generator_timeout_detector_decorator_clean_exit() -> None: i = None # Test regular timeout with clean exit - @python_utils.aio_generator_timeout_detector_decorator( + generator_clean = python_utils.aio_generator_timeout_detector_decorator( timeout=0.05, on_timeout=None - ) - async def generator_clean() -> types.AsyncGenerator[int, None]: - """Yield with increasing delays to trip the timeout.""" - for i in range(10): - await asyncio.sleep(i / 100.0) - yield i + )(stalling_generator) async for i in generator_clean(): pass @@ -168,17 +220,16 @@ async def generator_clean() -> types.AsyncGenerator[int, None]: @pytest.mark.asyncio -async def test_aio_generator_timeout_detector_decorator_reraise_total() -> ( - None -): +async def test_aio_generator_timeout_detector_decorator_reraise_total( + fake_clock: clock.FakeClock, +) -> None: """Reraise ``TimeoutError`` on a total timeout.""" # Test total timeout with reraise - @python_utils.aio_generator_timeout_detector_decorator(total_timeout=0.1) + @python_utils.aio_generator_timeout_detector_decorator(total_timeout=0.45) async def generator_reraise() -> types.AsyncGenerator[int, None]: - """Yield with increasing delays to trip the timeout.""" - for i in range(10): - await asyncio.sleep(i / 100.0) + """Let the fake clock pass the total timeout while yielding.""" + async for i in ticking_generator(fake_clock): yield i with pytest.raises(asyncio.TimeoutError): @@ -187,19 +238,20 @@ async def generator_reraise() -> types.AsyncGenerator[int, None]: @pytest.mark.asyncio -async def test_aio_generator_timeout_detector_decorator_clean_total() -> None: +async def test_aio_generator_timeout_detector_decorator_clean_total( + fake_clock: clock.FakeClock, +) -> None: """Exit cleanly on total timeout when ``on_timeout`` is ``None``.""" # Make pyright happy i = None # Test total timeout with clean exit @python_utils.aio_generator_timeout_detector_decorator( - total_timeout=0.1, on_timeout=None + total_timeout=0.45, on_timeout=None ) async def generator_clean_total() -> types.AsyncGenerator[int, None]: - """Yield with increasing delays to trip the timeout.""" - for i in range(10): - await asyncio.sleep(i / 100.0) + """Let the fake clock pass the total timeout while yielding.""" + async for i in ticking_generator(fake_clock): yield i async for i in generator_clean_total(): diff --git a/conftest.py b/conftest.py new file mode 100644 index 0000000..4a22f5b --- /dev/null +++ b/conftest.py @@ -0,0 +1,8 @@ +""" +Load the shared fixtures for the tests and the doctests. + +The fixtures live in the tests package. They are loaded from the repository +root because the doctests in ``python_utils`` need them as well. +""" + +pytest_plugins: tuple[str, ...] = ('_python_utils_tests.clock',) diff --git a/pyproject.toml b/pyproject.toml index c875f23..ce32e4d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -8,7 +8,7 @@ module-root = '' module-name = 'python_utils' # Keep the tests and tox config in the sdist (parity with the old MANIFEST.in) # so downstream packagers can build and test from the source distribution. -source-include = ['_python_utils_tests/**/*.py', 'tox.ini'] +source-include = ['_python_utils_tests/**/*.py', 'conftest.py', 'tox.ini'] [project] name = 'python-utils' diff --git a/python_utils/containers.py b/python_utils/containers.py index 9318d95..48bc14e 100644 --- a/python_utils/containers.py +++ b/python_utils/containers.py @@ -61,10 +61,13 @@ import abc import collections import collections.abc +import contextlib +import operator import typing if typing.TYPE_CHECKING: import _typeshed # noqa: F401 + import typing_extensions #: A type alias for a type that can be used as a key in a dictionary. KT = typing.TypeVar('KT') @@ -82,20 +85,65 @@ T = typing.TypeVar('T') #: Argument shapes accepted when updating a casted dict: a mapping, an iterable -#: of key/value pairs, an iterable of mappings, or a keys-and-getitem object. +#: of key/value pairs, or a keys-and-getitem object. # Kept as `typing.Union` (not PEP 604 `|`): one member is a string forward # reference, and `|` evaluates its operands eagerly, raising `TypeError` on a # `str` operand at runtime. `typing.Union` accepts it as a lazy `ForwardRef`. DictUpdateArgs = typing.Union[ collections.abc.Mapping[KT, VT], collections.abc.Iterable[tuple[KT, VT]], - collections.abc.Iterable[collections.abc.Mapping[KT, VT]], '_typeshed.SupportsKeysAndGetItem[KT, VT]', ] #: Policy for ``UniqueList`` duplicates: silently ``'ignore'`` or ``'raise'``. OnDuplicate = typing.Literal['ignore', 'raise'] +#: Marks the state that ``CastedDictBase.__reduce_ex__`` writes, to tell it +#: apart from the default state of an instance with slots. +_CASTED_DICT_STATE: str = 'python_utils.containers.CastedDictBase' + + +def _restore_attributes( + instance: object, + state: typing.Any, + inherited: collections.abc.Callable[[typing.Any], None] | None, +) -> None: + """ + Restores the attributes of an instance from its default pickle state. + + The default state is the instance dictionary, or a tuple of that + dictionary and the slot values when the class has slots. `pickle` and + `copy` apply both forms themselves, but only for a class without a + `__setstate__` of its own. + + A `__setstate__` that follows the calling class in the method resolution + order gets the state when there is one. It would have received it if the + calling class did not define `__setstate__`. + + Args: + instance (object): The instance to restore. + state (typing.Any): The default state of the instance. + inherited (Callable[[typing.Any], None] | None): The `__setstate__` + that follows the calling class, or None without one. + """ + if inherited is not None: + inherited(state) + return + + slots: dict[str, typing.Any] | None = None + if isinstance(state, tuple): + state, slots = typing.cast( + 'tuple[dict[str, typing.Any] | None, dict[str, typing.Any]]', + state, + ) + + if state: + vars(instance).update(state) + + if slots: + for name, value in slots.items(): + setattr(instance, name, value) + class CastedDictBase(dict[KT, VT], abc.ABC): """ @@ -119,8 +167,10 @@ class CastedDictBase(dict[KT, VT], abc.ABC): callable is provided. """ - _key_cast: KT_cast[KT] - _value_cast: VT_cast[VT] + # The casts default to `None` on the class so that a dict which `pickle` + # created without calling `__init__` can already store items. + _key_cast: KT_cast[KT] = None + _value_cast: VT_cast[VT] = None def __init__( self, @@ -160,12 +210,73 @@ def update( the dictionary. **kwargs (typing.Any): Keyword arguments to update the dictionary. """ - if args: - kwargs.update(*args) + # Keyword arguments are applied last so that they win, as they do + # for `dict.update`. + positional: dict[typing.Any, typing.Any] = {} + positional.update(*args) + for key, value in positional.items(): + self[key] = value + + for key, value in kwargs.items(): + self[key] = value + + def setdefault(self, key: typing.Any, default: typing.Any = None) -> VT: + """ + Stores the default through the casts when the key is missing. + + A default of `None` is stored as `None` without casting it, so that + `setdefault(key)` works for every value cast. + + Args: + key (typing.Any): The key to look up, before casting. + default (typing.Any, optional): The value to store when the key + is missing. Defaults to None. + + Returns: + VT: The value that is stored under the key. + """ + if self._key_cast is not None: + key = self._key_cast(key) + + if not super().__contains__(key): + stored: typing.Any = default + if default is not None: + stored = self._cast_stored_value(default) + + super().__setitem__(key, stored) + + return super().__getitem__(key) + + def _cast_stored_value(self, value: typing.Any) -> typing.Any: + """ + Casts a value on its way into the dictionary. + + The base class stores values as they are. A subclass that casts when + it stores a value overrides this method. + + Args: + value (typing.Any): The value to store. + + Returns: + typing.Any: The value as it is stored. + """ + return value + + def __ior__( # type: ignore[override, misc] + self, other: DictUpdateArgs[typing.Any, typing.Any] + ) -> 'typing_extensions.Self': + """ + Updates the dictionary in place through the casts, see `update`. - if kwargs: - for key, value in kwargs.items(): - self[key] = value + Args: + other (DictUpdateArgs[typing.Any, typing.Any]): The items to + merge into the dictionary. + + Returns: + typing_extensions.Self: The dictionary itself. + """ + self.update(other) + return self def __setitem__(self, key: typing.Any, value: typing.Any) -> None: """ @@ -181,6 +292,81 @@ def __setitem__(self, key: typing.Any, value: typing.Any) -> None: return super().__setitem__(key, value) + def __reduce_ex__(self, protocol: typing.SupportsIndex) -> typing.Any: + """ + Describes the dictionary for `pickle` and `copy` with its raw items. + + By default the items are restored through `__setitem__`. `pickle` + does that before the casts are back and `copy` does it after, which + casts every value a second time. The items travel inside the state + to avoid both. + + The default description is returned unchanged where it already + works or where other code relies on its form, see the comment below. + + Args: + protocol (typing.SupportsIndex): The pickle protocol. + + Returns: + typing.Any: The default description, with the items moved into + the state for protocol 2 and up. + """ + reduced: typing.Any = super().__reduce_ex__(protocol) + if ( + operator.index(protocol) < 2 + or len(self) == 0 + or type(self).__reduce__ is not object.__reduce__ + or type(self).__reduce_ex__ is not _BASE_REDUCE_EX + or type(self).__setstate__ is not _BASE_SETSTATE + ): + # Protocol 0 and 1 restore the items through `dict` itself and + # an empty dictionary has none. A subclass with its own + # `__reduce__` describes itself. One with its own `__reduce_ex__` + # or `__setstate__` expects the default state and items. + return reduced + + state: tuple[str, typing.Any, dict[KT, VT]] = ( + _CASTED_DICT_STATE, + reduced[2], + super().copy(), + ) + return reduced[0], reduced[1], state, None, None + + def __setstate__(self, state: object) -> None: + """ + Restores the attributes and the raw items without casting them. + + Args: + state (object): The state from `__reduce_ex__`, or the default + state for a pickle that already restored its items. + """ + parts: tuple[object, ...] = () + if isinstance(state, tuple): + parts = typing.cast('tuple[object, ...]', state) + + # A mixin after this class can have a `__setstate__` of its own. + inherited: typing.Any = getattr(super(), '__setstate__', None) + if len(parts) == 3 and parts[0] == _CASTED_DICT_STATE: + _restore_attributes(self, parts[1], inherited) + super().update(typing.cast('dict[KT, VT]', parts[2])) + else: + _restore_attributes(self, state, inherited) + + +#: The `__setstate__` that understands the state of +#: `CastedDictBase.__reduce_ex__`. A subclass that replaces it gets the default +#: state. +# Both are taken from the class dictionary. `CastedDictBase[...].__reduce_ex__` +# would be the method of the generic alias itself. +_BASE_SETSTATE: collections.abc.Callable[..., None] = vars(CastedDictBase)[ + '__setstate__' +] +#: The `__reduce_ex__` that writes that state. A subclass that replaces it +#: gets the default description to build on. +_BASE_REDUCE_EX: collections.abc.Callable[..., typing.Any] = vars( + CastedDictBase +)['__reduce_ex__'] + class CastedDict(CastedDictBase[KT, VT]): """ @@ -223,10 +409,22 @@ def __setitem__(self, key: typing.Any, value: typing.Any) -> None: The key itself is cast by ``CastedDictBase.__setitem__`` when a key cast is configured. """ + super().__setitem__(key, self._cast_stored_value(value)) + + def _cast_stored_value(self, value: typing.Any) -> typing.Any: + """ + Casts a value on its way into the dictionary. + + Args: + value (typing.Any): The value to store. + + Returns: + typing.Any: The value, cast when a value cast is set. + """ if self._value_cast is not None: value = self._value_cast(value) - super().__setitem__(key, value) + return value class LazyCastedDict(CastedDictBase[KT, VT]): @@ -278,15 +476,14 @@ class LazyCastedDict(CastedDictBase[KT, VT]): def __setitem__(self, key: typing.Any, value: typing.Any) -> None: """ Sets the item in the dictionary, casting the key if a key cast - callable is provided. + callable is provided. The value is stored as it is. Args: key (typing.Any): The key to set in the dictionary. value (typing.Any): The value to set in the dictionary. """ - if self._key_cast is not None: - key = self._key_cast(key) - + # The base class casts the key. Casting it here as well would cast + # it twice. super().__setitem__(key, value) def __getitem__(self, key: typing.Any) -> VT: @@ -357,6 +554,16 @@ class UniqueList(list[HT]): >>> l [5, 10, 2, 3, 4] + A value that was removed or replaced can be added again: + + >>> l = UniqueList(1, 2, 3) + >>> l[0] = 4 + >>> l.pop() + 3 + >>> l.extend([1, 2, 3]) + >>> l + [4, 2, 1, 3] + >>> l = UniqueList(1, 2, 3, on_duplicate='raise') >>> l.append(4) >>> l.append(4) @@ -378,6 +585,30 @@ class UniqueList(list[HT]): """ _set: set[HT] + #: The default lives on the class, so that a subclass can set its own + #: and a list that `pickle` created without `__init__` has one. + on_duplicate: OnDuplicate = 'ignore' + + def __new__( + cls, *args: typing.Any, **kwargs: typing.Any + ) -> 'typing_extensions.Self': + """ + Creates the list with an empty membership set. + + `pickle` and `copy` create the list through `__new__` and refill it + through `extend` or `append` without calling `__init__`, so the + membership set has to exist before `__init__` runs. + + Args: + *args (typing.Any): Ignored, handled by `__init__`. + **kwargs (typing.Any): Ignored, handled by `__init__`. + + Returns: + typing_extensions.Self: The new, empty list. + """ + instance: typing_extensions.Self = super().__new__(cls) + instance._set = set() + return instance def __init__( self, @@ -398,6 +629,44 @@ def __init__( for arg in args: self.append(arg) + def __setstate__(self, state: typing.Any) -> None: + """ + Restores the attributes and rebuilds the membership from the items. + + `copy` restores the attributes before the items and `pickle` restores + them after the items. A stored membership set is only right in the + second case, so the set is always derived from the items that are in + the list at this point. + + Args: + state (typing.Any): The instance attributes, with the slot values + when a subclass has slots. + """ + # A mixin after this class can have a `__setstate__` of its own. + inherited: typing.Any = getattr(super(), '__setstate__', None) + _restore_attributes(self, state, inherited) + self._set = set(self) + + def _release(self, value: HT) -> None: + """ + Drops a value that left the list from the membership set. + + The set is rebuilt from the list when the value cannot be removed + from it. That happens when the value changed its hash after it was + added, and when its `__hash__` or `__eq__` raises. + + The value has left the list by now, so nothing here may raise. When + the rebuild fails as well the set stays as it is. + + Args: + value (HT): The value that was removed from the list. + """ + try: + self._set.remove(value) + except Exception: # noqa: BLE001 + with contextlib.suppress(Exception): + self._set = set(self) + def insert(self, index: typing.SupportsIndex, value: HT) -> None: """ Inserts a value at the specified index, ensuring uniqueness. @@ -416,8 +685,15 @@ def insert(self, index: typing.SupportsIndex, value: HT) -> None: else: return + # The membership first, right after the check above. Another thread + # that inserts the same value in between would get past its own + # check as well. A failed insert takes the membership back. self._set.add(value) - super().insert(index, value) + try: + super().insert(index, value) + except BaseException: + self._set.discard(value) + raise def append(self, value: HT) -> None: """ @@ -439,6 +715,133 @@ def append(self, value: HT) -> None: self._set.add(value) super().append(value) + def extend(self, values: collections.abc.Iterable[HT]) -> None: + """ + Extends the list with the values that are not in it yet. + + Args: + values (Iterable[HT]): The values to append. + + Raises: + ValueError: If `on_duplicate` is set to 'raise' and a value is + already in the list or occurs more than once in `values`. + The list is left unchanged in that case. + """ + new_values: list[HT] = list(values) + if self.on_duplicate == 'raise': + duplicates: set[HT] = self._find_duplicates(new_values) + if duplicates: + raise ValueError(f'Duplicate values: {duplicates}') + + for value in new_values: + self.append(value) + + # `list.__iadd__` accepts any iterable while `list.__add__` only accepts + # a list. Typeshed ignores the same mismatch. + def __iadd__( # type: ignore[misc, override] + self, values: collections.abc.Iterable[HT] + ) -> 'typing_extensions.Self': + """ + Extends the list in place, see `extend`. + + Args: + values (Iterable[HT]): The values to append. + + Returns: + typing_extensions.Self: The list itself. + """ + self.extend(values) + return self + + def __imul__( + self, value: typing.SupportsIndex + ) -> 'typing_extensions.Self': + """ + Multiplies the list in place without ever repeating an item. + + A count below 1 empties the list, as it does for a regular list. A + count above 1 would repeat every item, so the repeat goes through + `extend` and follows `on_duplicate`. + + Args: + value (typing.SupportsIndex): The number of times to repeat. + + Returns: + typing_extensions.Self: The list itself. + + Raises: + ValueError: If `on_duplicate` is set to 'raise' and a non-empty + list is multiplied by more than 1. + """ + count: int = operator.index(value) + if count < 1: + self.clear() + elif count > 1: + self.extend(self) + + return self + + def pop(self, index: typing.SupportsIndex = -1) -> HT: + """ + Removes and returns the item at the given index. + + Args: + index (typing.SupportsIndex, optional): The index to pop. + Defaults to the last item. + + Returns: + HT: The removed item. + """ + value: HT = super().pop(index) + self._release(value) + return value + + def remove(self, value: HT) -> None: + """ + Removes a value from the list. + + Args: + value (HT): The value to remove. + + Raises: + ValueError: If the value is not in the list. + """ + # One list operation, so that another thread cannot get between + # finding the value and removing it. + super().remove(value) + self._release(value) + + def clear(self) -> None: + """Removes all items from the list.""" + super().clear() + self._set.clear() + + def _find_duplicates( + self, + values: list[HT], + replaced: collections.abc.Set[HT] = frozenset(), + ) -> set[HT]: + """ + Finds the values that would break uniqueness when added. + + Args: + values (list[HT]): The values to add. + replaced (Set[HT], optional): Items that leave the list in the + same operation, so that adding them again is allowed. + + Returns: + set[HT]: The values that occur more than once in `values` or + that are already in the list and not in `replaced`. + """ + seen: set[HT] = set() + duplicates: set[HT] = set() + for value in values: + if value in seen or (value in self._set and value not in replaced): + duplicates.add(value) + seen.add(value) + + return duplicates + def __contains__(self, item: HT) -> bool: # type: ignore[override] """ Checks if the list contains the specified item. @@ -449,7 +852,12 @@ def __contains__(self, item: HT) -> bool: # type: ignore[override] Returns: bool: True if the item is in the list, False otherwise. """ - return item in self._set + try: + return item in self._set + except TypeError: + # An unhashable item cannot be in the set, but it can still be + # equal to an item in the list. + return super().__contains__(item) @typing.overload def __setitem__(self, indices: typing.SupportsIndex, values: HT) -> None: @@ -480,31 +888,93 @@ def __setitem__( set to 'raise'. """ if isinstance(indices, slice): - values = typing.cast(collections.abc.Iterable[HT], values) - if self.on_duplicate == 'ignore': - raise RuntimeError( - 'ignore mode while setting slices introduces ambiguous ' - 'behaviour and is therefore not supported' - ) - - duplicates: set[HT] = set(values) & self._set - if duplicates and values != list(self[indices]): - raise ValueError(f'Duplicate values: {duplicates}') - - self._set.update(values) + self._set_slice( + indices, typing.cast(collections.abc.Iterable[HT], values) + ) else: - values = typing.cast(HT, values) - if values in self._set and values != self[indices]: - if self.on_duplicate == 'raise': - raise ValueError(f'Duplicate value: {values}') - else: - return + self._set_index(indices, typing.cast(HT, values)) + + def _set_slice( + self, indices: slice, values: collections.abc.Iterable[HT] + ) -> None: + """ + Replaces a slice of the list and keeps the membership in sync. + + The new values can reuse the items they replace. They cannot repeat + each other or an item that stays in the list. - self._set.add(values) + Args: + indices (slice): The slice to replace. + values (Iterable[HT]): The values to store. + + Raises: + RuntimeError: If `on_duplicate` is 'ignore'. + ValueError: If storing the values would create a duplicate. + """ + if self.on_duplicate == 'ignore': + raise RuntimeError( + 'ignore mode while setting slices introduces ambiguous ' + 'behaviour and is therefore not supported' + ) - super().__setitem__( - typing.cast(slice, indices), typing.cast(list[HT], values) + new_values: list[HT] = list(values) + old_values: list[HT] = self[indices] + duplicates: set[HT] = self._find_duplicates( + new_values, replaced=set(old_values) ) + if duplicates: + raise ValueError(f'Duplicate values: {duplicates}') + + # The membership first, for the same reason as in `insert`. + fresh: set[HT] = set(new_values).difference(self._set) + self._set.update(fresh) + try: + super().__setitem__(indices, new_values) + except BaseException: + self._set.difference_update(fresh) + raise + + # Only what really left the list. A value that the slice reuses + # stays a member the whole time. + self._set.difference_update(set(old_values).difference(new_values)) + + def _set_index(self, index: typing.SupportsIndex, value: HT) -> None: + """ + Replaces a single item and keeps the membership in sync. + + Args: + index (typing.SupportsIndex): The index to replace. + value (HT): The value to store. + + Raises: + ValueError: If the value is a duplicate of another item and + `on_duplicate` is set to 'raise'. + """ + old_value: HT = self[index] + known: bool = value in self._set + if known and value != old_value: + if self.on_duplicate == 'raise': + raise ValueError(f'Duplicate value: {value}') + else: + return + + # The membership first, for the same reason as in `insert`. + self._set.add(value) + try: + super().__setitem__(index, value) + except BaseException: + if not known: + self._set.discard(value) + + raise + + if not known: + # A known value replaces itself and stays a member the whole + # time. For a new value the old one leaves. + self._release(old_value) + # The release dropped the new value as well when it is the old + # object with a changed hash. + self._set.add(value) def __delitem__(self, index: typing.SupportsIndex | slice) -> None: """ @@ -514,13 +984,16 @@ def __delitem__(self, index: typing.SupportsIndex | slice) -> None: index (typing.SupportsIndex | slice): The index or slice to delete the item(s) at. """ + removed: list[HT] if isinstance(index, slice): - for value in self[index]: - self._set.remove(value) + removed = self[index] else: - self._set.remove(self[index]) + removed = [self[index]] + # The list first, then the membership, like every other mutator. super().__delitem__(index) + for value in removed: + self._release(value) # Type hinting `collections.deque` does not work consistently between Python @@ -595,10 +1068,33 @@ def __eq__(self, other: typing.Any) -> bool: elif isinstance(other, tuple): return tuple(self) == other elif isinstance(other, set): - return set(self) == other + try: + return set(self) == other + except TypeError: + # An unhashable item cannot be in a set. + return False else: return super().__eq__(other) + def __ne__(self, other: typing.Any) -> bool: + """ + Checks inequality as the opposite of `__eq__`. + + `collections.deque` has its own `__ne__`, which does not know about + the lists, tuples and sets that `__eq__` accepts. + + Args: + other (typing.Any): The object to compare with. + + Returns: + bool: False if the objects are equal, True otherwise. + """ + equal: typing.Any = self.__eq__(other) + if equal is NotImplemented: + return NotImplemented + + return not equal + def pop(self, index: int = -1) -> T: """ Removes and returns the item at the given index. Only supports index 0