codekingpro/portable-devtools
114k
1"""
2Integration test for the Collection statistics wrapper methods
3"""
4
5import json
6import time
7from typing import Any
8
9import pytest
10
11from chromadb.api.client import Client as ClientCreator
12from chromadb.base_types import SparseVector
13from chromadb.config import System
14from chromadb.test.conftest import skip_if_not_cluster
15from chromadb.test.utils.wait_for_version_increase import (
16 get_collection_version,
17 wait_for_version_increase,
18)
19from chromadb.utils.statistics import (
20 attach_statistics_function,
21 detach_statistics_function,
22 get_statistics,
23 get_statistics_fn_name,
24)
25
26pytestmark = [skip_if_not_cluster()]
27
28
29def test_statistics_wrapper(basic_http_client: System) -> None:
30 """Test the statistics wrapper methods on Collection"""
31 client = ClientCreator.from_system(basic_http_client)
32 client.reset()
33
34 # Create a collection
35 collection = client.get_or_create_collection(
36 name="test_collection",
37 metadata={"description": "Test collection for statistics"},
38 )
39
40 # Enable statistics
41 attached_fn, created = attach_statistics_function(
42 collection, "test_collection_statistics"
43 )
44 assert attached_fn is not None
45 assert created is True
46 assert attached_fn.function_name == "statistics"
47 assert attached_fn.output_collection == "test_collection_statistics"
48
49 initial_version = get_collection_version(client, collection.name)
50
51 # Add some documents with metadata
52 collection.add(
53 ids=["doc1", "doc2", "doc3"],
54 documents=["test document 1", "test document 2", "test document 3"],
55 metadatas=[
56 {"category": "A", "score": 10, "active": True},
57 {"category": "B", "score": 10, "active": False},
58 {"category": "A", "score": 20, "active": True},
59 ],
60 )
61
62 # Wait for statistics to be computed
63 wait_for_version_increase(client, collection.name, initial_version)
64 time.sleep(60)
65
66 # Get statistics
67 stats = get_statistics(collection, "test_collection_statistics")
68 print("\nStatistics output:")
69 print(json.dumps(stats, indent=2))
70
71 # Verify the structure
72 assert "statistics" in stats
73 assert "summary" in stats
74
75 # Verify summary
76 assert stats["summary"]["total_count"] == 3
77
78 # Verify category statistics
79 assert "category" in stats["statistics"]
80 assert "A" in stats["statistics"]["category"]
81 assert "B" in stats["statistics"]["category"]
82 assert stats["statistics"]["category"]["A"]["count"] == 2
83 assert stats["statistics"]["category"]["B"]["count"] == 1
84
85 # Verify score statistics
86 assert "score" in stats["statistics"]
87 assert "10" in stats["statistics"]["score"]
88 assert "20" in stats["statistics"]["score"]
89 assert stats["statistics"]["score"]["10"]["count"] == 2
90 assert stats["statistics"]["score"]["20"]["count"] == 1
91
92 # Verify active statistics
93 assert "active" in stats["statistics"]
94 assert "true" in stats["statistics"]["active"]
95 assert "false" in stats["statistics"]["active"]
96 assert stats["statistics"]["active"]["true"]["count"] == 2
97 assert stats["statistics"]["active"]["false"]["count"] == 1
98
99 # Test get_attached_function
100 stats_fn = collection.get_attached_function(get_statistics_fn_name(collection))
101 assert stats_fn.function_name == "statistics"
102
103 # Disable statistics (keep the collection)
104 success = detach_statistics_function(collection, delete_stats_collection=False)
105 assert success is True
106
107 # Verify the statistics collection still exists
108 stats_collection = client.get_collection("test_collection_statistics")
109 assert stats_collection is not None
110
111
112def test_backfill_statistics(basic_http_client: System) -> None:
113 """Test backfill statistics"""
114 client = ClientCreator.from_system(basic_http_client)
115 client.reset()
116
117 collection = client.create_collection(name="my_collection")
118
119 initial_version = get_collection_version(client, collection.name)
120
121 # Add some documents with metadata
122 collection.add(
123 ids=["doc1", "doc2", "doc3"],
124 documents=["test document 1", "test document 2", "test document 3"],
125 metadatas=[
126 {"category": "A", "score": 10, "active": True},
127 {"category": "B", "score": 10, "active": False},
128 {"category": "A", "score": 20, "active": True},
129 ],
130 )
131
132 # Let this all be compacted
133 wait_for_version_increase(client, collection.name, initial_version)
134 initial_version = get_collection_version(client, collection.name)
135
136 # Enable statistics
137 attached_fn, created = attach_statistics_function(
138 collection, "my_collection_statistics"
139 )
140 assert created is True
141 assert attached_fn.function_name == "statistics"
142 assert attached_fn.output_collection == "my_collection_statistics"
143
144 # Wait for statistics to be computed
145 wait_for_version_increase(client, collection.name, initial_version)
146
147 stats = get_statistics(collection, "my_collection_statistics")
148 assert stats is not None
149 assert "statistics" in stats
150 assert "summary" in stats
151
152 # Verify summary
153 assert stats["summary"]["total_count"] == 3
154
155 # Verify category statistics
156 assert "category" in stats["statistics"]
157 assert "A" in stats["statistics"]["category"]
158 assert "B" in stats["statistics"]["category"]
159 assert stats["statistics"]["category"]["A"]["count"] == 2
160 assert stats["statistics"]["category"]["B"]["count"] == 1
161
162 # Verify score statistics
163 assert "score" in stats["statistics"]
164 assert "10" in stats["statistics"]["score"]
165 assert "20" in stats["statistics"]["score"]
166 assert stats["statistics"]["score"]["10"]["count"] == 2
167 assert stats["statistics"]["score"]["20"]["count"] == 1
168
169 # Verify active statistics
170 assert "active" in stats["statistics"]
171 assert "true" in stats["statistics"]["active"]
172 assert "false" in stats["statistics"]["active"]
173 assert stats["statistics"]["active"]["true"]["count"] == 2
174 assert stats["statistics"]["active"]["false"]["count"] == 1
175
176 # Disable statistics
177 success = detach_statistics_function(collection, delete_stats_collection=True)
178 assert success is True
179
180
181def test_statistics_wrapper_custom_output_collection(basic_http_client: System) -> None:
182 """Test statistics with custom output collection name"""
183 client = ClientCreator.from_system(basic_http_client)
184 client.reset()
185
186 collection = client.create_collection(name="my_collection")
187
188 # Enable statistics with custom output collection name
189 attached_fn, created = attach_statistics_function(
190 collection, stats_collection_name="my_custom_stats"
191 )
192 assert created is True
193 assert attached_fn.output_collection == "my_custom_stats"
194
195 initial_version = get_collection_version(client, collection.name)
196
197 # Add data
198 collection.add(
199 ids=["id1"],
200 documents=["doc1"],
201 metadatas=[{"key": "value"}],
202 )
203
204 wait_for_version_increase(client, collection.name, initial_version)
205
206 # Get statistics
207 stats = get_statistics(collection, "my_custom_stats")
208 assert "statistics" in stats
209 assert "key" in stats["statistics"]
210
211 # Disable and delete the custom collection
212 detach_statistics_function(collection, delete_stats_collection=True)
213
214
215def test_statistics_wrapper_key_filter(basic_http_client: System) -> None:
216 """Test get_statistics with key filter parameter"""
217 client = ClientCreator.from_system(basic_http_client)
218 client.reset()
219
220 collection = client.create_collection(name="key_filter_test")
221
222 # Enable statistics
223 _, created = attach_statistics_function(collection, "key_filter_test_statistics")
224 assert created is True
225
226 initial_version = get_collection_version(client, collection.name)
227
228 # Add documents with multiple metadata keys
229 collection.add(
230 ids=["doc1", "doc2", "doc3"],
231 documents=["test document 1", "test document 2", "test document 3"],
232 metadatas=[
233 {"category": "A", "score": 10, "active": True},
234 {"category": "B", "score": 10, "active": False},
235 {"category": "A", "score": 20, "active": True},
236 ],
237 )
238
239 wait_for_version_increase(client, collection.name, initial_version)
240 time.sleep(60)
241
242 # Get all statistics (no key filter)
243 all_stats = get_statistics(collection, "key_filter_test_statistics")
244 assert "category" in all_stats["statistics"]
245 assert "score" in all_stats["statistics"]
246 assert "active" in all_stats["statistics"]
247
248 # Get statistics filtered by "category" key only
249 category_stats = get_statistics(
250 collection, "key_filter_test_statistics", keys=["category"]
251 )
252 assert "category" in category_stats["statistics"]
253 assert "score" not in category_stats["statistics"]
254 assert "active" not in category_stats["statistics"]
255 assert category_stats["statistics"]["category"]["A"]["count"] == 2
256 assert category_stats["statistics"]["category"]["B"]["count"] == 1
257 # Summary should still be present when filtering by key
258 assert "summary" in category_stats
259 assert category_stats["summary"]["total_count"] == 3
260
261 # Get statistics filtered by "score" key only
262 score_stats = get_statistics(
263 collection, "key_filter_test_statistics", keys=["score"]
264 )
265 assert "score" in score_stats["statistics"]
266 assert "category" not in score_stats["statistics"]
267 assert "active" not in score_stats["statistics"]
268 assert score_stats["statistics"]["score"]["10"]["count"] == 2
269 assert score_stats["statistics"]["score"]["20"]["count"] == 1
270 # Summary should still be present when filtering by key
271 assert "summary" in score_stats
272 assert score_stats["summary"]["total_count"] == 3
273
274 # Cleanup
275 detach_statistics_function(collection, delete_stats_collection=True)
276
277
278def test_statistics_wrapper_key_filter_too_many_keys(basic_http_client: System) -> None:
279 """Test that get_statistics raises ValueError when more than 30 keys are provided"""
280 client = ClientCreator.from_system(basic_http_client)
281 client.reset()
282
283 collection = client.create_collection(name="too_many_keys_test")
284
285 # Enable statistics
286 attach_statistics_function(collection, "too_many_keys_test_statistics")
287
288 # Generate more than 30 keys
289 too_many_keys = [f"key_{i}" for i in range(31)]
290
291 # Should raise ValueError when more than 30 keys are provided
292 with pytest.raises(ValueError) as exc_info:
293 get_statistics(collection, "too_many_keys_test_statistics", keys=too_many_keys)
294
295 assert "Too many keys provided: 31" in str(exc_info.value)
296 assert "Maximum allowed is 30" in str(exc_info.value)
297
298 # Cleanup
299 detach_statistics_function(collection, delete_stats_collection=True)
300
301
302# commenting out for now as waiting for query cache invalidateion slows down the test suite
303def test_statistics_wrapper_incremental_updates(basic_http_client: System) -> None:
304 """Test that statistics are updated incrementally"""
305 client = ClientCreator.from_system(basic_http_client)
306 client.reset()
307
308 collection = client.create_collection(name="incremental_test")
309 _, created = attach_statistics_function(collection, "incremental_test_statistics")
310 assert created is True
311
312 initial_version = get_collection_version(client, collection.name)
313
314 # Add initial batch
315 collection.add(
316 ids=["id1", "id2"],
317 documents=["doc1", "doc2"],
318 metadatas=[{"category": "A"}, {"category": "A"}],
319 )
320
321 wait_for_version_increase(client, collection.name, initial_version)
322 next_version = get_collection_version(client, collection.name)
323
324 # Check initial statistics
325 stats = get_statistics(collection, "incremental_test_statistics")
326 assert stats["statistics"]["category"]["A"]["count"] == 2
327 assert stats["summary"]["total_count"] == 2
328
329 # Add more data
330 collection.add(
331 ids=["id3", "id4"],
332 documents=["doc3", "doc4"],
333 metadatas=[{"category": "B"}, {"category": "A"}],
334 )
335
336 wait_for_version_increase(client, collection.name, next_version)
337 # TODO(tanujnay112): Remove this sleep once query cache invalidation is solidified
338 # or figure out a different testing harness where we don't have to wait for query cache invalidation
339 time.sleep(70)
340
341 # Check updated statistics
342 stats = get_statistics(collection, "incremental_test_statistics")
343 assert stats["statistics"]["category"]["A"]["count"] == 3
344 assert stats["statistics"]["category"]["B"]["count"] == 1
345 assert stats["summary"]["total_count"] == 4
346
347 detach_statistics_function(collection, delete_stats_collection=True)
348
349
350def test_sparse_vector_statistics(basic_http_client: System) -> None:
351 """Test statistics with sparse vector that includes labels"""
352 client = ClientCreator.from_system(basic_http_client)
353 client.reset()
354
355 collection = client.create_collection(name="sparse_vector_test1")
356
357 # Create sparse vectors with labels
358 sparse_vec1 = SparseVector(
359 indices=[100, 200, 300],
360 values=[1.0, 2.0, 3.0],
361 labels=["apple", "banana", "cherry"],
362 )
363 sparse_vec2 = SparseVector(
364 indices=[100, 400], values=[1.5, 2.5], labels=["apple", "date"]
365 )
366 sparse_vec3 = SparseVector(
367 indices=[200, 300], values=[2.0, 3.0], labels=["banana", "cherry"]
368 )
369
370 # Add data with sparse vectors
371 collection.add(
372 ids=["id1", "id2", "id3"],
373 documents=["doc1", "doc2", "doc3"],
374 metadatas=[
375 {"category": "A", "vec": sparse_vec1},
376 {"category": "B", "vec": sparse_vec2},
377 {"category": "A", "vec": sparse_vec3},
378 ],
379 )
380 _, created = attach_statistics_function(
381 collection, "sparse_vector_test1_statistics"
382 )
383 assert created is True
384
385 initial_version = get_collection_version(client, collection.name)
386
387 wait_for_version_increase(client, collection.name, initial_version)
388
389 # Get statistics
390 stats = get_statistics(collection, "sparse_vector_test1_statistics")
391 print("\nSparse vector statistics output:")
392 print(json.dumps(stats, indent=2))
393
394 assert "statistics" in stats
395 assert "summary" in stats
396 assert stats["summary"]["total_count"] == 3
397
398 # Verify category statistics
399 assert "category" in stats["statistics"]
400 assert "A" in stats["statistics"]["category"]
401 assert "B" in stats["statistics"]["category"]
402 assert stats["statistics"]["category"]["A"]["count"] == 2
403 assert stats["statistics"]["category"]["B"]["count"] == 1
404
405 # Verify sparse vector statistics use labels instead of hash IDs
406 assert "vec" in stats["statistics"]
407 assert "apple" in stats["statistics"]["vec"], "Should use label 'apple' not hash ID"
408 assert (
409 "banana" in stats["statistics"]["vec"]
410 ), "Should use label 'banana' not hash ID"
411 assert (
412 "cherry" in stats["statistics"]["vec"]
413 ), "Should use label 'cherry' not hash ID"
414 assert "date" in stats["statistics"]["vec"], "Should use label 'date' not hash ID"
415
416 # Verify counts
417 assert stats["statistics"]["vec"]["apple"]["count"] == 2 # in id1 and id2
418 assert stats["statistics"]["vec"]["banana"]["count"] == 2 # in id1 and id3
419 assert stats["statistics"]["vec"]["cherry"]["count"] == 2 # in id1 and id3
420 assert stats["statistics"]["vec"]["date"]["count"] == 1 # in id2 only
421
422
423def test_statistics_high_cardinality(basic_http_client: System) -> None:
424 """Test statistics with high cardinality metadata"""
425 client = ClientCreator.from_system(basic_http_client)
426 client.reset()
427
428 collection = client.create_collection(name="high_cardinality_test")
429
430 # Generate 500 documents with 10 metadata fields each
431 num_docs = 500
432 num_fields = 10
433 ids = [f"id{i}" for i in range(num_docs)]
434 documents = [f"doc{i}" for i in range(num_docs)]
435
436 metadatas: list[dict[str, Any]] = []
437 for i in range(num_docs):
438 meta: dict[str, Any] = {}
439 for j in range(num_fields):
440 meta[f"field_{j}"] = f"value_{j}_{i}"
441 metadatas.append(meta)
442
443 # Add in batches to avoid hitting request size limits
444 batch_size = 100
445 initial_version = get_collection_version(client, collection.name)
446
447 for i in range(0, num_docs, batch_size):
448 collection.add(
449 ids=ids[i : i + batch_size],
450 documents=documents[i : i + batch_size],
451 metadatas=metadatas[i : i + batch_size], # type: ignore[arg-type]
452 )
453
454 # Let all data be compacted
455 wait_for_version_increase(client, collection.name, initial_version)
456 initial_version = get_collection_version(client, collection.name)
457
458 # Enable statistics
459 _, created = attach_statistics_function(
460 collection, "high_cardinality_test_statistics"
461 )
462 assert created is True
463
464 # Wait for statistics to be computed
465 wait_for_version_increase(client, collection.name, initial_version)
466
467 # Get statistics
468 stats = get_statistics(collection, "high_cardinality_test_statistics")
469
470 assert "statistics" in stats
471
472 # Verify we have stats for all fields
473 for j in range(num_fields):
474 field_key = f"field_{j}"
475 assert field_key in stats["statistics"]
476
477 field_stats = stats["statistics"][field_key]
478 assert len(field_stats) == num_docs
479
480 # Verify each value has count 1
481 for i in range(num_docs):
482 value = f"value_{j}_{i}"
483 assert value in field_stats
484 assert field_stats[value]["count"] == 1
485
486 # Verify total count
487 assert stats["summary"]["total_count"] == num_docs
488
489 detach_statistics_function(collection, delete_stats_collection=True)
490 