diff --git a/deepdiff/deephash.py b/deepdiff/deephash.py index e0d60da2..0546a079 100644 --- a/deepdiff/deephash.py +++ b/deepdiff/deephash.py @@ -100,6 +100,10 @@ class BoolObj(Enum): FALSE = 0 +class _NumberHashKey: + """Pickle-safe namespace for numeric cache keys, separate from user tuples.""" + + def prepare_string_for_hashing( obj: Union[str, bytes, memoryview], ignore_string_type_changes: bool = False, @@ -382,9 +386,9 @@ def get_key(hashes: Dict[Any, Any], key: Any, default: Any = None, extract_index @staticmethod def _unwrap_hash_key(key: Any) -> Any: - """Unwrap a (type, value) hash key back to the original value for public API.""" - if isinstance(key, tuple) and len(key) == 2 and isinstance(key[0], type) and isinstance(key[1], only_numbers): - return key[1] + """Unwrap an internal numeric hash key for the public API.""" + if isinstance(key, tuple) and len(key) == 3 and key[0] is _NumberHashKey: + return key[2] return key def _get_objects_to_hashes_dict(self, extract_index: Optional[int] = 0) -> Dict[Any, Any]: @@ -619,18 +623,18 @@ def _make_hash_key(self, obj: Any) -> Any: In Python, 1 == 1.0 and hash(1) == hash(1.0), so int and float values collide as dict keys. When ignore_numeric_type_changes is False, we wrap - numeric objects as (type, value) tuples so that each type gets its own - cache entry and its own hash. + numeric objects as (_NumberHashKey, type, value) tuples so that each type + gets its own cache entry without colliding with ordinary (type, value) tuples. """ if not self.ignore_numeric_type_changes and isinstance(obj, only_numbers): - return (type(obj), obj) + return (_NumberHashKey, type(obj), obj) return obj @staticmethod def _make_hash_key_for_lookup(obj: Any, ignore_numeric_type_changes: bool = False) -> Any: """Static version of _make_hash_key for use in static accessor methods.""" if not ignore_numeric_type_changes and isinstance(obj, only_numbers): - return (type(obj), obj) + return (_NumberHashKey, type(obj), obj) return obj def _hash(self, obj: Any, parent: str, parents_ids: frozenset = EMPTY_FROZENSET) -> HashTuple: diff --git a/deepdiff/serialization.py b/deepdiff/serialization.py index 07be29bd..7ab7df92 100644 --- a/deepdiff/serialization.py +++ b/deepdiff/serialization.py @@ -91,6 +91,7 @@ class UnsupportedFormatErr(TypeError): 'collections.OrderedDict', 're.Pattern', 'deepdiff.helper.Opcode', + 'deepdiff.deephash._NumberHashKey', 'ipaddress.IPv4Interface', 'ipaddress.IPv6Interface', 'ipaddress.IPv4Network', diff --git a/tests/test_hash.py b/tests/test_hash.py index f1e2e912..57c1d035 100755 --- a/tests/test_hash.py +++ b/tests/test_hash.py @@ -6,6 +6,9 @@ import logging import datetime import ipaddress +import pickle +from copy import deepcopy +from decimal import Decimal from typing import Union from pathlib import Path from collections import namedtuple @@ -16,6 +19,7 @@ prepare_string_for_hashing, unprocessed, UNPROCESSED_KEY, BoolObj, HASH_LOOKUP_ERR_MSG, combine_hashes_lists) from deepdiff.helper import pypy3, get_id, number_to_string, np, py_major_version, py_minor_version +from deepdiff.serialization import pickle_dump, pickle_load from tests import CustomClass2 logging.disable(logging.CRITICAL) @@ -404,6 +408,51 @@ def test_number_type_change(self): result2 = DeepHashPrep(obj2, ignore_numeric_type_changes=True) assert result1[obj1] == result2[obj2] + @pytest.mark.parametrize('number', [1, 1.0, 1j, Decimal('1'), np.int64(1)]) + @pytest.mark.parametrize('reverse', [False, True]) + def test_numeric_cache_key_does_not_collide_with_tuple(self, number, reverse): + item = (type(number), number) + objects = [item, number] if reverse else [number, item] + result = DeepHash(objects) + number_hash = DeepHash(number)[number] + tuple_hash = DeepHash(item)[item] + + assert result[number] == number_hash + assert result[item] == tuple_hash + assert result[number] != result[item] + assert result.get(item) == tuple_hash + assert DeepHash.get_key(result.hashes, number) == number_hash + assert DeepHash.get_key(result.hashes, item) == tuple_hash + assert number in result + assert item in result + assert item in set(result.keys()) + assert dict(result.items())[item] == tuple_hash + assert result._get_objects_to_hashes_dict()[item] == tuple_hash + + def test_numeric_cache_does_not_contain_unhashed_tuple(self): + result = DeepHash(1) + item = (int, 1) + + assert item not in result + assert result.get(item) is None + with pytest.raises(KeyError): + result[item] + + @pytest.mark.parametrize('copy_cache', [ + deepcopy, + pytest.param(lambda value: pickle.loads(pickle.dumps(value)), id='pickle'), + pytest.param(lambda value: pickle_load(pickle_dump(value)), id='restricted_pickle'), + ]) + def test_numeric_cache_key_survives_copy(self, copy_cache): + original = DeepHash([1, 1.0]) + hashes = copy_cache(original.hashes) + + for item in [1, 1.0]: + assert DeepHash.get_key(hashes, item) == original[item] + result = DeepHash((int, 1), hashes=hashes) + assert result[1] != result[(int, 1)] + assert result[1] != result[1.0] + def test_prep_str_fail_if_deephash_leaks_results(self): """ This test fails if DeepHash is getting a mutable copy of hashes diff --git a/tests/test_ignore_order.py b/tests/test_ignore_order.py index aaceeb71..8dfeac64 100644 --- a/tests/test_ignore_order.py +++ b/tests/test_ignore_order.py @@ -1430,6 +1430,23 @@ def test_int_vs_float_in_list_of_dicts(self): assert result["type_changes"]["root[0]['a']"]["old_type"] is int assert result["type_changes"]["root[0]['a']"]["new_type"] is float + @pytest.mark.parametrize('number', [1, 1.0, 1j]) + @pytest.mark.parametrize('reverse', [False, True]) + def test_numeric_cache_key_does_not_hide_tuple_change(self, number, reverse): + item = (type(number), number) + t1, t2 = ([item], [number]) if reverse else ([number], [item]) + + assert DeepDiff(t1, t2, ignore_order=True) == { + 'type_changes': { + 'root[0]': { + 'old_type': type(t1[0]), + 'new_type': type(t2[0]), + 'old_value': t1[0], + 'new_value': t2[0], + }, + }, + } + def test_ignore_numeric_type_changes_suppresses_report(self): """When ignore_numeric_type_changes=True the type change must be hidden.""" result = DeepDiff(