summa-client-python 2.0.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- summa_client_python/__init__.py +75 -0
- summa_client_python/client.py +1187 -0
- summa_client_python/summa_pb2.py +212 -0
- summa_client_python/summa_pb2_grpc.py +845 -0
- summa_client_python/types.py +305 -0
- summa_client_python-2.0.0.dist-info/METADATA +193 -0
- summa_client_python-2.0.0.dist-info/RECORD +8 -0
- summa_client_python-2.0.0.dist-info/WHEEL +4 -0
|
@@ -0,0 +1,1187 @@
|
|
|
1
|
+
"""Async Summa client implementation.
|
|
2
|
+
|
|
3
|
+
All search types mirror the proto API structure exactly.
|
|
4
|
+
See types.py for Query, Reranker definitions.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
import json
|
|
10
|
+
from collections.abc import AsyncIterator
|
|
11
|
+
from typing import Any
|
|
12
|
+
|
|
13
|
+
import grpc
|
|
14
|
+
from google.protobuf.json_format import MessageToDict
|
|
15
|
+
from grpc import aio
|
|
16
|
+
|
|
17
|
+
from . import summa_pb2 as pb
|
|
18
|
+
from . import summa_pb2_grpc as pb_grpc
|
|
19
|
+
from .types import (
|
|
20
|
+
CandidateScores,
|
|
21
|
+
DocAddress,
|
|
22
|
+
Document,
|
|
23
|
+
DocumentMutationResult,
|
|
24
|
+
FusionCandidate,
|
|
25
|
+
FusionCandidateList,
|
|
26
|
+
IndexInfo,
|
|
27
|
+
OrdinalScore,
|
|
28
|
+
PassageScores,
|
|
29
|
+
QueryTrace,
|
|
30
|
+
RrfContribution,
|
|
31
|
+
SearchHit,
|
|
32
|
+
SearchResponse,
|
|
33
|
+
SearchTimings,
|
|
34
|
+
SearchTrace,
|
|
35
|
+
ShardSearchTrace,
|
|
36
|
+
VectorFieldStats,
|
|
37
|
+
)
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
class SummaClient:
|
|
41
|
+
"""Async client for Summa search server.
|
|
42
|
+
|
|
43
|
+
All search types mirror the proto API structure exactly.
|
|
44
|
+
|
|
45
|
+
Example:
|
|
46
|
+
async with SummaClient("localhost:50051") as client:
|
|
47
|
+
# Create index
|
|
48
|
+
await client.create_index("articles", '''
|
|
49
|
+
index articles {
|
|
50
|
+
title: text indexed stored
|
|
51
|
+
body: text indexed stored
|
|
52
|
+
}
|
|
53
|
+
''')
|
|
54
|
+
|
|
55
|
+
# Index documents
|
|
56
|
+
await client.index_documents("articles", [
|
|
57
|
+
{"title": "Hello", "body": "World"},
|
|
58
|
+
{"title": "Foo", "body": "Bar"},
|
|
59
|
+
])
|
|
60
|
+
await client.commit("articles")
|
|
61
|
+
|
|
62
|
+
# Search
|
|
63
|
+
results = await client.search("articles",
|
|
64
|
+
query={"match": {"field": "title", "text": "hello"}})
|
|
65
|
+
for hit in results.hits:
|
|
66
|
+
print(hit.address, hit.score)
|
|
67
|
+
"""
|
|
68
|
+
|
|
69
|
+
def __init__(
|
|
70
|
+
self,
|
|
71
|
+
address: str = "localhost:50051",
|
|
72
|
+
default_timeout: float | None = None,
|
|
73
|
+
):
|
|
74
|
+
"""Initialize client.
|
|
75
|
+
|
|
76
|
+
Args:
|
|
77
|
+
address: Server address in format "host:port"
|
|
78
|
+
default_timeout: Default per-RPC deadline in seconds, applied to
|
|
79
|
+
every call unless overridden by the call's ``timeout=``
|
|
80
|
+
argument. ``None`` = no deadline (current behaviour). On
|
|
81
|
+
expiry the call raises ``grpc.aio.AioRpcError`` with
|
|
82
|
+
``DEADLINE_EXCEEDED``.
|
|
83
|
+
"""
|
|
84
|
+
self.address = address
|
|
85
|
+
self.default_timeout = default_timeout
|
|
86
|
+
self._channel: aio.Channel | None = None
|
|
87
|
+
self._index_stub: pb_grpc.IndexServiceStub | None = None
|
|
88
|
+
self._search_stub: pb_grpc.SearchServiceStub | None = None
|
|
89
|
+
|
|
90
|
+
def _deadline(self, timeout: float | None) -> float | None:
|
|
91
|
+
"""Effective per-call gRPC deadline in seconds."""
|
|
92
|
+
return timeout if timeout is not None else self.default_timeout
|
|
93
|
+
|
|
94
|
+
async def connect(self) -> None:
|
|
95
|
+
"""Connect to the server."""
|
|
96
|
+
# Increase message size limits for large responses (e.g., loading content fields)
|
|
97
|
+
options = [
|
|
98
|
+
("grpc.max_receive_message_length", 50 * 1024 * 1024), # 50MB
|
|
99
|
+
(
|
|
100
|
+
"grpc.max_send_message_length",
|
|
101
|
+
200 * 1024 * 1024,
|
|
102
|
+
), # singleton upsert ceiling
|
|
103
|
+
]
|
|
104
|
+
# Enable gzip compression for smaller message sizes over the wire
|
|
105
|
+
self._channel = aio.insecure_channel(
|
|
106
|
+
self.address,
|
|
107
|
+
options=options,
|
|
108
|
+
compression=grpc.Compression.Gzip,
|
|
109
|
+
)
|
|
110
|
+
self._index_stub = pb_grpc.IndexServiceStub(self._channel)
|
|
111
|
+
self._search_stub = pb_grpc.SearchServiceStub(self._channel)
|
|
112
|
+
|
|
113
|
+
async def close(self) -> None:
|
|
114
|
+
"""Close the connection."""
|
|
115
|
+
if self._channel:
|
|
116
|
+
await self._channel.close()
|
|
117
|
+
self._channel = None
|
|
118
|
+
self._index_stub = None
|
|
119
|
+
self._search_stub = None
|
|
120
|
+
|
|
121
|
+
async def __aenter__(self) -> SummaClient:
|
|
122
|
+
await self.connect()
|
|
123
|
+
return self
|
|
124
|
+
|
|
125
|
+
async def __aexit__(self, exc_type, exc_val, exc_tb) -> None:
|
|
126
|
+
await self.close()
|
|
127
|
+
|
|
128
|
+
def _ensure_connected(self) -> None:
|
|
129
|
+
if self._index_stub is None or self._search_stub is None:
|
|
130
|
+
raise RuntimeError(
|
|
131
|
+
"Client not connected. Use 'async with' or call connect() first."
|
|
132
|
+
)
|
|
133
|
+
|
|
134
|
+
# =========================================================================
|
|
135
|
+
# Index Management
|
|
136
|
+
# =========================================================================
|
|
137
|
+
|
|
138
|
+
async def create_index(
|
|
139
|
+
self, index_name: str, schema: str, timeout: float | None = None
|
|
140
|
+
) -> bool:
|
|
141
|
+
"""Create a new index.
|
|
142
|
+
|
|
143
|
+
Args:
|
|
144
|
+
index_name: Name of the index
|
|
145
|
+
schema: Schema definition in SDL or JSON format
|
|
146
|
+
|
|
147
|
+
Returns:
|
|
148
|
+
True if successful
|
|
149
|
+
|
|
150
|
+
Example SDL schema:
|
|
151
|
+
index myindex {
|
|
152
|
+
field title: text [indexed, stored]
|
|
153
|
+
field body: text [indexed, stored]
|
|
154
|
+
field score: f64 [stored]
|
|
155
|
+
}
|
|
156
|
+
|
|
157
|
+
Example JSON schema:
|
|
158
|
+
{
|
|
159
|
+
"fields": [
|
|
160
|
+
{"name": "title", "type": "text", "indexed": true, "stored": true}
|
|
161
|
+
]
|
|
162
|
+
}
|
|
163
|
+
"""
|
|
164
|
+
self._ensure_connected()
|
|
165
|
+
request = pb.CreateIndexRequest(index_name=index_name, schema=schema)
|
|
166
|
+
response = await self._index_stub.CreateIndex(
|
|
167
|
+
request, timeout=self._deadline(timeout)
|
|
168
|
+
)
|
|
169
|
+
return response.success
|
|
170
|
+
|
|
171
|
+
async def delete_index(self, index_name: str, timeout: float | None = None) -> bool:
|
|
172
|
+
"""Delete an index.
|
|
173
|
+
|
|
174
|
+
Args:
|
|
175
|
+
index_name: Name of the index to delete
|
|
176
|
+
|
|
177
|
+
Returns:
|
|
178
|
+
True if successful
|
|
179
|
+
"""
|
|
180
|
+
self._ensure_connected()
|
|
181
|
+
request = pb.DeleteIndexRequest(index_name=index_name)
|
|
182
|
+
response = await self._index_stub.DeleteIndex(
|
|
183
|
+
request, timeout=self._deadline(timeout)
|
|
184
|
+
)
|
|
185
|
+
return response.success
|
|
186
|
+
|
|
187
|
+
async def list_indexes(self, timeout: float | None = None) -> list[str]:
|
|
188
|
+
"""List all indexes on the server.
|
|
189
|
+
|
|
190
|
+
Returns:
|
|
191
|
+
List of index names
|
|
192
|
+
"""
|
|
193
|
+
self._ensure_connected()
|
|
194
|
+
request = pb.ListIndexesRequest()
|
|
195
|
+
response = await self._index_stub.ListIndexes(
|
|
196
|
+
request, timeout=self._deadline(timeout)
|
|
197
|
+
)
|
|
198
|
+
return list(response.index_names)
|
|
199
|
+
|
|
200
|
+
async def get_index_info(
|
|
201
|
+
self, index_name: str, timeout: float | None = None
|
|
202
|
+
) -> IndexInfo:
|
|
203
|
+
"""Get information about an index.
|
|
204
|
+
|
|
205
|
+
Args:
|
|
206
|
+
index_name: Name of the index
|
|
207
|
+
|
|
208
|
+
Returns:
|
|
209
|
+
IndexInfo with document count, segments, and schema
|
|
210
|
+
"""
|
|
211
|
+
self._ensure_connected()
|
|
212
|
+
request = pb.GetIndexInfoRequest(index_name=index_name)
|
|
213
|
+
response = await self._search_stub.GetIndexInfo(
|
|
214
|
+
request, timeout=self._deadline(timeout)
|
|
215
|
+
)
|
|
216
|
+
vector_stats = [
|
|
217
|
+
VectorFieldStats(
|
|
218
|
+
field_name=vs.field_name,
|
|
219
|
+
vector_type=vs.vector_type,
|
|
220
|
+
total_vectors=vs.total_vectors,
|
|
221
|
+
dimension=vs.dimension,
|
|
222
|
+
)
|
|
223
|
+
for vs in response.vector_stats
|
|
224
|
+
]
|
|
225
|
+
return IndexInfo(
|
|
226
|
+
index_name=response.index_name,
|
|
227
|
+
num_docs=response.num_docs,
|
|
228
|
+
num_segments=response.num_segments,
|
|
229
|
+
schema=response.schema,
|
|
230
|
+
physical_num_docs=response.physical_num_docs,
|
|
231
|
+
num_deleted_docs=response.num_deleted_docs,
|
|
232
|
+
deleted_ratio=response.deleted_ratio,
|
|
233
|
+
candidate_scoring_version=response.candidate_scoring_version,
|
|
234
|
+
unprepared_candidate_fields=list(response.unprepared_candidate_fields),
|
|
235
|
+
vector_stats=vector_stats,
|
|
236
|
+
)
|
|
237
|
+
|
|
238
|
+
# =========================================================================
|
|
239
|
+
# Document Indexing
|
|
240
|
+
# =========================================================================
|
|
241
|
+
|
|
242
|
+
async def index_documents(
|
|
243
|
+
self,
|
|
244
|
+
index_name: str,
|
|
245
|
+
documents: list[dict[str, Any]],
|
|
246
|
+
timeout: float | None = None,
|
|
247
|
+
) -> tuple[int, int, list[dict[str, Any]]]:
|
|
248
|
+
"""Index multiple documents in batch.
|
|
249
|
+
|
|
250
|
+
Args:
|
|
251
|
+
index_name: Name of the index
|
|
252
|
+
documents: List of documents (dicts with field names as keys)
|
|
253
|
+
|
|
254
|
+
Returns:
|
|
255
|
+
Tuple of (indexed_count, error_count, errors) where errors is a list
|
|
256
|
+
of dicts with 'index' (0-based position) and 'error' (message) keys.
|
|
257
|
+
"""
|
|
258
|
+
self._ensure_connected()
|
|
259
|
+
|
|
260
|
+
named_docs = []
|
|
261
|
+
for doc in documents:
|
|
262
|
+
fields = _to_field_entries(doc)
|
|
263
|
+
named_docs.append(pb.NamedDocument(fields=fields))
|
|
264
|
+
|
|
265
|
+
request = pb.BatchIndexDocumentsRequest(
|
|
266
|
+
index_name=index_name, documents=named_docs
|
|
267
|
+
)
|
|
268
|
+
response = await self._index_stub.BatchIndexDocuments(
|
|
269
|
+
request, timeout=self._deadline(timeout)
|
|
270
|
+
)
|
|
271
|
+
errors = [{"index": e.index, "error": e.error} for e in response.errors]
|
|
272
|
+
return response.indexed_count, response.error_count, errors
|
|
273
|
+
|
|
274
|
+
async def index_document(
|
|
275
|
+
self,
|
|
276
|
+
index_name: str,
|
|
277
|
+
document: dict[str, Any],
|
|
278
|
+
timeout: float | None = None,
|
|
279
|
+
) -> None:
|
|
280
|
+
"""Index a single document.
|
|
281
|
+
|
|
282
|
+
Args:
|
|
283
|
+
index_name: Name of the index
|
|
284
|
+
document: Document as dict with field names as keys
|
|
285
|
+
timeout: Per-call deadline in seconds (overrides ``default_timeout``)
|
|
286
|
+
"""
|
|
287
|
+
await self.index_documents(index_name, [document], timeout=timeout)
|
|
288
|
+
|
|
289
|
+
async def delete_documents(
|
|
290
|
+
self,
|
|
291
|
+
index_name: str,
|
|
292
|
+
primary_keys: list[str],
|
|
293
|
+
timeout: float | None = None,
|
|
294
|
+
) -> DocumentMutationResult:
|
|
295
|
+
"""Stage exact-key deletions of whole documents, including all chunks.
|
|
296
|
+
|
|
297
|
+
Missing keys are accepted. Inspect per-item errors; commit publishes
|
|
298
|
+
accepted operations. Do not blindly retry after an uncertain RPC outcome.
|
|
299
|
+
"""
|
|
300
|
+
self._ensure_connected()
|
|
301
|
+
if (
|
|
302
|
+
len(primary_keys) > 100_000
|
|
303
|
+
or sum(len(key.encode("utf-8")) for key in primary_keys) > 8 * 1024 * 1024
|
|
304
|
+
):
|
|
305
|
+
raise ValueError(
|
|
306
|
+
"deletion request exceeds 100000 keys or 8 MiB of key bytes"
|
|
307
|
+
)
|
|
308
|
+
response = await self._index_stub.DeleteDocuments(
|
|
309
|
+
pb.DeleteDocumentsRequest(index_name=index_name, primary_keys=primary_keys),
|
|
310
|
+
timeout=self._deadline(timeout),
|
|
311
|
+
)
|
|
312
|
+
return DocumentMutationResult(
|
|
313
|
+
response.accepted_count,
|
|
314
|
+
[{"index": error.index, "error": error.error} for error in response.errors],
|
|
315
|
+
)
|
|
316
|
+
|
|
317
|
+
async def delete_document(
|
|
318
|
+
self,
|
|
319
|
+
index_name: str,
|
|
320
|
+
primary_key: str,
|
|
321
|
+
timeout: float | None = None,
|
|
322
|
+
) -> None:
|
|
323
|
+
"""Stage one deletion; raises if rejected. Call commit to publish."""
|
|
324
|
+
result = await self.delete_documents(index_name, [primary_key], timeout=timeout)
|
|
325
|
+
if result.errors or result.accepted_count != 1:
|
|
326
|
+
raise RuntimeError(
|
|
327
|
+
result.errors[0]["error"]
|
|
328
|
+
if result.errors
|
|
329
|
+
else "deletion was not accepted"
|
|
330
|
+
)
|
|
331
|
+
|
|
332
|
+
async def upsert_documents(
|
|
333
|
+
self,
|
|
334
|
+
index_name: str,
|
|
335
|
+
documents: list[dict[str, Any]],
|
|
336
|
+
timeout: float | None = None,
|
|
337
|
+
) -> DocumentMutationResult:
|
|
338
|
+
"""Stage complete replacements, inserting missing keys. No partial patches.
|
|
339
|
+
|
|
340
|
+
Each document must contain its primary key. Inspect per-item errors;
|
|
341
|
+
commit publishes the latest accepted replacement for each key,
|
|
342
|
+
including multiple replacements across calls.
|
|
343
|
+
"""
|
|
344
|
+
self._ensure_connected()
|
|
345
|
+
if len(documents) > 1_000:
|
|
346
|
+
raise ValueError("upsert request exceeds 1000 documents")
|
|
347
|
+
request = pb.UpsertDocumentsRequest(
|
|
348
|
+
index_name=index_name,
|
|
349
|
+
documents=[
|
|
350
|
+
pb.NamedDocument(fields=_to_field_entries(doc)) for doc in documents
|
|
351
|
+
],
|
|
352
|
+
)
|
|
353
|
+
limit_mib = 200 if len(documents) == 1 else 32
|
|
354
|
+
if request.ByteSize() > limit_mib * 1024 * 1024:
|
|
355
|
+
raise ValueError(f"upsert request exceeds {limit_mib} MiB encoded bytes")
|
|
356
|
+
response = await self._index_stub.UpsertDocuments(
|
|
357
|
+
request, timeout=self._deadline(timeout)
|
|
358
|
+
)
|
|
359
|
+
return DocumentMutationResult(
|
|
360
|
+
response.accepted_count,
|
|
361
|
+
[{"index": error.index, "error": error.error} for error in response.errors],
|
|
362
|
+
)
|
|
363
|
+
|
|
364
|
+
async def upsert_document(
|
|
365
|
+
self,
|
|
366
|
+
index_name: str,
|
|
367
|
+
document: dict[str, Any],
|
|
368
|
+
timeout: float | None = None,
|
|
369
|
+
) -> None:
|
|
370
|
+
"""Stage one complete replacement; raises if rejected. Commit to publish."""
|
|
371
|
+
result = await self.upsert_documents(index_name, [document], timeout=timeout)
|
|
372
|
+
if result.errors or result.accepted_count != 1:
|
|
373
|
+
raise RuntimeError(
|
|
374
|
+
result.errors[0]["error"]
|
|
375
|
+
if result.errors
|
|
376
|
+
else "upsert was not accepted"
|
|
377
|
+
)
|
|
378
|
+
|
|
379
|
+
async def index_documents_stream(
|
|
380
|
+
self,
|
|
381
|
+
index_name: str,
|
|
382
|
+
documents: AsyncIterator[dict[str, Any]],
|
|
383
|
+
timeout: float | None = None,
|
|
384
|
+
) -> tuple[int, list[dict[str, Any]]]:
|
|
385
|
+
"""Stream documents for indexing.
|
|
386
|
+
|
|
387
|
+
Args:
|
|
388
|
+
index_name: Name of the index
|
|
389
|
+
documents: Async iterator of documents
|
|
390
|
+
|
|
391
|
+
Returns:
|
|
392
|
+
Tuple of (indexed_count, errors) where errors is a list of dicts
|
|
393
|
+
with 'index' and 'error' keys.
|
|
394
|
+
"""
|
|
395
|
+
self._ensure_connected()
|
|
396
|
+
|
|
397
|
+
async def request_iterator():
|
|
398
|
+
async for doc in documents:
|
|
399
|
+
fields = _to_field_entries(doc)
|
|
400
|
+
yield pb.IndexDocumentRequest(index_name=index_name, fields=fields)
|
|
401
|
+
|
|
402
|
+
response = await self._index_stub.IndexDocuments(
|
|
403
|
+
request_iterator(), timeout=self._deadline(timeout)
|
|
404
|
+
)
|
|
405
|
+
errors = [{"index": e.index, "error": e.error} for e in response.errors]
|
|
406
|
+
return response.indexed_count, errors
|
|
407
|
+
|
|
408
|
+
async def commit(self, index_name: str, timeout: float | None = None) -> int:
|
|
409
|
+
"""Commit pending changes.
|
|
410
|
+
|
|
411
|
+
Args:
|
|
412
|
+
index_name: Name of the index
|
|
413
|
+
|
|
414
|
+
Returns:
|
|
415
|
+
Total number of documents in the index
|
|
416
|
+
"""
|
|
417
|
+
self._ensure_connected()
|
|
418
|
+
request = pb.CommitRequest(index_name=index_name)
|
|
419
|
+
response = await self._index_stub.Commit(
|
|
420
|
+
request, timeout=self._deadline(timeout)
|
|
421
|
+
)
|
|
422
|
+
return response.num_docs
|
|
423
|
+
|
|
424
|
+
async def force_merge(
|
|
425
|
+
self, index_name: str, timeout: float | None = None, *, compact: bool = False
|
|
426
|
+
) -> int:
|
|
427
|
+
"""Force merge all segments; compact=True physically removes deleted rows.
|
|
428
|
+
|
|
429
|
+
Args:
|
|
430
|
+
index_name: Name of the index
|
|
431
|
+
|
|
432
|
+
Returns:
|
|
433
|
+
Number of segments after merge
|
|
434
|
+
"""
|
|
435
|
+
self._ensure_connected()
|
|
436
|
+
request = pb.ForceMergeRequest(index_name=index_name, compact=compact)
|
|
437
|
+
response = await self._index_stub.ForceMerge(
|
|
438
|
+
request, timeout=self._deadline(timeout)
|
|
439
|
+
)
|
|
440
|
+
return response.num_segments
|
|
441
|
+
|
|
442
|
+
async def reorder(self, index_name: str, timeout: float | None = None) -> int:
|
|
443
|
+
"""Run configured text reordering and bounded vector maintenance.
|
|
444
|
+
|
|
445
|
+
Coalesces binary ANN runs and repairs Seismic nomination fragmentation.
|
|
446
|
+
Sparse maintenance uses pending debt independently of text reorder flags;
|
|
447
|
+
a bounded pass may leave more work for subsequent background passes.
|
|
448
|
+
|
|
449
|
+
Args:
|
|
450
|
+
index_name: Name of the index
|
|
451
|
+
|
|
452
|
+
Returns:
|
|
453
|
+
Number of segments after reorder
|
|
454
|
+
"""
|
|
455
|
+
self._ensure_connected()
|
|
456
|
+
request = pb.ReorderRequest(index_name=index_name)
|
|
457
|
+
response = await self._index_stub.Reorder(
|
|
458
|
+
request, timeout=self._deadline(timeout)
|
|
459
|
+
)
|
|
460
|
+
return response.num_segments
|
|
461
|
+
|
|
462
|
+
async def retrain_vector_index(
|
|
463
|
+
self, index_name: str, timeout: float | None = None
|
|
464
|
+
) -> bool:
|
|
465
|
+
"""Retrain vector index centroids/codebooks from current data.
|
|
466
|
+
|
|
467
|
+
Resets the trained ANN structures and rebuilds them from scratch.
|
|
468
|
+
Use this when significant new data has been added and you want
|
|
469
|
+
better centroids, or when the vector distribution has changed.
|
|
470
|
+
|
|
471
|
+
Args:
|
|
472
|
+
index_name: Name of the index
|
|
473
|
+
|
|
474
|
+
Returns:
|
|
475
|
+
True if successful
|
|
476
|
+
"""
|
|
477
|
+
self._ensure_connected()
|
|
478
|
+
request = pb.RetrainVectorIndexRequest(index_name=index_name)
|
|
479
|
+
response = await self._index_stub.RetrainVectorIndex(
|
|
480
|
+
request, timeout=self._deadline(timeout)
|
|
481
|
+
)
|
|
482
|
+
return response.success
|
|
483
|
+
|
|
484
|
+
# =========================================================================
|
|
485
|
+
# Search
|
|
486
|
+
# =========================================================================
|
|
487
|
+
|
|
488
|
+
async def search(
|
|
489
|
+
self,
|
|
490
|
+
index_name: str,
|
|
491
|
+
*,
|
|
492
|
+
query: dict[str, Any],
|
|
493
|
+
limit: int = 10,
|
|
494
|
+
offset: int = 0,
|
|
495
|
+
fields_to_load: list[str] | None = None,
|
|
496
|
+
reranker: dict[str, Any] | None = None,
|
|
497
|
+
candidate_limit: int = 0,
|
|
498
|
+
l1: dict[str, Any] | None = None,
|
|
499
|
+
score_export: dict[str, Any] | None = None,
|
|
500
|
+
include_rrf_scores: bool = False,
|
|
501
|
+
tracing: bool = False,
|
|
502
|
+
timeout: float | None = None,
|
|
503
|
+
) -> SearchResponse:
|
|
504
|
+
"""Search for documents.
|
|
505
|
+
|
|
506
|
+
All parameters mirror the proto SearchRequest structure exactly.
|
|
507
|
+
``query`` is a dict with exactly one key matching the proto Query oneof.
|
|
508
|
+
|
|
509
|
+
Args:
|
|
510
|
+
index_name: Name of the index
|
|
511
|
+
query: Query dict with one key: "term", "match", "boolean",
|
|
512
|
+
"sparse_vector", "dense_vector", "binary_dense_vector",
|
|
513
|
+
"boost", "range", "prefix", "all", or "fusion".
|
|
514
|
+
Fusion (hybrid union of sub-queries, top-level only):
|
|
515
|
+
{"fusion": {"queries": [{"query": {...}, "weight": 1.0}, ...],
|
|
516
|
+
"method": "rrf" | "normalized_weighted_sum",
|
|
517
|
+
"rrf_k": 60}}
|
|
518
|
+
limit: Maximum number of results
|
|
519
|
+
offset: Offset for pagination
|
|
520
|
+
fields_to_load: List of fields to include in results
|
|
521
|
+
reranker: Reranker dict matching proto Reranker message
|
|
522
|
+
candidate_limit: Shared first-stage candidate pool. Zero uses the
|
|
523
|
+
result window; explicit values are capped at 2x that window.
|
|
524
|
+
include_rrf_scores: Add independent RRF scores and per-branch votes
|
|
525
|
+
to fusion hits, preserving the requested ranking. Default false.
|
|
526
|
+
tracing: Preserve every shard's bounded branch nominations, query
|
|
527
|
+
provenance and selected results in response.trace. Default false.
|
|
528
|
+
|
|
529
|
+
Returns:
|
|
530
|
+
SearchResponse with hits
|
|
531
|
+
|
|
532
|
+
Examples:
|
|
533
|
+
# Term query (exact single token)
|
|
534
|
+
results = await client.search("articles",
|
|
535
|
+
query={"term": {"field": "title", "term": "hello"}})
|
|
536
|
+
|
|
537
|
+
# Match query (full-text, tokenized server-side)
|
|
538
|
+
results = await client.search("articles",
|
|
539
|
+
query={"match": {"field": "title", "text": "what is hemoglobin"}})
|
|
540
|
+
|
|
541
|
+
# Boolean query
|
|
542
|
+
results = await client.search("articles",
|
|
543
|
+
query={"boolean": {
|
|
544
|
+
"must": [{"match": {"field": "title", "text": "hello"}}],
|
|
545
|
+
"should": [{"match": {"field": "body", "text": "world"}}],
|
|
546
|
+
}})
|
|
547
|
+
|
|
548
|
+
# Sparse text query (server-side tokenization) with pruning
|
|
549
|
+
results = await client.search("docs",
|
|
550
|
+
query={"sparse_vector": {
|
|
551
|
+
"field": "embedding",
|
|
552
|
+
"text": "machine learning",
|
|
553
|
+
"pruning": 0.5,
|
|
554
|
+
}},
|
|
555
|
+
fields_to_load=["title", "body"])
|
|
556
|
+
|
|
557
|
+
# Sparse vector query (pre-computed)
|
|
558
|
+
results = await client.search("docs",
|
|
559
|
+
query={"sparse_vector": {
|
|
560
|
+
"field": "embedding",
|
|
561
|
+
"indices": [1, 5, 10],
|
|
562
|
+
"values": [0.5, 0.3, 0.2],
|
|
563
|
+
}})
|
|
564
|
+
|
|
565
|
+
# Dense vector query with reranker
|
|
566
|
+
results = await client.search("docs",
|
|
567
|
+
query={"dense_vector": {
|
|
568
|
+
"field": "embedding",
|
|
569
|
+
"vector": [0.1, 0.2, 0.3],
|
|
570
|
+
"nprobe": 10,
|
|
571
|
+
}},
|
|
572
|
+
reranker={
|
|
573
|
+
"field": "embedding",
|
|
574
|
+
"vector": [0.1, 0.2, 0.3],
|
|
575
|
+
},
|
|
576
|
+
limit=50,
|
|
577
|
+
candidate_limit=100,
|
|
578
|
+
fields_to_load=["title"])
|
|
579
|
+
|
|
580
|
+
"""
|
|
581
|
+
self._ensure_connected()
|
|
582
|
+
|
|
583
|
+
pb_query = _build_query(query)
|
|
584
|
+
pb_reranker = _build_reranker(reranker) if reranker else None
|
|
585
|
+
|
|
586
|
+
request = pb.SearchRequest(
|
|
587
|
+
index_name=index_name,
|
|
588
|
+
query=pb_query,
|
|
589
|
+
limit=limit,
|
|
590
|
+
offset=offset,
|
|
591
|
+
fields_to_load=fields_to_load or [],
|
|
592
|
+
reranker=pb_reranker,
|
|
593
|
+
candidate_limit=candidate_limit,
|
|
594
|
+
l1=pb.L1Ranking(**l1) if l1 is not None else None,
|
|
595
|
+
include_rrf_scores=include_rrf_scores,
|
|
596
|
+
tracing=tracing,
|
|
597
|
+
score_export=pb.ScoreExport(**score_export)
|
|
598
|
+
if score_export is not None
|
|
599
|
+
else None,
|
|
600
|
+
)
|
|
601
|
+
|
|
602
|
+
response = await self._search_stub.Search(
|
|
603
|
+
request, timeout=self._deadline(timeout)
|
|
604
|
+
)
|
|
605
|
+
|
|
606
|
+
expected_l1 = "formula_v1"
|
|
607
|
+
if (
|
|
608
|
+
score_export
|
|
609
|
+
and score_export.get("seed_document_passages")
|
|
610
|
+
and not response.seeded_document_passages
|
|
611
|
+
):
|
|
612
|
+
raise RuntimeError(
|
|
613
|
+
"Backend did not acknowledge document passage seeding; upgrade the broker and all backends"
|
|
614
|
+
)
|
|
615
|
+
if l1 is not None and response.ranking_method != expected_l1:
|
|
616
|
+
raise RuntimeError(
|
|
617
|
+
f"L1 requires a backend with {expected_l1} ranking semantics; "
|
|
618
|
+
f"received {response.ranking_method!r}"
|
|
619
|
+
)
|
|
620
|
+
|
|
621
|
+
if tracing and not response.HasField("trace"):
|
|
622
|
+
raise RuntimeError("Backend omitted requested search trace; upgrade Summa")
|
|
623
|
+
if include_rrf_scores and any(
|
|
624
|
+
not hit.HasField("rrf_score") for hit in response.hits
|
|
625
|
+
):
|
|
626
|
+
raise RuntimeError(
|
|
627
|
+
"Backend omitted requested RRF diagnostics; upgrade Summa"
|
|
628
|
+
)
|
|
629
|
+
|
|
630
|
+
hits = [
|
|
631
|
+
SearchHit(
|
|
632
|
+
address=DocAddress(
|
|
633
|
+
segment_id=hit.address.segment_id,
|
|
634
|
+
doc_id=hit.address.doc_id,
|
|
635
|
+
),
|
|
636
|
+
score=hit.score,
|
|
637
|
+
rrf_score=hit.rrf_score if hit.HasField("rrf_score") else None,
|
|
638
|
+
rrf_contributions=[
|
|
639
|
+
RrfContribution(
|
|
640
|
+
query_index=vote.query_index,
|
|
641
|
+
query_name=vote.query_name,
|
|
642
|
+
rank=vote.rank,
|
|
643
|
+
score=vote.score,
|
|
644
|
+
ordinal=vote.ordinal if vote.HasField("ordinal") else None,
|
|
645
|
+
)
|
|
646
|
+
for vote in hit.rrf_contributions
|
|
647
|
+
],
|
|
648
|
+
fields={k: _from_field_value_list(v) for k, v in hit.fields.items()},
|
|
649
|
+
candidate_scores=CandidateScores(
|
|
650
|
+
document=dict(hit.candidate_scores.document),
|
|
651
|
+
passages=[
|
|
652
|
+
PassageScores(
|
|
653
|
+
ordinal=row.ordinal,
|
|
654
|
+
scores=dict(row.scores),
|
|
655
|
+
l1_score=row.l1_score if row.HasField("l1_score") else None,
|
|
656
|
+
)
|
|
657
|
+
for row in hit.candidate_scores.passages
|
|
658
|
+
],
|
|
659
|
+
scored_passages=hit.candidate_scores.scored_passages,
|
|
660
|
+
)
|
|
661
|
+
if hit.HasField("candidate_scores")
|
|
662
|
+
else None,
|
|
663
|
+
ordinal_scores=[
|
|
664
|
+
OrdinalScore(ordinal=os.ordinal, score=os.score)
|
|
665
|
+
for os in hit.ordinal_scores
|
|
666
|
+
],
|
|
667
|
+
)
|
|
668
|
+
for hit in response.hits
|
|
669
|
+
]
|
|
670
|
+
|
|
671
|
+
timings = None
|
|
672
|
+
if response.HasField("timings"):
|
|
673
|
+
t = response.timings
|
|
674
|
+
timings = SearchTimings(
|
|
675
|
+
search_us=t.search_us,
|
|
676
|
+
rerank_us=t.rerank_us,
|
|
677
|
+
load_us=t.load_us,
|
|
678
|
+
total_us=t.total_us,
|
|
679
|
+
candidate_scoring_us=t.candidate_scoring_us,
|
|
680
|
+
)
|
|
681
|
+
|
|
682
|
+
return SearchResponse(
|
|
683
|
+
hits=hits,
|
|
684
|
+
total_hits=response.total_hits,
|
|
685
|
+
trace=_from_search_trace(response.trace)
|
|
686
|
+
if response.HasField("trace")
|
|
687
|
+
else None,
|
|
688
|
+
took_ms=response.took_ms,
|
|
689
|
+
timings=timings,
|
|
690
|
+
ranking_method=response.ranking_method,
|
|
691
|
+
seeded_document_passages=response.seeded_document_passages,
|
|
692
|
+
truncated=response.truncated,
|
|
693
|
+
fusion_candidates=[
|
|
694
|
+
FusionCandidateList(
|
|
695
|
+
query_index=branch.query_index,
|
|
696
|
+
candidates=[
|
|
697
|
+
FusionCandidate(
|
|
698
|
+
address=DocAddress(
|
|
699
|
+
segment_id=hit.address.segment_id,
|
|
700
|
+
doc_id=hit.address.doc_id,
|
|
701
|
+
),
|
|
702
|
+
score=hit.score,
|
|
703
|
+
ordinal_scores=[
|
|
704
|
+
OrdinalScore(ordinal=score.ordinal, score=score.score)
|
|
705
|
+
for score in hit.ordinal_scores
|
|
706
|
+
],
|
|
707
|
+
)
|
|
708
|
+
for hit in branch.candidates
|
|
709
|
+
],
|
|
710
|
+
)
|
|
711
|
+
for branch in response.fusion_candidates
|
|
712
|
+
],
|
|
713
|
+
)
|
|
714
|
+
|
|
715
|
+
async def get_document(
|
|
716
|
+
self,
|
|
717
|
+
index_name: str,
|
|
718
|
+
address: DocAddress,
|
|
719
|
+
timeout: float | None = None,
|
|
720
|
+
) -> Document | None:
|
|
721
|
+
"""Get a document by address.
|
|
722
|
+
|
|
723
|
+
Args:
|
|
724
|
+
index_name: Name of the index
|
|
725
|
+
address: DocAddress from a SearchHit
|
|
726
|
+
|
|
727
|
+
Returns:
|
|
728
|
+
Document or None if not found
|
|
729
|
+
"""
|
|
730
|
+
self._ensure_connected()
|
|
731
|
+
request = pb.GetDocumentRequest(
|
|
732
|
+
index_name=index_name,
|
|
733
|
+
address=pb.DocAddress(segment_id=address.segment_id, doc_id=address.doc_id),
|
|
734
|
+
)
|
|
735
|
+
try:
|
|
736
|
+
response = await self._search_stub.GetDocument(
|
|
737
|
+
request, timeout=self._deadline(timeout)
|
|
738
|
+
)
|
|
739
|
+
fields = {k: _from_field_value_list(v) for k, v in response.fields.items()}
|
|
740
|
+
return Document(fields=fields)
|
|
741
|
+
except grpc.RpcError as e:
|
|
742
|
+
if e.code() == grpc.StatusCode.NOT_FOUND:
|
|
743
|
+
return None
|
|
744
|
+
raise
|
|
745
|
+
|
|
746
|
+
|
|
747
|
+
# =============================================================================
|
|
748
|
+
# Helper functions
|
|
749
|
+
# =============================================================================
|
|
750
|
+
|
|
751
|
+
|
|
752
|
+
def _is_sparse_vector(value: list) -> bool:
|
|
753
|
+
"""Check if list is a sparse vector: list of (int, float) pairs."""
|
|
754
|
+
if not value:
|
|
755
|
+
return False
|
|
756
|
+
for item in value:
|
|
757
|
+
if not isinstance(item, (list, tuple)) or len(item) != 2:
|
|
758
|
+
return False
|
|
759
|
+
idx, val = item
|
|
760
|
+
if not isinstance(idx, int) or not isinstance(val, (int, float)):
|
|
761
|
+
return False
|
|
762
|
+
return True
|
|
763
|
+
|
|
764
|
+
|
|
765
|
+
def _is_multi_sparse_vector(value: list) -> bool:
|
|
766
|
+
"""Check if list is a multi-value sparse vector: list of sparse vectors."""
|
|
767
|
+
if not value:
|
|
768
|
+
return False
|
|
769
|
+
# All items must be lists and each must be a valid sparse vector
|
|
770
|
+
if not all(isinstance(item, list) for item in value):
|
|
771
|
+
return False
|
|
772
|
+
return all(_is_sparse_vector(item) for item in value)
|
|
773
|
+
|
|
774
|
+
|
|
775
|
+
def _is_dense_vector(value: list) -> bool:
|
|
776
|
+
"""Check if list is a dense vector: flat list of numeric values."""
|
|
777
|
+
if not value:
|
|
778
|
+
return False
|
|
779
|
+
return all(isinstance(v, (int, float)) and not isinstance(v, bool) for v in value)
|
|
780
|
+
|
|
781
|
+
|
|
782
|
+
def _is_multi_dense_vector(value: list) -> bool:
|
|
783
|
+
"""Check if list is a multi-value dense vector: list of dense vectors."""
|
|
784
|
+
if not value:
|
|
785
|
+
return False
|
|
786
|
+
# All items must be lists and each must be a valid dense vector
|
|
787
|
+
if not all(isinstance(item, list) for item in value):
|
|
788
|
+
return False
|
|
789
|
+
return all(_is_dense_vector(item) for item in value)
|
|
790
|
+
|
|
791
|
+
|
|
792
|
+
def _sparse_vector_to_proto(value: list) -> pb.SparseVector:
|
|
793
|
+
"""Build the protobuf representation used by fields and queries."""
|
|
794
|
+
return pb.SparseVector(
|
|
795
|
+
indices=[int(item[0]) for item in value],
|
|
796
|
+
values=[float(item[1]) for item in value],
|
|
797
|
+
)
|
|
798
|
+
|
|
799
|
+
|
|
800
|
+
def _dense_vector_to_proto(value: list) -> pb.DenseVector:
|
|
801
|
+
"""Build a dense protobuf vector with normalized float values."""
|
|
802
|
+
return pb.DenseVector(values=[float(item) for item in value])
|
|
803
|
+
|
|
804
|
+
|
|
805
|
+
def _to_field_entries(doc: dict[str, Any]) -> list[pb.FieldEntry]:
|
|
806
|
+
"""Convert document dict to list of FieldEntry for multi-value field support.
|
|
807
|
+
|
|
808
|
+
Multi-value fields (list of sparse vectors or list of dense vectors) are
|
|
809
|
+
expanded into multiple FieldEntry with the same name.
|
|
810
|
+
"""
|
|
811
|
+
entries = []
|
|
812
|
+
for name, value in doc.items():
|
|
813
|
+
if isinstance(value, list):
|
|
814
|
+
# Check for multi-value sparse vectors: [[(idx, val), ...], ...]
|
|
815
|
+
if _is_multi_sparse_vector(value):
|
|
816
|
+
for sv in value:
|
|
817
|
+
fv = pb.FieldValue(sparse_vector=_sparse_vector_to_proto(sv))
|
|
818
|
+
entries.append(pb.FieldEntry(name=name, value=fv))
|
|
819
|
+
continue
|
|
820
|
+
# Check for multi-value dense vectors: [[f1, f2, ...], ...]
|
|
821
|
+
if _is_multi_dense_vector(value):
|
|
822
|
+
for dv in value:
|
|
823
|
+
fv = pb.FieldValue(dense_vector=_dense_vector_to_proto(dv))
|
|
824
|
+
entries.append(pb.FieldEntry(name=name, value=fv))
|
|
825
|
+
continue
|
|
826
|
+
# Single sparse vector: [(idx, val), ...]
|
|
827
|
+
if _is_sparse_vector(value):
|
|
828
|
+
fv = pb.FieldValue(sparse_vector=_sparse_vector_to_proto(value))
|
|
829
|
+
entries.append(pb.FieldEntry(name=name, value=fv))
|
|
830
|
+
continue
|
|
831
|
+
# Single dense vector: [f1, f2, ...]
|
|
832
|
+
if _is_dense_vector(value):
|
|
833
|
+
fv = pb.FieldValue(dense_vector=_dense_vector_to_proto(value))
|
|
834
|
+
entries.append(pb.FieldEntry(name=name, value=fv))
|
|
835
|
+
continue
|
|
836
|
+
# Multi-value plain field: ["val1", "val2", ...] -> separate entries
|
|
837
|
+
for item in value:
|
|
838
|
+
entries.append(pb.FieldEntry(name=name, value=_to_field_value(item)))
|
|
839
|
+
continue
|
|
840
|
+
# Single value - use standard conversion
|
|
841
|
+
entries.append(pb.FieldEntry(name=name, value=_to_field_value(value)))
|
|
842
|
+
return entries
|
|
843
|
+
|
|
844
|
+
|
|
845
|
+
def _to_field_value(value: Any) -> pb.FieldValue:
|
|
846
|
+
"""Convert Python value to protobuf FieldValue.
|
|
847
|
+
|
|
848
|
+
Special handling for vector types:
|
|
849
|
+
- list[(int, float)] -> SparseVector (list of (index, value) tuples)
|
|
850
|
+
- list[float] -> DenseVector (flat list of numeric values)
|
|
851
|
+
- Other lists/dicts -> JSON
|
|
852
|
+
"""
|
|
853
|
+
if isinstance(value, str):
|
|
854
|
+
return pb.FieldValue(text=value)
|
|
855
|
+
elif isinstance(value, bool):
|
|
856
|
+
return pb.FieldValue(u64=1 if value else 0)
|
|
857
|
+
elif isinstance(value, int):
|
|
858
|
+
if value >= 0:
|
|
859
|
+
return pb.FieldValue(u64=value)
|
|
860
|
+
else:
|
|
861
|
+
return pb.FieldValue(i64=value)
|
|
862
|
+
elif isinstance(value, float):
|
|
863
|
+
return pb.FieldValue(f64=value)
|
|
864
|
+
elif isinstance(value, bytes):
|
|
865
|
+
return pb.FieldValue(bytes_value=value)
|
|
866
|
+
elif isinstance(value, dict):
|
|
867
|
+
# Dicts are always JSON
|
|
868
|
+
return pb.FieldValue(json_value=json.dumps(value))
|
|
869
|
+
elif isinstance(value, list):
|
|
870
|
+
# Check if it's a sparse vector: list of (index, value) pairs
|
|
871
|
+
if _is_sparse_vector(value):
|
|
872
|
+
return pb.FieldValue(sparse_vector=_sparse_vector_to_proto(value))
|
|
873
|
+
# Check if it's a dense vector: flat list of numeric values
|
|
874
|
+
if _is_dense_vector(value):
|
|
875
|
+
return pb.FieldValue(dense_vector=_dense_vector_to_proto(value))
|
|
876
|
+
# Otherwise treat as JSON
|
|
877
|
+
return pb.FieldValue(json_value=json.dumps(value))
|
|
878
|
+
else:
|
|
879
|
+
return pb.FieldValue(text=str(value))
|
|
880
|
+
|
|
881
|
+
|
|
882
|
+
def _from_field_value(fv: pb.FieldValue) -> Any:
|
|
883
|
+
"""Convert protobuf FieldValue to Python value."""
|
|
884
|
+
which = fv.WhichOneof("value")
|
|
885
|
+
if which == "text":
|
|
886
|
+
return fv.text
|
|
887
|
+
elif which == "u64":
|
|
888
|
+
return fv.u64
|
|
889
|
+
elif which == "i64":
|
|
890
|
+
return fv.i64
|
|
891
|
+
elif which == "f64":
|
|
892
|
+
return fv.f64
|
|
893
|
+
elif which == "bytes_value":
|
|
894
|
+
return fv.bytes_value
|
|
895
|
+
elif which == "json_value":
|
|
896
|
+
return json.loads(fv.json_value)
|
|
897
|
+
elif which == "sparse_vector":
|
|
898
|
+
return {
|
|
899
|
+
"indices": list(fv.sparse_vector.indices),
|
|
900
|
+
"values": list(fv.sparse_vector.values),
|
|
901
|
+
}
|
|
902
|
+
elif which == "dense_vector":
|
|
903
|
+
return list(fv.dense_vector.values)
|
|
904
|
+
elif which == "binary_dense_vector":
|
|
905
|
+
return bytes(fv.binary_dense_vector)
|
|
906
|
+
return None
|
|
907
|
+
|
|
908
|
+
|
|
909
|
+
def _from_field_value_list(fvl: pb.FieldValueList) -> Any:
|
|
910
|
+
"""Convert protobuf FieldValueList to Python value.
|
|
911
|
+
|
|
912
|
+
Single-value fields are unwrapped to a scalar.
|
|
913
|
+
Multi-value fields are returned as a list.
|
|
914
|
+
"""
|
|
915
|
+
values = [_from_field_value(v) for v in fvl.values]
|
|
916
|
+
if len(values) == 1:
|
|
917
|
+
return values[0]
|
|
918
|
+
return values
|
|
919
|
+
|
|
920
|
+
|
|
921
|
+
_COMBINER_MAP: dict[str, int] = {
|
|
922
|
+
"log_sum_exp": 0,
|
|
923
|
+
"max": 1,
|
|
924
|
+
"avg": 2,
|
|
925
|
+
"sum": 3,
|
|
926
|
+
"weighted_top_k": 4,
|
|
927
|
+
}
|
|
928
|
+
|
|
929
|
+
|
|
930
|
+
def _combiner_to_proto(combiner: str | None) -> int:
|
|
931
|
+
"""Convert combiner string to proto MultiValueCombiner enum value."""
|
|
932
|
+
if combiner is None:
|
|
933
|
+
return 0 # LOG_SUM_EXP default
|
|
934
|
+
return _COMBINER_MAP.get(combiner.lower(), 0)
|
|
935
|
+
|
|
936
|
+
|
|
937
|
+
def _build_query(q: dict[str, Any]) -> pb.Query:
|
|
938
|
+
"""Recursively convert a Query dict to protobuf Query.
|
|
939
|
+
|
|
940
|
+
The dict must have exactly one key matching the proto Query oneof:
|
|
941
|
+
"term", "match", "phrase", "boolean", "sparse_vector", "dense_vector",
|
|
942
|
+
"binary_dense_vector", "boost", "range", "prefix", "all", or "fusion".
|
|
943
|
+
"""
|
|
944
|
+
if "term" in q:
|
|
945
|
+
t = q["term"]
|
|
946
|
+
return pb.Query(
|
|
947
|
+
term=pb.TermQuery(
|
|
948
|
+
field=t["field"],
|
|
949
|
+
term=t["term"],
|
|
950
|
+
tokenizer_hint=t.get("tokenizer_hint", ""),
|
|
951
|
+
)
|
|
952
|
+
)
|
|
953
|
+
|
|
954
|
+
if "match" in q:
|
|
955
|
+
m = q["match"]
|
|
956
|
+
return pb.Query(
|
|
957
|
+
match=pb.MatchQuery(
|
|
958
|
+
field=m["field"],
|
|
959
|
+
text=m["text"],
|
|
960
|
+
tokenizer_hint=m.get("tokenizer_hint", ""),
|
|
961
|
+
)
|
|
962
|
+
)
|
|
963
|
+
|
|
964
|
+
if "phrase" in q:
|
|
965
|
+
p = q["phrase"]
|
|
966
|
+
return pb.Query(
|
|
967
|
+
phrase=pb.PhraseQuery(
|
|
968
|
+
field=p["field"],
|
|
969
|
+
text=p["text"],
|
|
970
|
+
slop=p.get("slop", 0),
|
|
971
|
+
tokenizer_hint=p.get("tokenizer_hint", ""),
|
|
972
|
+
)
|
|
973
|
+
)
|
|
974
|
+
|
|
975
|
+
if "boolean" in q:
|
|
976
|
+
b = q["boolean"]
|
|
977
|
+
return pb.Query(
|
|
978
|
+
boolean=pb.BooleanQuery(
|
|
979
|
+
must=[_build_query(sq) for sq in b.get("must", [])],
|
|
980
|
+
should=[_build_query(sq) for sq in b.get("should", [])],
|
|
981
|
+
must_not=[_build_query(sq) for sq in b.get("must_not", [])],
|
|
982
|
+
)
|
|
983
|
+
)
|
|
984
|
+
|
|
985
|
+
if "sparse_vector" in q:
|
|
986
|
+
sv = q["sparse_vector"]
|
|
987
|
+
sparse_vector = {
|
|
988
|
+
"field": sv["field"],
|
|
989
|
+
"indices": sv.get("indices", []),
|
|
990
|
+
"values": sv.get("values", []),
|
|
991
|
+
"text": sv.get("text", ""),
|
|
992
|
+
"combiner": _combiner_to_proto(sv.get("combiner")),
|
|
993
|
+
"heap_factor": sv.get("heap_factor", 0),
|
|
994
|
+
"combiner_temperature": sv.get("combiner_temperature", 0),
|
|
995
|
+
"combiner_top_k": sv.get("combiner_top_k", 0),
|
|
996
|
+
"combiner_decay": sv.get("combiner_decay", 0),
|
|
997
|
+
"weight_threshold": sv.get("weight_threshold", 0),
|
|
998
|
+
"max_query_dims": sv.get("max_query_dims", 0),
|
|
999
|
+
"pruning": sv.get("pruning", 0),
|
|
1000
|
+
}
|
|
1001
|
+
for option in ("lsp_gamma", "seismic_cut", "seismic_factor", "exhaustive"):
|
|
1002
|
+
if option in sv:
|
|
1003
|
+
sparse_vector[option] = sv[option]
|
|
1004
|
+
return pb.Query(sparse_vector=pb.SparseVectorQuery(**sparse_vector))
|
|
1005
|
+
|
|
1006
|
+
if "dense_vector" in q:
|
|
1007
|
+
dv = q["dense_vector"]
|
|
1008
|
+
return pb.Query(
|
|
1009
|
+
dense_vector=pb.DenseVectorQuery(
|
|
1010
|
+
field=dv["field"],
|
|
1011
|
+
vector=dv["vector"],
|
|
1012
|
+
nprobe=dv.get("nprobe", 0),
|
|
1013
|
+
combiner=_combiner_to_proto(dv.get("combiner")),
|
|
1014
|
+
combiner_temperature=dv.get("combiner_temperature", 0),
|
|
1015
|
+
combiner_top_k=dv.get("combiner_top_k", 0),
|
|
1016
|
+
combiner_decay=dv.get("combiner_decay", 0),
|
|
1017
|
+
)
|
|
1018
|
+
)
|
|
1019
|
+
|
|
1020
|
+
if "binary_dense_vector" in q:
|
|
1021
|
+
bv = q["binary_dense_vector"]
|
|
1022
|
+
return pb.Query(
|
|
1023
|
+
binary_dense_vector=pb.BinaryDenseVectorQuery(
|
|
1024
|
+
field=bv["field"],
|
|
1025
|
+
vector=bv["vector"],
|
|
1026
|
+
combiner=_combiner_to_proto(bv.get("combiner")),
|
|
1027
|
+
combiner_temperature=bv.get("combiner_temperature", 0),
|
|
1028
|
+
combiner_top_k=bv.get("combiner_top_k", 0),
|
|
1029
|
+
combiner_decay=bv.get("combiner_decay", 0),
|
|
1030
|
+
)
|
|
1031
|
+
)
|
|
1032
|
+
|
|
1033
|
+
if "boost" in q:
|
|
1034
|
+
bq = q["boost"]
|
|
1035
|
+
return pb.Query(
|
|
1036
|
+
boost=pb.BoostQuery(
|
|
1037
|
+
query=_build_query(bq["query"]),
|
|
1038
|
+
boost=bq["boost"],
|
|
1039
|
+
)
|
|
1040
|
+
)
|
|
1041
|
+
|
|
1042
|
+
if "range" in q:
|
|
1043
|
+
rq = q["range"]
|
|
1044
|
+
kwargs: dict[str, Any] = {"field": rq["field"]}
|
|
1045
|
+
if "min_u64" in rq:
|
|
1046
|
+
kwargs["min_u64"] = int(rq["min_u64"])
|
|
1047
|
+
if "max_u64" in rq:
|
|
1048
|
+
kwargs["max_u64"] = int(rq["max_u64"])
|
|
1049
|
+
if "min_i64" in rq:
|
|
1050
|
+
kwargs["min_i64"] = int(rq["min_i64"])
|
|
1051
|
+
if "max_i64" in rq:
|
|
1052
|
+
kwargs["max_i64"] = int(rq["max_i64"])
|
|
1053
|
+
if "min_f64" in rq:
|
|
1054
|
+
kwargs["min_f64"] = float(rq["min_f64"])
|
|
1055
|
+
if "max_f64" in rq:
|
|
1056
|
+
kwargs["max_f64"] = float(rq["max_f64"])
|
|
1057
|
+
return pb.Query(range=pb.RangeQuery(**kwargs))
|
|
1058
|
+
|
|
1059
|
+
if "prefix" in q:
|
|
1060
|
+
p = q["prefix"]
|
|
1061
|
+
return pb.Query(prefix=pb.PrefixQuery(field=p["field"], prefix=p["prefix"]))
|
|
1062
|
+
|
|
1063
|
+
if "all" in q:
|
|
1064
|
+
return pb.Query(all=pb.AllQuery())
|
|
1065
|
+
|
|
1066
|
+
if "fusion" in q:
|
|
1067
|
+
f = q["fusion"]
|
|
1068
|
+
method = f.get("method", "rrf")
|
|
1069
|
+
if method == "rrf":
|
|
1070
|
+
pb_method = pb.FusionMethod.FUSION_RRF
|
|
1071
|
+
elif method == "candidates":
|
|
1072
|
+
pb_method = pb.FusionMethod.FUSION_CANDIDATES
|
|
1073
|
+
elif method == "normalized_weighted_sum":
|
|
1074
|
+
pb_method = pb.FusionMethod.FUSION_NORMALIZED_WEIGHTED_SUM
|
|
1075
|
+
else:
|
|
1076
|
+
raise ValueError(
|
|
1077
|
+
f"Unknown fusion method {method!r}: "
|
|
1078
|
+
"expected 'rrf', 'normalized_weighted_sum' or 'candidates'"
|
|
1079
|
+
)
|
|
1080
|
+
return pb.Query(
|
|
1081
|
+
fusion=pb.FusionQuery(
|
|
1082
|
+
queries=[
|
|
1083
|
+
pb.WeightedQuery(
|
|
1084
|
+
query=_build_query(wq["query"]),
|
|
1085
|
+
weight=wq.get("weight", 0.0 if wq.get("name") else 1.0),
|
|
1086
|
+
name=wq.get("name", ""),
|
|
1087
|
+
scope={
|
|
1088
|
+
None: pb.SCORE_SCOPE_UNSPECIFIED,
|
|
1089
|
+
"document": pb.SCORE_SCOPE_DOCUMENT,
|
|
1090
|
+
"chunk": pb.SCORE_SCOPE_CHUNK,
|
|
1091
|
+
}[wq.get("scope")],
|
|
1092
|
+
score_only=wq.get("score_only", False),
|
|
1093
|
+
)
|
|
1094
|
+
for wq in f["queries"]
|
|
1095
|
+
],
|
|
1096
|
+
filters=[_build_query(q) for q in f.get("filters", [])],
|
|
1097
|
+
candidate_depth=f.get("candidate_depth", 0),
|
|
1098
|
+
method=pb_method,
|
|
1099
|
+
rrf_k=f.get("rrf_k", 0),
|
|
1100
|
+
# Chunk combiner; unset -> MAX server-side (chunk-level fusion)
|
|
1101
|
+
combiner=_combiner_to_proto(f.get("combiner")),
|
|
1102
|
+
)
|
|
1103
|
+
)
|
|
1104
|
+
|
|
1105
|
+
# No recognized query key found
|
|
1106
|
+
valid_keys = [
|
|
1107
|
+
"term",
|
|
1108
|
+
"match",
|
|
1109
|
+
"boolean",
|
|
1110
|
+
"sparse_vector",
|
|
1111
|
+
"dense_vector",
|
|
1112
|
+
"binary_dense_vector",
|
|
1113
|
+
"boost",
|
|
1114
|
+
"range",
|
|
1115
|
+
"prefix",
|
|
1116
|
+
"all",
|
|
1117
|
+
"fusion",
|
|
1118
|
+
]
|
|
1119
|
+
raise ValueError(
|
|
1120
|
+
f"Unrecognized query key(s): {set(q.keys()) - set(valid_keys)}. "
|
|
1121
|
+
f"Valid keys: {valid_keys}"
|
|
1122
|
+
)
|
|
1123
|
+
|
|
1124
|
+
|
|
1125
|
+
def _build_reranker(r: dict[str, Any]) -> pb.Reranker:
|
|
1126
|
+
"""Convert a Reranker dict to protobuf Reranker."""
|
|
1127
|
+
return pb.Reranker(
|
|
1128
|
+
field=r["field"],
|
|
1129
|
+
vector=r.get("vector", []),
|
|
1130
|
+
combiner=_combiner_to_proto(r.get("combiner")),
|
|
1131
|
+
combiner_temperature=r.get("combiner_temperature", 0),
|
|
1132
|
+
combiner_top_k=r.get("combiner_top_k", 0),
|
|
1133
|
+
combiner_decay=r.get("combiner_decay", 0),
|
|
1134
|
+
matryoshka_dims=r.get("matryoshka_dims", 0),
|
|
1135
|
+
binary_vector=r.get("binary_vector", b""),
|
|
1136
|
+
rrf_k=r.get("rrf_k", 0),
|
|
1137
|
+
)
|
|
1138
|
+
|
|
1139
|
+
|
|
1140
|
+
def _from_trace_candidate(candidate: pb.FusionCandidate) -> FusionCandidate:
|
|
1141
|
+
return FusionCandidate(
|
|
1142
|
+
address=DocAddress(candidate.address.segment_id, candidate.address.doc_id),
|
|
1143
|
+
score=candidate.score,
|
|
1144
|
+
ordinal_scores=[
|
|
1145
|
+
OrdinalScore(row.ordinal, row.score) for row in candidate.ordinal_scores
|
|
1146
|
+
],
|
|
1147
|
+
)
|
|
1148
|
+
|
|
1149
|
+
|
|
1150
|
+
def _from_search_trace(trace: pb.SearchTrace) -> SearchTrace:
|
|
1151
|
+
return SearchTrace(
|
|
1152
|
+
shards=[
|
|
1153
|
+
ShardSearchTrace(
|
|
1154
|
+
shard_id=shard.shard_id,
|
|
1155
|
+
backend_id=shard.backend_id,
|
|
1156
|
+
index_name=shard.index_name,
|
|
1157
|
+
ranking_method=shard.ranking_method,
|
|
1158
|
+
truncated=shard.truncated,
|
|
1159
|
+
filters=[
|
|
1160
|
+
MessageToDict(query, preserving_proto_field_name=True)
|
|
1161
|
+
for query in shard.filters
|
|
1162
|
+
],
|
|
1163
|
+
selected=[
|
|
1164
|
+
_from_trace_candidate(candidate) for candidate in shard.selected
|
|
1165
|
+
],
|
|
1166
|
+
queries=[
|
|
1167
|
+
QueryTrace(
|
|
1168
|
+
query_index=branch.query_index,
|
|
1169
|
+
query_name=branch.query_name,
|
|
1170
|
+
query=MessageToDict(
|
|
1171
|
+
branch.query, preserving_proto_field_name=True
|
|
1172
|
+
),
|
|
1173
|
+
scope=branch.scope,
|
|
1174
|
+
score_only=branch.score_only,
|
|
1175
|
+
candidate_depth=branch.candidate_depth,
|
|
1176
|
+
total_seen=branch.total_seen,
|
|
1177
|
+
candidates=[
|
|
1178
|
+
_from_trace_candidate(candidate)
|
|
1179
|
+
for candidate in branch.candidates
|
|
1180
|
+
],
|
|
1181
|
+
)
|
|
1182
|
+
for branch in shard.queries
|
|
1183
|
+
],
|
|
1184
|
+
)
|
|
1185
|
+
for shard in trace.shards
|
|
1186
|
+
]
|
|
1187
|
+
)
|