Skip to content

Commit 8f57386

Browse files
committed
feat: enable multiple targets for nearText query
1 parent 355ac56 commit 8f57386

2 files changed

Lines changed: 196 additions & 19 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: 67 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,6 @@
99
import io.weaviate.client6.v1.internal.grpc.ByteStringUtil;
1010
import io.weaviate.client6.v1.internal.grpc.protocol.WeaviateProtoBase;
1111
import io.weaviate.client6.v1.internal.grpc.protocol.WeaviateProtoBaseSearch;
12-
import io.weaviate.client6.v1.internal.grpc.protocol.WeaviateProtoBaseSearch.Targets.Builder;
1312

1413
public interface Target {
1514

@@ -171,23 +170,15 @@ public void appendVectors(WeaviateProtoBaseSearch.NearVector.Builder req) {
171170
}
172171
}
173172

174-
record TextTarget(String vectorName, Float weight, List<String> query) implements Target {
173+
record TextTarget(VectorWeight weight, List<String> query) implements Target {
175174

176-
@Override
177-
public boolean appendTargets(Builder req) {
178-
if (vectorName == null) {
179-
return false;
180-
}
181-
req.addTargetVectors(vectorName);
182-
183-
var weightsForTarget = WeaviateProtoBaseSearch.WeightsForTarget.newBuilder()
184-
.setTarget(vectorName);
185-
if (weight != null) {
186-
weightsForTarget.setWeight(weight);
187-
}
188-
req.addWeightsForTargets(weightsForTarget);
189-
return true;
175+
private TextTarget(String vectorName, Float weight, List<String> query) {
176+
this(new VectorWeight(vectorName, weight), query);
177+
}
190178

179+
@Override
180+
public boolean appendTargets(WeaviateProtoBaseSearch.Targets.Builder req) {
181+
return weight.appendTargets(req);
191182
}
192183
}
193184

@@ -203,12 +194,34 @@ static TextTarget text(String vectorName, float weight, String... text) {
203194
return new TextTarget(vectorName, weight, Arrays.asList(text));
204195
}
205196

206-
record CombinedTextTarget(List<String> query, CombinationMethod combinationMethod, List<TextTarget> targets)
197+
/**
198+
* Weight to be applied to the vector distance. Used for text-based
199+
* queries where only a single input is allowed.
200+
*/
201+
record VectorWeight(String vectorName, Float weight) implements Target {
202+
@Override
203+
public boolean appendTargets(WeaviateProtoBaseSearch.Targets.Builder req) {
204+
if (vectorName == null) {
205+
return false;
206+
}
207+
req.addTargetVectors(vectorName);
208+
209+
var weightsForTarget = WeaviateProtoBaseSearch.WeightsForTarget.newBuilder()
210+
.setTarget(vectorName);
211+
if (weight != null) {
212+
weightsForTarget.setWeight(weight);
213+
}
214+
req.addWeightsForTargets(weightsForTarget);
215+
return true;
216+
}
217+
}
218+
219+
record CombinedTextTarget(List<String> query, CombinationMethod combinationMethod, List<VectorWeight> vectorWeights)
207220
implements Target {
208221

209222
@Override
210223
public boolean appendTargets(WeaviateProtoBaseSearch.Targets.Builder req) {
211-
if (targets.isEmpty()) {
224+
if (vectorWeights.isEmpty()) {
212225
return false;
213226
}
214227
switch (combinationMethod) {
@@ -228,8 +241,43 @@ public boolean appendTargets(WeaviateProtoBaseSearch.Targets.Builder req) {
228241
req.setCombination(WeaviateProtoBaseSearch.CombinationMethod.COMBINATION_METHOD_TYPE_MANUAL);
229242
break;
230243
}
231-
targets.forEach(t -> t.appendTargets(req));
244+
vectorWeights.forEach(t -> t.appendTargets(req));
232245
return true;
233246
}
234247
}
248+
249+
static VectorWeight weight(String vectorName, float weight) {
250+
return new VectorWeight(vectorName, weight);
251+
}
252+
253+
static Target combine(List<String> query, CombinationMethod combinationMethod, VectorWeight... vectorWeights) {
254+
return new CombinedTextTarget(query, combinationMethod, Arrays.asList(vectorWeights));
255+
}
256+
257+
static Target combine(List<String> query, CombinationMethod combinationMethod, String... targetVectors) {
258+
var vectorWeights = Arrays.stream(targetVectors)
259+
.map(vw -> new VectorWeight(vw, null))
260+
.toArray(VectorWeight[]::new);
261+
return combine(query, combinationMethod, vectorWeights);
262+
}
263+
264+
static Target sum(List<String> query, String... targetVectors) {
265+
return combine(query, CombinationMethod.SUM, targetVectors);
266+
}
267+
268+
static Target min(List<String> query, String... targetVectors) {
269+
return combine(query, CombinationMethod.MIN, targetVectors);
270+
}
271+
272+
static Target average(List<String> query, String... targetVectors) {
273+
return combine(query, CombinationMethod.AVERAGE, targetVectors);
274+
}
275+
276+
static Target relativeScore(List<String> query, VectorWeight... weights) {
277+
return combine(query, CombinationMethod.RELATIVE_SCORE, weights);
278+
}
279+
280+
static Target manualWeights(List<String> query, VectorWeight... weights) {
281+
return combine(query, CombinationMethod.MANUAL_WEIGHTS, weights);
282+
}
235283
}
Lines changed: 129 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,129 @@
1+
package io.weaviate.client6.v1.api.collections.query;
2+
3+
import java.util.List;
4+
5+
import org.assertj.core.api.Assertions;
6+
import org.junit.Test;
7+
import org.junit.runner.RunWith;
8+
9+
import com.google.gson.JsonParser;
10+
import com.google.protobuf.InvalidProtocolBufferException;
11+
import com.google.protobuf.MessageOrBuilder;
12+
import com.google.protobuf.util.JsonFormat;
13+
import com.jparams.junit4.JParamsTestRunner;
14+
import com.jparams.junit4.data.DataMethod;
15+
16+
import io.weaviate.client6.v1.internal.grpc.protocol.WeaviateProtoBaseSearch;
17+
18+
@RunWith(JParamsTestRunner.class)
19+
public class TargetTest {
20+
21+
public static Object[][] appendTargetsTestCases() {
22+
return new Object[][] {
23+
{
24+
Target.vector(new float[] { 1, 2, 3 }),
25+
null,
26+
},
27+
{
28+
Target.average(
29+
Target.vector("title_vec", new float[] { 1, 2, 3 }),
30+
Target.vector("body_vec", new float[] { 4, 5, 6 })),
31+
"""
32+
{
33+
"combination": "COMBINATION_METHOD_TYPE_AVERAGE",
34+
"targetVectors": ["title_vec", "body_vec"],
35+
"weightsForTargets": [
36+
{"target": "title_vec"},
37+
{"target": "body_vec"}
38+
]
39+
}
40+
""",
41+
42+
},
43+
{
44+
Target.manualWeights(
45+
Target.vector("title_vec", .2f, new float[] { 1, 2, 3 }),
46+
Target.vector("title_vec", .3f, new float[] { 1, 2, 3 }),
47+
Target.vector("body_vec", .5f, new float[] { 4, 5, 6 })),
48+
"""
49+
{
50+
"combination": "COMBINATION_METHOD_TYPE_MANUAL",
51+
"targetVectors": ["title_vec", "title_vec", "body_vec"],
52+
"weightsForTargets": [
53+
{"target": "title_vec", "weight": 0.2},
54+
{"target": "title_vec", "weight": 0.3},
55+
{"target": "body_vec", "weight": 0.5}
56+
]
57+
}
58+
""",
59+
60+
},
61+
{
62+
Target.min(
63+
List.of("day", "night"),
64+
"title_vec", "body_vec"),
65+
"""
66+
{
67+
"combination": "COMBINATION_METHOD_TYPE_MIN",
68+
"targetVectors": ["title_vec", "body_vec"],
69+
"weightsForTargets": [
70+
{"target": "title_vec"},
71+
{"target": "body_vec"}
72+
]
73+
}
74+
""",
75+
76+
},
77+
{
78+
Target.relativeScore(
79+
List.of("one", "two", "three"),
80+
Target.weight("title_vec", 1),
81+
Target.weight("title_vec", 2),
82+
Target.weight("body_vec", 3)),
83+
"""
84+
{
85+
"combination": "COMBINATION_METHOD_TYPE_RELATIVE_SCORE",
86+
"targetVectors": ["title_vec", "title_vec", "body_vec"],
87+
"weightsForTargets": [
88+
{"target": "title_vec", "weight": 1.0},
89+
{"target": "title_vec", "weight": 2.0},
90+
{"target": "body_vec", "weight": 3.0}
91+
]
92+
}
93+
""",
94+
95+
},
96+
};
97+
}
98+
99+
@Test
100+
@DataMethod(source = TargetTest.class, method = "appendTargetsTestCases")
101+
public void test_appendTargets(Target target, String want) {
102+
var req = WeaviateProtoBaseSearch.Targets.newBuilder();
103+
var appended = target.appendTargets(req);
104+
if (want == null) {
105+
Assertions.assertThat(appended).as("should not append targets").isFalse();
106+
return;
107+
}
108+
109+
var got = proto2json(req);
110+
assertEqualJson(want, got);
111+
}
112+
113+
private static final String proto2json(MessageOrBuilder proto) {
114+
String out;
115+
try {
116+
out = JsonFormat.printer().print(proto);
117+
} catch (InvalidProtocolBufferException e) {
118+
out = e.getMessage();
119+
}
120+
121+
return out;
122+
}
123+
124+
private static void assertEqualJson(String want, String got) {
125+
var wantJson = JsonParser.parseString(want);
126+
var gotJson = JsonParser.parseString(got);
127+
Assertions.assertThat(gotJson).isEqualTo(wantJson);
128+
}
129+
}

0 commit comments

Comments
 (0)