55from vortexdb .connection import GRPCConnection
66from vortexdb .models import DenseVector , Payload , Similarity , ContentType , Point
77from vortexdb .exceptions import InvalidArgumentError
8+ from vortexdb .models import SearchQuery
89
910
1011
@@ -45,7 +46,6 @@ def test_insert_success(client, mock_connection):
4546 assert point_id == "point-123"
4647
4748
48-
4949def test_insert_rejects_invalid_vector (client ):
5050 with pytest .raises (TypeError ):
5151 client .insert (
@@ -54,6 +54,41 @@ def test_insert_rejects_invalid_vector(client):
5454 )
5555
5656
57+ # Batch Insert
58+
59+ def test_batch_insert_success (client , mock_connection ):
60+ response = Mock ()
61+ response .ids = [
62+ Mock (id = Mock (value = "p1" )),
63+ Mock (id = Mock (value = "p2" )),
64+ ]
65+ mock_connection .call .return_value = response
66+ items = [
67+ (DenseVector ([1 , 2 , 3 ]), Payload .text ("a" )),
68+ (DenseVector ([4 , 5 , 6 ]), Payload .text ("b" )),
69+ ]
70+ result = client .batch_insert (items = items )
71+ assert result == ["p1" , "p2" ]
72+
73+ def test_batch_insert_invalid_items_type (client ):
74+ with pytest .raises (TypeError ):
75+ client .batch_insert (items = "not-a-list" )
76+
77+ def test_batch_insert_invalid_tuple_structure (client ):
78+ items = [
79+ (DenseVector ([1 , 2 , 3 ]),), # only one element
80+ ]
81+ with pytest .raises (TypeError ):
82+ client .batch_insert (items = items )
83+
84+ def test_batch_insert_invalid_vector (client ):
85+ items = [
86+ ([1 , 2 , 3 ], Payload .text ("a" )), # not DenseVector
87+ ]
88+ with pytest .raises (TypeError ):
89+ client .batch_insert (items = items )
90+
91+
5792# Get
5893
5994def test_get_point_success (client , mock_connection ):
@@ -118,6 +153,80 @@ def test_search_invalid_vector(client):
118153 )
119154
120155
156+ # Batch Search
157+
158+ def test_batch_search_full_tuple (client , mock_connection ):
159+ mock_connection .call .return_value = Mock (
160+ results = [
161+ Mock (result_point_ids = [Mock (id = Mock (value = "p1" ))]),
162+ Mock (result_point_ids = [Mock (id = Mock (value = "p2" ))]),
163+ ]
164+ )
165+ queries = [
166+ (DenseVector ([1 , 2 , 3 ]), Similarity .COSINE , 2 ),
167+ (DenseVector ([4 , 5 , 6 ]), Similarity .EUCLIDEAN , 1 ),
168+ ]
169+ result = client .batch_search (queries = queries )
170+ assert result == [["p1" ], ["p2" ]]
171+
172+ def test_batch_search_vectors_with_global_params (client , mock_connection ):
173+ mock_connection .call .return_value = Mock (
174+ results = [
175+ Mock (result_point_ids = [Mock (id = Mock (value = "p1" ))]),
176+ ]
177+ )
178+ queries = [DenseVector ([1 , 2 , 3 ])]
179+ result = client .batch_search (
180+ queries = queries ,
181+ similarity = Similarity .MANHATTAN ,
182+ limit = 2 ,
183+ )
184+ assert result == [["p1" ]]
185+
186+ def test_batch_search_vector_similarity_with_global_limit (client , mock_connection ):
187+ mock_connection .call .return_value = Mock (
188+ results = [
189+ Mock (result_point_ids = [Mock (id = Mock (value = "p1" ))]),
190+ ]
191+ )
192+ queries = [
193+ (DenseVector ([1 , 2 , 3 ]), Similarity .COSINE ),
194+ ]
195+ result = client .batch_search (
196+ queries = queries ,
197+ limit = 2 ,
198+ )
199+ assert result == [["p1" ]]
200+
201+ def test_batch_search_searchquery_objects (client , mock_connection ):
202+ mock_connection .call .return_value = Mock (
203+ results = [
204+ Mock (result_point_ids = [Mock (id = Mock (value = "p1" ))]),
205+ ]
206+ )
207+ queries = [
208+ SearchQuery (DenseVector ([1 , 2 , 3 ]), Similarity .COSINE , 2 ),
209+ ]
210+ result = client .batch_search (queries = queries )
211+ assert result == [["p1" ]]
212+
213+ def test_batch_search_missing_globals_for_vector (client ):
214+ queries = [DenseVector ([1 , 2 , 3 ])]
215+ with pytest .raises (ValueError ):
216+ client .batch_search (queries = queries )
217+
218+ def test_batch_search_missing_limit (client ):
219+ queries = [
220+ (DenseVector ([1 , 2 , 3 ]), Similarity .COSINE ),
221+ ]
222+ with pytest .raises (ValueError ):
223+ client .batch_search (queries = queries )
224+
225+ def test_batch_search_invalid_format (client ):
226+ queries = ["invalid" ]
227+ with pytest .raises (TypeError ):
228+ client .batch_search (queries = queries )
229+
121230# Close
122231
123232def test_close_closes_connection (client , mock_connection ):
0 commit comments