Skip to content

Commit 8eb1ece

Browse files
committed
feat: introduce target vectors to hybrid query
1 parent b7320a9 commit 8eb1ece

4 files changed

Lines changed: 96 additions & 19 deletions

File tree

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

Lines changed: 57 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -248,6 +248,29 @@ public ResponseT hybrid(String query, Function<Hybrid.Builder, ObjectBuilder<Hyb
248248
return hybrid(Hybrid.of(query, fn));
249249
}
250250

251+
/**
252+
* Query collection objects using hybrid search.
253+
*
254+
* @param searchTarget Query target.
255+
* @throws WeaviateApiException in case the server returned with an
256+
* error status code.
257+
*/
258+
public ResponseT hybrid(Target searchTarget) {
259+
return hybrid(Hybrid.of(searchTarget));
260+
}
261+
262+
/**
263+
* Query collection objects using hybrid search.
264+
*
265+
* @param searchTarget Query target.
266+
* @param fn Lambda expression for optional parameters.
267+
* @throws WeaviateApiException in case the server returned with an
268+
* error status code.
269+
*/
270+
public ResponseT hybrid(Target searchTarget, Function<Hybrid.Builder, ObjectBuilder<Hybrid>> fn) {
271+
return hybrid(Hybrid.of(searchTarget, fn));
272+
}
273+
251274
/**
252275
* Query collection objects using hybrid search.
253276
*
@@ -292,6 +315,40 @@ public GroupedResponseT hybrid(String query, Function<Hybrid.Builder, ObjectBuil
292315
return hybrid(Hybrid.of(query, fn), groupBy);
293316
}
294317

318+
/**
319+
* Query collection objects using hybrid search.
320+
*
321+
* @param searchTarget Query target.
322+
* @param groupBy Group-by clause.
323+
* @return Grouped query result.
324+
* @throws WeaviateApiException in case the server returned with an
325+
* error status code.
326+
*
327+
* @see GroupBy
328+
* @see QueryResponseGrouped
329+
*/
330+
public GroupedResponseT hybrid(Target searchTarget, GroupBy groupBy) {
331+
return hybrid(Hybrid.of(searchTarget), groupBy);
332+
}
333+
334+
/**
335+
* Query collection objects using hybrid search.
336+
*
337+
* @param searchTarget Query target.
338+
* @param fn Lambda expression for optional parameters.
339+
* @param groupBy Group-by clause.
340+
* @return Grouped query result.
341+
* @throws WeaviateApiException in case the server returned with an
342+
* error status code.
343+
*
344+
* @see GroupBy
345+
* @see QueryResponseGrouped
346+
*/
347+
public GroupedResponseT hybrid(Target searchTarget, Function<Hybrid.Builder, ObjectBuilder<Hybrid>> fn,
348+
GroupBy groupBy) {
349+
return hybrid(Hybrid.of(searchTarget, fn), groupBy);
350+
}
351+
295352
/**
296353
* Query collection objects using hybrid search.
297354
*

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

Lines changed: 30 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -6,13 +6,15 @@
66
import java.util.function.Function;
77

88
import io.weaviate.client6.v1.api.collections.aggregate.AggregateObjectFilter;
9+
import io.weaviate.client6.v1.api.collections.query.Target.CombinedTextTarget;
10+
import io.weaviate.client6.v1.api.collections.query.Target.TextTarget;
911
import io.weaviate.client6.v1.internal.ObjectBuilder;
1012
import io.weaviate.client6.v1.internal.grpc.protocol.WeaviateProtoAggregate;
1113
import io.weaviate.client6.v1.internal.grpc.protocol.WeaviateProtoBaseSearch;
1214
import io.weaviate.client6.v1.internal.grpc.protocol.WeaviateProtoSearchGet;
1315

1416
public record Hybrid(
15-
String query,
17+
Target searchTarget,
1618
List<String> queryProperties,
1719
SearchOperator searchOperator,
1820
Float alpha,
@@ -27,16 +29,24 @@ public static enum FusionType {
2729
}
2830

2931
public static final Hybrid of(String query) {
30-
return of(query, ObjectBuilder.identity());
32+
return of(Target.text(List.of(query)));
3133
}
3234

3335
public static final Hybrid of(String query, Function<Builder, ObjectBuilder<Hybrid>> fn) {
34-
return fn.apply(new Builder(query)).build();
36+
return of(Target.text(List.of(query)), fn);
37+
}
38+
39+
public static final Hybrid of(Target searchTarget) {
40+
return of(searchTarget, ObjectBuilder.identity());
41+
}
42+
43+
public static final Hybrid of(Target searchTarget, Function<Builder, ObjectBuilder<Hybrid>> fn) {
44+
return fn.apply(new Builder(searchTarget)).build();
3545
}
3646

3747
public Hybrid(Builder builder) {
3848
this(
39-
builder.query,
49+
builder.searchTarget,
4050
builder.queryProperties,
4151
builder.searchOperator,
4252
builder.alpha,
@@ -48,7 +58,7 @@ public Hybrid(Builder builder) {
4858

4959
public static class Builder extends BaseQueryOptions.Builder<Builder, Hybrid> {
5060
// Required query parameters.
51-
private final String query;
61+
private final Target searchTarget;
5262

5363
// Optional query parameters.
5464
List<String> queryProperties = new ArrayList<>();
@@ -58,8 +68,8 @@ public static class Builder extends BaseQueryOptions.Builder<Builder, Hybrid> {
5868
FusionType fusionType;
5969
Float maxVectorDistance;
6070

61-
public Builder(String query) {
62-
this.query = query;
71+
public Builder(Target searchTarget) {
72+
this.searchTarget = searchTarget;
6373
}
6474

6575
/** Select properties to be included in the results scoring. */
@@ -155,9 +165,14 @@ public final void appendTo(WeaviateProtoSearchGet.SearchRequest.Builder req) {
155165

156166
private WeaviateProtoBaseSearch.Hybrid.Builder protoBuilder() {
157167
var hybrid = WeaviateProtoBaseSearch.Hybrid.newBuilder()
158-
.setQuery(query)
159168
.addAllProperties(queryProperties);
160169

170+
if (searchTarget instanceof TextTarget text) {
171+
hybrid.setQuery(text.query().get(0));
172+
} else if (searchTarget instanceof CombinedTextTarget combined) {
173+
hybrid.setQuery(combined.query().get(0));
174+
}
175+
161176
if (alpha != null) {
162177
hybrid.setAlpha(alpha);
163178
}
@@ -177,12 +192,17 @@ private WeaviateProtoBaseSearch.Hybrid.Builder protoBuilder() {
177192

178193
if (near != null) {
179194
if (near instanceof NearVector nv) {
180-
hybrid.setNearVector(nv.protoBuilder());
195+
hybrid.setNearVector(nv.protoBuilder(false));
181196
} else if (near instanceof NearText nt) {
182-
hybrid.setNearText(nt.protoBuilder());
197+
hybrid.setNearText(nt.protoBuilder(false));
183198
}
184199
}
185200

201+
var targets = WeaviateProtoBaseSearch.Targets.newBuilder();
202+
if (searchTarget.appendTargets(targets)) {
203+
hybrid.setTargets(targets);
204+
}
205+
186206
return hybrid;
187207
}
188208
}

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

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -117,19 +117,19 @@ public final void appendTo(WeaviateProtoBaseSearch.NearTextSearch.Move.Builder m
117117
@Override
118118
public void appendTo(WeaviateProtoSearchGet.SearchRequest.Builder req) {
119119
common.appendTo(req);
120-
req.setNearText(protoBuilder());
120+
req.setNearText(protoBuilder(true));
121121
}
122122

123123
@Override
124124
public void appendTo(WeaviateProtoAggregate.AggregateRequest.Builder req) {
125125
if (common.limit() != null) {
126126
req.setLimit(common.limit());
127127
}
128-
req.setNearText(protoBuilder());
128+
req.setNearText(protoBuilder(true));
129129
}
130130

131131
// Package-private for Hybrid to see.
132-
WeaviateProtoBaseSearch.NearTextSearch.Builder protoBuilder() {
132+
WeaviateProtoBaseSearch.NearTextSearch.Builder protoBuilder(boolean withTargets) {
133133
var nearText = WeaviateProtoBaseSearch.NearTextSearch.newBuilder();
134134

135135
if (searchTarget instanceof TextTarget text) {
@@ -139,7 +139,7 @@ WeaviateProtoBaseSearch.NearTextSearch.Builder protoBuilder() {
139139
}
140140

141141
var targets = WeaviateProtoBaseSearch.Targets.newBuilder();
142-
if (searchTarget.appendTargets(targets)) {
142+
if (withTargets && searchTarget.appendTargets(targets)) {
143143
nearText.setTargets(targets);
144144
}
145145

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

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -40,24 +40,24 @@ public final NearVector build() {
4040
@Override
4141
public final void appendTo(WeaviateProtoSearchGet.SearchRequest.Builder req) {
4242
common.appendTo(req);
43-
req.setNearVector(protoBuilder());
43+
req.setNearVector(protoBuilder(true));
4444
}
4545

4646
@Override
4747
public void appendTo(WeaviateProtoAggregate.AggregateRequest.Builder req) {
4848
if (common.limit() != null) {
4949
req.setLimit(common.limit());
5050
}
51-
req.setNearVector(protoBuilder());
51+
req.setNearVector(protoBuilder(true));
5252
}
5353

54-
// This is made package-private for Hybrid to see. Should we refactor?
55-
WeaviateProtoBaseSearch.NearVector.Builder protoBuilder() {
54+
WeaviateProtoBaseSearch.NearVector.Builder protoBuilder(boolean withTargets) {
5655
var nearVector = WeaviateProtoBaseSearch.NearVector.newBuilder();
5756

5857
searchTarget.appendVectors(nearVector);
58+
5959
var targets = WeaviateProtoBaseSearch.Targets.newBuilder();
60-
if (searchTarget.appendTargets(targets)) {
60+
if (withTargets && searchTarget.appendTargets(targets)) {
6161
nearVector.setTargets(targets);
6262
}
6363

0 commit comments

Comments
 (0)