diff --git a/datasketch/hnsw.py b/datasketch/hnsw.py index 9fecf693..ac2fc9dd 100644 --- a/datasketch/hnsw.py +++ b/datasketch/hnsw.py @@ -446,7 +446,7 @@ def setdefault(self, key: Hashable, default: np.ndarray) -> np.ndarray: raise ValueError("Default value cannot be None.") if key not in self._nodes or self._nodes[key].is_deleted: self.insert(key, default) - return self._nodes[key] + return self._nodes[key].point def insert( self, diff --git a/test/test_hnsw.py b/test/test_hnsw.py index e437fffb..74e39b4c 100644 --- a/test/test_hnsw.py +++ b/test/test_hnsw.py @@ -4,7 +4,7 @@ import numpy as np import pytest -from datasketch.hnsw import HNSW +from datasketch.hnsw import HNSW, _Node from datasketch.minhash import MinHash @@ -132,6 +132,35 @@ def test_copy(self): self.assertTrue(0 not in hnsw) self.assertTrue(0 in hnsw2) + def test_setdefault(self): + # setdefault must return the point (the mapping value), matching + # MutableMapping semantics and its documented "-> np.ndarray" contract, + # not the internal _Node wrapper. + data = self._create_random_points(n=12) + hnsw = self._create_index(data[:10]) + + # Existing, non-deleted key: return its point without overwriting. + returned = hnsw.setdefault(0, data[10]) + self.assertNotIsInstance(returned, _Node) + self.assertTrue(np.array_equal(returned, data[0])) + self.assertTrue(np.array_equal(hnsw[0], data[0])) + + # Absent key: insert the default and return it. + returned = hnsw.setdefault(100, data[10]) + self.assertNotIsInstance(returned, _Node) + self.assertTrue(np.array_equal(returned, data[10])) + self.assertIn(100, hnsw) + self.assertTrue(np.array_equal(hnsw[100], data[10])) + + # Soft-removed key: reinsert the default and return it. + hnsw.remove(1) + self.assertNotIn(1, hnsw) + returned = hnsw.setdefault(1, data[11]) + self.assertNotIsInstance(returned, _Node) + self.assertTrue(np.array_equal(returned, data[11])) + self.assertIn(1, hnsw) + self.assertTrue(np.array_equal(hnsw[1], data[11])) + def test_soft_remove_and_pop_and_clean(self): data = self._create_random_points() hnsw = self._create_index(data)