codekingpro/portable-devtools
114k
1import numpy as np
2from typing import List, Dict, Any, cast, Union
3from chromadb.utils.results import (
4 _transform_embeddings,
5 _add_query_fields,
6 _add_get_fields,
7 query_result_to_dfs,
8 get_result_to_df,
9)
10from chromadb.api.types import (
11 QueryResult,
12 GetResult,
13)
14from numpy.typing import NDArray
15
16
17def test_transform_embeddings() -> None:
18 # Test with None input
19 assert _transform_embeddings(None) is None
20
21 # Test with numpy arrays
22 embeddings = cast(
23 List[NDArray[Union[np.int32, np.float32]]],
24 [np.array([1.0, 2.0]), np.array([3.0, 4.0])],
25 )
26 transformed = _transform_embeddings(embeddings)
27 assert isinstance(transformed, list)
28 assert transformed == [[1.0, 2.0], [3.0, 4.0]]
29
30 # Test with list of lists
31 embeddings = cast(
32 List[NDArray[Union[np.int32, np.float32]]],
33 [np.array([1.0, 2.0]), np.array([3.0, 4.0])],
34 )
35 transformed = _transform_embeddings(embeddings)
36 assert transformed == [[1.0, 2.0], [3.0, 4.0]]
37
38
39def test_add_query_fields() -> None:
40 data_dict: Dict[str, Any] = {}
41 query_result: QueryResult = {
42 "ids": [["id1"], ["id2"]],
43 "embeddings": [[np.array([1.0, 2.0])], [np.array([3.0, 4.0])]],
44 "documents": [["doc1"], ["doc2"]],
45 "metadatas": [[{"key": "value1"}], [{"key": "value2"}]],
46 "distances": [[0.1], [0.2]],
47 "uris": [["uri1", "uri2"]],
48 "data": [
49 [np.array([1, 2, 3]), np.array([4, 5, 6])]
50 ], # Using numpy arrays as Image type
51 "included": ["embeddings", "documents", "metadatas", "distances"],
52 }
53
54 _add_query_fields(data_dict, query_result, 0)
55 assert np.array_equal(data_dict["embedding"], [np.array([1.0, 2.0])])
56 assert data_dict["document"] == ["doc1"]
57 assert data_dict["metadata"] == [{"key": "value1"}]
58 assert data_dict["distance"] == [0.1]
59
60
61def test_add_get_fields() -> None:
62 data_dict: Dict[str, Any] = {}
63 get_result: GetResult = {
64 "ids": ["id1", "id2"],
65 "embeddings": [np.array([1.0, 2.0]), np.array([3.0, 4.0])],
66 "documents": ["doc1", "doc2"],
67 "metadatas": [{"key": "value1"}, {"key": "value2"}],
68 "uris": ["uri1", "uri2"],
69 "data": [
70 np.array([1, 2, 3]),
71 np.array([4, 5, 6]),
72 ], # Using numpy arrays as Image type
73 "included": ["embeddings", "documents", "metadatas"],
74 }
75
76 _add_get_fields(data_dict, get_result)
77 assert all(
78 np.array_equal(a, b)
79 for a, b in zip(
80 data_dict["embedding"], [np.array([1.0, 2.0]), np.array([3.0, 4.0])]
81 )
82 )
83 assert data_dict["document"] == ["doc1", "doc2"]
84 assert data_dict["metadata"] == [{"key": "value1"}, {"key": "value2"}]
85
86
87def test_query_result_to_dfs() -> None:
88 query_result: QueryResult = {
89 "ids": [["id1", "id2"]],
90 "embeddings": [[np.array([1.0, 2.0]), np.array([3.0, 4.0])]],
91 "documents": [["doc1", "doc2"]],
92 "metadatas": [[{"key": "value1"}, {"key": "value2"}]],
93 "distances": [[0.1, 0.2]],
94 "uris": [["uri1", "uri2"]],
95 "data": [
96 [np.array([1, 2, 3]), np.array([4, 5, 6])]
97 ], # Using numpy arrays as Image type
98 "included": ["embeddings", "documents", "metadatas", "distances"],
99 }
100
101 dfs = query_result_to_dfs(query_result)
102 assert len(dfs) == 1 # Only one query
103
104 # Test DataFrame
105 df = dfs[0]
106 assert df.index[0] == "id1"
107 assert df["document"].iloc[0] == "doc1"
108 assert df["metadata"].iloc[0] == {"key": "value1"}
109 assert np.array_equal(df["embedding"].iloc[0], np.array([1.0, 2.0]))
110 assert df["distance"].iloc[0] == 0.1
111
112 # Test column order
113 assert list(df.columns) == ["embedding", "document", "metadata", "distance"]
114
115
116def test_get_result_to_df() -> None:
117 get_result: GetResult = {
118 "ids": ["id1", "id2"],
119 "embeddings": [np.array([1.0, 2.0]), np.array([3.0, 4.0])],
120 "documents": ["doc1", "doc2"],
121 "metadatas": [{"key": "value1"}, {"key": "value2"}],
122 "uris": ["uri1", "uri2"],
123 "data": [
124 np.array([1, 2, 3]),
125 np.array([4, 5, 6]),
126 ], # Using numpy arrays as Image type
127 "included": ["embeddings", "documents", "metadatas"],
128 }
129
130 df = get_result_to_df(get_result)
131 assert len(df) == 2
132 assert list(df.index) == ["id1", "id2"]
133 assert df["document"].tolist() == ["doc1", "doc2"]
134 assert df["metadata"].tolist() == [{"key": "value1"}, {"key": "value2"}]
135 assert all(
136 np.array_equal(a, b)
137 for a, b in zip(
138 df["embedding"].tolist(), [np.array([1.0, 2.0]), np.array([3.0, 4.0])]
139 )
140 )
141
142 # Test column order
143 assert list(df.columns) == ["embedding", "document", "metadata"]
144
145
146def test_query_result_to_dfs_with_missing_fields() -> None:
147 query_result: QueryResult = {
148 "ids": [["id1"]],
149 "documents": [["doc1"]],
150 "embeddings": [[]], # type:ignore
151 "metadatas": [[]],
152 "distances": [[]],
153 "uris": [[]],
154 "data": [[]],
155 "included": ["documents"],
156 }
157
158 dfs = query_result_to_dfs(query_result)
159 assert len(dfs) == 1
160 df = dfs[0]
161 assert df.index[0] == "id1"
162 assert df["document"].iloc[0] == "doc1"
163 assert "metadata" not in df.columns
164 assert "embedding" not in df.columns
165 assert "distance" not in df.columns
166
167
168def test_get_result_to_df_with_missing_fields() -> None:
169 get_result: GetResult = {
170 "ids": ["id1", "id2"],
171 "documents": ["doc1", "doc2"],
172 "embeddings": [],
173 "metadatas": [],
174 "uris": [],
175 "data": [],
176 "included": ["documents"],
177 }
178
179 df = get_result_to_df(get_result)
180 assert len(df) == 2
181 assert list(df.index) == ["id1", "id2"]
182 assert df["document"].tolist() == ["doc1", "doc2"]
183 assert "metadata" not in df.columns
184 assert "embedding" not in df.columns
185 