Team Ai
Datasetpublic

codekingpro/portable-devtools

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