codekingpro/portable-devtools
114k
1from dataclasses import dataclass, field
2from typing import List, Dict, Any, Union, Set, Optional
3
4from chromadb.execution.expression.operator import (
5 KNN,
6 Filter,
7 GroupBy,
8 Limit,
9 Projection,
10 Scan,
11 Rank,
12 Select,
13 Where,
14 Key,
15)
16
17
18@dataclass
19class CountPlan:
20 scan: Scan
21
22
23@dataclass
24class GetPlan:
25 scan: Scan
26 filter: Filter = field(default_factory=Filter)
27 limit: Limit = field(default_factory=Limit)
28 projection: Projection = field(default_factory=Projection)
29
30
31@dataclass
32class KNNPlan:
33 scan: Scan
34 knn: KNN
35 filter: Filter = field(default_factory=Filter)
36 projection: Projection = field(default_factory=Projection)
37
38
39class Search:
40 """Payload for hybrid search operations.
41
42 Can be constructed by directly providing the parameters, or by using the builder pattern.
43
44 Examples:
45 Direct construction with expressions:
46 Search(
47 where=Key("status") == "active",
48 rank=Knn(query=[0.1, 0.2]),
49 limit=Limit(limit=10),
50 select=Select(keys={Key.DOCUMENT}),
51 )
52
53 Direct construction with dicts:
54 Search(
55 where={"status": "active"},
56 rank={"$knn": {"query": [0.1, 0.2]}},
57 limit=10,
58 select=["#document", "#score"],
59 )
60
61 Builder pattern:
62 (Search()
63 .where(Key("status") == "active")
64 .rank(Knn(query=[0.1, 0.2]))
65 .limit(10)
66 .select(Key.DOCUMENT))
67 """
68
69 def __init__(
70 self,
71 where: Optional[Union[Where, Dict[str, Any]]] = None,
72 rank: Optional[Union[Rank, Dict[str, Any]]] = None,
73 group_by: Optional[Union[GroupBy, Dict[str, Any]]] = None,
74 limit: Optional[Union[Limit, Dict[str, Any], int]] = None,
75 select: Optional[Union[Select, Dict[str, Any], List[str], Set[str]]] = None,
76 ):
77 """Initialize a Search payload.
78
79 Args:
80 where: Where expression or dict for filtering results (defaults to None - no filtering)
81 Dict will be converted using Where.from_dict()
82 rank: Rank expression or dict for scoring (defaults to None - no ranking)
83 Dict will be converted using Rank.from_dict()
84 Note: Primitive numbers are not accepted - use {"$val": number} for constant ranks
85 group_by: GroupBy configuration for grouping and aggregating results (defaults to None)
86 Dict will be converted using GroupBy.from_dict()
87 limit: Limit configuration for pagination (defaults to no limit)
88 Can be a Limit object, a dict for Limit.from_dict(), or an int
89 When passing an int, it creates Limit(limit=value, offset=0)
90 select: Select configuration for keys (defaults to empty selection)
91 Can be a Select object, a dict for Select.from_dict(),
92 or a list/set of strings (e.g., ["#document", "#score"])
93 """
94 # Handle where parameter
95 if where is None:
96 self._where = None
97 elif isinstance(where, Where):
98 self._where = where
99 elif isinstance(where, dict):
100 self._where = Where.from_dict(where)
101 else:
102 raise TypeError(
103 f"where must be a Where object, dict, or None, got {type(where).__name__}"
104 )
105
106 # Handle rank parameter
107 if rank is None:
108 self._rank = None
109 elif isinstance(rank, Rank):
110 self._rank = rank
111 elif isinstance(rank, dict):
112 self._rank = Rank.from_dict(rank)
113 else:
114 raise TypeError(
115 f"rank must be a Rank object, dict, or None, got {type(rank).__name__}"
116 )
117
118 # Handle group_by parameter
119 if group_by is None:
120 self._group_by = GroupBy()
121 elif isinstance(group_by, GroupBy):
122 self._group_by = group_by
123 elif isinstance(group_by, dict):
124 self._group_by = GroupBy.from_dict(group_by)
125 else:
126 raise TypeError(
127 f"group_by must be a GroupBy object, dict, or None, got {type(group_by).__name__}"
128 )
129
130 # Handle limit parameter
131 if limit is None:
132 self._limit = Limit()
133 elif isinstance(limit, Limit):
134 self._limit = limit
135 elif isinstance(limit, int):
136 self._limit = Limit.from_dict({"limit": limit, "offset": 0})
137 elif isinstance(limit, dict):
138 self._limit = Limit.from_dict(limit)
139 else:
140 raise TypeError(
141 f"limit must be a Limit object, dict, int, or None, got {type(limit).__name__}"
142 )
143
144 # Handle select parameter
145 if select is None:
146 self._select = Select()
147 elif isinstance(select, Select):
148 self._select = select
149 elif isinstance(select, dict):
150 self._select = Select.from_dict(select)
151 elif isinstance(select, (list, set)):
152 # Convert list/set of strings to Select object
153 self._select = Select.from_dict({"keys": list(select)})
154 else:
155 raise TypeError(
156 f"select must be a Select object, dict, list, set, or None, got {type(select).__name__}"
157 )
158
159 def to_dict(self) -> Dict[str, Any]:
160 """Return a JSON-serializable dictionary representation."""
161 return {
162 "filter": self._where.to_dict() if self._where is not None else None,
163 "rank": self._rank.to_dict() if self._rank is not None else None,
164 "group_by": self._group_by.to_dict(),
165 "limit": self._limit.to_dict(),
166 "select": self._select.to_dict(),
167 }
168
169 # Builder methods for chaining
170 def select_all(self) -> "Search":
171 """Select all predefined keys (document, embedding, metadata, score)."""
172 new_select = Select(keys={Key.DOCUMENT, Key.EMBEDDING, Key.METADATA, Key.SCORE})
173 return Search(
174 where=self._where,
175 rank=self._rank,
176 group_by=self._group_by,
177 limit=self._limit,
178 select=new_select,
179 )
180
181 def select(self, *keys: Union[Key, str]) -> "Search":
182 """Select specific keys to return.
183
184 Args:
185 *keys: Key objects or string key names.
186
187 Returns:
188 Search: A new Search with updated selection.
189 """
190 new_select = Select(keys=set(keys))
191 return Search(
192 where=self._where,
193 rank=self._rank,
194 group_by=self._group_by,
195 limit=self._limit,
196 select=new_select,
197 )
198
199 def where(self, where: Optional[Union[Where, Dict[str, Any]]]) -> "Search":
200 """Set the where clause for filtering.
201
202 Args:
203 where: Where expression, dict, or None.
204
205 Example:
206 search.where((Key("status") == "active") & (Key("score") > 0.5))
207 search.where({"status": "active"})
208 search.where({"$and": [{"status": "active"}, {"score": {"$gt": 0.5}}]})
209 """
210 return Search(
211 where=where,
212 rank=self._rank,
213 group_by=self._group_by,
214 limit=self._limit,
215 select=self._select,
216 )
217
218 def rank(self, rank_expr: Optional[Union[Rank, Dict[str, Any]]]) -> "Search":
219 """Set the ranking expression.
220
221 Args:
222 rank_expr: A Rank expression, dict, or None for scoring
223 Dicts will be converted using Rank.from_dict()
224 Note: Primitive numbers are not accepted - use {"$val": number} for constant ranks
225
226 Example:
227 search.rank(Knn(query=[0.1, 0.2]) * 0.8 + Val(0.5) * 0.2)
228 search.rank({"$knn": {"query": [0.1, 0.2]}})
229 search.rank({"$sum": [{"$knn": {"query": [0.1, 0.2]}}, {"$val": 0.5}]})
230 """
231 return Search(
232 where=self._where,
233 rank=rank_expr,
234 group_by=self._group_by,
235 limit=self._limit,
236 select=self._select,
237 )
238
239 def group_by(self, group_by: Optional[Union[GroupBy, Dict[str, Any]]]) -> "Search":
240 """Set the group_by configuration for grouping and aggregating results
241
242 Args:
243 group_by: A GroupBy object, dict, or None for grouping
244 Dicts will be converted using GroupBy.from_dict()
245
246 Example:
247 search.group_by(GroupBy(
248 keys=[Key("category")],
249 aggregate=MinK(keys=[Key.SCORE], k=3)
250 ))
251 search.group_by({
252 "keys": ["category"],
253 "aggregate": {"$min_k": {"keys": ["#score"], "k": 3}}
254 })
255 """
256 return Search(
257 where=self._where,
258 rank=self._rank,
259 group_by=group_by,
260 limit=self._limit,
261 select=self._select,
262 )
263
264 def limit(self, limit: int, offset: int = 0) -> "Search":
265 """Set the limit and offset for pagination
266
267 Args:
268 limit: Maximum number of results to return
269 offset: Number of results to skip (default: 0)
270
271 Example:
272 search.limit(20, offset=10)
273 """
274 new_limit = Limit(offset=offset, limit=limit)
275 return Search(
276 where=self._where,
277 rank=self._rank,
278 group_by=self._group_by,
279 limit=new_limit,
280 select=self._select,
281 )
282 