diff --git a/services/graph/pkg/service/v0/searchquery.go b/services/graph/pkg/service/v0/searchquery.go index 4f71d6ae98..1afa7b8076 100644 --- a/services/graph/pkg/service/v0/searchquery.go +++ b/services/graph/pkg/service/v0/searchquery.go @@ -19,6 +19,7 @@ import ( searchmsg "github.com/opencloud-eu/opencloud/protogen/gen/opencloud/messages/search/v0" searchsvc "github.com/opencloud-eu/opencloud/protogen/gen/opencloud/services/search/v0" "github.com/opencloud-eu/opencloud/services/graph/pkg/errorcode" + "github.com/opencloud-eu/opencloud/services/search/pkg/aggregation" "github.com/opencloud-eu/opencloud/services/search/pkg/search" ) @@ -90,9 +91,10 @@ func (g Graph) runSingleSearch(ctx context.Context, sr libregraph.SearchRequest, } rsp, err := g.searchService.Search(ctx, &searchsvc.SearchRequest{ - Query: sr.Query.QueryString, - PageSize: pageSize, - Aggregations: libregraphAggregationsToSearch(sr.Aggregations), + Query: sr.Query.QueryString, + PageSize: pageSize, + Aggregations: libregraphAggregationsToSearch(sr.Aggregations), + AggregationFilters: sr.AggregationFilters, }) if err != nil { return libregraph.SearchResponse{}, err @@ -293,6 +295,9 @@ func searchAggregationsToLibregraph(in []*searchsvc.AggregationResult, defs []li Key: &key, Count: &count, } + if token := aggregationTokenForBucket(key, def); token != "" { + lb.AggregationFilterToken = &token + } if subs := b.GetSubAggregations(); len(subs) > 0 { lb.LibreGraphSubAggregations = searchAggregationsToLibregraph(subs, def.LibreGraphSubAggregations) } @@ -306,6 +311,27 @@ func searchAggregationsToLibregraph(in []*searchsvc.AggregationResult, defs []li return out } +// range buckets are matched to their definition via the from-to merge key +func aggregationTokenForBucket(key string, def libregraph.AggregationOption) string { + if def.BucketDefinition != nil && len(def.BucketDefinition.Ranges) > 0 { + for _, r := range def.BucketDefinition.Ranges { + from, to := ptrStr(r.From), ptrStr(r.To) + if from+"-"+to == key { + return aggregation.EncodeRangeToken(from, to) + } + } + return "" + } + return aggregation.EncodeTermsToken(key) +} + +func ptrStr(s *string) string { + if s == nil { + return "" + } + return *s +} + func (g Graph) renderSearchError(w http.ResponseWriter, r *http.Request, err error) { e := merrors.Parse(err.Error()) switch e.Code { diff --git a/services/search/pkg/aggregation/aggregation_suite_test.go b/services/search/pkg/aggregation/aggregation_suite_test.go new file mode 100644 index 0000000000..60b9affb61 --- /dev/null +++ b/services/search/pkg/aggregation/aggregation_suite_test.go @@ -0,0 +1,13 @@ +package aggregation_test + +import ( + "testing" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +func TestAggregation(t *testing.T) { + RegisterFailHandler(Fail) + RunSpecs(t, "Aggregation Suite") +} diff --git a/services/search/pkg/aggregation/token.go b/services/search/pkg/aggregation/token.go new file mode 100644 index 0000000000..a3083df364 --- /dev/null +++ b/services/search/pkg/aggregation/token.go @@ -0,0 +1,141 @@ +// Package aggregation encodes and decodes the MS Graph aggregationFilterToken: +// quoted hex terms tokens behind a U+01C2 pair, and range(from, to) range tokens. +package aggregation + +import ( + "encoding/hex" + "fmt" + "strings" +) + +// U+01C2 LATIN LETTER ALVEOLAR CLICK, twice +const termPrefix = "ǂǂ" + +// EncodeTermsToken hex-encodes the bucket key; the quotes are part of the token value. +func EncodeTermsToken(key string) string { + return `"` + termPrefix + hex.EncodeToString([]byte(key)) + `"` +} + +// EncodeRangeToken uses the MS Graph spelling: open bounds are min/max, an open +// upper bound carries to="le". +func EncodeRangeToken(from, to string) string { + if from == "" { + from = "min" + } + if to == "" { + return "range(" + from + ", max, to=\"le\")" + } + return "range(" + from + ", " + to + ")" +} + +// DecodeAggregationFilter turns {field}:{token} into a KQL fragment; tokens +// that are not server-shaped are rejected. +func DecodeAggregationFilter(filter string) (string, error) { + field, token, ok := strings.Cut(filter, ":") + if !ok || field == "" || token == "" { + return "", fmt.Errorf("invalid aggregation filter %q", filter) + } + switch { + case strings.HasPrefix(token, "or("): + return decodeOr(field, token) + case strings.HasPrefix(token, "range("): + return decodeRange(field, token) + default: + v, err := decodeTerm(token) + if err != nil { + return "", err + } + frag, err := kqlTerm(field, v) + if err != nil { + return "", err + } + return frag, nil + } +} + +func decodeTerm(token string) (string, error) { + if len(token) < 2 || token[0] != '"' || token[len(token)-1] != '"' { + return "", fmt.Errorf("invalid terms token %q", token) + } + inner := token[1 : len(token)-1] + if !strings.HasPrefix(inner, termPrefix) { + return "", fmt.Errorf("invalid terms token %q", token) + } + b, err := hex.DecodeString(strings.TrimPrefix(inner, termPrefix)) + if err != nil { + return "", fmt.Errorf("invalid terms token %q: %w", token, err) + } + return string(b), nil +} + +// open bounds (min/max) are dropped +func decodeRange(field, token string) (string, error) { + inner, ok := trimCall(token, "range") + if !ok { + return "", fmt.Errorf("invalid range token %q", token) + } + segments := strings.Split(inner, ",") + if len(segments) == 3 && strings.TrimSpace(segments[2]) == `to="le"` { + // MS Graph appends to="le" to an open upper bound; the bound itself + // is still max, so the marker carries no extra information. + segments = segments[:2] + } + if len(segments) != 2 { + return "", fmt.Errorf("invalid range token %q", token) + } + from, to := strings.TrimSpace(segments[0]), strings.TrimSpace(segments[1]) + var parts []string + if from != "" && from != "min" { + parts = append(parts, field+">="+from) + } + if to != "" && to != "max" { + // bucket ranges are from-inclusive, to-exclusive; match that so the + // drilldown returns exactly the matches the bucket counted + parts = append(parts, field+"<"+to) + } + if len(parts) == 0 { + return "", fmt.Errorf("range token %q has no bounds", token) + } + return "(" + strings.Join(parts, " AND ") + ")", nil +} + +func decodeOr(field, token string) (string, error) { + inner, ok := trimCall(token, "or") + if !ok { + return "", fmt.Errorf("invalid or token %q", token) + } + // terms tokens are quote-wrapped lowercase hex, so they never contain a + // comma; a plain split is safe. + parts := make([]string, 0) + for _, t := range strings.Split(inner, ",") { + v, err := decodeTerm(strings.TrimSpace(t)) + if err != nil { + return "", err + } + frag, err := kqlTerm(field, v) + if err != nil { + return "", err + } + parts = append(parts, frag) + } + if len(parts) == 0 { + return "", fmt.Errorf("empty or token %q", token) + } + return "(" + strings.Join(parts, " OR ") + ")", nil +} + +func trimCall(token, name string) (string, bool) { + if !strings.HasPrefix(token, name+"(") || !strings.HasSuffix(token, ")") { + return "", false + } + return token[len(name)+1 : len(token)-1], true +} + +// KQL quoted strings have no escape syntax, so a value containing a double +// quote is rejected rather than emitted as broken KQL. +func kqlTerm(field, value string) (string, error) { + if strings.Contains(value, `"`) { + return "", fmt.Errorf("aggregation value %q contains an unsupported double quote", value) + } + return field + `:"` + value + `"`, nil +} diff --git a/services/search/pkg/aggregation/token_test.go b/services/search/pkg/aggregation/token_test.go new file mode 100644 index 0000000000..88225110c4 --- /dev/null +++ b/services/search/pkg/aggregation/token_test.go @@ -0,0 +1,70 @@ +package aggregation_test + +import ( + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + + "github.com/opencloud-eu/opencloud/services/search/pkg/aggregation" +) + +var _ = Describe("Token", func() { + Describe("EncodeTermsToken", func() { + It("encodes the key as quoted ǂǂ-prefixed lowercase hex", func() { + Expect(aggregation.EncodeTermsToken("And the Bands Played On")).To(Equal(`"ǂǂ416e64207468652042616e647320506c61796564204f6e"`)) + }) + }) + + Describe("EncodeRangeToken", func() { + DescribeTable("bounds", + func(from, to, want string) { + Expect(aggregation.EncodeRangeToken(from, to)).To(Equal(want)) + }, + Entry("closed", "0", "100", "range(0, 100)"), + Entry("open lower", "", "100", "range(min, 100)"), + Entry("open upper", "0", "", `range(0, max, to="le")`), + ) + }) + + Describe("DecodeAggregationFilter", func() { + DescribeTable("valid tokens", + func(filter, want string) { + got, err := aggregation.DecodeAggregationFilter(filter) + Expect(err).ToNot(HaveOccurred()) + Expect(got).To(Equal(want)) + }, + Entry("terms with a space", `audio.artist:"ǂǂ5361786f6e"`, `audio.artist:"Saxon"`), + Entry("closed range", "Size:range(0,100)", "(Size>=0 AND Size<100)"), + Entry("closed range with spaces", "Size:range(0, 100)", "(Size>=0 AND Size<100)"), + Entry("open lower range", "Size:range(min,100)", "(Size<100)"), + Entry("open upper range", "Size:range(0,max)", "(Size>=0)"), + Entry("open upper range with le marker", `Size:range(0, max, to="le")`, "(Size>=0)"), + Entry("or of two terms", + `audio.artist:or("ǂǂ5361786f6e","ǂǂ49726f6e204d616964656e")`, + `(audio.artist:"Saxon" OR audio.artist:"Iron Maiden")`), + ) + + DescribeTable("rejected tokens", + func(filter string) { + _, err := aggregation.DecodeAggregationFilter(filter) + Expect(err).To(HaveOccurred()) + }, + Entry("no colon", `audio.artist"ǂǂ00"`), + Entry("empty field", `:"ǂǂ00"`), + Entry("missing ǂǂ prefix", `audio.artist:"deadbeef"`), + Entry("odd hex", `audio.artist:"ǂǂabc"`), + Entry("range without bounds", "Size:range(min,max)"), + ) + + It("round-trips a terms key through encode+decode", func() { + got, err := aggregation.DecodeAggregationFilter("audio.artist:" + aggregation.EncodeTermsToken("AC/DC")) + Expect(err).ToNot(HaveOccurred()) + Expect(got).To(Equal(`audio.artist:"AC/DC"`)) + }) + + It("rejects a decoded value containing a double quote", func() { + // 22 is a double quote; it cannot be expressed in a KQL string. + _, err := aggregation.DecodeAggregationFilter(`audio.artist:"ǂǂ22"`) + Expect(err).To(HaveOccurred()) + }) + }) +}) diff --git a/services/search/pkg/bleve/backend.go b/services/search/pkg/bleve/backend.go index de86a79ca9..6dfef81568 100644 --- a/services/search/pkg/bleve/backend.go +++ b/services/search/pkg/bleve/backend.go @@ -43,7 +43,7 @@ func NewBackend(index bleve.Index, queryCreator searchQuery.Creator[query.Query] // Search executes a search request operation within the index. // Returns a SearchIndexResponse object or an error. func (b *Backend) Search(ctx context.Context, sir *searchService.SearchIndexRequest) (*searchService.SearchIndexResponse, error) { - createdQuery, err := b.queryCreator.Create(sir.Query) + createdQuery, err := b.queryCreator.CreateWithFilters(sir.Query, sir.GetAggregationFilters()) if err != nil { if kql.IsValidationError(err) { return nil, errtypes.BadRequest(err.Error()) diff --git a/services/search/pkg/opensearch/backend.go b/services/search/pkg/opensearch/backend.go index a7f43d392a..c3d7a3e710 100644 --- a/services/search/pkg/opensearch/backend.go +++ b/services/search/pkg/opensearch/backend.go @@ -73,7 +73,7 @@ func NewBackend(ctx context.Context, name string, client *opensearchgoAPI.Client } func (b *Backend) Search(ctx context.Context, sir *searchService.SearchIndexRequest) (*searchService.SearchIndexResponse, error) { - boolQuery, err := convert.KQLToOpenSearchBoolQuery(sir.Query) + boolQuery, err := convert.KQLToOpenSearchBoolQueryWithFilters(sir.Query, sir.GetAggregationFilters()) switch { case kql.IsValidationError(err): return nil, errtypes.BadRequest(err.Error()) diff --git a/services/search/pkg/opensearch/internal/convert/kql_query.go b/services/search/pkg/opensearch/internal/convert/kql_query.go index 711a67b4b4..7ee23024ff 100644 --- a/services/search/pkg/opensearch/internal/convert/kql_query.go +++ b/services/search/pkg/opensearch/internal/convert/kql_query.go @@ -13,14 +13,18 @@ var ( ) func KQLToOpenSearchBoolQuery(kqlQuery string) (*osu.BoolQuery, error) { - kqlAst, err := kql.Builder{}.Build(kqlQuery) + return KQLToOpenSearchBoolQueryWithFilters(kqlQuery, nil) +} + +// KQLToOpenSearchBoolQueryWithFilters ANDs the decoded aggregation filters in +// as exact case-sensitive matches. +func KQLToOpenSearchBoolQueryWithFilters(kqlQuery string, filters []string) (*osu.BoolQuery, error) { + // shared lowering: field resolution, media-type expansion, value lowercasing + kqlAst, err := query.MergeFilters(kql.Builder{}, kqlQuery, filters) if err != nil { return nil, err } - // shared lowering: field resolution, media-type expansion, value lowercasing. - kqlAst = query.Normalize(kqlAst, query.ResolveField) - builder, err := TranspileKQLToOpenSearch(kqlAst.Nodes) if err != nil { return nil, fmt.Errorf("failed to compile query: %w", err) diff --git a/services/search/pkg/query/bleve/bleve.go b/services/search/pkg/query/bleve/bleve.go index 41e260b81c..ecde178687 100644 --- a/services/search/pkg/query/bleve/bleve.go +++ b/services/search/pkg/query/bleve/bleve.go @@ -34,5 +34,15 @@ func (c Creator[T]) Create(qs string) (T, error) { return t, nil } +// CreateWithFilters implements the Creator interface. +func (c Creator[T]) CreateWithFilters(qs string, filters []string) (T, error) { + var t T + merged, err := query.MergeFilters(c.builder, qs, filters) + if err != nil { + return t, err + } + return c.compiler.Compile(merged) +} + // DefaultCreator exposes a kql to bleve query creator. var DefaultCreator = Creator[bQuery.Query]{kql.Builder{}, Compiler{}} diff --git a/services/search/pkg/query/casesensitive.go b/services/search/pkg/query/casesensitive.go new file mode 100644 index 0000000000..fa4135bbd3 --- /dev/null +++ b/services/search/pkg/query/casesensitive.go @@ -0,0 +1,26 @@ +package query + +import "github.com/opencloud-eu/opencloud/pkg/ast" + +// ForceCaseSensitive runs on decoded aggregation filters after Normalize: the +// values are exact bucket keys, so they must match the case-preserving base +// field, not the lowercased sibling. +func ForceCaseSensitive(a *ast.Ast) *ast.Ast { + if a == nil { + return a + } + forceCaseSensitiveNodes(a.Nodes) + return a +} + +func forceCaseSensitiveNodes(nodes []ast.Node) { + for _, n := range nodes { + switch node := n.(type) { + case *ast.StringNode: + node.Exact = true + node.CaseInsensitive = false + case *ast.GroupNode: + forceCaseSensitiveNodes(node.Nodes) + } + } +} diff --git a/services/search/pkg/query/merge.go b/services/search/pkg/query/merge.go new file mode 100644 index 0000000000..c08efc4ab9 --- /dev/null +++ b/services/search/pkg/query/merge.go @@ -0,0 +1,40 @@ +package query + +import ( + "github.com/opencloud-eu/opencloud/pkg/ast" + "github.com/opencloud-eu/opencloud/pkg/kql" +) + +// MergeFilters puts each part in its own group so the AND binds across whole +// queries, not into their operator precedence. +func MergeFilters(b Builder, qs string, filters []string) (*ast.Ast, error) { + main, err := b.Build(qs) + if err != nil { + return nil, err + } + main = Normalize(main, ResolveField) + if len(filters) == 0 { + return main, nil + } + + nodes := make([]ast.Node, 0, 2*len(filters)+1) + if len(main.Nodes) > 0 { + nodes = append(nodes, &ast.GroupNode{Base: &ast.Base{}, Nodes: main.Nodes}) + } + for _, f := range filters { + fa, err := b.Build(f) + if err != nil { + return nil, err + } + fa = Normalize(fa, ResolveField) + ForceCaseSensitive(fa) + if len(fa.Nodes) == 0 { + continue + } + if len(nodes) > 0 { + nodes = append(nodes, &ast.OperatorNode{Value: kql.BoolAND}) + } + nodes = append(nodes, &ast.GroupNode{Base: &ast.Base{}, Nodes: fa.Nodes}) + } + return &ast.Ast{Nodes: nodes}, nil +} diff --git a/services/search/pkg/query/merge_test.go b/services/search/pkg/query/merge_test.go new file mode 100644 index 0000000000..2f68633ba1 --- /dev/null +++ b/services/search/pkg/query/merge_test.go @@ -0,0 +1,62 @@ +package query_test + +import ( + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + + "github.com/opencloud-eu/opencloud/pkg/ast" + "github.com/opencloud-eu/opencloud/pkg/kql" + "github.com/opencloud-eu/opencloud/services/search/pkg/query" +) + +func collectStringNodes(nodes []ast.Node) []*ast.StringNode { + var out []*ast.StringNode + for _, n := range nodes { + switch node := n.(type) { + case *ast.StringNode: + out = append(out, node) + case *ast.GroupNode: + out = append(out, collectStringNodes(node.Nodes)...) + } + } + return out +} + +var _ = Describe("MergeFilters", func() { + It("returns the normalized main query unchanged when there are no filters", func() { + a, err := query.MergeFilters(kql.Builder{}, `name:"hello"`, nil) + Expect(err).ToNot(HaveOccurred()) + Expect(a.Nodes).To(HaveLen(1)) + }) + + It("ANDs a filter in as an exact, case-sensitive match", func() { + a, err := query.MergeFilters(kql.Builder{}, `name:"hello"`, []string{`Tags:"Saxon"`}) + Expect(err).ToNot(HaveOccurred()) + + strs := collectStringNodes(a.Nodes) + var forced *ast.StringNode + for _, s := range strs { + if s.Value == "Saxon" { + forced = s + } + } + Expect(forced).ToNot(BeNil(), "the decoded filter node should be present") + Expect(forced.Exact).To(BeTrue()) + Expect(forced.CaseInsensitive).To(BeFalse()) + }) + + It("forces every node of an OR filter", func() { + a, err := query.MergeFilters(kql.Builder{}, `name:"hello"`, []string{`(Tags:"a" OR Tags:"b")`}) + Expect(err).ToNot(HaveOccurred()) + forced := 0 + for _, s := range collectStringNodes(a.Nodes) { + if s.Value != "a" && s.Value != "b" { + continue + } + forced++ + Expect(s.Exact).To(BeTrue()) + Expect(s.CaseInsensitive).To(BeFalse()) + } + Expect(forced).To(Equal(2)) + }) +}) diff --git a/services/search/pkg/query/query.go b/services/search/pkg/query/query.go index bd4091637e..1af2a28718 100644 --- a/services/search/pkg/query/query.go +++ b/services/search/pkg/query/query.go @@ -16,4 +16,7 @@ type Compiler[T any] interface { // Creator is the interface that wraps the basic Create method. type Creator[T any] interface { Create(qs string) (T, error) + // CreateWithFilters ANDs the decoded aggregation filters in as exact + // case-sensitive matches. + CreateWithFilters(qs string, filters []string) (T, error) } diff --git a/services/search/pkg/search/service.go b/services/search/pkg/search/service.go index b7f985b622..a290b0988d 100644 --- a/services/search/pkg/search/service.go +++ b/services/search/pkg/search/service.go @@ -34,6 +34,7 @@ import ( searchmsg "github.com/opencloud-eu/opencloud/protogen/gen/opencloud/messages/search/v0" searchsvc "github.com/opencloud-eu/opencloud/protogen/gen/opencloud/services/search/v0" "github.com/opencloud-eu/opencloud/services/graph/pkg/unifiedrole" + "github.com/opencloud-eu/opencloud/services/search/pkg/aggregation" "github.com/opencloud-eu/opencloud/services/search/pkg/config" "github.com/opencloud-eu/opencloud/services/search/pkg/content" "github.com/opencloud-eu/opencloud/services/search/pkg/metrics" @@ -139,6 +140,20 @@ func (s *Service) Search(ctx context.Context, req *searchsvc.SearchRequest) (*se } req.Query = query + // decode the aggregation filters once, up front, so a malformed token fails + // the whole request instead of silently dropping + if raw := req.GetAggregationFilters(); len(raw) > 0 { + decoded := make([]string, 0, len(raw)) + for _, f := range raw { + frag, err := aggregation.DecodeAggregationFilter(f) + if err != nil { + return nil, errtypes.BadRequest(err.Error()) + } + decoded = append(decoded, frag) + } + req.AggregationFilters = decoded + } + if len(scope) > 0 { scopedID, err := storagespace.ParseID(scope) if err != nil { @@ -602,8 +617,9 @@ func (s *Service) searchIndex(ctx context.Context, req *searchsvc.SearchRequest, } searchRequest := &searchsvc.SearchIndexRequest{ - Query: req.Query, - Aggregations: req.GetAggregations(), + Query: req.Query, + Aggregations: req.GetAggregations(), + AggregationFilters: req.GetAggregationFilters(), Ref: &searchmsg.Reference{ ResourceId: searchRootID, Path: searchPathPrefix, diff --git a/services/search/pkg/search/service_test.go b/services/search/pkg/search/service_test.go index 5db825bc9e..a1d8be8d7b 100644 --- a/services/search/pkg/search/service_test.go +++ b/services/search/pkg/search/service_test.go @@ -258,6 +258,159 @@ var _ = Describe("Searchprovider", func() { Expect(match.Entity.Ref.ResourceId.OpaqueId).To(Equal(personalSpace.Root.OpaqueId)) Expect(match.Entity.Ref.Path).To(Equal("./path/to/Foo.pdf")) }) + + It("forwards aggregations to the engine", func() { + _, err := s.Search(ctx, &searchsvc.SearchRequest{ + Query: "foo", + Aggregations: []*searchsvc.AggregationOption{ + {Field: "audio.artist", Size: 10}, + }, + }) + Expect(err).ToNot(HaveOccurred()) + indexClient.AssertCalled(GinkgoT(), "Search", mock.Anything, mock.MatchedBy(func(req *searchsvc.SearchIndexRequest) bool { + return len(req.Aggregations) == 1 && + req.Aggregations[0].Field == "audio.artist" && + req.Aggregations[0].Size == 10 + })) + }) + }) + + Context("with two personal spaces returning aggregations", func() { + var ( + spaceA = &sprovider.StorageSpace{ + Id: &sprovider.StorageSpaceId{OpaqueId: "storageid$a!a"}, + Root: &sprovider.ResourceId{StorageId: "storageid", SpaceId: "a", OpaqueId: "a"}, + Name: "space-a", + SpaceType: "personal", + } + spaceB = &sprovider.StorageSpace{ + Id: &sprovider.StorageSpaceId{OpaqueId: "storageid$b!b"}, + Root: &sprovider.ResourceId{StorageId: "storageid", SpaceId: "b", OpaqueId: "b"}, + Name: "space-b", + SpaceType: "personal", + } + ) + + BeforeEach(func() { + gatewayClient.On("ListStorageSpaces", mock.Anything, mock.Anything).Return(&sprovider.ListStorageSpacesResponse{ + Status: status.NewOK(ctx), + StorageSpaces: []*sprovider.StorageSpace{spaceA, spaceB}, + }, nil) + indexClient.On("Search", mock.Anything, mock.MatchedBy(func(req *searchsvc.SearchIndexRequest) bool { + return req.Ref != nil && req.Ref.ResourceId.SpaceId == "a" + })).Return(&searchsvc.SearchIndexResponse{ + TotalMatches: 2, + Aggregations: []*searchsvc.AggregationResult{{ + Field: "audio.artist", + Buckets: []*searchsvc.Bucket{ + {Key: "Saxon", Count: 2}, + {Key: "Motörhead", Count: 1}, + }, + }}, + }, nil) + indexClient.On("Search", mock.Anything, mock.MatchedBy(func(req *searchsvc.SearchIndexRequest) bool { + return req.Ref != nil && req.Ref.ResourceId.SpaceId == "b" + })).Return(&searchsvc.SearchIndexResponse{ + TotalMatches: 3, + Aggregations: []*searchsvc.AggregationResult{{ + Field: "audio.artist", + Buckets: []*searchsvc.Bucket{ + {Key: "Saxon", Count: 3}, + {Key: "Led Zeppelin", Count: 1}, + }, + }}, + }, nil) + }) + + It("merges bucket counts across spaces", func() { + res, err := s.Search(ctx, &searchsvc.SearchRequest{ + Query: "mediatype:audio", + Aggregations: []*searchsvc.AggregationOption{ + {Field: "audio.artist", Size: 10}, + }, + }) + Expect(err).ToNot(HaveOccurred()) + Expect(res.Aggregations).To(HaveLen(1)) + agg := res.Aggregations[0] + Expect(agg.Field).To(Equal("audio.artist")) + + counts := map[string]int64{} + for _, b := range agg.Buckets { + counts[b.Key] = b.Count + } + Expect(counts).To(HaveKeyWithValue("Saxon", int64(5))) + Expect(counts).To(HaveKeyWithValue("Motörhead", int64(1))) + Expect(counts).To(HaveKeyWithValue("Led Zeppelin", int64(1))) + }) + + It("sorts buckets by count descending by default", func() { + res, err := s.Search(ctx, &searchsvc.SearchRequest{ + Query: "mediatype:audio", + Aggregations: []*searchsvc.AggregationOption{ + {Field: "audio.artist"}, + }, + }) + Expect(err).ToNot(HaveOccurred()) + keys := []string{} + for _, b := range res.Aggregations[0].Buckets { + keys = append(keys, b.Key) + } + // Saxon:5, Motörhead:1, Led Zeppelin:1 (count desc) + Expect(keys[0]).To(Equal("Saxon")) + }) + + It("sorts buckets alphabetically ascending with sortBy keyAsString", func() { + res, err := s.Search(ctx, &searchsvc.SearchRequest{ + Query: "mediatype:audio", + Aggregations: []*searchsvc.AggregationOption{ + { + Field: "audio.artist", + BucketDefinition: &searchsvc.BucketDefinition{ + SortBy: "keyAsString", + }, + }, + }, + }) + Expect(err).ToNot(HaveOccurred()) + keys := []string{} + for _, b := range res.Aggregations[0].Buckets { + keys = append(keys, b.Key) + } + Expect(keys).To(Equal([]string{"Led Zeppelin", "Motörhead", "Saxon"})) + }) + + It("applies minimumCount filter and size cap", func() { + res, err := s.Search(ctx, &searchsvc.SearchRequest{ + Query: "mediatype:audio", + Aggregations: []*searchsvc.AggregationOption{ + { + Field: "audio.artist", + Size: 5, + BucketDefinition: &searchsvc.BucketDefinition{ + SortBy: "count", + IsDescending: true, + MinimumCount: 2, + }, + }, + }, + }) + Expect(err).ToNot(HaveOccurred()) + // only Saxon has count >= 2 + Expect(res.Aggregations[0].Buckets).To(HaveLen(1)) + Expect(res.Aggregations[0].Buckets[0].Key).To(Equal("Saxon")) + }) + + It("trims the bucket list to Size", func() { + res, err := s.Search(ctx, &searchsvc.SearchRequest{ + Query: "mediatype:audio", + Aggregations: []*searchsvc.AggregationOption{ + {Field: "audio.artist", Size: 1}, + }, + }) + Expect(err).ToNot(HaveOccurred()) + Expect(res.Aggregations[0].Buckets).To(HaveLen(1)) + Expect(res.Aggregations[0].Buckets[0].Key).To(Equal("Saxon")) + }) }) Context("with a personal space with a filter", func() {