Skip to content

Commit a2d8bfb

Browse files
committed
test: write unit tests for marshalling targets / vectors
1 parent 8f57386 commit a2d8bfb

2 files changed

Lines changed: 139 additions & 3 deletions

File tree

  • src
    • main/java/io/weaviate/client6/v1/api/collections/query
    • test/java/io/weaviate/client6/v1/api/collections/query

src/main/java/io/weaviate/client6/v1/api/collections/query/Target.java

Lines changed: 23 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
package io.weaviate.client6.v1.api.collections.query;
22

33
import java.util.Arrays;
4+
import java.util.LinkedHashMap;
45
import java.util.List;
56
import java.util.stream.Collectors;
67

@@ -80,14 +81,26 @@ static VectorTarget vector(float[] vector) {
8081
return new VectorTarget(null, null, vector);
8182
}
8283

84+
static VectorTarget vector(float[][] vector) {
85+
return new VectorTarget(null, null, vector);
86+
}
87+
8388
static VectorTarget vector(String vectorName, float[] vector) {
8489
return new VectorTarget(vectorName, null, vector);
8590
}
8691

92+
static VectorTarget vector(String vectorName, float[][] vector) {
93+
return new VectorTarget(vectorName, null, vector);
94+
}
95+
8796
static VectorTarget vector(String vectorName, float weight, float[] vector) {
8897
return new VectorTarget(vectorName, weight, vector);
8998
}
9099

100+
static VectorTarget vector(String vectorName, float weight, float[][] vector) {
101+
return new VectorTarget(vectorName, weight, vector);
102+
}
103+
91104
static Target combine(CombinationMethod combinationMethod, VectorTarget... vectorTargets) {
92105
return new CombinedVectorTarget(combinationMethod, Arrays.asList(vectorTargets));
93106
}
@@ -156,16 +169,23 @@ public void appendVectors(WeaviateProtoBaseSearch.NearVector.Builder req) {
156169
return;
157170
}
158171

172+
// We use LinkedHashMap to preserve insertion order.
173+
// This has negligble performance penalty, if any,
174+
// but allows for a predictable output in tests.
159175
targets
160176
.stream()
161-
.collect(Collectors.groupingBy(VectorTarget::vectorName, Collectors.toList()))
177+
.collect(Collectors.groupingBy(
178+
VectorTarget::vectorName,
179+
LinkedHashMap::new,
180+
Collectors.toList()))
162181
.entrySet()
163182
.forEach(target -> {
164-
var vectorForTarget = WeaviateProtoBaseSearch.VectorForTarget.newBuilder()
183+
var vectorForTargets = WeaviateProtoBaseSearch.VectorForTarget.newBuilder()
165184
.setName(target.getKey());
166185
target.getValue().forEach(vt -> {
167-
vectorForTarget.addVectors(vt.encodeVectors());
186+
vectorForTargets.addVectors(vt.encodeVectors());
168187
});
188+
req.addVectorForTargets(vectorForTargets);
169189
});
170190
}
171191
}

src/test/java/io/weaviate/client6/v1/api/collections/query/TargetTest.java

Lines changed: 116 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -110,6 +110,122 @@ public void test_appendTargets(Target target, String want) {
110110
assertEqualJson(want, got);
111111
}
112112

113+
public static Object[][] appendVectorsTestCases() {
114+
return new Object[][] {
115+
{
116+
Target.vector(new float[] { 1, 2, 3 }),
117+
"""
118+
{
119+
"vectors": [{
120+
"vectorBytes": "AACAPwAAAEAAAEBA",
121+
"type": "VECTOR_TYPE_SINGLE_FP32"
122+
}]
123+
}
124+
""",
125+
},
126+
{
127+
Target.vector(new float[][] { { 1, 2, 3 }, { 4, 5, 6 } }),
128+
"""
129+
{
130+
"vectors": [{
131+
"vectorBytes": "AwAAAIA/AAAAQAAAQEAAAIBAAACgQAAAwEA=",
132+
"type": "VECTOR_TYPE_MULTI_FP32"
133+
}]
134+
}
135+
""",
136+
},
137+
{
138+
Target.vector("title_vec", new float[] { 1, 2, 3 }),
139+
"""
140+
{
141+
"vectorForTargets": [{
142+
"name": "title_vec",
143+
"vectors": [{
144+
"vectorBytes": "AACAPwAAAEAAAEBA",
145+
"type": "VECTOR_TYPE_SINGLE_FP32"
146+
}]
147+
}]
148+
}
149+
""",
150+
},
151+
{
152+
Target.vector("title_vec", new float[][] { { 1, 2, 3 }, { 4, 5, 6 } }),
153+
"""
154+
{
155+
"vectorForTargets": [{
156+
"name": "title_vec",
157+
"vectors": [{
158+
"vectorBytes": "AwAAAIA/AAAAQAAAQEAAAIBAAACgQAAAwEA=",
159+
"type": "VECTOR_TYPE_MULTI_FP32"
160+
}]
161+
}]
162+
}
163+
""",
164+
},
165+
{
166+
Target.average(
167+
Target.vector("title_vec", new float[] { 1, 2, 3 }),
168+
Target.vector("title_vec", new float[] { 4, 5, 6 }),
169+
Target.vector("lyrics_vec", new float[] { 7, 8, 9 })),
170+
"""
171+
{
172+
"vectorForTargets": [
173+
{
174+
"name": "title_vec",
175+
"vectors": [
176+
{"vectorBytes": "AACAPwAAAEAAAEBA", "type": "VECTOR_TYPE_SINGLE_FP32" },
177+
{"vectorBytes": "AACAQAAAoEAAAMBA", "type": "VECTOR_TYPE_SINGLE_FP32" }
178+
]
179+
},
180+
{
181+
"name": "lyrics_vec",
182+
"vectors": [
183+
{"vectorBytes": "AADgQAAAAEEAABBB", "type": "VECTOR_TYPE_SINGLE_FP32" }
184+
]
185+
}
186+
]
187+
}
188+
""",
189+
},
190+
{
191+
Target.average(
192+
Target.vector("title_vec", new float[][] { { 1, 2, 3 }, { 4, 5, 6 } }),
193+
Target.vector("title_vec", new float[][] { { 4, 5, 6 }, { 7, 8, 9 } }),
194+
Target.vector("lyrics_vec", new float[][] { { 7, 8, 9 }, { 1, 2, 3 } })),
195+
"""
196+
{
197+
"vectorForTargets": [
198+
{
199+
"name": "title_vec",
200+
"vectors": [
201+
{"vectorBytes": "AwAAAIA/AAAAQAAAQEAAAIBAAACgQAAAwEA=", "type": "VECTOR_TYPE_MULTI_FP32" },
202+
{"vectorBytes": "AwAAAIBAAACgQAAAwEAAAOBAAAAAQQAAEEE=", "type": "VECTOR_TYPE_MULTI_FP32" }
203+
]
204+
},
205+
{
206+
"name": "lyrics_vec",
207+
"vectors": [
208+
{"vectorBytes": "AwAAAOBAAAAAQQAAEEEAAIA/AAAAQAAAQEA=", "type": "VECTOR_TYPE_MULTI_FP32" }
209+
]
210+
}
211+
]
212+
}
213+
""",
214+
},
215+
};
216+
}
217+
218+
@Test
219+
@DataMethod(source = TargetTest.class, method = "appendVectorsTestCases")
220+
public void test_appendVectors(NearVectorTarget target, String want) {
221+
var req = WeaviateProtoBaseSearch.NearVector.newBuilder();
222+
223+
target.appendVectors(req);
224+
225+
var got = proto2json(req);
226+
assertEqualJson(want, got);
227+
}
228+
113229
private static final String proto2json(MessageOrBuilder proto) {
114230
String out;
115231
try {

0 commit comments

Comments
 (0)