Skip to content

Commit 1a469a0

Browse files
authored
Merge pull request #166 from weaviate/rod/hfresh-followup
Fix distance_metric enum conversion and add hfresh unit tests
2 parents af49288 + 31d4998 commit 1a469a0

3 files changed

Lines changed: 156 additions & 38 deletions

File tree

.claude/skills/operating-weaviate-cli/references/collections.md

Lines changed: 23 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -26,7 +26,7 @@ weaviate-cli create collection \
2626
- `--collection` -- Name (default: "Movies")
2727
- `--replication_factor` -- Number of replicas (default: 3)
2828
- `--async_enabled` -- Enable async replication
29-
- `--vector_index` -- Index type: hnsw, flat, dynamic, hnsw_pq, hnsw_bq, hnsw_sq, hnsw_rq, hnsw_acorn, hnsw_multivector, flat_bq, dynamic_*
29+
- `--vector_index` -- Index type: hnsw, flat, dynamic, hnsw_pq, hnsw_bq, hnsw_sq, hnsw_rq, hnsw_acorn, hnsw_multivector, flat_bq, dynamic_*, hfresh
3030
- `--inverted_index` -- Inverted index: timestamp, null, length
3131
- `--training_limit` -- PQ/SQ training limit (default: 10000)
3232
- `--multitenant` -- Enable multi-tenancy
@@ -43,6 +43,28 @@ weaviate-cli create collection \
4343
- `--object_ttl_time` -- Time to live in seconds (default: None, TTL disabled when omitted)
4444
- `--object_ttl_filter_expired` -- Filter expired-but-not-yet-deleted objects from queries
4545
- `--object_ttl_property_name` -- Date property name for TTL when `object_ttl_type=property` (default: "releaseDate"). **Only valid when `--object_ttl_type=property`**; rejected otherwise.
46+
- `--hfresh_max_posting_size_kb` -- (hfresh only) Max posting list size in KB (default: None, uses server default)
47+
- `--hfresh_replicas` -- (hfresh only) Number of replicas per element across posting lists (default: None, uses server default)
48+
- `--hfresh_search_probe` -- (hfresh only) Search probe size (default: None, uses server default)
49+
- `--distance_metric` -- Distance metric: cosine, dot, l2-squared, hamming, manhattan (default: None, uses server default). Applies to all vector index types.
50+
- `--rescore_limit` -- Rescore limit for quantized indexes (default: None, uses server default)
51+
52+
**hfresh examples:**
53+
```bash
54+
# Basic hfresh collection
55+
weaviate-cli create collection --collection Movies --vector_index hfresh --json
56+
57+
# hfresh with all tuning parameters
58+
weaviate-cli create collection \
59+
--collection Movies \
60+
--vector_index hfresh \
61+
--hfresh_max_posting_size_kb 64 \
62+
--hfresh_replicas 2 \
63+
--hfresh_search_probe 100 \
64+
--distance_metric cosine \
65+
--rescore_limit 200 \
66+
--json
67+
```
4668

4769
**Object TTL examples:**
4870
```bash

test/unittests/test_managers/test_collection_manager.py

Lines changed: 83 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -777,3 +777,86 @@ def test_update_collection_with_ttl_disable_type(mock_client, mock_wvc_object_tt
777777
assert update_call_kwargs["object_ttl_config"] is not None
778778
# Verify the disable method was called
779779
mock_wvc_object_ttl["reconfigure"].disable.assert_called_once()
780+
781+
782+
def test_create_collection_with_hfresh_defaults(mock_client, mock_wvc_object_ttl):
783+
"""Test creating a collection with hfresh vector index using default parameters."""
784+
mock_collections = MagicMock()
785+
mock_client.collections = mock_collections
786+
mock_collections.exists.side_effect = [False, True]
787+
788+
manager = CollectionManager(mock_client)
789+
790+
manager.create_collection(
791+
collection="TestCollection",
792+
vector_index="hfresh",
793+
)
794+
795+
mock_collections.create.assert_called_once()
796+
create_call_kwargs = mock_collections.create.call_args.kwargs
797+
assert create_call_kwargs["name"] == "TestCollection"
798+
assert create_call_kwargs["vector_index_config"] is not None
799+
800+
801+
def test_create_collection_with_hfresh_all_params(mock_client, mock_wvc_object_ttl):
802+
"""Test creating a collection with hfresh vector index with all parameters set."""
803+
mock_collections = MagicMock()
804+
mock_client.collections = mock_collections
805+
mock_collections.exists.side_effect = [False, True]
806+
807+
manager = CollectionManager(mock_client)
808+
809+
manager.create_collection(
810+
collection="TestCollection",
811+
vector_index="hfresh",
812+
hfresh_max_posting_size_kb=64,
813+
hfresh_replicas=2,
814+
hfresh_search_probe=100,
815+
distance_metric="cosine",
816+
rescore_limit=200,
817+
)
818+
819+
mock_collections.create.assert_called_once()
820+
create_call_kwargs = mock_collections.create.call_args.kwargs
821+
assert create_call_kwargs["name"] == "TestCollection"
822+
assert create_call_kwargs["vector_index_config"] is not None
823+
824+
825+
def test_create_collection_with_hfresh_valid_distance_metrics(
826+
mock_client, mock_wvc_object_ttl
827+
):
828+
"""Test creating an hfresh collection with each valid distance metric."""
829+
valid_metrics = ["cosine", "dot", "l2-squared", "hamming", "manhattan"]
830+
for metric in valid_metrics:
831+
mock_collections = MagicMock()
832+
mock_client.collections = mock_collections
833+
mock_collections.exists.side_effect = [False, True]
834+
835+
manager = CollectionManager(mock_client)
836+
manager.create_collection(
837+
collection="TestCollection",
838+
vector_index="hfresh",
839+
distance_metric=metric,
840+
)
841+
mock_collections.create.assert_called_once()
842+
843+
844+
def test_create_collection_with_hfresh_invalid_distance_metric(
845+
mock_client, mock_wvc_object_ttl
846+
):
847+
"""Test that an unsupported distance metric raises ValueError."""
848+
mock_collections = MagicMock()
849+
mock_client.collections = mock_collections
850+
mock_collections.exists.return_value = False
851+
852+
manager = CollectionManager(mock_client)
853+
854+
with pytest.raises(ValueError) as exc_info:
855+
manager.create_collection(
856+
collection="TestCollection",
857+
vector_index="hfresh",
858+
distance_metric="invalid_metric",
859+
)
860+
861+
assert "Invalid distance_metric: 'invalid_metric'" in str(exc_info.value)
862+
mock_collections.create.assert_not_called()

weaviate_cli/managers/collection_manager.py

Lines changed: 50 additions & 37 deletions
Original file line numberDiff line numberDiff line change
@@ -147,35 +147,42 @@ def _print_text():
147147
def get_all_collections(self) -> dict[str, _CollectionConfigSimple]:
148148
return self.client.collections.list_all()
149149

150+
_DISTANCE_METRIC_MAP = {
151+
"cosine": wvc.VectorDistances.COSINE,
152+
"dot": wvc.VectorDistances.DOT,
153+
"l2-squared": wvc.VectorDistances.L2_SQUARED,
154+
"hamming": wvc.VectorDistances.HAMMING,
155+
"manhattan": wvc.VectorDistances.MANHATTAN,
156+
}
157+
158+
def _resolve_distance_metric(
159+
self, distance_metric: Optional[str]
160+
) -> Optional[wvc.VectorDistances]:
161+
"""Convert a distance metric string to its VectorDistances enum value."""
162+
if distance_metric is None:
163+
return None
164+
if distance_metric not in self._DISTANCE_METRIC_MAP:
165+
raise ValueError(
166+
f"Invalid distance_metric: '{distance_metric}'. "
167+
f"Must be one of: {list(self._DISTANCE_METRIC_MAP.keys())}"
168+
)
169+
return self._DISTANCE_METRIC_MAP[distance_metric]
170+
150171
def _build_hfresh_config(
151172
self,
152173
max_posting_size_kb: Optional[int] = None,
153-
distance_metric: Optional[str] = "cosine",
174+
distance_metric: Optional[wvc.VectorDistances] = None,
154175
rescore_limit: Optional[int] = None,
155176
replicas: Optional[int] = None,
156177
search_probe: Optional[int] = None,
157178
):
158179
"""Build hfresh configuration with provided parameters."""
159-
# Explicit mapping of distance metric strings to enum values
160-
distance_metric_map = {
161-
"cosine": wvc.VectorDistances.COSINE,
162-
"dot": wvc.VectorDistances.DOT,
163-
"l2-squared": wvc.VectorDistances.L2_SQUARED,
164-
"hamming": wvc.VectorDistances.HAMMING,
165-
"manhattan": wvc.VectorDistances.MANHATTAN,
166-
}
167-
168180
kwargs = {}
169181

170182
if max_posting_size_kb is not None:
171183
kwargs["max_posting_size_kb"] = max_posting_size_kb
172184
if distance_metric is not None:
173-
if distance_metric not in distance_metric_map:
174-
raise ValueError(
175-
f"Invalid distance_metric: '{distance_metric}'. "
176-
f"Must be one of: {list(distance_metric_map.keys())}"
177-
)
178-
kwargs["distance_metric"] = distance_metric_map[distance_metric]
185+
kwargs["distance_metric"] = distance_metric
179186
if replicas is not None:
180187
kwargs["replicas"] = replicas
181188
if search_probe is not None:
@@ -248,56 +255,62 @@ def create_collection(
248255
"Error: Named vector name is only supported with named vectors. Please use --named_vector to enable named vectors."
249256
)
250257

258+
distance_metric_enum = self._resolve_distance_metric(distance_metric)
259+
251260
vector_index_map: Dict[str, wvc.VectorIndexConfig] = {
252-
"hnsw": wvc.Configure.VectorIndex.hnsw(distance_metric=distance_metric),
253-
"flat": wvc.Configure.VectorIndex.flat(distance_metric=distance_metric),
261+
"hnsw": wvc.Configure.VectorIndex.hnsw(
262+
distance_metric=distance_metric_enum
263+
),
264+
"flat": wvc.Configure.VectorIndex.flat(
265+
distance_metric=distance_metric_enum
266+
),
254267
"dynamic": wvc.Configure.VectorIndex.dynamic(),
255268
"dynamic_flat_bq": wvc.Configure.VectorIndex.dynamic(
256269
flat=wvc.Configure.VectorIndex.flat(
257270
quantizer=wvc.Configure.VectorIndex.Quantizer.bq(),
258-
distance_metric=distance_metric,
271+
distance_metric=distance_metric_enum,
259272
)
260273
),
261274
"dynamic_flat_bq_hnsw_pq": wvc.Configure.VectorIndex.dynamic(
262275
flat=wvc.Configure.VectorIndex.flat(
263276
quantizer=wvc.Configure.VectorIndex.Quantizer.bq(
264277
rescore_limit=rescore_limit
265278
),
266-
distance_metric=distance_metric,
279+
distance_metric=distance_metric_enum,
267280
),
268281
hnsw=wvc.Configure.VectorIndex.hnsw(
269282
quantizer=wvc.Configure.VectorIndex.Quantizer.pq(
270283
training_limit=training_limit
271284
),
272-
distance_metric=distance_metric,
285+
distance_metric=distance_metric_enum,
273286
),
274287
),
275288
"dynamic_flat_bq_hnsw_sq": wvc.Configure.VectorIndex.dynamic(
276289
flat=wvc.Configure.VectorIndex.flat(
277290
quantizer=wvc.Configure.VectorIndex.Quantizer.bq(
278291
rescore_limit=rescore_limit
279292
),
280-
distance_metric=distance_metric,
293+
distance_metric=distance_metric_enum,
281294
),
282295
hnsw=wvc.Configure.VectorIndex.hnsw(
283296
quantizer=wvc.Configure.VectorIndex.Quantizer.sq(
284297
rescore_limit=rescore_limit, training_limit=training_limit
285298
),
286-
distance_metric=distance_metric,
299+
distance_metric=distance_metric_enum,
287300
),
288301
),
289302
"dynamic_flat_bq_hnsw_bq": wvc.Configure.VectorIndex.dynamic(
290303
flat=wvc.Configure.VectorIndex.flat(
291304
quantizer=wvc.Configure.VectorIndex.Quantizer.bq(
292305
rescore_limit=rescore_limit
293306
),
294-
distance_metric=distance_metric,
307+
distance_metric=distance_metric_enum,
295308
),
296309
hnsw=wvc.Configure.VectorIndex.hnsw(
297310
quantizer=wvc.Configure.VectorIndex.Quantizer.bq(
298311
rescore_limit=rescore_limit
299312
),
300-
distance_metric=distance_metric,
313+
distance_metric=distance_metric_enum,
301314
),
302315
),
303316
"dynamic_hnsw_pq": wvc.Configure.VectorIndex.dynamic(
@@ -312,70 +325,70 @@ def create_collection(
312325
quantizer=wvc.Configure.VectorIndex.Quantizer.sq(
313326
rescore_limit=rescore_limit, training_limit=training_limit
314327
),
315-
distance_metric=distance_metric,
328+
distance_metric=distance_metric_enum,
316329
)
317330
),
318331
"dynamic_hnsw_bq": wvc.Configure.VectorIndex.dynamic(
319332
hnsw=wvc.Configure.VectorIndex.hnsw(
320333
quantizer=wvc.Configure.VectorIndex.Quantizer.bq(
321334
rescore_limit=rescore_limit
322335
),
323-
distance_metric=distance_metric,
336+
distance_metric=distance_metric_enum,
324337
)
325338
),
326339
"hnsw_pq": wvc.Configure.VectorIndex.hnsw(
327340
quantizer=wvc.Configure.VectorIndex.Quantizer.pq(
328341
training_limit=training_limit
329342
),
330-
distance_metric=distance_metric,
343+
distance_metric=distance_metric_enum,
331344
),
332345
"hnsw_bq": wvc.Configure.VectorIndex.hnsw(
333346
quantizer=wvc.Configure.VectorIndex.Quantizer.bq(
334347
rescore_limit=rescore_limit
335348
),
336-
distance_metric=distance_metric,
349+
distance_metric=distance_metric_enum,
337350
),
338351
"hnsw_bq_cache": wvc.Configure.VectorIndex.hnsw(
339352
quantizer=wvc.Configure.VectorIndex.Quantizer.bq(
340353
cache=True, rescore_limit=rescore_limit
341354
),
342-
distance_metric=distance_metric,
355+
distance_metric=distance_metric_enum,
343356
),
344357
"hnsw_sq": wvc.Configure.VectorIndex.hnsw(
345358
quantizer=wvc.Configure.VectorIndex.Quantizer.sq(
346359
rescore_limit=rescore_limit, training_limit=training_limit
347360
),
348-
distance_metric=distance_metric,
361+
distance_metric=distance_metric_enum,
349362
),
350363
"hnsw_rq": wvc.Configure.VectorIndex.hnsw(
351364
quantizer=wvc.Configure.VectorIndex.Quantizer.rq(
352365
rescore_limit=rescore_limit
353366
),
354-
distance_metric=distance_metric,
367+
distance_metric=distance_metric_enum,
355368
),
356369
"hnsw_acorn": wvc.Configure.VectorIndex.hnsw(
357370
filter_strategy=VectorFilterStrategy.ACORN,
358-
distance_metric=distance_metric,
371+
distance_metric=distance_metric_enum,
359372
),
360373
"hnsw_multivector": wvc.Configure.VectorIndex.hnsw(
361374
multi_vector=wvc.Configure.VectorIndex.MultiVector.multi_vector(),
362-
distance_metric=distance_metric,
375+
distance_metric=distance_metric_enum,
363376
),
364377
"flat_bq": wvc.Configure.VectorIndex.flat(
365378
quantizer=wvc.Configure.VectorIndex.Quantizer.bq(
366379
rescore_limit=rescore_limit
367380
),
368-
distance_metric=distance_metric,
381+
distance_metric=distance_metric_enum,
369382
),
370383
"flat_bq_cache": wvc.Configure.VectorIndex.flat(
371384
quantizer=wvc.Configure.VectorIndex.Quantizer.bq(
372385
cache=True, rescore_limit=rescore_limit
373386
),
374-
distance_metric=distance_metric,
387+
distance_metric=distance_metric_enum,
375388
),
376389
"hfresh": self._build_hfresh_config(
377390
max_posting_size_kb=hfresh_max_posting_size_kb,
378-
distance_metric=distance_metric,
391+
distance_metric=distance_metric_enum,
379392
rescore_limit=rescore_limit,
380393
replicas=hfresh_replicas,
381394
search_probe=hfresh_search_probe,

0 commit comments

Comments
 (0)