diff --git a/src/main/java/io/weaviate/client/v1/graphql/query/argument/HybridArgument.java b/src/main/java/io/weaviate/client/v1/graphql/query/argument/HybridArgument.java index 4c4d241e1..549906a12 100644 --- a/src/main/java/io/weaviate/client/v1/graphql/query/argument/HybridArgument.java +++ b/src/main/java/io/weaviate/client/v1/graphql/query/argument/HybridArgument.java @@ -4,6 +4,7 @@ import java.util.Set; import io.weaviate.client.v1.graphql.query.util.Serializer; +import java.util.stream.Collectors; import lombok.AccessLevel; import lombok.Builder; import lombok.EqualsAndHashCode; @@ -25,6 +26,7 @@ public class HybridArgument implements Argument { String fusionType; String[] properties; String[] targetVectors; + Searches searches; @Override @@ -47,7 +49,26 @@ public String build() { if (ArrayUtils.isNotEmpty(targetVectors)) { arg.add(String.format("targetVectors:%s", Serializer.arrayWithQuotes(targetVectors))); } + if (searches != null && (searches.nearVector != null || searches.nearText != null)) { + Set searchesArgs = new LinkedHashSet<>(); + if (searches.nearVector != null) { + searchesArgs.add(searches.nearVector.build()); + } + if (searches.nearText != null) { + searchesArgs.add(searches.nearText.build()); + } + arg.add(String.format("searches:{%s}", String.join(" ", searchesArgs))); + } return String.format("hybrid:{%s}", String.join(" ", arg)); } + + @Getter + @Builder + @ToString + @FieldDefaults(makeFinal = true, level = AccessLevel.PRIVATE) + public static class Searches { + NearVectorArgument nearVector; + NearTextArgument nearText; + } } diff --git a/src/test/java/io/weaviate/client/v1/graphql/query/argument/HybridArgumentTest.java b/src/test/java/io/weaviate/client/v1/graphql/query/argument/HybridArgumentTest.java index 9de0a8f7f..4962c0161 100644 --- a/src/test/java/io/weaviate/client/v1/graphql/query/argument/HybridArgumentTest.java +++ b/src/test/java/io/weaviate/client/v1/graphql/query/argument/HybridArgumentTest.java @@ -96,4 +96,38 @@ public void shouldCreateArgumentWithProperties() { assertThat(str).isEqualTo("hybrid:{query:\"I'm a simple string\" " + "properties:[\"prop1\",\"prop2\"]}"); } + + @Test + public void shouldCreateArgumentWithNearVectorSearches() { + NearVectorArgument nearVector = NearVectorArgument.builder() + .vector(new Float[]{ .1f, .2f, .3f }) + .certainty(0.9f) + .build(); + + HybridArgument hybrid = HybridArgument.builder() + .query("I'm a simple string") + .searches(HybridArgument.Searches.builder().nearVector(nearVector).build()) + .build(); + + String str = hybrid.build(); + + assertThat(str).isEqualTo("hybrid:{query:\"I'm a simple string\" searches:{nearVector:{vector:[0.1,0.2,0.3] certainty:0.9}}}"); + } + + @Test + public void shouldCreateArgumentWithNearTextSearches() { + NearTextArgument nearText = NearTextArgument.builder() + .concepts(new String[]{"concept"}) + .certainty(0.9f) + .build(); + + HybridArgument hybrid = HybridArgument.builder() + .query("I'm a simple string") + .searches(HybridArgument.Searches.builder().nearText(nearText).build()) + .build(); + + String str = hybrid.build(); + + assertThat(str).isEqualTo("hybrid:{query:\"I'm a simple string\" searches:{nearText:{concepts:[\"concept\"] certainty:0.9}}}"); + } } diff --git a/src/test/java/io/weaviate/integration/client/graphql/ClientGraphQLTest.java b/src/test/java/io/weaviate/integration/client/graphql/ClientGraphQLTest.java index 769c3efba..68704bb78 100644 --- a/src/test/java/io/weaviate/integration/client/graphql/ClientGraphQLTest.java +++ b/src/test/java/io/weaviate/integration/client/graphql/ClientGraphQLTest.java @@ -1551,7 +1551,7 @@ public void testGraphQLGetWithGroupBy() { Config config = new Config("http", address); WeaviateClient client = new WeaviateClient(config); WeaviateTestGenerics.DocumentPassageSchema testData = new WeaviateTestGenerics.DocumentPassageSchema(); - ; + List ofDocumentA = Collections.singletonList( new GroupHitOfDocument(new AdditionalOfDocument(testData.DOCUMENT_IDS[0])) ); @@ -1630,6 +1630,68 @@ public void testGraphQLGetWithGroupBy() { checkGroupElements(expectedHits2, groups.get(1).getHits()); } + @Test + public void testGraphQLGetWithGroupByWithHybrid() { + // given + Config config = new Config("http", address); + WeaviateClient client = new WeaviateClient(config); + WeaviateTestGenerics.DocumentPassageSchema testData = new WeaviateTestGenerics.DocumentPassageSchema(); + // hits + Field[] hits = new Field[]{ + Field.builder().name("content").build(), + Field.builder().name("_additional{id distance}").build(), + }; + // group + Field group = Field.builder() + .name("group") + .fields(new Field[]{ + Field.builder().name("id").build(), + Field.builder().name("groupedBy") + .fields(new Field[]{ + Field.builder().name("value").build(), + Field.builder().name("path").build(), + }).build(), + Field.builder().name("count").build(), + Field.builder().name("maxDistance").build(), + Field.builder().name("minDistance").build(), + Field.builder().name("hits").fields(hits).build(), + }).build(); + // _additional + Field _additional = Field.builder().name("_additional").fields(new Field[]{group}).build(); + // filter arguments + GroupByArgument groupBy = client.graphQL().arguments().groupByArgBuilder() + .path(new String[]{"content"}).groups(3).objectsPerGroup(10).build(); + NearTextArgument nearText = NearTextArgument.builder().concepts(new String[]{"Passage content 2"}).build(); + HybridArgument hybrid = HybridArgument.builder() + .searches(HybridArgument.Searches.builder().nearText(nearText).build()) + .query("Passage content 2") + .alpha(0.9f) + .build(); + // when + testData.createAndInsertData(client); + Result groupByResult = client.graphQL().get() + .withClassName(testData.PASSAGE) + .withHybrid(hybrid) + .withGroupBy(groupBy) + .withFields(_additional).run(); + testData.cleanupWeaviate(client); + // then + assertThat(groupByResult).isNotNull(); + assertThat(groupByResult.getError()).isNull(); + assertThat(groupByResult.getResult()).isNotNull(); + List> result = extractResult(groupByResult, testData.PASSAGE); + assertThat(result).isNotNull().hasSize(3); + List groups = getGroups(result); + assertThat(groups).isNotNull().hasSize(3); + for (int i = 0; i < 3; i++) { + if (i == 0) { + assertThat(groups.get(i).groupedBy.value).isEqualTo("Passage content 2"); + } + assertThat(groups.get(i).minDistance).isEqualTo(groups.get(i).getHits().get(0).get_additional().getDistance()); + assertThat(groups.get(i).maxDistance).isEqualTo(groups.get(i).getHits().get(groups.get(i).getHits().size() - 1).get_additional().getDistance()); + } + } + private void checkGroupElements(List expected, List actual) { assertThat(expected).hasSameSizeAs(actual); for (int i = 0; i < actual.size(); i++) {