From ecb2c4f3db31e40ffd8ba103e101b26fe213edb8 Mon Sep 17 00:00:00 2001 From: xblwh <64543322+xblwh@users.noreply.github.com> Date: Tue, 15 Sep 2026 13:54:02 +0800 Subject: [PATCH 1/2] fix: reject mismatched target vector names --- test/collection/test_target_vectors.py | 87 ++++++++++++++++++++++++++ weaviate/collections/grpc/shared.py | 10 +-- 2 files changed, 92 insertions(+), 5 deletions(-) create mode 100644 test/collection/test_target_vectors.py diff --git a/test/collection/test_target_vectors.py b/test/collection/test_target_vectors.py new file mode 100644 index 000000000..f3b160ce8 --- /dev/null +++ b/test/collection/test_target_vectors.py @@ -0,0 +1,87 @@ +import pytest + +from weaviate.classes.query import HybridVector, NearVector, TargetVectors +from weaviate.collections.classes.grpc import NearVectorInputType, TargetVectorJoinType +from weaviate.collections.grpc.query import _QueryGRPC +from weaviate.exceptions import WeaviateInvalidInputError +from weaviate.proto.v1 import search_get_pb2 +from weaviate.util import _ServerVersion + + +def _request( + version: str, + query_type: str, + vector: NearVectorInputType, + target_vector: TargetVectorJoinType, +) -> search_get_pb2.SearchRequest: + query = _QueryGRPC( + weaviate_version=_ServerVersion.from_string(version), + name="Documents", + tenant=None, + consistency_level=None, + validate_arguments=True, + uses_125_api=True, + uses_127_api=True, + ) + if query_type == "near_vector": + return query.near_vector(near_vector=vector, target_vector=target_vector) + return query.hybrid( + query="example", + vector=HybridVector.near_vector(vector) if query_type == "hybrid_near_vector" else vector, + target_vector=target_vector, + ) + + +@pytest.mark.parametrize("version", ["1.26.0", "1.27.0", "1.29.0"]) +@pytest.mark.parametrize("query_type", ["near_vector", "hybrid", "hybrid_near_vector"]) +@pytest.mark.parametrize("weighted", [False, True]) +def test_mismatched_target_vector_names(version: str, query_type: str, weighted: bool) -> None: + target_vector = ( + TargetVectors.manual_weights({"title": 1.0, "summary": 2.0}) + if weighted + else ["title", "summary"] + ) + with pytest.raises(WeaviateInvalidInputError): + _request(version, query_type, {"title": [1.0, 0.0], "body": [0.0, 1.0]}, target_vector) + + +@pytest.mark.parametrize("version", ["1.26.0", "1.27.0", "1.29.0"]) +@pytest.mark.parametrize("query_type", ["near_vector", "hybrid", "hybrid_near_vector"]) +def test_target_vector_names_can_be_in_a_different_order(version: str, query_type: str) -> None: + request = _request( + version, + query_type, + {"summary": [0.0, 1.0], "title": [1.0, 0.0]}, + TargetVectors.manual_weights({"title": 1.0, "summary": 2.0}), + ) + targets = ( + request.near_vector.targets + if query_type == "near_vector" + else request.hybrid_search.targets + ) + assert set(targets.target_vectors) == {"title", "summary"} + assert {weight.target: weight.weight for weight in targets.weights_for_targets} == { + "title": 1.0, + "summary": 2.0, + } + + +@pytest.mark.parametrize("version", ["1.27.0", "1.29.0"]) +@pytest.mark.parametrize("query_type", ["near_vector", "hybrid", "hybrid_near_vector"]) +def test_repeated_target_vector_names_with_multiple_weights(version: str, query_type: str) -> None: + request = _request( + version, + query_type, + {"title": NearVector.list_of_vectors([1.0, 0.0], [0.0, 1.0])}, + TargetVectors.manual_weights({"title": [1.0, 2.0]}), + ) + targets = ( + request.near_vector.targets + if query_type == "near_vector" + else request.hybrid_search.targets + ) + assert list(targets.target_vectors) == ["title", "title"] + assert [(weight.target, weight.weight) for weight in targets.weights_for_targets] == [ + ("title", 1.0), + ("title", 2.0), + ] diff --git a/weaviate/collections/grpc/shared.py b/weaviate/collections/grpc/shared.py index 3b39a611b..7bd607757 100644 --- a/weaviate/collections/grpc/shared.py +++ b/weaviate/collections/grpc/shared.py @@ -155,6 +155,10 @@ def _vector_per_target( raise WeaviateInvalidInputError( "The number of target vectors must be equal to the number of vectors." ) + if set(targets.target_vectors) != vector.keys(): + raise WeaviateInvalidInputError( + "The vector dictionary keys must match the target vector names." + ) vector_per_target: Dict[str, bytes] = {} for key, value in vector.items(): @@ -292,11 +296,7 @@ def add_list_of_vectors(value: _ListOfVectorsQuery, key: str) -> None: target_vectors.append(key) if isinstance(vector, dict): - if ( - len(vector) == 0 - or targets is None - or len(set(targets.target_vectors)) != len(vector) - ): + if len(vector) == 0 or targets is None or set(targets.target_vectors) != vector.keys(): raise invalid_nv_exception for key, value in vector.items(): if _is_1d_vector(value): From 5475a295ea168d01765a1f31ca8cfa799c66689f Mon Sep 17 00:00:00 2001 From: Ivan Despot <66276597+g-despot@users.noreply.github.com> Date: Thu, 17 Sep 2026 11:29:54 +0200 Subject: [PATCH 2/2] Name the mismatched keys in the error and simplify the tests --- test/collection/test_target_vectors.py | 111 +++++++++++++------------ weaviate/collections/grpc/shared.py | 18 ++-- 2 files changed, 72 insertions(+), 57 deletions(-) diff --git a/test/collection/test_target_vectors.py b/test/collection/test_target_vectors.py index f3b160ce8..07b4ee74e 100644 --- a/test/collection/test_target_vectors.py +++ b/test/collection/test_target_vectors.py @@ -1,85 +1,92 @@ +from typing import Callable + import pytest from weaviate.classes.query import HybridVector, NearVector, TargetVectors from weaviate.collections.classes.grpc import NearVectorInputType, TargetVectorJoinType from weaviate.collections.grpc.query import _QueryGRPC from weaviate.exceptions import WeaviateInvalidInputError -from weaviate.proto.v1 import search_get_pb2 +from weaviate.proto.v1 import base_search_pb2 from weaviate.util import _ServerVersion -def _request( - version: str, - query_type: str, - vector: NearVectorInputType, - target_vector: TargetVectorJoinType, -) -> search_get_pb2.SearchRequest: - query = _QueryGRPC( - weaviate_version=_ServerVersion.from_string(version), +def _query(version: str) -> _QueryGRPC: + weaviate_version = _ServerVersion.from_string(version) + return _QueryGRPC( + weaviate_version=weaviate_version, name="Documents", tenant=None, consistency_level=None, validate_arguments=True, - uses_125_api=True, - uses_127_api=True, - ) - if query_type == "near_vector": - return query.near_vector(near_vector=vector, target_vector=target_vector) - return query.hybrid( - query="example", - vector=HybridVector.near_vector(vector) if query_type == "hybrid_near_vector" else vector, - target_vector=target_vector, + uses_125_api=weaviate_version.is_at_least(1, 25, 0), + uses_127_api=weaviate_version.is_at_least(1, 27, 0), ) -@pytest.mark.parametrize("version", ["1.26.0", "1.27.0", "1.29.0"]) -@pytest.mark.parametrize("query_type", ["near_vector", "hybrid", "hybrid_near_vector"]) -@pytest.mark.parametrize("weighted", [False, True]) -def test_mismatched_target_vector_names(version: str, query_type: str, weighted: bool) -> None: - target_vector = ( - TargetVectors.manual_weights({"title": 1.0, "summary": 2.0}) - if weighted - else ["title", "summary"] +def _near_vector( + version: str, vector: NearVectorInputType, target_vector: TargetVectorJoinType +) -> base_search_pb2.Targets: + request = _query(version).near_vector(near_vector=vector, target_vector=target_vector) + return request.near_vector.targets + + +def _hybrid( + version: str, vector: NearVectorInputType, target_vector: TargetVectorJoinType +) -> base_search_pb2.Targets: + request = _query(version).hybrid(query="example", vector=vector, target_vector=target_vector) + return request.hybrid_search.targets + + +def _hybrid_near_vector( + version: str, vector: NearVectorInputType, target_vector: TargetVectorJoinType +) -> base_search_pb2.Targets: + request = _query(version).hybrid( + query="example", vector=HybridVector.near_vector(vector), target_vector=target_vector ) - with pytest.raises(WeaviateInvalidInputError): - _request(version, query_type, {"title": [1.0, 0.0], "body": [0.0, 1.0]}, target_vector) + return request.hybrid_search.targets -@pytest.mark.parametrize("version", ["1.26.0", "1.27.0", "1.29.0"]) -@pytest.mark.parametrize("query_type", ["near_vector", "hybrid", "hybrid_near_vector"]) -def test_target_vector_names_can_be_in_a_different_order(version: str, query_type: str) -> None: - request = _request( +Search = Callable[[str, NearVectorInputType, TargetVectorJoinType], base_search_pb2.Targets] +SEARCHES = [_near_vector, _hybrid, _hybrid_near_vector] +VERSIONS = ["1.26.0", "1.27.0", "1.29.0"] + + +@pytest.mark.parametrize("search", SEARCHES) +@pytest.mark.parametrize("version", VERSIONS) +@pytest.mark.parametrize( + "target_vector", + [["title", "summary"], TargetVectors.manual_weights({"title": 1.0, "summary": 2.0})], + ids=["list", "weights"], +) +def test_mismatched_vector_names( + search: Search, version: str, target_vector: TargetVectorJoinType +) -> None: + with pytest.raises(WeaviateInvalidInputError, match="must match the target vectors"): + search(version, {"title": [1.0, 0.0], "body": [0.0, 1.0]}, target_vector) + + +@pytest.mark.parametrize("search", SEARCHES) +@pytest.mark.parametrize("version", VERSIONS) +def test_vector_names_in_a_different_order(search: Search, version: str) -> None: + targets = search( version, - query_type, {"summary": [0.0, 1.0], "title": [1.0, 0.0]}, TargetVectors.manual_weights({"title": 1.0, "summary": 2.0}), ) - targets = ( - request.near_vector.targets - if query_type == "near_vector" - else request.hybrid_search.targets - ) - assert set(targets.target_vectors) == {"title", "summary"} - assert {weight.target: weight.weight for weight in targets.weights_for_targets} == { - "title": 1.0, - "summary": 2.0, - } + weights = {weight.target: weight.weight for weight in targets.weights_for_targets} + assert weights == {"title": 1.0, "summary": 2.0} + if version != "1.26.0": + assert list(targets.target_vectors) == ["summary", "title"] +@pytest.mark.parametrize("search", SEARCHES) @pytest.mark.parametrize("version", ["1.27.0", "1.29.0"]) -@pytest.mark.parametrize("query_type", ["near_vector", "hybrid", "hybrid_near_vector"]) -def test_repeated_target_vector_names_with_multiple_weights(version: str, query_type: str) -> None: - request = _request( +def test_repeated_target_vector_names_with_multiple_weights(search: Search, version: str) -> None: + targets = search( version, - query_type, {"title": NearVector.list_of_vectors([1.0, 0.0], [0.0, 1.0])}, TargetVectors.manual_weights({"title": [1.0, 2.0]}), ) - targets = ( - request.near_vector.targets - if query_type == "near_vector" - else request.hybrid_search.targets - ) assert list(targets.target_vectors) == ["title", "title"] assert [(weight.target, weight.weight) for weight in targets.weights_for_targets] == [ ("title", 1.0), diff --git a/weaviate/collections/grpc/shared.py b/weaviate/collections/grpc/shared.py index 7bd607757..dcfa31fb7 100644 --- a/weaviate/collections/grpc/shared.py +++ b/weaviate/collections/grpc/shared.py @@ -124,6 +124,16 @@ def _recompute_target_vector_to_grpc( target_vector = target_vectors_tmp return self.__target_vector_to_grpc(target_vector) + @staticmethod + def __check_vector_keys( + vector: Dict[str, Any], targets: base_search_pb2.Targets, argument_name: str + ) -> None: + target_names = set(targets.target_vectors) + if target_names != vector.keys(): + raise WeaviateInvalidInputError( + f"The {argument_name} keys {sorted(vector.keys())} must match the target vectors {sorted(target_names)}" + ) + def __target_vector_to_grpc( self, target_vector: Optional[TargetVectorJoinType] ) -> Tuple[Optional[base_search_pb2.Targets], Optional[List[str]]]: @@ -155,10 +165,7 @@ def _vector_per_target( raise WeaviateInvalidInputError( "The number of target vectors must be equal to the number of vectors." ) - if set(targets.target_vectors) != vector.keys(): - raise WeaviateInvalidInputError( - "The vector dictionary keys must match the target vector names." - ) + self.__check_vector_keys(vector, targets, argument_name) vector_per_target: Dict[str, bytes] = {} for key, value in vector.items(): @@ -296,8 +303,9 @@ def add_list_of_vectors(value: _ListOfVectorsQuery, key: str) -> None: target_vectors.append(key) if isinstance(vector, dict): - if len(vector) == 0 or targets is None or set(targets.target_vectors) != vector.keys(): + if len(vector) == 0 or targets is None: raise invalid_nv_exception + self.__check_vector_keys(vector, targets, argument_name) for key, value in vector.items(): if _is_1d_vector(value): add_1d_vector(value, key)