codekingpro/portable-devtools
114k
1from typing import Dict, Generator, List, Optional, Sequence, Union
2import numpy as np
3from numpy.typing import NDArray
4
5import pytest
6import chromadb
7from chromadb.api.types import URI, DataLoader, Documents, IDs, Image, URIs
8from chromadb.api import ClientAPI
9from chromadb.test.conftest import reset
10from chromadb.test.ef.test_multimodal_ef import hashing_multimodal_ef
11
12
13def encode_data(data: str) -> NDArray[np.uint8]:
14 return np.array(data.encode())
15
16
17class DefaultDataLoader(DataLoader[List[Optional[Image]]]):
18 def __call__(self, uris: Sequence[Optional[URI]]) -> List[Optional[Image]]:
19 # Convert each URI to a numpy array
20 return [None if uri is None else encode_data(uri) for uri in uris]
21
22
23def record_set_with_uris(n: int = 3) -> Dict[str, Union[IDs, Documents, URIs]]:
24 return {
25 "ids": [f"{i}" for i in range(n)],
26 "documents": [f"document_{i}" for i in range(n)],
27 "uris": [f"uri_{i}" for i in range(n)],
28 }
29
30
31@pytest.fixture()
32def collection_with_data_loader(
33 client: ClientAPI,
34) -> Generator[chromadb.Collection, None, None]:
35 reset(client)
36 collection = client.create_collection(
37 name="collection_with_data_loader",
38 data_loader=DefaultDataLoader(),
39 embedding_function=hashing_multimodal_ef(),
40 )
41 yield collection
42 client.delete_collection(collection.name)
43
44
45@pytest.fixture
46def collection_without_data_loader(
47 client: ClientAPI,
48) -> Generator[chromadb.Collection, None, None]:
49 reset(client)
50 collection = client.create_collection(
51 name="collection_without_data_loader",
52 embedding_function=hashing_multimodal_ef(),
53 )
54 yield collection
55 client.delete_collection(collection.name)
56
57
58def test_without_data_loader(
59 collection_without_data_loader: chromadb.Collection,
60 n_examples: int = 3,
61) -> None:
62 record_set = record_set_with_uris(n=n_examples)
63
64 # Can't embed data in URIs without a data loader
65 with pytest.raises(ValueError):
66 collection_without_data_loader.add(
67 ids=record_set["ids"],
68 uris=record_set["uris"],
69 )
70
71 # Can't get data from URIs without a data loader
72 with pytest.raises(ValueError):
73 collection_without_data_loader.get(include=["data"])
74
75
76def test_without_uris(
77 collection_with_data_loader: chromadb.Collection, n_examples: int = 3
78) -> None:
79 record_set = record_set_with_uris(n=n_examples)
80
81 collection_with_data_loader.add(
82 ids=record_set["ids"],
83 documents=record_set["documents"],
84 )
85
86 get_result = collection_with_data_loader.get(include=["data"])
87
88 assert get_result["data"] is not None
89 for data in get_result["data"]:
90 assert data is None
91
92
93def test_data_loader(
94 collection_with_data_loader: chromadb.Collection, n_examples: int = 3
95) -> None:
96 record_set = record_set_with_uris(n=n_examples)
97
98 collection_with_data_loader.add(
99 ids=record_set["ids"],
100 uris=record_set["uris"],
101 )
102
103 # Get with "data"
104 get_result = collection_with_data_loader.get(include=["data"])
105
106 assert get_result["data"] is not None
107 for i, data in enumerate(get_result["data"]):
108 assert data is not None
109 assert data == encode_data(record_set["uris"][i])
110
111 # Query by URI
112 query_result = collection_with_data_loader.query(
113 query_uris=record_set["uris"],
114 n_results=len(record_set["uris"][0]),
115 include=["data", "uris"],
116 )
117
118 assert query_result["data"] is not None
119 for i, data in enumerate(query_result["data"][0]):
120 assert data is not None
121 assert query_result["uris"] is not None
122 assert data == encode_data(query_result["uris"][0][i])
123 