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
2 changes: 1 addition & 1 deletion datasketch/hnsw.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
31 changes: 30 additions & 1 deletion test/test_hnsw.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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)
Expand Down