Skip to content

Commit ea4ef65

Browse files
fix: reject ragged multi-vector embeddings
1 parent c307d97 commit ea4ef65

2 files changed

Lines changed: 53 additions & 3 deletions

File tree

test/collection/test_byteops.py

Lines changed: 32 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,9 @@
1-
from weaviate.collections.grpc.shared import _ByteOps
1+
import pytest
2+
3+
from weaviate.collections.grpc.query import _QueryGRPC
4+
from weaviate.collections.grpc.shared import _ByteOps, _Pack, _Unpack
5+
from weaviate.exceptions import WeaviateInvalidInputError
6+
from weaviate.util import _ServerVersion
27

38

49
def test_decode_float32s():
@@ -22,3 +27,29 @@ def test_decode_int64s():
2227
assert _ByteOps.decode_int64s(
2328
b"\x01\x00\x00\x00\x00\x00\x00\x00\x02\x00\x00\x00\x00\x00\x00\x00"
2429
) == [1, 2]
30+
31+
32+
def test_multi_vector_pack_round_trip():
33+
vector = [[1.0, 2.0], [3.0, 4.0]]
34+
35+
assert _Unpack.multi(_Pack.multi(vector)) == vector
36+
37+
38+
def test_multi_vector_pack_rejects_ragged_vectors():
39+
with pytest.raises(WeaviateInvalidInputError, match="consistent dimensions"):
40+
_Pack.multi([[1.0, 2.0], [3.0]])
41+
42+
43+
def test_near_vector_request_rejects_ragged_multi_vector():
44+
query = _QueryGRPC(
45+
_ServerVersion.from_string("1.29.0"),
46+
"TestCollection",
47+
None,
48+
None,
49+
True,
50+
True,
51+
True,
52+
)
53+
54+
with pytest.raises(WeaviateInvalidInputError, match="consistent dimensions"):
55+
query.near_vector(near_vector=[[1.0, 2.0], [3.0]])

weaviate/collections/grpc/shared.py

Lines changed: 21 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -802,8 +802,27 @@ def single(vector: OneDimensionalVectorType) -> bytes:
802802

803803
@staticmethod
804804
def multi(vector: TwoDimensionalVectorType) -> bytes:
805-
vector_list = [item for sublist in vector for item in sublist]
806-
return struct.pack("<H", len(vector[0])) + struct.pack(
805+
if len(vector) == 0:
806+
raise WeaviateInvalidInputError("Multi-vector embeddings must not be empty.")
807+
808+
first_vector = _get_vector_v4(vector[0])
809+
dimension = len(first_vector)
810+
if dimension == 0:
811+
raise WeaviateInvalidInputError(
812+
"Multi-vector embeddings must not contain empty vectors."
813+
)
814+
815+
vector_list: List[float] = []
816+
for subvector in vector:
817+
subvector_list = _get_vector_v4(subvector)
818+
if len(subvector_list) != dimension:
819+
raise WeaviateInvalidInputError(
820+
"Multi-vector embeddings must have consistent dimensions. "
821+
f"Expected dimension {dimension}, got {len(subvector_list)}."
822+
)
823+
vector_list.extend(subvector_list)
824+
825+
return struct.pack("<H", dimension) + struct.pack(
807826
"{}f".format(len(vector_list)), *vector_list
808827
)
809828

0 commit comments

Comments
 (0)