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.
@@ -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
+ )