44import logging
55import os
66import uuid
7- from typing import Any , Dict , List , Optional
7+ from typing import Any , ClassVar , Dict , List , Optional , Sequence
88
99import httpx
1010from moss_core import (
1111 CLOUD_API_MANAGE_URL ,
12- ManageClient ,
1312 DocumentInfo ,
1413 GetDocumentsOptions ,
1514 IndexInfo ,
1615 IndexManager ,
16+ JobStatusResponse ,
17+ ManageClient ,
1718 MutationOptions ,
1819 MutationResult ,
19- JobStatusResponse ,
2020 QueryResultDocumentInfo ,
2121 SearchResult ,
2222)
2323
2424logger = logging .getLogger (__name__ )
2525
26- from typing import Sequence
2726
2827class QueryOptions :
2928 """Options for search queries."""
29+
3030 def __init__ (
3131 self ,
3232 embedding : Optional [Sequence [float ]] = None ,
@@ -37,16 +37,31 @@ def __init__(
3737 rerank_top_k : Optional [int ] = None ,
3838 rerank_model : Optional [str ] = None ,
3939 ):
40+ if top_k is not None and (not isinstance (top_k , int ) or top_k < 1 ):
41+ raise ValueError ("top_k must be an integer >= 1" )
42+ if alpha is not None and (
43+ not isinstance (alpha , (int , float )) or not (0.0 <= alpha <= 1.0 )
44+ ):
45+ raise ValueError ("alpha must be a float between 0.0 and 1.0" )
46+ if embedding is not None :
47+ try :
48+ embedding = [float (x ) for x in embedding ]
49+ except (TypeError , ValueError ):
50+ raise ValueError ("embedding must be a sequence of numbers" )
51+ if rerank_top_k is not None and (
52+ not isinstance (rerank_top_k , int ) or rerank_top_k < 1
53+ ):
54+ raise ValueError ("rerank_top_k must be an integer >= 1" )
55+
4056 self .embedding = embedding
4157 self .top_k = top_k
42- self .alpha = alpha
58+ self .alpha = float ( alpha ) if alpha is not None else None
4359 self .filter = filter
44- self .rerank = rerank
60+ self .rerank = bool ( rerank )
4561 self .rerank_top_k = rerank_top_k
4662 self .rerank_model = rerank_model
4763
4864
49-
5065def _get_manage_url () -> str :
5166 """Manage URL, overridable via env for local development."""
5267 return os .getenv ("MOSS_CLOUD_API_MANAGE_URL" , CLOUD_API_MANAGE_URL )
@@ -83,6 +98,7 @@ class MossClient:
8398 """
8499
85100 DEFAULT_MODEL_ID = "moss-minilm"
101+ _cross_encoder_cache : ClassVar [Dict [str , Any ]] = {}
86102
87103 def __init__ (self , project_id : str , project_key : str ) -> None :
88104 self ._project_id = project_id
@@ -215,8 +231,10 @@ async def query(
215231 """
216232 is_loaded = await asyncio .to_thread (self ._manager .has_index , name )
217233
218- rerank = getattr (options , "rerank" , False )
219- override_top_k = getattr (options , "rerank_top_k" , 50 ) if rerank else None
234+ rerank = getattr (options , "rerank" , False ) is True
235+ override_top_k = (
236+ (getattr (options , "rerank_top_k" , None ) or 50 ) if rerank else None
237+ )
220238
221239 if is_loaded :
222240 result = await self ._query_local (name , query , options , override_top_k )
@@ -228,10 +246,10 @@ async def query(
228246 name ,
229247 )
230248 result = await self ._query_cloud (name , query , options , override_top_k )
231-
249+
232250 if rerank :
233251 result = await self ._rerank_results (query , result , options )
234-
252+
235253 return result
236254
237255 # -- Internal ---------------------------------------------------
@@ -243,7 +261,11 @@ async def _query_local(
243261 options : Optional [QueryOptions ],
244262 override_top_k : Optional [int ] = None ,
245263 ) -> SearchResult :
246- top_k = override_top_k if override_top_k is not None else (getattr (options , "top_k" , None ) or 5 )
264+ top_k = (
265+ override_top_k
266+ if override_top_k is not None
267+ else (getattr (options , "top_k" , None ) or 5 )
268+ )
247269 alpha = getattr (options , "alpha" , None )
248270 if alpha is None :
249271 alpha = 0.8
@@ -286,7 +308,11 @@ async def _query_cloud(
286308 override_top_k : Optional [int ] = None ,
287309 ) -> SearchResult :
288310 """Fallback: query via the cloud API when the index is not loaded locally."""
289- top_k = override_top_k if override_top_k is not None else (getattr (options , "top_k" , None ) or 10 )
311+ top_k = (
312+ override_top_k
313+ if override_top_k is not None
314+ else (getattr (options , "top_k" , None ) or 10 )
315+ )
290316 query_embedding = getattr (options , "embedding" , None )
291317
292318 request_body : Dict [str , Any ] = {
@@ -346,32 +372,37 @@ async def _rerank_results(
346372 "Install it with: pip install 'moss[rerank]'"
347373 )
348374
349- model_name = getattr (options , "rerank_model" , None ) or "cross-encoder/ms-marco-MiniLM-L-6-v2"
375+ model_name = (
376+ getattr (options , "rerank_model" , None )
377+ or "cross-encoder/ms-marco-MiniLM-L-6-v2"
378+ )
350379
351- def do_rerank ():
380+ def do_rerank () -> SearchResult :
352381 if not hasattr (self .__class__ , "_cross_encoder_cache" ):
353382 self .__class__ ._cross_encoder_cache = {}
354383
355384 if model_name not in self .__class__ ._cross_encoder_cache :
356- self .__class__ ._cross_encoder_cache [model_name ] = CrossEncoder (model_name )
385+ self .__class__ ._cross_encoder_cache [model_name ] = CrossEncoder (
386+ model_name
387+ )
357388
358389 model = self .__class__ ._cross_encoder_cache [model_name ]
359390
360- pairs = [[query , doc .text ] for doc in search_result .docs ]
391+ local_docs = search_result .docs
392+ pairs = [[query , doc .text ] for doc in local_docs ]
361393 scores = model .predict (pairs )
362394
363- for doc , score in zip (search_result . docs , scores ):
395+ for doc , score in zip (local_docs , scores ):
364396 doc .score = float (score )
365397
366- search_result . docs .sort (key = lambda d : d .score , reverse = True )
398+ local_docs .sort (key = lambda d : d .score , reverse = True )
367399
368400 original_top_k = getattr (options , "top_k" , None ) or 5
369- search_result .docs = search_result . docs [:original_top_k ]
401+ search_result .docs = local_docs [:original_top_k ]
370402 return search_result
371403
372404 return await asyncio .to_thread (do_rerank )
373405
374-
375406 def _resolve_model_id (
376407 self ,
377408 docs : List [DocumentInfo ],
0 commit comments