Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
test_statistics_wrapper.py490 linesDownload Raw Back to distributed
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 
codekingpro/portable-devtools · Team Ai