Skip to content

Commit bee582e

Browse files
authored
Expose neural filters to query visitors (#1992)
* Expose neural filters to query visitors Signed-off-by: Cédric Pelvet <cedric.pelvet@gmail.com> * Add changelog entry for query visitor fix Signed-off-by: Cédric Pelvet <cedric.pelvet@gmail.com> --------- Signed-off-by: Cédric Pelvet <cedric.pelvet@gmail.com>
1 parent 3e52b21 commit bee582e

3 files changed

Lines changed: 70 additions & 0 deletions

File tree

CHANGELOG.md

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@ The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/),
1111
- Add `model_selection` (language_option/model_type) parameter to semantic field to resolve the model id from cluster settings ([#1918](https://github.com/opensearch-project/neural-search/issues/1918))
1212

1313
### Bug Fixes
14+
* [Neural Query] Expose embedded filters to `QueryBuilderVisitor` traversal ([#1992](https://github.com/opensearch-project/neural-search/pull/1992))
1415
* [Hybrid Query] Fix NoSuchElementException in hybrid query with sort/search_after when a shard returns no results ([#1939](https://github.com/opensearch-project/neural-search/pull/1939))
1516
* [SemanticHighlighter] Fix SemanticHighlighterExtBuilder.toXContent ([#1906](https://github.com/opensearch-project/neural-search/issues/1906)) (query-insights [#651](https://github.com/opensearch-project/query-insights/issues/651))
1617
* [Sparse ANN] Fold sparse vector tokens into the signed-short range (modulus 32768) so folded tokens are never sign-extended to a negative value when stored in short[] ([#1926](https://github.com/opensearch-project/neural-search/pull/1926))

src/main/java/org/opensearch/neuralsearch/query/NeuralQueryBuilder.java

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -57,6 +57,7 @@
5757
import lombok.experimental.Accessors;
5858
import org.apache.commons.lang3.StringUtils;
5959
import org.apache.commons.lang3.builder.EqualsBuilder;
60+
import org.apache.lucene.search.BooleanClause;
6061
import org.apache.lucene.search.MatchNoDocsQuery;
6162
import org.apache.lucene.search.Query;
6263
import org.apache.lucene.search.join.ScoreMode;
@@ -77,6 +78,7 @@
7778
import org.opensearch.index.mapper.RankFeaturesFieldMapper;
7879
import org.opensearch.index.query.NestedQueryBuilder;
7980
import org.opensearch.index.query.QueryBuilder;
81+
import org.opensearch.index.query.QueryBuilderVisitor;
8082
import org.opensearch.index.query.QueryCoordinatorContext;
8183
import org.opensearch.index.query.WithFieldName;
8284
import org.opensearch.index.query.QueryRewriteContext;
@@ -525,6 +527,19 @@ public QueryBuilder filter(QueryBuilder filterToBeAdded) {
525527

526528
}
527529

530+
/**
531+
* Exposes the embedded query filter as a {@link BooleanClause.Occur#FILTER} child. The default
532+
* {@link QueryBuilder#visit(QueryBuilderVisitor)} implementation visits only this builder, so composite query
533+
* builders must override it to make their complete query tree available to visitors.
534+
*/
535+
@Override
536+
public void visit(QueryBuilderVisitor visitor) {
537+
visitor.accept(this);
538+
if (queryfilter != null) {
539+
queryfilter.visit(visitor.getChildVisitor(BooleanClause.Occur.FILTER));
540+
}
541+
}
542+
528543
@Override
529544
protected void doXContent(XContentBuilder xContentBuilder, Params params) throws IOException {
530545
xContentBuilder.startObject(NAME);

src/test/java/org/opensearch/neuralsearch/query/NeuralQueryBuilderTests.java

Lines changed: 54 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@
66

77
import lombok.SneakyThrows;
88
import org.apache.commons.lang3.tuple.Pair;
9+
import org.apache.lucene.search.BooleanClause;
910
import org.apache.lucene.search.MatchNoDocsQuery;
1011
import org.apache.lucene.search.Query;
1112
import org.junit.Before;
@@ -28,6 +29,7 @@
2829
import org.opensearch.index.query.MatchAllQueryBuilder;
2930
import org.opensearch.index.query.MatchNoneQueryBuilder;
3031
import org.opensearch.index.query.QueryBuilder;
32+
import org.opensearch.index.query.QueryBuilderVisitor;
3133
import org.opensearch.index.query.QueryShardContext;
3234
import org.opensearch.index.query.TermQueryBuilder;
3335
import org.opensearch.knn.index.query.KNNQueryBuilder;
@@ -39,6 +41,7 @@
3941
import org.opensearch.test.OpenSearchTestCase;
4042

4143
import java.io.IOException;
44+
import java.util.ArrayList;
4245
import java.util.List;
4346
import java.util.Map;
4447
import java.util.function.Supplier;
@@ -826,6 +829,57 @@ public void testFilter_whenAddBoolQueryBuilderToNeuralQueryBuilder_thenFilterSuc
826829
assertEquals(TEST_FILTER, neuralQueryBuilder.queryfilter());
827830
}
828831

832+
public void testVisit_whenFilterPresent_thenVisitsFilterClause() {
833+
NeuralQueryBuilder neuralQueryBuilder = getBaselineNeuralQueryBuilder();
834+
List<QueryBuilder> visitedQueries = new ArrayList<>();
835+
List<BooleanClause.Occur> visitedOccurrences = new ArrayList<>();
836+
QueryBuilderVisitor visitor = new QueryBuilderVisitor() {
837+
@Override
838+
public void accept(QueryBuilder queryBuilder) {
839+
visitedQueries.add(queryBuilder);
840+
}
841+
842+
@Override
843+
public QueryBuilderVisitor getChildVisitor(BooleanClause.Occur occur) {
844+
visitedOccurrences.add(occur);
845+
return this;
846+
}
847+
};
848+
849+
neuralQueryBuilder.visit(visitor);
850+
851+
assertEquals(List.of(neuralQueryBuilder, TEST_FILTER), visitedQueries);
852+
assertEquals(List.of(BooleanClause.Occur.FILTER), visitedOccurrences);
853+
}
854+
855+
public void testVisit_whenFilterAbsent_thenVisitsOnlyNeuralQuery() {
856+
NeuralQueryBuilder neuralQueryBuilder = NeuralQueryBuilder.builder()
857+
.fieldName(FIELD_NAME)
858+
.queryText(QUERY_TEXT)
859+
.modelId(MODEL_ID)
860+
.k(K)
861+
.build();
862+
List<QueryBuilder> visitedQueries = new ArrayList<>();
863+
List<BooleanClause.Occur> visitedOccurrences = new ArrayList<>();
864+
QueryBuilderVisitor visitor = new QueryBuilderVisitor() {
865+
@Override
866+
public void accept(QueryBuilder queryBuilder) {
867+
visitedQueries.add(queryBuilder);
868+
}
869+
870+
@Override
871+
public QueryBuilderVisitor getChildVisitor(BooleanClause.Occur occur) {
872+
visitedOccurrences.add(occur);
873+
return this;
874+
}
875+
};
876+
877+
neuralQueryBuilder.visit(visitor);
878+
879+
assertEquals(List.of(neuralQueryBuilder), visitedQueries);
880+
assertTrue(visitedOccurrences.isEmpty());
881+
}
882+
829883
public void testQueryCreation_whenCreateQueryWithDoToQuery_thenFail() {
830884
setUpClusterService();
831885
NeuralQueryBuilder neuralQueryBuilder = NeuralQueryBuilder.builder()

0 commit comments

Comments
 (0)