diff --git a/python/tests/detail/test_collection_dql.py b/python/tests/detail/test_collection_dql.py index bccce61..8eb04e3 100644 --- a/python/tests/detail/test_collection_dql.py +++ b/python/tests/detail/test_collection_dql.py @@ -824,7 +824,7 @@ class TestCollectionQuery: for k, v in DEFAULT_VECTOR_FIELD_NAME.items(): multi_query_vectors.append(Query(field_name=v, vector=doc_vectors[v])) - rrf_reranker = RrfReRanker(topn=3) + rrf_reranker = RrfReRanker() multi_query_result = full_collection.query( multi_query_vectors, reranker=rrf_reranker, @@ -876,8 +876,8 @@ class TestCollectionQuery: batchdoc_and_check(full_collection, multiple_docs, doc_num, operator="insert") doc_fields, doc_vectors = generate_vectordict_random(full_collection.schema) - metrics = {field: MetricType.IP for field in weights} - weighted_reranker = WeightedReRanker(topn=3, weights=weights, metrics=metrics) + weight_list = [weights[v] for v in DEFAULT_VECTOR_FIELD_NAME.values()] + weighted_reranker = WeightedReRanker(weights=weight_list) single_query_results = {} for k, v in DEFAULT_VECTOR_FIELD_NAME.items(): @@ -1165,13 +1165,6 @@ class TestCollectionQuery: ), # Non-existent ID "Expected exception for non-existent document ID", ), - ( - "Both vector and id specified (invalid combination)", - lambda ref_dense_vector: Query( - field_name="vector_fp32_field", vector=ref_dense_vector, id="5" - ), - "Expected exception for specifying both vector and id", - ), ( "Neither vector nor id specified", lambda ref_dense_vector: Query( diff --git a/python/tests/test_collection.py b/python/tests/test_collection.py index cf94a4e..b16e2ee 100644 --- a/python/tests/test_collection.py +++ b/python/tests/test_collection.py @@ -976,18 +976,6 @@ class TestCollectionQuery: result = collection_with_multiple_docs.query(filter="id in (1)", topk=100) assert len(result) == 1 - def test_collection_query_with_vector_and_id( - self, collection_with_single_doc: Collection, single_doc: Doc - ): - with pytest.raises(ValueError): - collection_with_single_doc.query( - Query( - field_name="dense", - id=single_doc.id, - vector=single_doc.vector("dense"), - ) - ) - def test_collection_query_with_filter_not_in( self, collection_with_multiple_docs: Collection, multiple_docs ): @@ -1013,30 +1001,6 @@ class TestCollectionQuery: ) assert len(result) == 10 - def test_collection_query_multi_vector_with_same_field( - self, collection_with_multiple_docs: Collection, multiple_docs - ): - # Multi-vector query on same field without reranker should raise ValueError - with pytest.raises(ValueError, match="Reranker is required"): - collection_with_multiple_docs.query( - [ - Query(field_name="dense", vector=multiple_docs[0].vector("dense")), - Query(field_name="dense", vector=multiple_docs[1].vector("dense")), - ] - ) - - # Same field name with reranker should also raise ValueError - reranker = RrfReRanker(topn=10, rank_constant=60) - with pytest.raises(ValueError, match="appears more than once"): - collection_with_multiple_docs.query( - [ - Query(field_name="dense", vector=multiple_docs[0].vector("dense")), - Query(field_name="dense", vector=multiple_docs[1].vector("dense")), - ], - topk=10, - reranker=reranker, - ) - def test_collection_query_by_dense_vector( self, collection_with_multiple_docs: Collection, multiple_docs ): @@ -1087,7 +1051,7 @@ class TestCollectionQuery: self, collection_with_multiple_docs: Collection, multiple_docs ): """Test multi-vector query with RRF reranker on multiple dense vectors.""" - reranker = RrfReRanker(topn=10, rank_constant=60) + reranker = RrfReRanker(rank_constant=60) result = collection_with_multiple_docs.query( [ Query(field_name="dense", vector=multiple_docs[0].vector("dense")), @@ -1106,7 +1070,7 @@ class TestCollectionQuery: self, collection_with_multiple_docs: Collection, multiple_docs ): """Test multi-vector query with RRF reranker on multiple sparse vectors.""" - reranker = RrfReRanker(topn=10, rank_constant=60) + reranker = RrfReRanker(rank_constant=60) result = collection_with_multiple_docs.query( [ Query(field_name="sparse", vector=multiple_docs[0].vector("sparse")), @@ -1125,7 +1089,7 @@ class TestCollectionQuery: self, collection_with_multiple_docs: Collection, multiple_docs ): """Test multi-vector query with RRF reranker combining dense + sparse.""" - reranker = RrfReRanker(topn=10, rank_constant=60) + reranker = RrfReRanker(rank_constant=60) result = collection_with_multiple_docs.query( [ Query(field_name="dense", vector=multiple_docs[0].vector("dense")), @@ -1141,9 +1105,7 @@ class TestCollectionQuery: self, collection_with_multiple_docs: Collection, multiple_docs ): """Test multi-vector query with Weighted reranker on multiple dense vectors.""" - metrics = {"dense": MetricType.IP, "dense2": MetricType.IP} - weights = {"dense": 0.6, "dense2": 0.4} - reranker = WeightedReRanker(topn=10, metrics=metrics, weights=weights) + reranker = WeightedReRanker(weights=[0.6, 0.4]) result = collection_with_multiple_docs.query( [ Query(field_name="dense", vector=multiple_docs[0].vector("dense")), @@ -1159,9 +1121,7 @@ class TestCollectionQuery: self, collection_with_multiple_docs: Collection, multiple_docs ): """Test multi-vector query with Weighted reranker on multiple sparse vectors.""" - metrics = {"sparse": MetricType.IP, "sparse2": MetricType.IP} - weights = {"sparse": 0.6, "sparse2": 0.4} - reranker = WeightedReRanker(topn=10, metrics=metrics, weights=weights) + reranker = WeightedReRanker(weights=[0.6, 0.4]) result = collection_with_multiple_docs.query( [ Query(field_name="sparse", vector=multiple_docs[0].vector("sparse")), @@ -1180,9 +1140,7 @@ class TestCollectionQuery: self, collection_with_multiple_docs: Collection, multiple_docs ): """Test multi-vector query with Weighted reranker combining dense + sparse.""" - metrics = {"dense": MetricType.IP, "sparse": MetricType.IP} - weights = {"dense": 0.7, "sparse": 0.3} - reranker = WeightedReRanker(topn=10, metrics=metrics, weights=weights) + reranker = WeightedReRanker(weights=[0.7, 0.3]) result = collection_with_multiple_docs.query( [ Query(field_name="dense", vector=multiple_docs[0].vector("dense")), @@ -1203,7 +1161,7 @@ class TestCollectionQuery: def my_rerank_callback(query_results, topn): callback_invoked.append(True) all_docs = [] - for docs in query_results.values(): + for docs in query_results: all_docs.extend(docs) seen = set() unique_docs = [] @@ -1214,7 +1172,7 @@ class TestCollectionQuery: unique_docs.sort(key=lambda d: d.score(), reverse=True) return unique_docs[:topn] - reranker = CallbackReRanker(callback=my_rerank_callback, topn=10) + reranker = CallbackReRanker(callback=my_rerank_callback) result = collection_with_multiple_docs.query( [ Query(field_name="dense", vector=multiple_docs[0].vector("dense")), @@ -1234,7 +1192,7 @@ class TestCollectionQuery: def my_rerank_callback(query_results, topn): all_docs = [] - for docs in query_results.values(): + for docs in query_results: all_docs.extend(docs) seen = set() unique_docs = [] @@ -1245,7 +1203,7 @@ class TestCollectionQuery: unique_docs.sort(key=lambda d: d.score(), reverse=True) return unique_docs[:topn] - reranker = CallbackReRanker(callback=my_rerank_callback, topn=5) + reranker = CallbackReRanker(callback=my_rerank_callback) result = collection_with_multiple_docs.query( [ Query(field_name="dense", vector=multiple_docs[0].vector("dense")), diff --git a/python/tests/test_collection_fts.py b/python/tests/test_collection_fts.py index 55832a1..6cc884b 100644 --- a/python/tests/test_collection_fts.py +++ b/python/tests/test_collection_fts.py @@ -172,7 +172,7 @@ class TestFtsOnlyCollectionLifecycle: class TestFtsOnlyCollectionQueryValidation: def test_vector_query_rejected(self, fts_collection: Collection): """Vector query on a no-vector collection must raise.""" - with pytest.raises(ValueError, match="vector or id"): + with pytest.raises(ValueError, match="No vector field found"): fts_collection.query( queries=Query(field_name="content", vector=[0.1, 0.2, 0.3]), topk=5, @@ -181,7 +181,7 @@ class TestFtsOnlyCollectionQueryValidation: def test_id_query_rejected(self, fts_collection: Collection): """ID-based query on a no-vector collection must raise.""" fts_collection.insert(_make_docs()[:1]) - with pytest.raises(ValueError, match="vector or id"): + with pytest.raises(ValueError, match="No vector field found"): fts_collection.query( queries=Query(field_name="content", id="pk_0"), topk=5, diff --git a/python/tests/test_collection_fts_vector_hybrid.py b/python/tests/test_collection_fts_vector_hybrid.py index f57a727..0478b89 100644 --- a/python/tests/test_collection_fts_vector_hybrid.py +++ b/python/tests/test_collection_fts_vector_hybrid.py @@ -29,7 +29,6 @@ from zvec import ( ) from zvec.extension.multi_vector_reranker import RrfReRanker, WeightedReRanker from zvec.model.param.query import Fts, Query -from zvec.typing import MetricType DIM = 16 @@ -166,7 +165,7 @@ class TestFtsVectorHybridQuery: def test_hybrid_fts_and_vector_basic(self, hybrid_collection_with_docs: Collection): """FTS + vector multi-query with RRF reranker returns results.""" - reranker = RrfReRanker(topn=10, rank_constant=60) + reranker = RrfReRanker(rank_constant=60) result = hybrid_collection_with_docs.query( queries=[ Query(field_name="content", fts=Fts(match_string="retrieval")), @@ -185,7 +184,7 @@ class TestFtsVectorHybridQuery: self, hybrid_collection_with_docs: Collection ): """Docs relevant in both FTS and vector should rank higher.""" - reranker = RrfReRanker(topn=10, rank_constant=60) + reranker = RrfReRanker(rank_constant=60) # FTS: "retrieval search" matches pk_3, pk_4 # Vector: ret_vec cluster matches pk_3, pk_4 # Both signals agree: pk_3 and pk_4 should rank top @@ -202,7 +201,7 @@ class TestFtsVectorHybridQuery: def test_hybrid_scores_descending(self, hybrid_collection_with_docs: Collection): """Hybrid query results must be sorted by score descending.""" - reranker = RrfReRanker(topn=10, rank_constant=60) + reranker = RrfReRanker(rank_constant=60) result = hybrid_collection_with_docs.query( queries=[ Query(field_name="content", fts=Fts(match_string="intelligence")), @@ -217,7 +216,7 @@ class TestFtsVectorHybridQuery: def test_hybrid_with_filter(self, hybrid_collection_with_docs: Collection): """Hybrid query respects SQL filter.""" - reranker = RrfReRanker(topn=10, rank_constant=60) + reranker = RrfReRanker(rank_constant=60) result = hybrid_collection_with_docs.query( queries=[ Query(field_name="content", fts=Fts(match_string="learning")), @@ -234,7 +233,7 @@ class TestFtsVectorHybridQuery: self, hybrid_collection_with_docs: Collection ): """When FTS matches nothing, vector results still appear.""" - reranker = RrfReRanker(topn=10, rank_constant=60) + reranker = RrfReRanker(rank_constant=60) result = hybrid_collection_with_docs.query( queries=[ Query( @@ -251,7 +250,7 @@ class TestFtsVectorHybridQuery: def test_hybrid_query_string_syntax(self, hybrid_collection_with_docs: Collection): """Hybrid query works with FTS query_string (advanced syntax).""" - reranker = RrfReRanker(topn=10, rank_constant=60) + reranker = RrfReRanker(rank_constant=60) result = hybrid_collection_with_docs.query( queries=[ Query( @@ -283,35 +282,35 @@ class TestFtsVectorHybridValidation: topk=5, ) - def test_duplicate_field_name_rejected( + def test_duplicate_field_name_allowed( self, hybrid_collection_with_docs: Collection ): - """Multi-query with duplicate field names should raise.""" - reranker = RrfReRanker(topn=10, rank_constant=60) - with pytest.raises(ValueError, match="appears more than once"): - hybrid_collection_with_docs.query( - queries=[ - Query(field_name="content", fts=Fts(match_string="hello")), - Query(field_name="content", fts=Fts(match_string="world")), - ], - topk=5, - reranker=reranker, - ) + """Multi-query with duplicate field names is allowed and returns results.""" + reranker = RrfReRanker(rank_constant=60) + result = hybrid_collection_with_docs.query( + queries=[ + Query(field_name="content", fts=Fts(match_string="learning")), + Query(field_name="content", fts=Fts(match_string="intelligence")), + ], + topk=5, + reranker=reranker, + ) + assert len(result) > 0 + assert len(result) <= 5 - def test_multiple_vectors_without_fts_rejected( - self, hybrid_collection_with_docs: Collection - ): - """Two vector queries on a single-vector-field collection should raise.""" - reranker = RrfReRanker(topn=10, rank_constant=60) - with pytest.raises(ValueError, match="cannot query with multiple vectors"): - hybrid_collection_with_docs.query( - queries=[ - Query(field_name="embedding", vector=[1.0] * DIM), - Query(field_name="embedding", vector=[0.5] * DIM), - ], - topk=5, - reranker=reranker, - ) + def test_multiple_vectors_allowed(self, hybrid_collection_with_docs: Collection): + """Two vector queries on the same field are allowed with a reranker.""" + reranker = RrfReRanker(rank_constant=60) + result = hybrid_collection_with_docs.query( + queries=[ + Query(field_name="embedding", vector=[1.0] * DIM), + Query(field_name="embedding", vector=[0.5] * DIM), + ], + topk=5, + reranker=reranker, + ) + assert len(result) > 0 + assert len(result) <= 5 class TestFtsVectorHybridWeightedReranker: @@ -321,9 +320,8 @@ class TestFtsVectorHybridWeightedReranker: self, hybrid_collection_with_docs: Collection ): """WeightedReranker correctly normalizes FTS scores alongside vector scores.""" - metrics = {"embedding": MetricType.IP} - weights = {"content": 0.5, "embedding": 0.5} - reranker = WeightedReRanker(topn=10, metrics=metrics, weights=weights) + weights = [0.5, 0.5] + reranker = WeightedReRanker(weights=weights) result = hybrid_collection_with_docs.query( queries=[ Query(field_name="content", fts=Fts(match_string="retrieval search")), @@ -341,9 +339,8 @@ class TestFtsVectorHybridWeightedReranker: self, hybrid_collection_with_docs: Collection ): """WeightedReranker hybrid results are sorted by score descending.""" - metrics = {"embedding": MetricType.IP} - weights = {"content": 0.4, "embedding": 0.6} - reranker = WeightedReRanker(topn=10, metrics=metrics, weights=weights) + weights = [0.4, 0.6] + reranker = WeightedReRanker(weights=weights) result = hybrid_collection_with_docs.query( queries=[ Query(field_name="content", fts=Fts(match_string="intelligence")), @@ -361,11 +358,8 @@ class TestFtsVectorHybridWeightedReranker: ): """Higher FTS weight should boost FTS-relevant docs in ranking.""" # High FTS weight: FTS signal dominates - metrics = {"embedding": MetricType.IP} - weights_fts_heavy = {"content": 0.9, "embedding": 0.1} - reranker_fts = WeightedReRanker( - topn=10, metrics=metrics, weights=weights_fts_heavy - ) + weights_fts_heavy = [0.9, 0.1] + reranker_fts = WeightedReRanker(weights=weights_fts_heavy) result_fts = hybrid_collection_with_docs.query( queries=[ Query(field_name="content", fts=Fts(match_string="retrieval")), @@ -376,10 +370,8 @@ class TestFtsVectorHybridWeightedReranker: ) # High vector weight: vector signal dominates - weights_vec_heavy = {"content": 0.1, "embedding": 0.9} - reranker_vec = WeightedReRanker( - topn=10, metrics=metrics, weights=weights_vec_heavy - ) + weights_vec_heavy = [0.1, 0.9] + reranker_vec = WeightedReRanker(weights=weights_vec_heavy) result_vec = hybrid_collection_with_docs.query( queries=[ Query(field_name="content", fts=Fts(match_string="retrieval")), diff --git a/python/tests/test_fts_query.py b/python/tests/test_fts_query.py index 74cca6a..16db8b4 100644 --- a/python/tests/test_fts_query.py +++ b/python/tests/test_fts_query.py @@ -110,11 +110,11 @@ class TestFtsQueryBinding: assert restored.query_string == "+vector search" assert restored.match_string == "" - def test_vector_query_fts_field(self): - """_VectorQuery should have fts field.""" - from _zvec.param import _Fts, _VectorQuery + def test_search_query_fts_field(self): + """_SearchQuery should have fts field.""" + from _zvec.param import _Fts, _SearchQuery - vq = _VectorQuery() + vq = _SearchQuery() # fts should be None by default (optional) assert vq.fts is None @@ -125,11 +125,11 @@ class TestFtsQueryBinding: assert vq.fts is not None assert vq.fts.query_string == "hello" - def test_vector_query_pickle_with_fts(self): - """_VectorQuery with fts should survive pickling.""" - from _zvec.param import _Fts, _VectorQuery + def test_search_query_pickle_with_fts(self): + """_SearchQuery with fts should survive pickling.""" + from _zvec.param import _Fts, _SearchQuery - vq = _VectorQuery() + vq = _SearchQuery() vq.topk = 10 vq.field_name = "embedding" fts = _Fts() @@ -143,11 +143,11 @@ class TestFtsQueryBinding: assert restored.fts is not None assert restored.fts.match_string == "test query" - def test_vector_query_pickle_without_fts(self): - """_VectorQuery without fts should survive pickling.""" - from _zvec.param import _VectorQuery + def test_search_query_pickle_without_fts(self): + """_SearchQuery without fts should survive pickling.""" + from _zvec.param import _SearchQuery - vq = _VectorQuery() + vq = _SearchQuery() vq.topk = 5 vq.field_name = "vec" diff --git a/python/tests/test_params.py b/python/tests/test_params.py index b969101..a519138 100644 --- a/python/tests/test_params.py +++ b/python/tests/test_params.py @@ -40,7 +40,7 @@ from zvec import ( VectorSchema, ) -from _zvec.param import _VectorQuery +from _zvec.param import _SearchQuery # ---------------------------- # Invert Index Param Test Case diff --git a/python/tests/test_query_executor.py b/python/tests/test_query_executor.py index 6b2266e..823e6ef 100644 --- a/python/tests/test_query_executor.py +++ b/python/tests/test_query_executor.py @@ -14,20 +14,16 @@ from __future__ import annotations from typing import Dict, Union -from unittest.mock import MagicMock +from unittest.mock import MagicMock, patch import numpy as np import math -from _zvec.param import _VectorQuery +from _zvec.param import _SearchQuery import pytest from zvec.executor.query_executor import ( - MultiVectorQueryExecutor, - NoVectorQueryExecutor, QueryContext, QueryExecutor, - QueryExecutorFactory, - SingleVectorQueryExecutor, ) from zvec import ( RrfReRanker, @@ -43,21 +39,6 @@ from zvec import ( from zvec.extension.multi_vector_reranker import CallbackReRanker -# ---------------------------- -# Mock Vector Schema -# ---------------------------- -class MockVectorSchema(VectorSchema): - def __init__(self, name="test_vector"): - self._name = name - - @property - def name(self): - return self._name - - def _get_object(self): - return MagicMock() - - # ---------------------------- # Mock Collection Schema # ---------------------------- @@ -110,7 +91,7 @@ class TestQuery: assert query.has_vector() def test_validate_dense_fp16_convert(self): - v = _VectorQuery() + v = _SearchQuery() schema = VectorSchema(name="test", data_type=DataType.VECTOR_FP16) vec = np.array([1.1, 2.1, 3.1], dtype=np.float16) v.set_vector(schema._get_object(), vec) @@ -118,7 +99,7 @@ class TestQuery: assert np.array_equal(vec, ret) def test_validate_dense_fp32_convert(self): - v = _VectorQuery() + v = _SearchQuery() schema = VectorSchema(name="test", data_type=DataType.VECTOR_FP32) vec = np.array([1.1, 2.1, 3.1], dtype=np.float32) v.set_vector(schema._get_object(), vec) @@ -126,7 +107,7 @@ class TestQuery: assert np.array_equal(vec, ret) def test_validate_dense_fp64_convert(self): - v = _VectorQuery() + v = _SearchQuery() schema = VectorSchema(name="test", data_type=DataType.VECTOR_FP64) vec = np.array([1.1, 2.1, 3.1], dtype=np.float64) v.set_vector(schema._get_object(), vec) @@ -134,7 +115,7 @@ class TestQuery: assert np.array_equal(vec, ret) def test_validate_dense_int8_convert(self): - v = _VectorQuery() + v = _SearchQuery() schema = VectorSchema(name="test", data_type=DataType.VECTOR_INT8) vec = np.array([1, 2, 3], dtype=np.int8) v.set_vector(schema._get_object(), vec) @@ -142,7 +123,7 @@ class TestQuery: assert np.array_equal(vec, ret) def test_validate_sparse_fp32_convert(self): - v = _VectorQuery() + v = _SearchQuery() schema = VectorSchema(name="test", data_type=DataType.SPARSE_VECTOR_FP32) vec = {1: 1.1, 2: 2.2, 3: 3.3} v.set_vector(schema._get_object(), vec) @@ -151,7 +132,7 @@ class TestQuery: assert math.isclose(vec[k], ret[k], abs_tol=1e-6) def test_validate_sparse_fp16_convert(self): - v = _VectorQuery() + v = _SearchQuery() schema = VectorSchema(name="test", data_type=DataType.SPARSE_VECTOR_FP16) vec = {1: 1.1, 2: 2.2, 3: 3.3} v.set_vector(schema._get_object(), vec) @@ -189,7 +170,6 @@ class TestQueryContext: assert ctx.reranker is None assert ctx.output_fields is None assert ctx.include_vector is False - assert ctx.core_vectors == [] def test_properties(self): queries = [Query(field_name="test")] @@ -215,9 +195,7 @@ class TestQueryContext: def test_properties_with_weighted_reranker(self): queries = [Query(field_name="test")] reranker = WeightedReRanker( - topn=10, - metrics={"test": MetricType.L2}, - weights={"test": 1.0}, + weights=[1.0], ) ctx = QueryContext( @@ -227,13 +205,12 @@ class TestQueryContext: ) assert ctx.reranker == reranker - assert ctx.reranker.weights == {"test": 1.0} - assert ctx.reranker.metrics == {"test": MetricType.L2} + assert ctx.reranker.weights == [1.0] def test_properties_with_callback_reranker(self): queries = [Query(field_name="test")] cb = lambda query_results, topn: [] - reranker = CallbackReRanker(callback=cb, topn=10) + reranker = CallbackReRanker(callback=cb) ctx = QueryContext( topk=5, @@ -243,141 +220,83 @@ class TestQueryContext: assert ctx.reranker == reranker - def test_core_vectors_setter(self): - ctx = QueryContext(topk=10) - core_vectors = [MagicMock()] - ctx.core_vectors = core_vectors - assert ctx.core_vectors == core_vectors - -class TestNoVectorQueryExecutor: +class TestQueryExecutor: def test_init(self): schema = MockCollectionSchema() - executor = NoVectorQueryExecutor(schema) + executor = QueryExecutor(schema) assert isinstance(executor, QueryExecutor) - def test_do_validate_with_queries(self): + def test_do_build_without_queries(self): + # When no queries are given, build a single vector-less query. schema = MockCollectionSchema() - executor = NoVectorQueryExecutor(schema) - ctx = QueryContext( - topk=10, queries=[Query(field_name="test", vector=[0.1, 0.2, 0.3])] - ) - - with pytest.raises( - ValueError, match="Collection does not support query with vector or id" - ): - executor._do_validate(ctx) - - def test_do_validate_without_queries(self): - schema = MockCollectionSchema() - executor = NoVectorQueryExecutor(schema) - ctx = QueryContext(topk=10) - - executor._do_validate(ctx) - - def test_do_build(self): - schema = MockCollectionSchema() - executor = NoVectorQueryExecutor(schema) + executor = QueryExecutor(schema) ctx = QueryContext(topk=5, filter="test_filter") - result = executor._do_build(ctx, MagicMock()) + result = executor._build_queries(ctx, MagicMock()) assert len(result) == 1 assert result[0].topk == 5 assert result[0].filter == "test_filter" - -class TestSingleVectorQueryExecutor: - def test_init(self): + def test_do_build_query_wo_vector(self): + # Vector-less core query should carry the context query params. schema = MockCollectionSchema() - executor = SingleVectorQueryExecutor(schema) - assert isinstance(executor, NoVectorQueryExecutor) + executor = QueryExecutor(schema) + ctx = QueryContext(topk=7, filter="f", include_vector=True) - def test_do_validate_multiple_queries(self): + core_vector = executor._build_base_search_query(ctx) + assert core_vector.topk == 7 + assert core_vector.filter == "f" + assert core_vector.include_vector is True + + def test_do_merge_rerank_results_single_without_reranker(self): + # A single result list without a reranker is returned as-is. schema = MockCollectionSchema() - executor = SingleVectorQueryExecutor(schema) - queries = [Query(field_name="test1"), Query(field_name="test2")] - ctx = QueryContext(topk=10, queries=queries) + executor = QueryExecutor(schema) + ctx = QueryContext(topk=5) + docs_list = [["doc1", "doc2"]] - with pytest.raises( - ValueError, - match="Collection has only one vector field, cannot query with multiple vectors", - ): - executor._do_validate(ctx) + result = executor._merge_and_rerank(ctx, docs_list) + assert result == ["doc1", "doc2"] - def test_do_build_without_queries(self): + def test_do_merge_rerank_results_empty(self): + # Empty results should raise an error. schema = MockCollectionSchema() - executor = SingleVectorQueryExecutor(schema) + executor = QueryExecutor(schema) ctx = QueryContext(topk=5) - result = executor._do_build(ctx, MagicMock()) - assert len(result) == 1 - assert result[0].topk == 5 + with pytest.raises(ValueError, match="Query results is empty"): + executor._merge_and_rerank(ctx, []) - -class TestMultiVectorQueryExecutor: - def test_init(self): + def test_do_merge_rerank_results_with_reranker(self): + # Multiple result lists are merged through the reranker. schema = MockCollectionSchema() - executor = MultiVectorQueryExecutor(schema) - assert isinstance(executor, SingleVectorQueryExecutor) - - def test_do_validate_multiple_queries_without_reranker(self): - schema = MockCollectionSchema() - executor = MultiVectorQueryExecutor(schema) - queries = [Query(field_name="test1"), Query(field_name="test2")] - ctx = QueryContext(topk=10, queries=queries) - - with pytest.raises(ValueError, match="Reranker is required for multi-query"): - executor._do_validate(ctx) - - def test_do_validate_multiple_queries_with_reranker(self): - schema = MockCollectionSchema() - executor = MultiVectorQueryExecutor(schema) - queries = [Query(field_name="test1"), Query(field_name="test2")] - reranker = RrfReRanker() - ctx = QueryContext(topk=10, queries=queries, reranker=reranker) - - executor._do_validate(ctx) - - def test_do_validate_multiple_queries_with_weighted_reranker(self): - schema = MockCollectionSchema() - executor = MultiVectorQueryExecutor(schema) - queries = [Query(field_name="test1"), Query(field_name="test2")] - reranker = WeightedReRanker( - topn=10, - metrics={"test1": MetricType.L2, "test2": MetricType.L2}, - weights={"test1": 0.7, "test2": 0.3}, + executor = QueryExecutor(schema) + reranker = MagicMock() + reranker.rerank.return_value = ["merged"] + ctx = QueryContext( + topk=5, + queries=[Query(field_name="test1"), Query(field_name="test2")], + reranker=reranker, ) - ctx = QueryContext(topk=10, queries=queries, reranker=reranker) + docs_list = [["d1"], ["d2"]] - executor._do_validate(ctx) + result = executor._merge_and_rerank(ctx, docs_list) + assert result == ["merged"] + reranker.rerank.assert_called_once_with(docs_list, ctx.topk) - def test_do_validate_multiple_queries_with_callback_reranker(self): + def test_execute_python_pipeline(self): + # Each query is executed serially and converted into a result list. schema = MockCollectionSchema() - executor = MultiVectorQueryExecutor(schema) - queries = [Query(field_name="test1"), Query(field_name="test2")] - reranker = CallbackReRanker( - callback=lambda query_results, topn: [], - topn=10, - ) - ctx = QueryContext(topk=10, queries=queries, reranker=reranker) + executor = QueryExecutor(schema) + collection = MagicMock() + collection.Query.side_effect = [["raw1"], ["raw2"]] + vectors = [MagicMock(), MagicMock()] - executor._do_validate(ctx) - - -class TestQueryExecutorFactory: - def test_create_no_vectors(self): - schema = MockCollectionSchema() - executor = QueryExecutorFactory.create(schema) - assert isinstance(executor, NoVectorQueryExecutor) - - def test_create_single_vector(self): - schema = MockCollectionSchema(vectors=MockVectorSchema()) - executor = QueryExecutorFactory.create(schema) - assert isinstance(executor, SingleVectorQueryExecutor) - - def test_create_multiple_vectors(self): - schema = MockCollectionSchema( - vectors={"test1": MockVectorSchema(), "test2": MockVectorSchema()} - ) - executor = QueryExecutorFactory.create(schema) - assert isinstance(executor, MultiVectorQueryExecutor) + with patch( + "zvec.executor.query_executor.convert_to_py_doc", + side_effect=lambda doc, schema: doc, + ): + results = executor._execute_python_pipeline(vectors, collection) + assert results == [["raw1"], ["raw2"]] + assert collection.Query.call_count == 2 diff --git a/python/tests/test_reranker.py b/python/tests/test_reranker.py index 044c042..19350e7 100644 --- a/python/tests/test_reranker.py +++ b/python/tests/test_reranker.py @@ -15,10 +15,9 @@ from __future__ import annotations from unittest.mock import patch, MagicMock import pytest -import math import os -from zvec import Doc, MetricType +from zvec import Doc from zvec.extension.multi_vector_reranker import ( CallbackReRanker, RrfReRanker, @@ -38,37 +37,27 @@ RUN_INTEGRATION_TESTS = os.environ.get("ZVEC_RUN_INTEGRATION_TESTS", "0") == "1" # ---------------------------- class TestRrfReRanker: def test_init(self): - reranker = RrfReRanker(topn=5, rerank_field="content", rank_constant=100) - assert reranker.topn == 5 - assert reranker.rerank_field == "content" + reranker = RrfReRanker(rank_constant=100) assert reranker.rank_constant == 100 - def test_rrf_score(self): - reranker = RrfReRanker(rank_constant=60) - # 根据公式 1.0 / (k + rank + 1),其中k=60 - assert reranker._rrf_score(0) == 1.0 / (60 + 0 + 1) - assert reranker._rrf_score(1) == 1.0 / (60 + 1 + 1) - assert reranker._rrf_score(10) == 1.0 / (60 + 10 + 1) - - def test_rerank(self): - reranker = RrfReRanker(topn=3) + def test_rerank_delegates_to_cpp(self): + """RrfReRanker.rerank() delegates to C++ (raises TypeError with Python Docs).""" + reranker = RrfReRanker() doc1 = Doc(id="1", score=0.8) doc2 = Doc(id="2", score=0.7) doc3 = Doc(id="3", score=0.9) doc4 = Doc(id="4", score=0.6) - query_results = {"vector1": [doc1, doc2, doc3], "vector2": [doc3, doc1, doc4]} + query_results = [[doc1, doc2, doc3], [doc3, doc1, doc4]] - results = reranker.rerank(query_results) + with pytest.raises((TypeError, RuntimeError)): + reranker.rerank(query_results, topn=3) - assert len(results) <= reranker.topn - - for doc in results: - assert hasattr(doc, "score") - - scores = [doc.score for doc in results] - assert scores == sorted(scores, reverse=True) + def test_get_object_returns_cpp_reranker(self): + """_get_object() returns a valid C++ reranker instance.""" + reranker = RrfReRanker() + assert reranker._get_object() is not None # ---------------------------- @@ -76,64 +65,30 @@ class TestRrfReRanker: # ---------------------------- class TestWeightedReRanker: def test_init(self): - metrics = {"vector1": MetricType.L2, "vector2": MetricType.COSINE} - weights = {"vector1": 0.7, "vector2": 0.3} + weights = [0.7, 0.3] reranker = WeightedReRanker( - topn=5, - rerank_field="content", - metrics=metrics, weights=weights, ) - assert reranker.topn == 5 - assert reranker.rerank_field == "content" - assert reranker.metrics == metrics - assert reranker.weights == weights + assert list(reranker.weights) == weights - def test_normalize_score(self): - reranker = WeightedReRanker() - - score = reranker._normalize_score(1.0, MetricType.L2) - expected = 1.0 - 2 * math.atan(1.0) / math.pi - assert score == expected - - score = reranker._normalize_score(1.0, MetricType.IP) - expected = 0.5 + math.atan(1.0) / math.pi - assert score == expected - - score = reranker._normalize_score(1.0, MetricType.COSINE) - expected = 1.0 - 1.0 / 2.0 - assert score == expected - - with pytest.raises(ValueError, match="Unsupported metric type"): - reranker._normalize_score(1.0, "unsupported_metric") - - def test_rerank(self): - metrics = {"vector1": MetricType.L2, "vector2": MetricType.L2} - weights = {"vector1": 0.7, "vector2": 0.3} - reranker = WeightedReRanker(topn=3, weights=weights, metrics=metrics) + def test_rerank_delegates_to_cpp(self): + """WeightedReRanker.rerank() delegates to C++ (raises TypeError with Python Docs).""" + weights = [0.7, 0.3] + reranker = WeightedReRanker(weights=weights) doc1 = Doc(id="1", score=0.8) doc2 = Doc(id="2", score=0.7) doc3 = Doc(id="3", score=0.9) - query_results = {"vector1": [doc1, doc2], "vector2": [doc2, doc3]} + query_results = [[doc1, doc2], [doc2, doc3]] - results = reranker.rerank(query_results) + with pytest.raises((TypeError, RuntimeError)): + reranker.rerank(query_results, topn=3) - assert len(results) <= reranker.topn - - for doc in results: - assert hasattr(doc, "score") - - def test_rerank_missing_metric_raises(self): - metrics = {"vector1": MetricType.L2} - reranker = WeightedReRanker(topn=3, metrics=metrics) - - doc1 = Doc(id="1", score=0.8) - query_results = {"vector1": [doc1], "vector2": [doc1]} - - with pytest.raises(ValueError, match="no metric type specified"): - reranker.rerank(query_results) + def test_get_object_returns_cpp_reranker(self): + """_get_object() returns a valid C++ reranker instance.""" + reranker = WeightedReRanker(weights=[0.5, 0.5]) + assert reranker._get_object() is not None # ---------------------------- @@ -144,27 +99,27 @@ class TestCallbackReRanker: def my_callback(query_results, topn): return [] - reranker = CallbackReRanker(callback=my_callback, topn=5) - assert reranker.topn == 5 + reranker = CallbackReRanker(callback=my_callback) + assert reranker._get_object() is not None def test_rerank(self): def my_callback(query_results, topn): all_docs = [] - for docs in query_results.values(): + for docs in query_results: all_docs.extend(docs) all_docs.sort(key=lambda d: d.score, reverse=True) return all_docs[:topn] - reranker = CallbackReRanker(callback=my_callback, topn=3) + reranker = CallbackReRanker(callback=my_callback) doc1 = Doc(id="1", score=0.8) doc2 = Doc(id="2", score=0.9) doc3 = Doc(id="3", score=0.7) doc4 = Doc(id="4", score=0.6) - query_results = {"vector1": [doc1, doc2], "vector2": [doc3, doc4]} + query_results = [[doc1, doc2], [doc3, doc4]] - results = reranker.rerank(query_results) + results = reranker.rerank(query_results, topn=3) assert len(results) == 3 scores = [doc.score for doc in results] @@ -177,8 +132,8 @@ class TestCallbackReRanker: received_topn.append(topn) return [] - reranker = CallbackReRanker(callback=my_callback, topn=7) - reranker.rerank({"v1": [Doc(id="1", score=0.5)]}) + reranker = CallbackReRanker(callback=my_callback) + reranker.rerank([[Doc(id="1", score=0.5)]], topn=7) assert received_topn == [7] @@ -237,12 +192,6 @@ class TestQwenReRanker: ) assert reranker.query == "test query" - def test_topn_property(self): - reranker = QwenReRanker( - query="test", topn=5, api_key="test_key", rerank_field="content" - ) - assert reranker.topn == 5 - def test_rerank_field_property(self): reranker = QwenReRanker(query="test", api_key="test_key", rerank_field="title") assert reranker.rerank_field == "title" @@ -251,7 +200,7 @@ class TestQwenReRanker: reranker = QwenReRanker( query="test", api_key="test_key", rerank_field="content" ) - results = reranker.rerank({}) + results = reranker.rerank([], topn=10) assert results == [] def test_rerank_no_valid_documents(self): @@ -259,22 +208,22 @@ class TestQwenReRanker: query="test", api_key="test_key", rerank_field="content" ) # Document without the rerank_field - query_results = {"vector1": [Doc(id="1")]} + query_results = [[Doc(id="1")]] with pytest.raises(ValueError, match="No documents to rerank"): - reranker.rerank(query_results) + reranker.rerank(query_results, topn=10) def test_rerank_skip_empty_content(self): reranker = QwenReRanker( query="test", api_key="test_key", rerank_field="content" ) - query_results = { - "vector1": [ + query_results = [ + [ Doc(id="1", fields={"content": ""}), Doc(id="2", fields={"content": " "}), ] - } + ] with pytest.raises(ValueError, match="No documents to rerank"): - reranker.rerank(query_results) + reranker.rerank(query_results, topn=10) @patch("zvec.extension.qwen_function.require_module") def test_rerank_success(self, mock_require_module): @@ -294,17 +243,17 @@ class TestQwenReRanker: mock_dashscope.TextReRank.call.return_value = mock_response reranker = QwenReRanker( - query="test query", topn=2, api_key="test_key", rerank_field="content" + query="test query", api_key="test_key", rerank_field="content" ) - query_results = { - "vector1": [ + query_results = [ + [ Doc(id="1", fields={"content": "Document 1"}), Doc(id="2", fields={"content": "Document 2"}), ] - } + ] - results = reranker.rerank(query_results) + results = reranker.rerank(query_results, topn=2) assert len(results) == 2 assert results[0].id == "1" @@ -338,14 +287,14 @@ class TestQwenReRanker: mock_dashscope.TextReRank.call.return_value = mock_response reranker = QwenReRanker( - query="test", topn=5, api_key="test_key", rerank_field="content" + query="test", api_key="test_key", rerank_field="content" ) # Same document in multiple vector results doc1 = Doc(id="1", fields={"content": "Document 1"}) - query_results = {"vector1": [doc1], "vector2": [doc1]} + query_results = [[doc1], [doc1]] - results = reranker.rerank(query_results) + results = reranker.rerank(query_results, topn=5) # Should only call API with document once call_args = mock_dashscope.TextReRank.call.call_args @@ -368,10 +317,10 @@ class TestQwenReRanker: query="test", api_key="test_key", rerank_field="content" ) - query_results = {"vector1": [Doc(id="1", fields={"content": "Document 1"})]} + query_results = [[Doc(id="1", fields={"content": "Document 1"})]] with pytest.raises(ValueError, match="DashScope API error"): - reranker.rerank(query_results) + reranker.rerank(query_results, topn=10) @patch("zvec.extension.qwen_function.require_module") def test_rerank_runtime_error(self, mock_require_module): @@ -384,10 +333,10 @@ class TestQwenReRanker: query="test", api_key="test_key", rerank_field="content" ) - query_results = {"vector1": [Doc(id="1", fields={"content": "Document 1"})]} + query_results = [[Doc(id="1", fields={"content": "Document 1"})]] with pytest.raises(RuntimeError, match="Failed to call DashScope API"): - reranker.rerank(query_results) + reranker.rerank(query_results, topn=10) @pytest.mark.skipif( not RUN_INTEGRATION_TESTS, @@ -403,14 +352,13 @@ class TestQwenReRanker: # Create reranker with real API reranker = QwenReRanker( query="What is machine learning?", - topn=3, rerank_field="content", model="gte-rerank-v2", ) # Prepare test documents - query_results = { - "vector1": [ + query_results = [ + [ Doc( id="1", score=0.8, @@ -433,7 +381,7 @@ class TestQwenReRanker: }, ), ], - "vector2": [ + [ Doc( id="4", score=0.6, @@ -449,10 +397,10 @@ class TestQwenReRanker: }, ), ], - } + ] # Call real API - results = reranker.rerank(query_results) + results = reranker.rerank(query_results, topn=3) # Verify results assert len(results) <= 3, "Should return at most topn documents" @@ -523,13 +471,11 @@ class TestDefaultLocalReRanker: reranker = DefaultLocalReRanker( query="test query", - topn=5, rerank_field="content", model_name="cross-encoder/ms-marco-MiniLM-L6-v2", ) assert reranker.query == "test query" - assert reranker.topn == 5 assert reranker.rerank_field == "content" assert reranker.model_name == "cross-encoder/ms-marco-MiniLM-L6-v2" assert reranker.model_source == "huggingface" @@ -551,7 +497,6 @@ class TestDefaultLocalReRanker: reranker = DefaultLocalReRanker( query="custom query", - topn=10, rerank_field="title", model_name="cross-encoder/ms-marco-MiniLM-L12-v2", model_source="modelscope", @@ -560,7 +505,6 @@ class TestDefaultLocalReRanker: ) assert reranker.query == "custom query" - assert reranker.topn == 10 assert reranker.rerank_field == "title" assert reranker.model_name == "cross-encoder/ms-marco-MiniLM-L12-v2" assert reranker.model_source == "modelscope" @@ -593,23 +537,6 @@ class TestDefaultLocalReRanker: reranker = DefaultLocalReRanker(query="test query", rerank_field="content") assert reranker.query == "test query" - def test_topn_property(self): - """Test topn property.""" - mock_model = MagicMock() - mock_model.predict = MagicMock() - - mock_st = MagicMock() - mock_st.CrossEncoder.return_value = mock_model - - with patch( - "zvec.extension.sentence_transformer_rerank_function.require_module", - return_value=mock_st, - ): - reranker = DefaultLocalReRanker( - query="test", topn=15, rerank_field="content" - ) - assert reranker.topn == 15 - def test_rerank_field_property(self): """Test rerank_field property.""" mock_model = MagicMock() @@ -655,7 +582,7 @@ class TestDefaultLocalReRanker: return_value=mock_st, ): reranker = DefaultLocalReRanker(query="test", rerank_field="content") - results = reranker.rerank({}) + results = reranker.rerank([], topn=10) assert results == [] def test_rerank_no_valid_documents(self): @@ -673,9 +600,9 @@ class TestDefaultLocalReRanker: reranker = DefaultLocalReRanker(query="test", rerank_field="content") # Document without the rerank_field - query_results = {"vector1": [Doc(id="1")]} + query_results = [[Doc(id="1")]] with pytest.raises(ValueError, match="No documents to rerank"): - reranker.rerank(query_results) + reranker.rerank(query_results, topn=10) def test_rerank_skip_empty_content(self): """Test rerank skips documents with empty content.""" @@ -691,14 +618,14 @@ class TestDefaultLocalReRanker: ): reranker = DefaultLocalReRanker(query="test", rerank_field="content") - query_results = { - "vector1": [ + query_results = [ + [ Doc(id="1", fields={"content": ""}), Doc(id="2", fields={"content": " "}), ] - } + ] with pytest.raises(ValueError, match="No documents to rerank"): - reranker.rerank(query_results) + reranker.rerank(query_results, topn=10) def test_rerank_success(self): """Test successful rerank with mocked model.""" @@ -720,19 +647,17 @@ class TestDefaultLocalReRanker: "zvec.extension.sentence_transformer_rerank_function.require_module", return_value=mock_st, ): - reranker = DefaultLocalReRanker( - query="test query", topn=3, rerank_field="content" - ) + reranker = DefaultLocalReRanker(query="test query", rerank_field="content") - query_results = { - "vector1": [ + query_results = [ + [ Doc(id="1", score=0.8, fields={"content": "Document 1"}), Doc(id="2", score=0.7, fields={"content": "Document 2"}), Doc(id="3", score=0.6, fields={"content": "Document 3"}), ] - } + ] - results = reranker.rerank(query_results) + results = reranker.rerank(query_results, topn=3) # Verify results assert len(results) == 3 @@ -771,21 +696,19 @@ class TestDefaultLocalReRanker: "zvec.extension.sentence_transformer_rerank_function.require_module", return_value=mock_st, ): - reranker = DefaultLocalReRanker( - query="test", topn=2, rerank_field="content" - ) + reranker = DefaultLocalReRanker(query="test", rerank_field="content") - query_results = { - "vector1": [ + query_results = [ + [ Doc(id="1", fields={"content": "Doc 1"}), Doc(id="2", fields={"content": "Doc 2"}), Doc(id="3", fields={"content": "Doc 3"}), Doc(id="4", fields={"content": "Doc 4"}), Doc(id="5", fields={"content": "Doc 5"}), ] - } + ] - results = reranker.rerank(query_results) + results = reranker.rerank(query_results, topn=2) # Should only return top 2 assert len(results) == 2 @@ -811,20 +734,18 @@ class TestDefaultLocalReRanker: "zvec.extension.sentence_transformer_rerank_function.require_module", return_value=mock_st, ): - reranker = DefaultLocalReRanker( - query="test", topn=5, rerank_field="content" - ) + reranker = DefaultLocalReRanker(query="test", rerank_field="content") # Same document in multiple vector results doc1 = Doc(id="1", fields={"content": "Document 1"}) doc2 = Doc(id="2", fields={"content": "Document 2"}) - query_results = { - "vector1": [doc1, doc2], - "vector2": [doc1], # doc1 appears in both - } + query_results = [ + [doc1, doc2], + [doc1], # doc1 appears in both + ] - results = reranker.rerank(query_results) + results = reranker.rerank(query_results, topn=5) # Should only process each document once assert len(results) == 2 @@ -852,19 +773,17 @@ class TestDefaultLocalReRanker: "zvec.extension.sentence_transformer_rerank_function.require_module", return_value=mock_st, ): - reranker = DefaultLocalReRanker( - query="test", topn=3, rerank_field="content" - ) + reranker = DefaultLocalReRanker(query="test", rerank_field="content") - query_results = { - "vector1": [ + query_results = [ + [ Doc(id="1", fields={"content": "Doc 1"}), Doc(id="2", fields={"content": "Doc 2"}), Doc(id="3", fields={"content": "Doc 3"}), ] - } + ] - results = reranker.rerank(query_results) + results = reranker.rerank(query_results, topn=3) # Should be sorted by score (descending) assert len(results) == 3 @@ -892,10 +811,10 @@ class TestDefaultLocalReRanker: ): reranker = DefaultLocalReRanker(query="test", rerank_field="content") - query_results = {"vector1": [Doc(id="1", fields={"content": "Document 1"})]} + query_results = [[Doc(id="1", fields={"content": "Document 1"})]] with pytest.raises(RuntimeError, match="Failed to compute rerank scores"): - reranker.rerank(query_results) + reranker.rerank(query_results, topn=10) def test_rerank_with_custom_batch_size(self): """Test rerank uses custom batch_size.""" @@ -918,14 +837,14 @@ class TestDefaultLocalReRanker: query="test", rerank_field="content", batch_size=64 ) - query_results = { - "vector1": [ + query_results = [ + [ Doc(id="1", fields={"content": "Doc 1"}), Doc(id="2", fields={"content": "Doc 2"}), ] - } + ] - reranker.rerank(query_results) + reranker.rerank(query_results, topn=10) # Verify batch_size is passed to predict call_args = mock_model.predict.call_args @@ -947,13 +866,12 @@ class TestDefaultLocalReRanker: # Create reranker with real model (using default lightweight model) reranker = DefaultLocalReRanker( query="What is machine learning?", - topn=3, rerank_field="content", ) # Prepare test documents - query_results = { - "vector1": [ + query_results = [ + [ Doc( id="1", score=0.8, @@ -976,7 +894,7 @@ class TestDefaultLocalReRanker: }, ), ], - "vector2": [ + [ Doc( id="4", score=0.6, @@ -992,10 +910,10 @@ class TestDefaultLocalReRanker: }, ), ], - } + ] # Call real model - results = reranker.rerank(query_results) + results = reranker.rerank(query_results, topn=3) # Verify results assert len(results) <= 3, "Should return at most topn documents" @@ -1030,3 +948,49 @@ class TestDefaultLocalReRanker: content = doc.field("content") if content: print(f" Content: {content[:80]}...") + + +# ---------------------------- +# DocList Type and Delegation Tests +# ---------------------------- +class TestDocList: + def test_type_alias(self): + """DocList is list[Doc].""" + from zvec.model.doc import DocList + from zvec import Doc, DocList as QR + + assert DocList == list[Doc] + assert QR == list[Doc] + + def test_rrf_reranker_delegates_to_cpp(self): + """RrfReRanker.rerank() delegates to C++ (raises TypeError with Python Docs).""" + reranker = RrfReRanker() + with pytest.raises(TypeError): + reranker.rerank([[Doc(id="1", score=0.5)]], topn=5) + + def test_weighted_reranker_delegates_to_cpp(self): + """WeightedReRanker.rerank() delegates to C++ (raises TypeError with Python Docs).""" + reranker = WeightedReRanker(weights=[0.7, 0.3]) + with pytest.raises(TypeError): + reranker.rerank( + [[Doc(id="1", score=0.5)], [Doc(id="2", score=0.3)]], topn=5 + ) + + def test_single_route_query_results(self): + """CallbackReRanker works with single-route (one element list).""" + + def cb(query_results, topn): + return query_results[0][:topn] + + reranker = CallbackReRanker(callback=cb) + results = reranker.rerank( + [ + [ + Doc(id="1", score=0.9), + Doc(id="2", score=0.8), + Doc(id="3", score=0.7), + ] + ], + topn=2, + ) + assert len(results) == 2 diff --git a/python/zvec/__init__.py b/python/zvec/__init__.py index 655535e..c1bd167 100644 --- a/python/zvec/__init__.py +++ b/python/zvec/__init__.py @@ -71,7 +71,7 @@ from .model import schema as schema # —— Core data structures —— from .model.collection import Collection -from .model.doc import Doc +from .model.doc import Doc, DocList # —— Query & index parameters —— # —— FTS params (C++ binding) —— @@ -127,6 +127,7 @@ __all__ = [ # Core classes "Collection", "Doc", + "DocList", # Schema "CollectionSchema", "FieldSchema", diff --git a/python/zvec/__init__.pyi b/python/zvec/__init__.pyi index dd50a75..66665f6 100644 --- a/python/zvec/__init__.pyi +++ b/python/zvec/__init__.pyi @@ -11,7 +11,7 @@ from .extension import ReRanker, RrfReRanker, WeightedReRanker from .extension.embedding import DenseEmbeddingFunction from .model import param, schema from .model.collection import Collection -from .model.doc import Doc +from .model.doc import Doc, DocList from .model.param import ( AddColumnOption, AlterColumnOption, @@ -52,8 +52,8 @@ __all__: list = [ "CollectionStats", "DataType", "DenseEmbeddingFunction", - "DenseEmbeddingFunction", "Doc", + "DocList", "FieldSchema", "FlatIndexParam", "HnswIndexParam", @@ -72,7 +72,6 @@ __all__: list = [ "QuantizeType", "Query", "ReRanker", - "ReRanker", "RrfReRanker", "Status", "StatusCode", @@ -127,7 +126,7 @@ class _Collection: def Optimize(self, arg0: param.OptimizeOption) -> None: ... def Options(self) -> param.CollectionOption: ... def Path(self) -> str: ... - def Query(self, arg0: param._VectorQuery) -> list[_Doc]: ... + def Query(self, arg0: param._SearchQuery) -> list[_Doc]: ... def Schema(self) -> schema._CollectionSchema: ... def Stats(self) -> schema.CollectionStats: ... def Update(self, arg0: collections.abc.Sequence[_Doc]) -> list[typing.Status]: ... diff --git a/python/zvec/executor/__init__.py b/python/zvec/executor/__init__.py index 2ed019e..96582e3 100644 --- a/python/zvec/executor/__init__.py +++ b/python/zvec/executor/__init__.py @@ -16,11 +16,9 @@ from __future__ import annotations from .query_executor import ( QueryContext, QueryExecutor, - QueryExecutorFactory, ) __all__ = [ "QueryContext", "QueryExecutor", - "QueryExecutorFactory", ] diff --git a/python/zvec/executor/query_executor.py b/python/zvec/executor/query_executor.py index 663ebc5..08c98ca 100644 --- a/python/zvec/executor/query_executor.py +++ b/python/zvec/executor/query_executor.py @@ -13,18 +13,15 @@ # limitations under the License. from __future__ import annotations -import os -from abc import ABC, abstractmethod -from concurrent.futures import ThreadPoolExecutor, as_completed -from typing import Optional, Union, final +from typing import Optional, Union import numpy as np from _zvec import _Collection, _MultiQuery -from _zvec.param import _Fts, _SubQuery, _VectorQuery +from _zvec.param import _Fts, _SearchQuery, _SubQuery -from ..extension import ReRanker, RrfReRanker, WeightedReRanker +from ..extension import ReRanker from ..model.convert import convert_to_py_doc -from ..model.doc import Doc +from ..model.doc import DocList from ..model.param.query import Query from ..model.schema import CollectionSchema from ..typing import DataType @@ -32,7 +29,6 @@ from ..typing import DataType __all__ = [ "QueryContext", "QueryExecutor", - "QueryExecutorFactory", ] DTYPE_MAP = { @@ -80,9 +76,6 @@ class QueryContext: # reranker self._reranker = reranker - # core vectors - self._core_vectors = [] - @property def topk(self): return self._topk @@ -107,61 +100,120 @@ class QueryContext: def include_vector(self): return self._include_vector - @property - def core_vectors(self): - return self._core_vectors - @core_vectors.setter - def core_vectors(self, core_vectors: list[_VectorQuery]): - self._core_vectors = core_vectors +class QueryExecutor: + """Unified query executor that routes based on query count and reranker type.""" - -class QueryExecutor(ABC): def __init__(self, schema: CollectionSchema): self._schema = schema - self._concurrency = max(1, int(os.getenv("ZVEC_QUERY_CONCURRENCY", "1"))) - @abstractmethod - def _do_validate(self, ctx: QueryContext) -> None: - pass - - @abstractmethod - def _do_build( + def _build_queries( self, ctx: QueryContext, collection: _Collection - ) -> list[_VectorQuery]: - pass + ) -> list[_SearchQuery]: + """Build query vector list (no validation, conversion only).""" + if not ctx.queries: + return [self._build_base_search_query(ctx)] + return [ + self._build_search_query(ctx, query, collection) for query in ctx.queries + ] - def _do_build_query_wo_vector(self, ctx: QueryContext) -> _VectorQuery: - core_vector = _VectorQuery() - core_vector.topk = ctx.topk - core_vector.include_vector = ctx.include_vector + def execute(self, ctx: QueryContext, collection: _Collection) -> DocList: + """Execute a query, routing by query count. + + A single (or vector-less) query is sent to C++ as a ``_SearchQuery``; + multiple queries are assembled into a ``_MultiQuery``. + """ + queries = self._build_queries(ctx, collection) + if not queries: + raise ValueError("No query to execute") + + if len(queries) == 1: + return self._execute_single_query(queries[0], collection) + return self._execute_multi_query(ctx, queries, collection) + + def _execute_single_query( + self, query: _SearchQuery, collection: _Collection + ) -> DocList: + """Single/vector-less query: send a ``_SearchQuery`` to C++.""" + docs = collection.Query(query) + return [convert_to_py_doc(doc, self._schema) for doc in docs] + + def _execute_multi_query( + self, ctx: QueryContext, queries: list[_SearchQuery], collection: _Collection + ) -> DocList: + """Multiple queries: send a ``_MultiQuery`` to C++. + + A Python-only reranker (``_get_object()`` returns None) cannot run + inside the C++ MultiQuery, so each route is executed individually and + merged by the reranker in Python. + """ + reranker = ctx.reranker + if reranker is not None and reranker._get_object() is None: + docs_list = self._execute_python_pipeline(queries, collection) + return self._merge_and_rerank(ctx, docs_list) + + multi_query = self._build_multi_query(ctx, queries) + docs = collection.Query(multi_query) + return [convert_to_py_doc(doc, self._schema) for doc in docs] + + def _build_multi_query( + self, ctx: QueryContext, queries: list[_SearchQuery] + ) -> _MultiQuery: + """Assemble a C++ ``_MultiQuery`` from per-route ``_SearchQuery`` objects.""" + multi_query = _MultiQuery() + multi_query.queries = [_SubQuery.from_search_query(query) for query in queries] + multi_query.topk = ctx.topk if ctx.filter: - core_vector.filter = ctx.filter - if ctx.output_fields: - core_vector.output_fields = ctx.output_fields - return core_vector + multi_query.filter = ctx.filter + multi_query.include_vector = ctx.include_vector + if ctx.output_fields is not None: + multi_query.output_fields = ctx.output_fields + if ctx.reranker is not None: + multi_query.reranker = ctx.reranker._get_object() + return multi_query - def _do_build_fts_query(self, query: Query, core_vector: _VectorQuery) -> None: - """Set FTS query on core_vector if the query has FTS parameters.""" + def _execute_python_pipeline( + self, vectors: list[_SearchQuery], collection: _Collection + ) -> list[DocList]: + """Execute queries serially for the Python-only reranker path.""" + return [self._execute_single_query(query, collection) for query in vectors] + + def _merge_and_rerank(self, ctx: QueryContext, docs_list: list[DocList]) -> DocList: + """Merge and rerank results from the Python pipeline path.""" + if not docs_list: + raise ValueError("Query results is empty") + if len(docs_list) == 1 and not ctx.reranker: + return docs_list[0] + return ctx.reranker.rerank(docs_list, ctx.topk) + + def _build_base_search_query(self, ctx: QueryContext) -> _SearchQuery: + search_query = _SearchQuery() + search_query.topk = ctx.topk + search_query.include_vector = ctx.include_vector + if ctx.filter: + search_query.filter = ctx.filter + if ctx.output_fields is not None: + search_query.output_fields = ctx.output_fields + return search_query + + def _apply_fts(self, query: Query, search_query: _SearchQuery) -> None: + """Set FTS query on search_query if the query has FTS parameters.""" if query.has_fts(): fts = _Fts() fts.query_string = query.fts.query_string or "" fts.match_string = query.fts.match_string or "" - core_vector.fts = fts + search_query.fts = fts - def _do_build_query_with_vector( + def _build_search_query( self, ctx: QueryContext, query: Query, collection: _Collection - ) -> _VectorQuery: - core_vector = self._do_build_query_wo_vector(ctx) - core_vector.field_name = query.field_name + ) -> _SearchQuery: + search_query = self._build_base_search_query(ctx) + search_query.field_name = query.field_name if query.param: - core_vector.query_params = query.param + search_query.query_params = query.param # set FTS query if provided - self._do_build_fts_query(query, core_vector) - - # set output_fields - core_vector.output_fields = ctx.output_fields + self._apply_fts(query, search_query) vector_schema = None if query.has_vector() or query.has_id(): @@ -181,189 +233,14 @@ class QueryExecutor(ABC): fetched = collection.Fetch([query.id]) doc = next(iter(fetched.values())) if not doc: - return core_vector + raise ValueError(f"Document with id '{query.id}' not found") vec_data = doc.get_any(vector_schema.name, vector_schema.data_type) else: - return core_vector + return search_query target_dtype = DTYPE_MAP.get(vector_schema.data_type.value) - core_vector.set_vector( + search_query.set_vector( vector_schema._get_object(), convert_to_numpy(vec_data, target_dtype) if target_dtype else vec_data, ) - return core_vector - - def _do_execute( - self, vectors: list[_VectorQuery], collection: _Collection - ) -> dict[str, list[Doc]]: - query_cnt = len(vectors) - if query_cnt == 0: - raise ValueError("No query to execute") - - if len(vectors) == 1 or self._concurrency == 1: - results = {} - for query in vectors: - docs = collection.Query(query) - results[query.field_name] = [ - convert_to_py_doc(doc, self._schema) for doc in docs - ] - return results - - results = {} - with ThreadPoolExecutor(max_workers=self._concurrency) as executor: - future_to_query = { - executor.submit(collection.Query, query): query.field_name - for query in vectors - } - - for future in as_completed(future_to_query): - field_name = future_to_query[future] - try: - docs = future.result() - results[field_name] = [ - convert_to_py_doc(doc, self._schema) for doc in docs - ] - except Exception as e: - raise e - return results - - def _do_merge_rerank_results( - self, ctx: QueryContext, docs_map: dict[str, list[Doc]] - ) -> list[Doc]: - query_result_cnt = len(docs_map) if docs_map else 0 - if query_result_cnt == 0: - raise ValueError("Query results is none and dost not to rerank") - if query_result_cnt == 1: - if not ctx.reranker or isinstance( - ctx.reranker, (RrfReRanker, WeightedReRanker) - ): - return next(iter(docs_map.values())) - return ctx.reranker.rerank(docs_map) - return ctx.reranker.rerank(docs_map) - - @final - def execute(self, ctx: QueryContext, collection: _Collection) -> list[Doc]: - # 1. validate query - self._do_validate(ctx) - # 2. build query vector - query_vectors = self._do_build(ctx, collection) - if not query_vectors: - raise ValueError("No query to execute") - # 3. execute query - docs = self._do_execute(query_vectors, collection) - # 4. merge and rerank result - return self._do_merge_rerank_results(ctx, docs) - - -class NoVectorQueryExecutor(QueryExecutor): - def __init__(self, schema: CollectionSchema): - super().__init__(schema) - - def _do_validate(self, ctx: QueryContext) -> None: - for query in ctx.queries: - if query.has_vector() or query.has_id(): - raise ValueError("Collection does not support query with vector or id") - query._validate() - - def _do_build( - self, ctx: QueryContext, collection: _Collection - ) -> list[_VectorQuery]: - if len(ctx.queries) == 0: - return [self._do_build_query_wo_vector(ctx)] - # FTS-only branch in _do_build_query_with_vector skips vector resolution. - return [ - self._do_build_query_with_vector(ctx, query, collection) - for query in ctx.queries - ] - - -class SingleVectorQueryExecutor(NoVectorQueryExecutor): - def __init__(self, schema: CollectionSchema) -> None: - super().__init__(schema) - - def _validate_multi_query(self, ctx: QueryContext) -> None: - """Shared validation for multi-query: reranker required + no duplicate fields.""" - if ctx.reranker is None: - raise ValueError("Reranker is required for multi-query") - seen_fields = set() - for query in ctx.queries: - query._validate() - if query.field_name in seen_fields: - raise ValueError( - f"Query field name '{query.field_name}' appears more than once" - ) - seen_fields.add(query.field_name) - - def _do_validate(self, ctx: QueryContext) -> None: - if len(ctx.queries) > 1: - # Allow FTS + vector hybrid multi-query (requires reranker) - if not any(q.has_fts() for q in ctx.queries): - raise ValueError( - "Collection has only one vector field, cannot query with multiple vectors" - ) - self._validate_multi_query(ctx) - return - for query in ctx.queries: - query._validate() - - def _do_build( - self, ctx: QueryContext, collection: _Collection - ) -> list[_VectorQuery]: - if len(ctx.queries) == 0: - return [self._do_build_query_wo_vector(ctx)] - vectors = [] - for query in ctx.queries: - vectors.append(self._do_build_query_with_vector(ctx, query, collection)) - return vectors - - def execute(self, ctx: QueryContext, collection: _Collection) -> list[Doc]: - # 1. validate query - self._do_validate(ctx) - # 2. build query vectors - query_vectors = self._do_build(ctx, collection) - if not query_vectors: - raise ValueError("No query to execute") - - # Multi-query fast path: route FTS + vector hybrid to C++ MultiQuery - if len(query_vectors) > 1 and ctx.reranker is not None: - cpp_reranker = ctx.reranker._get_object() - if cpp_reranker is not None: - mvq = _MultiQuery() - mvq.queries = [_SubQuery.from_vector_query(vq) for vq in query_vectors] - mvq.topk = ctx.topk - if ctx.filter: - mvq.filter = ctx.filter - mvq.include_vector = ctx.include_vector - if ctx.output_fields: - mvq.output_fields = ctx.output_fields - mvq.reranker = cpp_reranker - docs = collection.Query(mvq) - return [convert_to_py_doc(doc, self._schema) for doc in docs] - - # 3. execute query - docs = self._do_execute(query_vectors, collection) - # 4. merge and rerank result - return self._do_merge_rerank_results(ctx, docs) - - -class MultiVectorQueryExecutor(SingleVectorQueryExecutor): - def __init__(self, schema: CollectionSchema) -> None: - super().__init__(schema) - - def _do_validate(self, ctx: QueryContext) -> None: - if len(ctx.queries) > 1: - self._validate_multi_query(ctx) - return - for query in ctx.queries: - query._validate() - - -class QueryExecutorFactory: - @staticmethod - def create(schema: CollectionSchema) -> QueryExecutor: - vectors = schema.vectors - if len(vectors) == 0: - return NoVectorQueryExecutor(schema) - if len(vectors) == 1: - return SingleVectorQueryExecutor(schema) - return MultiVectorQueryExecutor(schema) + return search_query diff --git a/python/zvec/extension/multi_vector_reranker.py b/python/zvec/extension/multi_vector_reranker.py index ecb6407..e96de17 100644 --- a/python/zvec/extension/multi_vector_reranker.py +++ b/python/zvec/extension/multi_vector_reranker.py @@ -13,16 +13,12 @@ # limitations under the License. from __future__ import annotations -import heapq -import math -from collections import defaultdict from collections.abc import Callable from typing import Optional from _zvec import _CallbackReranker, _RrfReranker, _WeightedReranker -from ..model.doc import Doc -from ..typing import MetricType +from ..model.doc import DocList from .rerank_function import RerankFunction @@ -35,24 +31,15 @@ class RrfReRanker(RerankFunction): The RRF score for a document at rank ``r`` is: ``1 / (k + r + 1)``, where ``k`` is the rank constant. - Note: - This re-ranker is specifically designed for multi-vector scenarios where - query results from multiple vector fields need to be combined. - Args: - topn (int, optional): Number of top documents to return. Defaults to 10. - rerank_field (Optional[str], optional): Ignored by RRF. Defaults to None. rank_constant (int, optional): Smoothing constant ``k`` in RRF formula. Larger values reduce the impact of early ranks. Defaults to 60. """ def __init__( self, - topn: int = 10, - rerank_field: Optional[str] = None, rank_constant: int = 60, ): - super().__init__(topn=topn, rerank_field=rerank_field) self._rank_constant = rank_constant # Use C++ implementation for performance self._cpp_reranker = _RrfReranker(rank_constant) @@ -65,36 +52,18 @@ class RrfReRanker(RerankFunction): """Return the underlying C++ RrfReranker instance.""" return self._cpp_reranker - def _rrf_score(self, rank: int) -> float: - return 1.0 / (self._rank_constant + rank + 1) - - def rerank(self, query_results: dict[str, list[Doc]]) -> list[Doc]: - """Apply Reciprocal Rank Fusion to combine multiple query results. + def rerank(self, query_results: list[DocList], topn: int) -> DocList: + """Re-rank using C++ RRF implementation. Args: - query_results (dict[str, list[Doc]]): Results from one or more vector queries. + query_results (list[DocList]): Multi-route recall results, + positionally aligned with queries. + topn (int): Number of top documents to return. Returns: - list[Doc]: Re-ranked documents with RRF scores in the ``score`` field. + DocList: Re-ranked documents. """ - rrf_scores: dict[str, float] = defaultdict(float) - id_to_doc: dict[str, Doc] = {} - - for _, query_result in query_results.items(): - for rank, doc in enumerate(query_result): - doc_id = doc.id - rrf_score = self._rrf_score(rank) - rrf_scores[doc_id] += rrf_score - if doc_id not in id_to_doc: - id_to_doc[doc_id] = doc - - top_docs = heapq.nlargest(self.topn, rrf_scores.items(), key=lambda x: x[1]) - results: list[Doc] = [] - for doc_id, rrf_score in top_docs: - doc = id_to_doc[doc_id] - new_doc = doc._replace(score=rrf_score) - results.append(new_doc) - return results + return self._cpp_reranker.rerank(query_results, topn) class WeightedReRanker(RerankFunction): @@ -102,98 +71,41 @@ class WeightedReRanker(RerankFunction): Each vector field's relevance score is normalized based on its own metric type, then scaled by a user-provided weight. Final scores are summed across - fields. - - Note: - This re-ranker is specifically designed for multi-vector scenarios where - query results from multiple vector fields need to be combined with - configurable weights. + fields. The actual re-ranking logic lives in the C++ implementation. Args: - topn (int, optional): Number of top documents to return. Defaults to 10. - rerank_field (Optional[str], optional): Ignored. Defaults to None. - metrics (Optional[dict[str, MetricType]], optional): Per-field distance - metric used for score normalization. Every queried field must have - a metric specified; missing fields will raise an error at rerank time. - Defaults to None. - weights (Optional[dict[str, float]], optional): Weight per vector field. - Fields not listed use weight 1.0. Defaults to None. - - Note: - Supported metrics: L2, IP, COSINE. Scores are normalized to [0, 1]. + weights (Optional[list[float]], optional): Weight per vector field, + aligned by position with the queries supplied to ``collection.query()``. + Defaults to None (treated as an empty list). """ def __init__( self, - topn: int = 10, - rerank_field: Optional[str] = None, - metrics: Optional[dict[str, MetricType]] = None, - weights: Optional[dict[str, float]] = None, + weights: Optional[list[float]] = None, ): - super().__init__(topn=topn, rerank_field=rerank_field) - self._weights = weights or {} - self._metrics = metrics or {} + self._cpp_reranker = _WeightedReranker(weights or []) @property - def weights(self) -> dict[str, float]: - """dict[str, float]: Weight mapping for vector fields.""" - return self._weights - - @property - def metrics(self) -> dict[str, MetricType]: - """dict[str, MetricType]: Per-field metric type mapping.""" - return self._metrics + def weights(self) -> list[float]: + """list[float]: Weight list for vector fields, aligned with queries.""" + return self._cpp_reranker.weights def _get_object(self): - """Return a C++ WeightedReranker instance.""" - return _WeightedReranker(self._weights) + """Return the underlying C++ WeightedReranker instance.""" + return self._cpp_reranker - def rerank(self, query_results: dict[str, list[Doc]]) -> list[Doc]: - """Combine scores from multiple vector fields using weighted sum. + def rerank(self, query_results: list[DocList], topn: int) -> DocList: + """Re-rank using C++ Weighted implementation. Args: - query_results (dict[str, list[Doc]]): Results per vector field. + query_results (list[DocList]): Multi-route recall results, + positionally aligned with queries. + topn (int): Number of top documents to return. Returns: - list[Doc]: Re-ranked documents with combined scores in ``score`` field. + DocList: Re-ranked documents. """ - weighted_scores: dict[str, float] = defaultdict(float) - id_to_doc: dict[str, Doc] = {} - - for vector_name, query_result in query_results.items(): - if vector_name not in self._metrics: - raise ValueError( - f"WeightedReRanker: no metric type specified for field " - f"'{vector_name}'" - ) - metric = self._metrics[vector_name] - for _, doc in enumerate(query_result): - doc_id = doc.id - weighted_score = self._normalize_score( - doc.score, metric - ) * self.weights.get(vector_name, 1.0) - weighted_scores[doc_id] += weighted_score - if doc_id not in id_to_doc: - id_to_doc[doc_id] = doc - - top_docs = heapq.nlargest( - self.topn, weighted_scores.items(), key=lambda x: x[1] - ) - results: list[Doc] = [] - for doc_id, weighted_score in top_docs: - doc = id_to_doc[doc_id] - new_doc = doc._replace(score=weighted_score) - results.append(new_doc) - return results - - def _normalize_score(self, score: float, metric: MetricType) -> float: - if metric == MetricType.L2: - return 1.0 - 2 * math.atan(score) / math.pi - if metric == MetricType.IP: - return 0.5 + math.atan(score) / math.pi - if metric == MetricType.COSINE: - return 1.0 - score / 2.0 - raise ValueError("Unsupported metric type") + return self._cpp_reranker.rerank(query_results, topn) class CallbackReRanker(RerankFunction): @@ -202,21 +114,18 @@ class CallbackReRanker(RerankFunction): This bridges a Python callable into the C++ reranker interface, enabling custom re-ranking logic to be executed within the C++ MultiQuery path. - The callback receives the raw C++ Doc objects (as ``_Doc`` instances) grouped - by vector field name, and must return a list of ``_Doc`` instances. + The callback receives raw C++ ``_Doc`` objects grouped per query (as a + ``list[list[_Doc]]``) and must return a ``list[_Doc]``. Args: callback: A callable with signature - ``(query_results: dict[str, list[_Doc]], topn: int) -> list[_Doc]``. - topn (int, optional): Number of top documents to return. Defaults to 10. + ``(query_results: list[list[_Doc]], topn: int) -> list[_Doc]``. """ def __init__( self, callback: Callable, - topn: int = 10, ): - super().__init__(topn=topn) self._callback = callback self._cpp_reranker = _CallbackReranker(callback) @@ -224,13 +133,15 @@ class CallbackReRanker(RerankFunction): """Return the underlying C++ CallbackReranker instance.""" return self._cpp_reranker - def rerank(self, query_results: dict[str, list[Doc]]) -> list[Doc]: + def rerank(self, query_results: list[DocList], topn: int) -> DocList: """Invoke the callback to re-rank documents. Args: - query_results (dict[str, list[Doc]]): Results per vector field. + query_results (list[DocList]): Multi-route recall results, + positionally aligned with queries. + topn (int): Number of top documents to return. Returns: - list[Doc]: Re-ranked documents. + DocList: Re-ranked documents. """ - return self._callback(query_results, self.topn) + return self._callback(query_results, topn) diff --git a/python/zvec/extension/qwen_rerank_function.py b/python/zvec/extension/qwen_rerank_function.py index 9b4a66b..ead1f9e 100644 --- a/python/zvec/extension/qwen_rerank_function.py +++ b/python/zvec/extension/qwen_rerank_function.py @@ -15,7 +15,7 @@ from __future__ import annotations from typing import Optional -from ..model.doc import Doc +from ..model.doc import Doc, DocList from .qwen_function import QwenFunctionBase from .rerank_function import RerankFunction @@ -32,8 +32,6 @@ class QwenReRanker(QwenFunctionBase, RerankFunction): Args: query (str): Query text for semantic re-ranking. **Required**. - topn (int, optional): Maximum number of documents to return after re-ranking. - Defaults to 10. rerank_field (str): Document field name to use as re-ranking input text. **Required** (e.g., "content", "title", "body"). model (str, optional): DashScope re-ranking model identifier. @@ -53,7 +51,6 @@ class QwenReRanker(QwenFunctionBase, RerankFunction): Example: >>> reranker = QwenReRanker( ... query="machine learning algorithms", - ... topn=5, ... rerank_field="content", ... model="gte-rerank-v2", ... api_key="your-api-key" @@ -64,7 +61,6 @@ class QwenReRanker(QwenFunctionBase, RerankFunction): def __init__( self, query: Optional[str] = None, - topn: int = 10, rerank_field: Optional[str] = None, model: str = "gte-rerank-v2", api_key: Optional[str] = None, @@ -73,7 +69,6 @@ class QwenReRanker(QwenFunctionBase, RerankFunction): Args: query (Optional[str]): Query text for semantic matching. Required. - topn (int): Number of top results to return. rerank_field (Optional[str]): Document field for re-ranking input. model (str): DashScope model name. api_key (Optional[str]): API key or None to use environment variable. @@ -82,37 +77,45 @@ class QwenReRanker(QwenFunctionBase, RerankFunction): ValueError: If query is empty or API key is unavailable. """ QwenFunctionBase.__init__(self, model=model, api_key=api_key) - RerankFunction.__init__(self, topn=topn, rerank_field=rerank_field) + RerankFunction.__init__(self) if not query: raise ValueError("Query is required for QwenReRanker") self._query = query + self._rerank_field = rerank_field @property def query(self) -> str: """str: Query text used for semantic re-ranking.""" return self._query - def rerank(self, query_results: dict[str, list[Doc]]) -> list[Doc]: + @property + def rerank_field(self) -> Optional[str]: + """Optional[str]: Field name used as re-ranking input.""" + return self._rerank_field + + def rerank(self, query_results: list[DocList], topn: int) -> DocList: """Re-rank documents using Qwen's TextReRank API. Sends document texts to DashScope TextReRank service along with the query. Returns documents sorted by relevance scores from the cross-encoder model. Args: - query_results (dict[str, list[Doc]]): Mapping from vector field names - to lists of retrieved documents. Documents from all fields are + query_results (list[DocList]): Multi-route recall results, + positionally aligned with the queries supplied to + ``collection.query()``. Documents from all routes are deduplicated and re-ranked together. + topn (int): Maximum number of documents to return after re-ranking. Returns: - list[Doc]: Re-ranked documents (up to ``topn``) with updated ``score`` - fields containing relevance scores from the API. + DocList: Re-ranked documents (up to ``topn``) with updated + ``score`` fields containing relevance scores from the API. Raises: ValueError: If no valid documents are found or API call fails. Note: - - Duplicate documents (same ID) across fields are processed once + - Duplicate documents (same ID) across routes are processed once - Documents with empty/missing ``rerank_field`` content are skipped - Returned scores are relevance scores from the cross-encoder model """ @@ -124,7 +127,7 @@ class QwenReRanker(QwenFunctionBase, RerankFunction): doc_ids: list[str] = [] contents: list[str] = [] - for _, query_result in query_results.items(): + for query_result in query_results: for doc in query_result: doc_id = doc.id if doc_id in id_to_doc: @@ -147,11 +150,11 @@ class QwenReRanker(QwenFunctionBase, RerankFunction): output = self._call_rerank_api( query=self.query, documents=contents, - top_n=self.topn, + top_n=topn, ) # Build result list with updated scores - results: list[Doc] = [] + results: DocList = [] for item in output["results"]: idx = item["index"] doc_id = doc_ids[idx] diff --git a/python/zvec/extension/rerank_function.py b/python/zvec/extension/rerank_function.py index 0d8d002..54fe161 100644 --- a/python/zvec/extension/rerank_function.py +++ b/python/zvec/extension/rerank_function.py @@ -14,9 +14,8 @@ from __future__ import annotations from abc import ABC, abstractmethod -from typing import Optional -from ..model.doc import Doc +from ..model.doc import DocList class RerankFunction(ABC): @@ -26,44 +25,22 @@ class RerankFunction(ABC): a secondary scoring strategy. They are used in the ``query()`` method of ``Collection`` via the ``reranker`` parameter. - Args: - topn (int, optional): Number of top documents to return after re-ranking. - Defaults to 10. - rerank_field (Optional[str], optional): Field name used as input for - re-ranking (e.g., document title or body). Defaults to None. - Note: Subclasses must implement the ``rerank()`` method. """ - def __init__( - self, - topn: int = 10, - rerank_field: Optional[str] = None, - ): - self._topn = topn - self._rerank_field = rerank_field - - @property - def topn(self) -> int: - """int: Number of top documents to return after re-ranking.""" - return self._topn - - @property - def rerank_field(self) -> Optional[str]: - """Optional[str]: Field name used as re-ranking input.""" - return self._rerank_field - @abstractmethod - def rerank(self, query_results: dict[str, list[Doc]]) -> list[Doc]: - """Re-rank documents from one or more vector queries. + def rerank(self, query_results: list[DocList], topn: int) -> DocList: + """Re-rank documents from multi-route recall results. Args: - query_results (dict[str, list[Doc]]): Mapping from vector field name - to list of retrieved documents (sorted by relevance). + query_results (list[DocList]): List of query results from + multi-route recall. Each element corresponds to a Query in the + collection.query(queries=List[Query]) call, aligned by position. + topn (int): Number of top documents to return after re-ranking. Returns: - list[Doc]: Re-ranked list of documents (length ≤ ``topn``), + DocList: Re-ranked list of documents (length ≤ ``topn``), with updated ``score`` fields. """ ... diff --git a/python/zvec/extension/sentence_transformer_rerank_function.py b/python/zvec/extension/sentence_transformer_rerank_function.py index 58c5838..2e22d7c 100644 --- a/python/zvec/extension/sentence_transformer_rerank_function.py +++ b/python/zvec/extension/sentence_transformer_rerank_function.py @@ -15,7 +15,7 @@ from __future__ import annotations from typing import Literal, Optional -from ..model.doc import Doc +from ..model.doc import Doc, DocList from ..tool import require_module from .rerank_function import RerankFunction from .sentence_transformer_function import SentenceTransformerFunctionBase @@ -33,8 +33,6 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction): Args: query (str): Query text for semantic re-ranking. **Required**. - topn (int, optional): Maximum number of documents to return after re-ranking. - Defaults to 10. rerank_field (Optional[str], optional): Document field name to use as re-ranking input text. **Required** (e.g., "content", "title", "body"). model_name (str, optional): Cross-encoder model identifier or local path. @@ -56,7 +54,6 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction): Attributes: query (str): The query text used for re-ranking. - topn (int): Maximum number of documents to return. rerank_field (Optional[str]): Field name used for re-ranking input. model_name (str): The cross-encoder model being used. model_source (str): The model source ("huggingface" or "modelscope"). @@ -113,7 +110,6 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction): >>> >>> reranker = SentenceTransformerReRanker( ... query="machine learning algorithms", - ... topn=5, ... rerank_field="content" ... ) >>> @@ -127,7 +123,6 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction): >>> # Using ModelScope for users in China >>> reranker = SentenceTransformerReRanker( ... query="深度学习", - ... topn=10, ... rerank_field="content", ... model_source="modelscope" ... ) @@ -135,7 +130,6 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction): >>> # Using larger model for better quality >>> reranker = SentenceTransformerReRanker( ... query="neural networks", - ... topn=5, ... rerank_field="content", ... model_name="BAAI/bge-reranker-large", ... device="cuda", @@ -143,13 +137,13 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction): ... ) >>> # Direct rerank call (for testing) - >>> query_results = { - ... "vector1": [ + >>> query_results = [ + ... [ ... Doc(id="1", score=0.9, fields={"content": "Machine learning is..."}), ... Doc(id="2", score=0.8, fields={"content": "Deep learning is..."}), ... ] - ... } - >>> reranked = reranker.rerank(query_results) + ... ] + >>> reranked = reranker.rerank(query_results, topn=5) >>> for doc in reranked: ... print(f"ID: {doc.id}, Score: {doc.score:.4f}") ID: 2, Score: 0.9234 @@ -170,7 +164,6 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction): def __init__( self, query: Optional[str] = None, - topn: int = 10, rerank_field: Optional[str] = None, model_name: str = "cross-encoder/ms-marco-MiniLM-L6-v2", model_source: Literal["huggingface", "modelscope"] = "huggingface", @@ -181,7 +174,6 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction): Args: query (Optional[str]): Query text for semantic matching. Required. - topn (int): Number of top results to return. rerank_field (Optional[str]): Document field for re-ranking input. model_name (str): Cross-encoder model identifier. model_source (Literal["huggingface", "modelscope"]): Model source. @@ -197,12 +189,13 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction): ) # Initialize rerank function - RerankFunction.__init__(self, topn=topn, rerank_field=rerank_field) + RerankFunction.__init__(self) # Validate query if not query: raise ValueError("Query is required for DefaultLocalReRanker") self._query = query + self._rerank_field = rerank_field self._batch_size = batch_size # Load and validate cross-encoder model @@ -273,12 +266,17 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction): """str: Query text used for semantic re-ranking.""" return self._query + @property + def rerank_field(self) -> Optional[str]: + """Optional[str]: Field name used as re-ranking input.""" + return self._rerank_field + @property def batch_size(self) -> int: """int: Batch size for processing query-document pairs.""" return self._batch_size - def rerank(self, query_results: dict[str, list[Doc]]) -> list[Doc]: + def rerank(self, query_results: list[DocList], topn: int) -> DocList: """Re-rank documents using Sentence Transformer cross-encoder model. Evaluates each query-document pair using the cross-encoder model to compute @@ -286,19 +284,22 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction): results are returned. Args: - query_results (dict[str, list[Doc]]): Mapping from vector field names - to lists of retrieved documents. Documents from all fields are + query_results (list[DocList]): Multi-route recall results, + positionally aligned with the queries supplied to + ``collection.query()``. Documents from all routes are deduplicated and re-ranked together. + topn (int): Maximum number of documents to return after re-ranking. Returns: - list[Doc]: Re-ranked documents (up to ``topn``) with updated ``score`` - fields containing relevance scores from the cross-encoder model. + DocList: Re-ranked documents (up to ``topn``) with updated + ``score`` fields containing relevance scores from the + cross-encoder model. Raises: ValueError: If no valid documents are found or model inference fails. Note: - - Duplicate documents (same ID) across fields are processed once + - Duplicate documents (same ID) across routes are processed once - Documents with empty/missing ``rerank_field`` content are skipped - Returned scores are logits from the cross-encoder model - Higher scores indicate higher relevance @@ -310,13 +311,13 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction): ... topn=3, ... rerank_field="content" ... ) - >>> query_results = { - ... "vector1": [ + >>> query_results = [ + ... [ ... Doc(id="1", score=0.9, fields={"content": "ML basics"}), ... Doc(id="2", score=0.8, fields={"content": "DL tutorial"}), ... ] - ... } - >>> reranked = reranker.rerank(query_results) + ... ] + >>> reranked = reranker.rerank(query_results, topn=3) >>> len(reranked) <= 3 True """ @@ -328,7 +329,7 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction): doc_ids: list[str] = [] contents: list[str] = [] - for _, query_result in query_results.items(): + for query_result in query_results: for doc in query_result: doc_id = doc.id if doc_id in id_to_doc: @@ -373,10 +374,10 @@ class DefaultLocalReRanker(SentenceTransformerFunctionBase, RerankFunction): # Sort by score (descending) and take top-k scored_docs.sort(key=lambda x: x[2], reverse=True) - top_scored_docs = scored_docs[: self.topn] + top_scored_docs = scored_docs[:topn] # Build result list with updated scores - results: list[Doc] = [] + results: DocList = [] for _, doc, score in top_scored_docs: new_doc = doc._replace(score=score) results.append(new_doc) diff --git a/python/zvec/model/collection.py b/python/zvec/model/collection.py index 6262581..de16753 100644 --- a/python/zvec/model/collection.py +++ b/python/zvec/model/collection.py @@ -18,11 +18,11 @@ from typing import Optional, Union, overload from _zvec import _Collection -from ..executor import QueryContext, QueryExecutorFactory +from ..executor import QueryContext, QueryExecutor from ..extension import ReRanker from ..typing import Status from .convert import convert_to_cpp_doc, convert_to_py_doc -from .doc import Doc +from .doc import Doc, DocList from .param import ( AddColumnOption, AlterColumnOption, @@ -63,7 +63,7 @@ class Collection: inst._obj = core_collection schema = CollectionSchema._from_core(core_collection.Schema()) inst._schema = schema - inst._querier = QueryExecutorFactory.create(schema) + inst._querier = QueryExecutor(schema) return inst @property @@ -381,7 +381,7 @@ class Collection: include_vector: bool = False, output_fields: Optional[list[str]] = None, reranker: Optional[ReRanker] = None, - ) -> list[Doc]: + ) -> DocList: """Perform vector similarity search with optional filtering and re-ranking. At least one `Query` must be provided via `queries`. @@ -403,7 +403,7 @@ class Collection: Defaults to None. Returns: - list[Doc]: Top-k matching documents, sorted by relevance score. + DocList: Top-k matching documents, sorted by relevance score. Examples: >>> from zvec import Query diff --git a/python/zvec/model/doc.py b/python/zvec/model/doc.py index 722293d..175c946 100644 --- a/python/zvec/model/doc.py +++ b/python/zvec/model/doc.py @@ -20,6 +20,7 @@ from ..common import VectorType __all__ = [ "Doc", + "DocList", ] @@ -171,3 +172,7 @@ class Doc: else: obj.vectors = {} return obj + + +#: Type alias for query results: a list of documents returned by a single query route. +DocList = list[Doc] diff --git a/python/zvec/model/param/__init__.pyi b/python/zvec/model/param/__init__.pyi index 31b43f6..b408866 100644 --- a/python/zvec/model/param/__init__.pyi +++ b/python/zvec/model/param/__init__.pyi @@ -817,7 +817,7 @@ class VectorIndexParam(IndexParam): QuantizeType: Vector quantization type (e.g., FP16, INT8). """ -class _VectorQuery: +class _SearchQuery: field_name: str filter: str include_vector: bool diff --git a/src/binding/c/c_api.cc b/src/binding/c/c_api.cc index 714b711..df26bf9 100644 --- a/src/binding/c/c_api.cc +++ b/src/binding/c/c_api.cc @@ -5622,7 +5622,7 @@ zvec_error_code_t zvec_group_by_vector_query_set_flat_params( // Reranker Implementation // ============================================================================= -zvec_reranker_t *zvec_reranker_create_rrf(int rank_constant) { +zvec_reranker_t *zvec_create_rrf_reranker(int rank_constant) { ZVEC_TRY_RETURN_NULL("Failed to create RRF Reranker", auto *reranker = new zvec::Reranker::Ptr( @@ -5632,39 +5632,29 @@ zvec_reranker_t *zvec_reranker_create_rrf(int rank_constant) { return nullptr; } -zvec_reranker_t *zvec_reranker_create_weighted(const char **fields, - const double *weights, - size_t field_count) { - if ((!fields || !weights) && field_count > 0) { - set_last_error( - "Fields and weights pointers cannot be null when field_count > 0"); +zvec_reranker_t *zvec_create_weighted_reranker(const double *weights, + size_t weight_count) { + if (!weights && weight_count > 0) { + set_last_error("Weights pointer cannot be null when weight_count > 0"); return nullptr; } ZVEC_TRY_RETURN_NULL( "Failed to create Weighted Reranker", - std::map weight_map; - for (size_t i = 0; i < field_count; ++i) { - if (!fields[i]) { - set_last_error("Null field name at index " + std::to_string(i)); - return nullptr; - } - weight_map[fields[i]] = weights[i]; - } - auto *reranker = new zvec::Reranker::Ptr( - std::make_shared(weight_map)); + std::make_shared( + std::vector(weights, weights + weight_count))); return reinterpret_cast(reranker);) return nullptr; } -void zvec_reranker_destroy(zvec_reranker_t *reranker) { +void zvec_destroy_reranker(zvec_reranker_t *reranker) { if (reranker) { delete reinterpret_cast(reranker); } } -int zvec_reranker_get_rank_constant(const zvec_reranker_t *reranker) { +int zvec_get_reranker_rank_constant(const zvec_reranker_t *reranker) { if (!reranker) return -1; auto *ptr = reinterpret_cast(reranker); auto *rrf = dynamic_cast(ptr->get()); diff --git a/src/binding/python/model/param/python_param.cc b/src/binding/python/model/param/python_param.cc index 50bdfcc..3615433 100644 --- a/src/binding/python/model/param/python_param.cc +++ b/src/binding/python/model/param/python_param.cc @@ -1538,19 +1538,19 @@ void ZVecPyParams::bind_vector_query(py::module_ &m) { .def(py::init<>()) .def_readwrite("num_candidates", &SubQuery::num_candidates_) .def_static( - "from_vector_query", + "from_search_query", [](const SearchQuery &sq) { SubQuery sub; sub.num_candidates_ = sq.topk_; sub.target_ = sq.target_; return sub; }, - py::arg("vector_query"), + py::arg("search_query"), "Create a SubQuery from a single-target search query."); - // _VectorQuery is the historical Python class name; it now wraps the + // _SearchQuery is the Python class name; it wraps the // single-target SearchQuery so external Python code keeps working unchanged. - py::class_(m, "_VectorQuery") + py::class_(m, "_SearchQuery") .def(py::init<>()) // properties .def_readwrite("topk", &SearchQuery::topk_) @@ -1805,7 +1805,7 @@ void ZVecPyParams::bind_vector_query(py::module_ &m) { }, [](py::tuple t) { if (t.size() != 10) - throw std::runtime_error("Invalid pickle data for _VectorQuery"); + throw std::runtime_error("Invalid pickle data for _SearchQuery"); SearchQuery obj{}; obj.topk_ = t[0].cast(); diff --git a/src/binding/python/model/python_reranker.cc b/src/binding/python/model/python_reranker.cc index e8a3374..6543ecf 100644 --- a/src/binding/python/model/python_reranker.cc +++ b/src/binding/python/model/python_reranker.cc @@ -13,6 +13,7 @@ // limitations under the License. #include "python_reranker.h" +#include #include #include #include @@ -20,9 +21,40 @@ namespace zvec { +namespace { + +inline void reranker_throw_if_error(const Status &status) { + switch (status.code()) { + case StatusCode::OK: + return; + case StatusCode::NOT_FOUND: + throw py::key_error(status.message()); + case StatusCode::INVALID_ARGUMENT: + throw py::value_error(status.message()); + default: + throw std::runtime_error(status.message()); + } +} + +inline DocPtrList unwrap_rerank_result(Result result) { + if (!result.has_value()) { + reranker_throw_if_error(result.error()); + } + return std::move(result).value(); +} + +} // namespace + void ZVecPyReranker::Initialize(py::module_ &m) { // Bind Reranker base class (abstract, cannot be instantiated directly) - py::class_(m, "_Reranker"); + py::class_(m, "_Reranker") + .def( + "rerank", + [](const Reranker &self, const std::vector &query_results, + int topn) { + return unwrap_rerank_result(self.rerank(query_results, topn)); + }, + py::arg("query_results"), py::arg("topn") = 10); // Bind ScoreBasedReranker intermediate class py::class_>( @@ -37,7 +69,7 @@ void ZVecPyReranker::Initialize(py::module_ &m) { // Bind WeightedReranker py::class_>(m, "_WeightedReranker") - .def(py::init>(), py::arg("weights")) + .def(py::init>(), py::arg("weights")) .def_property_readonly("weights", &WeightedReranker::weights); // Bind CallbackReranker diff --git a/src/db/collection.cc b/src/db/collection.cc index 1b9dd69..e45b001 100644 --- a/src/db/collection.cc +++ b/src/db/collection.cc @@ -16,7 +16,6 @@ #include #include #include -#include #include #include #include @@ -1706,29 +1705,16 @@ Result CollectionImpl::Query(const MultiQuery &query) const { return DocPtrList(); } - struct PendingQuery { - std::string field_name; - SearchQuery query; - }; - // Convert each SubQuery to a SearchQuery and validate. - std::set seen_fields; - std::vector pending_queries; - pending_queries.reserve(query.queries.size()); + std::vector search_queries; + std::vector field_names; + search_queries.reserve(query.queries.size()); + field_names.reserve(query.queries.size()); for (const auto &sub : query.queries) { const auto &target = sub.target_; - auto [_, inserted] = seen_fields.insert(target.field_name_); - if (!inserted) { - return tl::make_unexpected(Status::InvalidArgument( - "Duplicate field name in multi-query: ", target.field_name_)); - } - // Use get_field uniformly; validate_and_sanitize checks type compatibility. + auto *field_schema = schema_->get_field(target.field_name_); - if (!field_schema) { - return tl::make_unexpected( - Status::InvalidArgument("Field not found: ", target.field_name_)); - } SearchQuery sq; sq.target_ = target; @@ -1740,43 +1726,44 @@ Result CollectionImpl::Query(const MultiQuery &query) const { auto s = sq.validate_and_sanitize(field_schema); CHECK_RETURN_STATUS_EXPECTED(s); - pending_queries.push_back({target.field_name_, std::move(sq)}); + field_names.push_back(target.field_name_); + search_queries.push_back(std::move(sq)); } - std::map query_results; - - auto execute_query = [&](PendingQuery &pending) -> Result { + // Execute sub-queries. + auto execute_query = [&](SearchQuery &sq) -> Result { auto engine = sqlengine::SQLEngine::create(std::make_shared()); - return engine->execute(schema_, std::move(pending.query), segments); + return engine->execute(schema_, std::move(sq), segments); }; - std::vector> results(pending_queries.size()); + std::vector> results(search_queries.size()); // Single-segment queries have no segment-level fanout; multi-segment queries // already use the query pool per sub-query. if (segments.size() == 1) { auto group = GlobalResource::Instance().query_thread_pool()->make_group(); - for (size_t i = 0; i < pending_queries.size(); ++i) { + for (size_t i = 0; i < search_queries.size(); ++i) { group->execute( - [&, i]() { results[i] = execute_query(pending_queries[i]); }); + [&, i]() { results[i] = execute_query(search_queries[i]); }); } group->wait_finish(); } else { - for (size_t i = 0; i < pending_queries.size(); ++i) { - results[i] = execute_query(pending_queries[i]); + for (size_t i = 0; i < search_queries.size(); ++i) { + results[i] = execute_query(search_queries[i]); } } - for (size_t i = 0; i < pending_queries.size(); ++i) { - if (!results[i]) { - return tl::make_unexpected(results[i].error()); + // Collect results and rerank. + std::vector query_results; + query_results.reserve(results.size()); + for (auto &result : results) { + if (!result) { + return tl::make_unexpected(result.error()); } - query_results[pending_queries[i].field_name] = - std::move(results[i].value()); + query_results.push_back(std::move(result.value())); } - // Merge and rerank results - query.reranker->bind_schema(schema_); + query.reranker->bind_schema(schema_, field_names); return query.reranker->rerank(query_results, query.topk); } diff --git a/src/db/reranker/reranker.cc b/src/db/reranker/reranker.cc index 7f55ef2..9fb49be 100644 --- a/src/db/reranker/reranker.cc +++ b/src/db/reranker/reranker.cc @@ -28,16 +28,22 @@ namespace zvec { // ==================== ScoreBasedReranker ==================== Result ScoreBasedReranker::rerank( - const std::map &query_results, int topn) const { + const std::vector &query_results, int topn) const { + if (topn <= 0) { + return DocPtrList(); + } + std::unordered_map scores; std::unordered_map id_to_doc; - for (const auto &[field_name, docs] : query_results) { + for (size_t query_index = 0; query_index < query_results.size(); + ++query_index) { + const auto &docs = query_results[query_index]; for (size_t rank = 0; rank < docs.size(); ++rank) { const auto &doc = docs[rank]; const std::string &doc_id = doc->pk(); auto rs = rescore(static_cast(doc->score()), - static_cast(rank), field_name); + static_cast(rank), static_cast(query_index)); if (!rs.has_value()) { return tl::make_unexpected(rs.error()); } @@ -79,18 +85,20 @@ Result ScoreBasedReranker::rerank( // ==================== RrfReranker ==================== Result RrfReranker::rescore(double /*score*/, int rank, - const std::string & /*field_name*/) const { + int /*query_index*/) const { return 1.0 / (static_cast(rank_constant_) + static_cast(rank) + 1.0); } // ==================== WeightedReranker ==================== -WeightedReranker::WeightedReranker(const std::map &weights) +WeightedReranker::WeightedReranker(const std::vector &weights) : weights_(weights) {} -void WeightedReranker::bind_schema(CollectionSchema::Ptr schema) { +void WeightedReranker::bind_schema( + CollectionSchema::Ptr schema, const std::vector &field_names) { schema_ = std::move(schema); + field_names_ = field_names; } Result WeightedReranker::normalize_score(double score, @@ -122,7 +130,18 @@ Result WeightedReranker::normalize_score(double score, } Result WeightedReranker::rescore(double score, int /*rank*/, - const std::string &field_name) const { + int query_index) const { + if (!schema_) { + return tl::make_unexpected( + Status::InvalidArgument("WeightedReranker: schema is null")); + } + if (query_index < 0 || + static_cast(query_index) >= field_names_.size()) { + return tl::make_unexpected( + Status::InvalidArgument("WeightedReranker: query_index out of range: ", + std::to_string(query_index))); + } + const auto &field_name = field_names_[query_index]; const auto *field = schema_->get_field(field_name); if (!field) { return tl::make_unexpected(Status::InvalidArgument( @@ -133,9 +152,8 @@ Result WeightedReranker::rescore(double score, int /*rank*/, return tl::make_unexpected(normalized.error()); } double weight = 1.0; - auto weight_it = weights_.find(field_name); - if (weight_it != weights_.end()) { - weight = weight_it->second; + if (static_cast(query_index) < weights_.size()) { + weight = weights_[query_index]; } return normalized.value() * weight; } diff --git a/src/include/zvec/c_api.h b/src/include/zvec/c_api.h index 64cec24..4adfaf9 100644 --- a/src/include/zvec/c_api.h +++ b/src/include/zvec/c_api.h @@ -1140,8 +1140,8 @@ typedef struct zvec_doc_t zvec_doc_t; /** * @brief Reranker structure (opaque pointer) * Aligned with zvec::Reranker - * Use zvec_reranker_create_rrf() or zvec_reranker_create_weighted() to create - * and zvec_reranker_destroy() to destroy + * Use zvec_create_rrf_reranker() or zvec_create_weighted_reranker() to create + * and zvec_destroy_reranker() to destroy */ typedef struct zvec_reranker_t zvec_reranker_t; typedef struct zvec_collection_schema_t zvec_collection_schema_t; @@ -1961,23 +1961,22 @@ zvec_group_by_vector_query_set_flat_params( * @return zvec_reranker_t* Pointer to the newly created reranker */ ZVEC_EXPORT zvec_reranker_t *ZVEC_CALL -zvec_reranker_create_rrf(int rank_constant); +zvec_create_rrf_reranker(int rank_constant); /** * @brief Create a Weighted reranker - * @param fields Array of field names - * @param weights Array of weights corresponding to fields - * @param field_count Number of field/weight entries + * @param weights Array of weights for each query + * @param weight_count Number of weight entries * @return zvec_reranker_t* Pointer to the newly created reranker */ -ZVEC_EXPORT zvec_reranker_t *ZVEC_CALL zvec_reranker_create_weighted( - const char **fields, const double *weights, size_t field_count); +ZVEC_EXPORT zvec_reranker_t *ZVEC_CALL +zvec_create_weighted_reranker(const double *weights, size_t weight_count); /** * @brief Destroy reranker * @param reranker Reranker pointer */ -ZVEC_EXPORT void ZVEC_CALL zvec_reranker_destroy(zvec_reranker_t *reranker); +ZVEC_EXPORT void ZVEC_CALL zvec_destroy_reranker(zvec_reranker_t *reranker); /** * @brief Get RRF rank constant (only valid for RRF reranker) @@ -1985,7 +1984,7 @@ ZVEC_EXPORT void ZVEC_CALL zvec_reranker_destroy(zvec_reranker_t *reranker); * @return int Rank constant, or -1 if not an RRF reranker */ ZVEC_EXPORT int ZVEC_CALL -zvec_reranker_get_rank_constant(const zvec_reranker_t *reranker); +zvec_get_reranker_rank_constant(const zvec_reranker_t *reranker); // ----------------------------------------------------------------------------- // zvec_multi_query_t (Multi Query) @@ -2100,7 +2099,7 @@ ZVEC_EXPORT zvec_error_code_t ZVEC_CALL zvec_multi_query_get_output_fields( * reranker) * @param query Multi-vector query pointer * @param reranker Reranker pointer (remains valid, caller must call - * zvec_reranker_destroy after use) + * zvec_destroy_reranker after use) * @return zvec_error_code_t Error code */ ZVEC_EXPORT zvec_error_code_t ZVEC_CALL zvec_multi_query_set_reranker( diff --git a/src/include/zvec/db/reranker.h b/src/include/zvec/db/reranker.h index 0b8b56e..c138f42 100644 --- a/src/include/zvec/db/reranker.h +++ b/src/include/zvec/db/reranker.h @@ -14,9 +14,9 @@ #pragma once #include -#include #include #include +#include #include #include #include @@ -32,16 +32,16 @@ class Reranker { Reranker() = default; virtual ~Reranker() = default; - virtual void bind_schema(CollectionSchema::Ptr) {} + virtual void bind_schema(CollectionSchema::Ptr /*schema*/, + const std::vector & /*field_names*/) {} //! Re-rank documents from one or more vector queries. - //! \param query_results Mapping from vector field name to list of retrieved - //! documents (sorted by relevance). + //! \param query_results Per-query lists of retrieved documents (sorted by + //! relevance), in the same order as the sub-queries supplied by the caller. //! \param topn Maximum number of documents to return. //! \return Re-ranked list of documents (length <= topn), with updated scores. virtual Result rerank( - const std::map &query_results, - int topn = 10) const = 0; + const std::vector &query_results, int topn = 10) const = 0; }; //! Intermediate base for rerankers that compute per-document scores. @@ -51,17 +51,18 @@ class Reranker { //! Subclasses only need to implement rescore(). class ScoreBasedReranker : public Reranker { public: - //! Compute the contribution score for a single document. - //! \param score The document's raw relevance score from the vector field. - //! \param rank The document's position (0-based) in the per-field result - //! list. \param field_name The name of the vector field this result came - //! from. \return The score contribution to be accumulated for this document. - virtual Result rescore(double score, int rank, - const std::string &field_name) const = 0; + Result rerank(const std::vector &query_results, + int topn = 10) const override; - Result rerank( - const std::map &query_results, - int topn = 10) const override; + private: + //! Compute the contribution score for a single document. + //! \param score The document's raw relevance score from the vector query. + //! \param rank The document's position (0-based) in the per-query result + //! list. \param query_index The index (0-based) of the sub-query this result + //! came from. \return The score contribution to be accumulated for this + //! document. + virtual Result rescore(double score, int rank, + int query_index) const = 0; }; //! Re-ranker using Reciprocal Rank Fusion (RRF) for multi-vector search. @@ -79,10 +80,10 @@ class RrfReranker : public ScoreBasedReranker { return rank_constant_; } - Result rescore(double score, int rank, - const std::string &field_name) const override; - private: + Result rescore(double score, int rank, + int query_index) const override; + int rank_constant_; }; @@ -91,24 +92,30 @@ class RrfReranker : public ScoreBasedReranker { //! Each vector field's relevance score is normalized based on its own metric //! type, then scaled by a user-provided weight. Final scores are summed across //! fields. Supported metrics: L2, IP, COSINE. +//! +//! @note NOT thread-safe. The bind_schema() and rerank() calls share mutable +//! state. Each concurrent query must use its own WeightedReranker instance or +//! serialize access externally. class WeightedReranker : public ScoreBasedReranker { public: - explicit WeightedReranker(const std::map &weights = {}); + explicit WeightedReranker(const std::vector &weights = {}); - void bind_schema(CollectionSchema::Ptr schema) override; + void bind_schema(CollectionSchema::Ptr schema, + const std::vector &field_names) override; - const std::map &weights() const { + const std::vector &weights() const { return weights_; } - Result rescore(double score, int rank, - const std::string &field_name) const override; - private: + Result rescore(double score, int rank, + int query_index) const override; + static Result normalize_score(double score, const FieldSchema &field); CollectionSchema::Ptr schema_; - std::map weights_; + std::vector field_names_; + std::vector weights_; }; //! Callback-based re-ranker for cross-language bridging. @@ -118,13 +125,16 @@ class WeightedReranker : public ScoreBasedReranker { class CallbackReranker : public Reranker { public: using Callback = - std::function &, int)>; + std::function &, int)>; explicit CallbackReranker(Callback fn) : callback_(std::move(fn)) {} - Result rerank( - const std::map &query_results, - int topn = 10) const override { + Result rerank(const std::vector &query_results, + int topn = 10) const override { + if (!callback_) { + return tl::make_unexpected( + Status::InvalidArgument("CallbackReranker: callback is empty")); + } return callback_(query_results, topn); } diff --git a/tests/c/c_api_test.c b/tests/c/c_api_test.c index 662405f..80013da 100644 --- a/tests/c/c_api_test.c +++ b/tests/c/c_api_test.c @@ -4370,41 +4370,40 @@ void test_reranker_functions(void) { TEST_START(); // Test 1: Create RRF reranker - zvec_reranker_t *rrf = zvec_reranker_create_rrf(60); + zvec_reranker_t *rrf = zvec_create_rrf_reranker(60); TEST_ASSERT(rrf != NULL); if (rrf) { - TEST_ASSERT(zvec_reranker_get_rank_constant(rrf) == 60); - zvec_reranker_destroy(rrf); + TEST_ASSERT(zvec_get_reranker_rank_constant(rrf) == 60); + zvec_destroy_reranker(rrf); } // Test 2: Create RRF reranker with different rank constant - zvec_reranker_t *rrf2 = zvec_reranker_create_rrf(100); + zvec_reranker_t *rrf2 = zvec_create_rrf_reranker(100); TEST_ASSERT(rrf2 != NULL); if (rrf2) { - TEST_ASSERT(zvec_reranker_get_rank_constant(rrf2) == 100); - zvec_reranker_destroy(rrf2); + TEST_ASSERT(zvec_get_reranker_rank_constant(rrf2) == 100); + zvec_destroy_reranker(rrf2); } // Test 3: Create Weighted reranker - const char *fields[] = {"embedding1", "embedding2"}; double weights[] = {0.7, 0.3}; - zvec_reranker_t *weighted = zvec_reranker_create_weighted(fields, weights, 2); + zvec_reranker_t *weighted = zvec_create_weighted_reranker(weights, 2); TEST_ASSERT(weighted != NULL); if (weighted) { - TEST_ASSERT(zvec_reranker_get_rank_constant(weighted) == -1); - zvec_reranker_destroy(weighted); + TEST_ASSERT(zvec_get_reranker_rank_constant(weighted) == -1); + zvec_destroy_reranker(weighted); } - // Test 4: Create Weighted reranker with no fields - zvec_reranker_t *weighted2 = zvec_reranker_create_weighted(NULL, NULL, 0); + // Test 4: Create Weighted reranker with no weights + zvec_reranker_t *weighted2 = zvec_create_weighted_reranker(NULL, 0); TEST_ASSERT(weighted2 != NULL); if (weighted2) { - zvec_reranker_destroy(weighted2); + zvec_destroy_reranker(weighted2); } // Test 5: NULL reranker operations - TEST_ASSERT(zvec_reranker_get_rank_constant(NULL) == -1); - zvec_reranker_destroy(NULL); // Should not crash + TEST_ASSERT(zvec_get_reranker_rank_constant(NULL) == -1); + zvec_destroy_reranker(NULL); // Should not crash TEST_END(); } @@ -4528,14 +4527,14 @@ void test_multi_vector_query_with_rrf_reranker(void) { multi_query_fixture_t f; TEST_ASSERT(setup_multi_query_fixture(&f, "zvec_test_mq_rrf", "mq_rrf")); - zvec_reranker_t *rrf = zvec_reranker_create_rrf(60); + zvec_reranker_t *rrf = zvec_create_rrf_reranker(60); TEST_ASSERT(rrf != NULL); int count = execute_multi_query_with_reranker(&f, rrf, 3, 3); TEST_ASSERT(count > 0); TEST_ASSERT(count <= 3); - zvec_reranker_destroy(rrf); + zvec_destroy_reranker(rrf); // MultiQuery property setters/getters zvec_multi_query_t *mvq2 = zvec_multi_query_create(); @@ -4589,16 +4588,15 @@ void test_multi_vector_query_with_weighted_reranker(void) { TEST_ASSERT( setup_multi_query_fixture(&f, "zvec_test_mq_weighted", "mq_weighted")); - const char *fields[] = {"embedding1", "embedding2"}; double weights[] = {0.7, 0.3}; - zvec_reranker_t *weighted = zvec_reranker_create_weighted(fields, weights, 2); + zvec_reranker_t *weighted = zvec_create_weighted_reranker(weights, 2); TEST_ASSERT(weighted != NULL); int count = execute_multi_query_with_reranker(&f, weighted, 3, 3); TEST_ASSERT(count > 0); TEST_ASSERT(count <= 3); - zvec_reranker_destroy(weighted); + zvec_destroy_reranker(weighted); teardown_multi_query_fixture(&f); TEST_END(); diff --git a/tests/db/collection_test.cc b/tests/db/collection_test.cc index 7f4e257..e23d47c 100644 --- a/tests/db/collection_test.cc +++ b/tests/db/collection_test.cc @@ -3804,30 +3804,6 @@ TEST_F(CollectionTest, Feature_MultiQuery_Validate) { EXPECT_EQ(result.error().code(), StatusCode::INVALID_ARGUMENT); } - // Test 4: Duplicate field names should fail - { - MultiQuery mvq; - mvq.topk = 10; - mvq.reranker = std::make_shared(60); - - SubQuery vq1; - vq1.num_candidates_ = 10; - vq1.target_.field_name_ = "dense_fp32"; - std::get(vq1.target_.clause_) - .query_vector_.assign(128 * sizeof(float), '\0'); - mvq.queries.push_back(vq1); - - SubQuery vq2; - vq2.num_candidates_ = 10; - vq2.target_.field_name_ = "dense_fp32"; - std::get(vq2.target_.clause_) - .query_vector_.assign(128 * sizeof(float), '\0'); - mvq.queries.push_back(vq2); - - auto result = collection->Query(mvq); - ASSERT_FALSE(result.has_value()); - EXPECT_EQ(result.error().code(), StatusCode::INVALID_ARGUMENT); - } } TEST_F(CollectionTest, Feature_MultiQuery_SingleFieldWithReranker) { @@ -3936,9 +3912,8 @@ TEST_F(CollectionTest, Feature_MultiQuery_MultiFieldWeighted) { MultiQuery mvq; mvq.topk = 10; - std::map weights = {{"dense_fp32", 0.7}, - {"sparse_fp32", 0.3}}; - mvq.reranker = std::make_shared(weights); + mvq.reranker = + std::make_shared(std::vector{0.7, 0.3}); // Query dense_fp32 field { @@ -4090,11 +4065,11 @@ TEST_F(CollectionTest, Feature_MultiQuery_CallbackReranker) { // Use CallbackReranker with a lambda that merges and sorts by score bool callback_invoked = false; auto callback_fn = [&callback_invoked]( - const std::map &query_results, + const std::vector &query_results, int topn) -> DocPtrList { callback_invoked = true; DocPtrList all_docs; - for (const auto &[_, docs] : query_results) { + for (const auto &docs : query_results) { for (const auto &doc : docs) { all_docs.push_back(doc); } diff --git a/tests/db/reranker_test.cc b/tests/db/reranker_test.cc index adbd3fd..b41c123 100644 --- a/tests/db/reranker_test.cc +++ b/tests/db/reranker_test.cc @@ -14,8 +14,8 @@ #define _USE_MATH_DEFINES #include -#include #include +#include #include #include #include @@ -55,11 +55,11 @@ TEST(RrfRerankerTest, BasicRRF) { RrfReranker reranker(/*rank_constant=*/60); // Two vector fields, each returning 3 documents with some overlap - std::map query_results; - query_results["vec1"] = {MakeDoc("a", 0.9f), MakeDoc("b", 0.8f), - MakeDoc("c", 0.7f)}; - query_results["vec2"] = {MakeDoc("b", 0.95f), MakeDoc("a", 0.85f), - MakeDoc("d", 0.75f)}; + std::vector query_results; + query_results.push_back( + {MakeDoc("a", 0.9f), MakeDoc("b", 0.8f), MakeDoc("c", 0.7f)}); + query_results.push_back( + {MakeDoc("b", 0.95f), MakeDoc("a", 0.85f), MakeDoc("d", 0.75f)}); auto result = reranker.rerank(query_results, /*topn=*/10); ASSERT_TRUE(result.has_value()); @@ -72,9 +72,9 @@ TEST(RrfRerankerTest, BasicRRF) { // So a and b should have equal scores and be at the top ASSERT_GE(results.size(), 3u); - // "a" and "b" should have the highest RRF scores - EXPECT_EQ(results[0]->pk(), "a"); - EXPECT_EQ(results[1]->pk(), "b"); + // "a" and "b" should have the highest RRF scores (equal, order unspecified) + std::set top2{results[0]->pk(), results[1]->pk()}; + EXPECT_EQ(top2, (std::set{"a", "b"})); // Verify scores are close (a and b have same RRF score) EXPECT_NEAR(results[0]->score(), results[1]->score(), 1e-10); } @@ -82,9 +82,9 @@ TEST(RrfRerankerTest, BasicRRF) { TEST(RrfRerankerTest, Topn) { RrfReranker reranker(/*rank_constant=*/60); - std::map query_results; - query_results["vec1"] = {MakeDoc("a", 0.9f), MakeDoc("b", 0.8f), - MakeDoc("c", 0.7f)}; + std::vector query_results; + query_results.push_back( + {MakeDoc("a", 0.9f), MakeDoc("b", 0.8f), MakeDoc("c", 0.7f)}); auto result = reranker.rerank(query_results, /*topn=*/2); ASSERT_TRUE(result.has_value()); @@ -94,8 +94,8 @@ TEST(RrfRerankerTest, Topn) { TEST(RrfRerankerTest, SingleField) { RrfReranker reranker(/*rank_constant=*/60); - std::map query_results; - query_results["vec1"] = {MakeDoc("a", 0.9f), MakeDoc("b", 0.8f)}; + std::vector query_results; + query_results.push_back({MakeDoc("a", 0.9f), MakeDoc("b", 0.8f)}); auto result = reranker.rerank(query_results); ASSERT_TRUE(result.has_value()); @@ -108,7 +108,7 @@ TEST(RrfRerankerTest, SingleField) { TEST(RrfRerankerTest, EmptyResults) { RrfReranker reranker(/*rank_constant=*/60); - std::map query_results; + std::vector query_results; auto result = reranker.rerank(query_results); ASSERT_TRUE(result.has_value()); EXPECT_TRUE(result.value().empty()); @@ -119,12 +119,12 @@ TEST(RrfRerankerTest, EmptyResults) { TEST(WeightedRerankerTest, BasicWeighted) { auto schema = MakeSchema({{"vec1", MetricType::L2}, {"vec2", MetricType::L2}}); - WeightedReranker reranker({{"vec1", 0.7}, {"vec2", 0.3}}); - reranker.bind_schema(schema); + WeightedReranker reranker({0.7, 0.3}); + reranker.bind_schema(schema, {"vec1", "vec2"}); - std::map query_results; - query_results["vec1"] = {MakeDoc("a", 0.5f), MakeDoc("b", 0.3f)}; - query_results["vec2"] = {MakeDoc("a", 0.8f), MakeDoc("c", 0.6f)}; + std::vector query_results; + query_results.push_back({MakeDoc("a", 0.5f), MakeDoc("b", 0.3f)}); + query_results.push_back({MakeDoc("a", 0.8f), MakeDoc("c", 0.6f)}); auto result = reranker.rerank(query_results); ASSERT_TRUE(result.has_value()); @@ -137,12 +137,12 @@ TEST(WeightedRerankerTest, BasicWeighted) { TEST(WeightedRerankerTest, MixedMetrics) { auto schema = MakeSchema({{"vec1", MetricType::L2}, {"vec2", MetricType::COSINE}}); - WeightedReranker reranker({{"vec1", 0.5}, {"vec2", 0.5}}); - reranker.bind_schema(schema); + WeightedReranker reranker({0.5, 0.5}); + reranker.bind_schema(schema, {"vec1", "vec2"}); - std::map query_results; - query_results["vec1"] = {MakeDoc("a", 0.5f)}; - query_results["vec2"] = {MakeDoc("a", 0.4f)}; + std::vector query_results; + query_results.push_back({MakeDoc("a", 0.5f)}); + query_results.push_back({MakeDoc("a", 0.4f)}); auto result = reranker.rerank(query_results); ASSERT_TRUE(result.has_value()); @@ -161,12 +161,12 @@ TEST(WeightedRerankerTest, MixedMetrics) { TEST(WeightedRerankerTest, MissingMetricError) { auto schema = MakeSchema({{"vec1", MetricType::L2}}); WeightedReranker reranker; - reranker.bind_schema(schema); - - std::map query_results; - query_results["vec1"] = {MakeDoc("a", 0.5f)}; - query_results["vec2"] = {MakeDoc("b", 0.3f)}; + // Binding a field that is absent from the schema should fail at rerank time. + reranker.bind_schema(schema, {"vec1", "vec2"}); + std::vector query_results; + query_results.push_back({MakeDoc("a", 0.5f)}); + query_results.push_back({MakeDoc("b", 0.3f)}); auto result = reranker.rerank(query_results); ASSERT_FALSE(result.has_value()); } @@ -174,10 +174,10 @@ TEST(WeightedRerankerTest, MissingMetricError) { TEST(WeightedRerankerTest, NormalizeL2) { auto schema = MakeSchema({{"vec1", MetricType::L2}}); WeightedReranker reranker; - reranker.bind_schema(schema); + reranker.bind_schema(schema, {"vec1"}); - std::map query_results; - query_results["vec1"] = {MakeDoc("a", 0.0f), MakeDoc("b", 1.0f)}; + std::vector query_results; + query_results.push_back({MakeDoc("a", 0.0f), MakeDoc("b", 1.0f)}); auto result = reranker.rerank(query_results); ASSERT_TRUE(result.has_value()); @@ -193,10 +193,10 @@ TEST(WeightedRerankerTest, NormalizeL2) { TEST(WeightedRerankerTest, NormalizeIP) { auto schema = MakeSchema({{"vec1", MetricType::IP}}); WeightedReranker reranker; - reranker.bind_schema(schema); + reranker.bind_schema(schema, {"vec1"}); - std::map query_results; - query_results["vec1"] = {MakeDoc("a", 0.0f), MakeDoc("b", 1.0f)}; + std::vector query_results; + query_results.push_back({MakeDoc("a", 0.0f), MakeDoc("b", 1.0f)}); auto result = reranker.rerank(query_results); ASSERT_TRUE(result.has_value()); @@ -211,11 +211,11 @@ TEST(WeightedRerankerTest, NormalizeIP) { TEST(WeightedRerankerTest, NormalizeCosine) { auto schema = MakeSchema({{"vec1", MetricType::COSINE}}); WeightedReranker reranker; - reranker.bind_schema(schema); + reranker.bind_schema(schema, {"vec1"}); - std::map query_results; - query_results["vec1"] = {MakeDoc("a", 0.0f), MakeDoc("b", 1.0f), - MakeDoc("c", 2.0f)}; + std::vector query_results; + query_results.push_back( + {MakeDoc("a", 0.0f), MakeDoc("b", 1.0f), MakeDoc("c", 2.0f)}); auto result = reranker.rerank(query_results); ASSERT_TRUE(result.has_value()); @@ -230,11 +230,11 @@ TEST(WeightedRerankerTest, NormalizeCosine) { TEST(WeightedRerankerTest, Topn) { auto schema = MakeSchema({{"vec1", MetricType::L2}}); WeightedReranker reranker; - reranker.bind_schema(schema); + reranker.bind_schema(schema, {"vec1"}); - std::map query_results; - query_results["vec1"] = {MakeDoc("a", 0.1f), MakeDoc("b", 0.2f), - MakeDoc("c", 0.3f)}; + std::vector query_results; + query_results.push_back( + {MakeDoc("a", 0.1f), MakeDoc("b", 0.2f), MakeDoc("c", 0.3f)}); auto result = reranker.rerank(query_results, /*topn=*/2); ASSERT_TRUE(result.has_value()); @@ -248,10 +248,9 @@ TEST(CallbackRerankerTest, BasicCallback) { // Simple callback that returns docs sorted by score descending, limited to // topn CallbackReranker::Callback cb = - [](const std::map &query_results, - int topn) -> DocPtrList { + [](const std::vector &query_results, int topn) -> DocPtrList { DocPtrList all_docs; - for (const auto &[_, docs] : query_results) { + for (const auto &docs : query_results) { for (const auto &doc : docs) { all_docs.push_back(doc); } @@ -268,9 +267,9 @@ TEST(CallbackRerankerTest, BasicCallback) { CallbackReranker reranker(cb); - std::map query_results; - query_results["vec1"] = {MakeDoc("a", 0.5f), MakeDoc("b", 0.9f)}; - query_results["vec2"] = {MakeDoc("c", 0.7f)}; + std::vector query_results; + query_results.push_back({MakeDoc("a", 0.5f), MakeDoc("b", 0.9f)}); + query_results.push_back({MakeDoc("c", 0.7f)}); auto result = reranker.rerank(query_results, /*topn=*/10); ASSERT_TRUE(result.has_value());