2222)
2323
2424from ..rerankers import get_reranker
25- from .models import QueryOptions
25+ from .models import QueryOptions , RerankOptions
2626
2727logger = logging .getLogger (__name__ )
2828
@@ -191,24 +191,35 @@ async def query(
191191 Otherwise, falls back to the cloud query API.
192192
193193 Args:
194- options: Query options (top_k, alpha, embedding, filter). Example filter:
195- QueryOptions(filter={"$and": [
196- {"field": "city", "condition": {"$eq": "NYC"}},
197- {"field": "price", "condition": {"$lt": "50"}},
198- ]})
194+ options: Query options (top_k, alpha, embedding, filter, rerank).
195+ Reranking is applied client-side after retrieval and works on
196+ both the local and cloud paths. Example filter:
197+ QueryOptions(filter={"$and": [
198+ {"field": "city", "condition": {"$eq": "NYC"}},
199+ {"field": "price", "condition": {"$lt": "50"}},
200+ ]})
199201 """
200202 is_loaded = await asyncio .to_thread (self ._manager .has_index , name )
201203
202204 if is_loaded :
203- return await self ._query_local (name , query , options )
205+ result = await self ._query_local (name , query , options )
206+ else :
207+ if getattr (options , "filter" , None ) is not None :
208+ logger .warning (
209+ "Metadata filter ignored: filtering is only supported for locally loaded indexes. "
210+ "Call load_index('%s') first." ,
211+ name ,
212+ )
213+ result = await self ._query_cloud (name , query , options )
204214
205- if getattr (options , "filter" , None ) is not None :
206- logger .warning (
207- "Metadata filter ignored: filtering is only supported for locally loaded indexes. "
208- "Call load_index('%s') first." ,
209- name ,
210- )
211- return await self ._query_cloud (name , query , options )
215+ rerank = getattr (options , "rerank" , None )
216+ if rerank :
217+ top_k = getattr (options , "top_k" , None )
218+ if top_k is None :
219+ top_k = 5
220+ result = await self ._apply_rerank (query , result , rerank , top_k )
221+
222+ return result
212223
213224 # -- Internal ---------------------------------------------------
214225
@@ -218,7 +229,9 @@ async def _query_local(
218229 query : str ,
219230 options : Optional [QueryOptions ],
220231 ) -> SearchResult :
221- top_k = getattr (options , "top_k" , None ) or 5
232+ top_k = getattr (options , "top_k" , None )
233+ if top_k is None :
234+ top_k = 5
222235 alpha = getattr (options , "alpha" , None )
223236 if alpha is None :
224237 alpha = 0.8
@@ -229,7 +242,7 @@ async def _query_local(
229242 fetch_k = top_k * 4 if rerank else top_k
230243
231244 if query_embedding is not None :
232- result = await asyncio .to_thread (
245+ return await asyncio .to_thread (
233246 self ._manager .query ,
234247 name ,
235248 query ,
@@ -238,41 +251,46 @@ async def _query_local(
238251 alpha ,
239252 filter ,
240253 )
241- else :
242- try :
243- result = await asyncio .to_thread (
244- self ._manager .query_text ,
245- name ,
246- query ,
247- fetch_k ,
248- alpha ,
249- filter ,
250- )
251- except RuntimeError as e :
252- if "requires explicit query embeddings" in str (e ):
253- raise ValueError (
254- "This index uses custom embeddings. "
255- "Query embeddings must be provided via QueryOptions.embedding."
256- ) from e
257- raise
258254
259- if rerank :
260- if rerank ._instance is None :
261- rerank ._instance = get_reranker (
262- rerank .provider , ** rerank .init_kwargs
263- )
264- final_n = rerank .top_n or top_k
265- reranked_docs = await rerank ._instance .rerank (
266- query , result .docs , top_k = final_n
267- )
268- result = SearchResult (
269- docs = reranked_docs ,
270- query = result .query ,
271- index_name = result .index_name ,
272- time_taken_ms = result .time_taken_ms ,
255+ try :
256+ return await asyncio .to_thread (
257+ self ._manager .query_text ,
258+ name ,
259+ query ,
260+ fetch_k ,
261+ alpha ,
262+ filter ,
273263 )
264+ except RuntimeError as e :
265+ if "requires explicit query embeddings" in str (e ):
266+ raise ValueError (
267+ "This index uses custom embeddings. "
268+ "Query embeddings must be provided via QueryOptions.embedding."
269+ ) from e
270+ raise
274271
275- return result
272+ @staticmethod
273+ async def _apply_rerank (
274+ query : str ,
275+ result : SearchResult ,
276+ rerank_opts : RerankOptions ,
277+ default_top_k : Optional [int ],
278+ ) -> SearchResult :
279+ """Rerank search results. Works on both local and cloud paths."""
280+ if rerank_opts ._instance is None :
281+ rerank_opts ._instance = get_reranker (
282+ rerank_opts .provider , ** rerank_opts .init_kwargs
283+ )
284+ final_n = rerank_opts .top_n or default_top_k
285+ reranked_docs = await rerank_opts ._instance .rerank (
286+ query , result .docs , top_k = final_n
287+ )
288+ return SearchResult (
289+ docs = reranked_docs ,
290+ query = result .query ,
291+ index_name = result .index_name ,
292+ time_taken_ms = result .time_taken_ms ,
293+ )
276294
277295 async def _query_cloud (
278296 self ,
@@ -281,15 +299,19 @@ async def _query_cloud(
281299 options : Optional [QueryOptions ],
282300 ) -> SearchResult :
283301 """Fallback: query via the cloud API when the index is not loaded locally."""
284- top_k = getattr (options , "top_k" , None ) or 10
302+ top_k = getattr (options , "top_k" , None )
303+ if top_k is None :
304+ top_k = 5
305+ rerank = getattr (options , "rerank" , None )
306+ fetch_k = top_k * 4 if rerank else top_k
285307 query_embedding = getattr (options , "embedding" , None )
286308
287309 request_body : Dict [str , Any ] = {
288310 "query" : query ,
289311 "indexName" : name ,
290312 "projectId" : self ._project_id ,
291313 "projectKey" : self ._project_key ,
292- "topK" : top_k ,
314+ "topK" : fetch_k ,
293315 }
294316 if query_embedding is not None :
295317 request_body ["queryEmbedding" ] = list (query_embedding )
0 commit comments