diff --git a/Makefile b/Makefile index 1bae8847..fd847ba6 100644 --- a/Makefile +++ b/Makefile @@ -108,7 +108,7 @@ ci-tests-race: test-deps # run diff lint like in pipeline .lint: $(info Running lint...) - GOBIN=$(LOCAL_BIN) go run github.com/golangci/golangci-lint/v2/cmd/golangci-lint@v2.12.1 run \ + GOBIN=$(LOCAL_BIN) go run github.com/golangci/golangci-lint/v2/cmd/golangci-lint@v2.14.0 run \ --config=.golangci.yaml ./... .PHONY: lint diff --git a/pkg/seqproxyapi/v1/mappings.go b/pkg/seqproxyapi/v1/mappings.go index 1390fea9..c4f7e76e 100644 --- a/pkg/seqproxyapi/v1/mappings.go +++ b/pkg/seqproxyapi/v1/mappings.go @@ -4,6 +4,7 @@ import ( "fmt" "github.com/ozontech/seq-db/asyncsearcher" + "github.com/ozontech/seq-db/query" "github.com/ozontech/seq-db/seq" ) @@ -128,3 +129,55 @@ func AsyncSearchStatusFromString(s string) (AsyncSearchStatus, error) { return 0, fmt.Errorf("unknown status") } + +var typeMappings = []DataType{ + query.DataTypeBytes: DataType_BYTES, + query.DataTypeSeqID: DataType_SEQ_ID, + query.DataTypeDocument: DataType_RAW_DOCUMENT, + query.DataTypeString: DataType_STRING, + query.DataTypeUint32: DataType_UINT32, + query.DataTypeUint64: DataType_UINT64, + query.DataTypeInt32: DataType_INT32, + query.DataTypeInt64: DataType_INT64, + query.DataTypeFloat64: DataType_FLOAT64, + query.DataTypeFloat64Array: DataType_FLOAT64_ARRAY, + query.DataTypeStringArray: DataType_STRING_ARRAY, +} + +var typeMappingsPb = func() []query.DataType { + mappings := make([]query.DataType, len(typeMappings)) + for from, to := range typeMappings { + mappings[to] = query.DataType(from) + } + return mappings +}() + +func (t DataType) ToQueryDataType() (query.DataType, error) { + if int(t) >= len(typeMappingsPb) || t < 0 { + return 0, fmt.Errorf("unknown data type: %d", t) + } + return typeMappingsPb[t], nil +} + +func (t DataType) MustQueryDataType() query.DataType { + v, err := t.ToQueryDataType() + if err != nil { + panic(err) + } + return v +} + +func ToProtoDataType(t query.DataType) (DataType, error) { + if int(t) >= len(typeMappings) { + return 0, fmt.Errorf("unknown data type: %d", t) + } + return typeMappings[t], nil +} + +func MustProtoDataType(t query.DataType) DataType { + v, err := ToProtoDataType(t) + if err != nil { + panic(err) + } + return v +} diff --git a/pkg/storeapi/mappings.go b/pkg/storeapi/mappings.go index 1114875d..25d4c46c 100644 --- a/pkg/storeapi/mappings.go +++ b/pkg/storeapi/mappings.go @@ -4,6 +4,7 @@ import ( "fmt" "github.com/ozontech/seq-db/asyncsearcher" + "github.com/ozontech/seq-db/query" "github.com/ozontech/seq-db/seq" ) @@ -142,3 +143,55 @@ func MustProtoAsyncSearchStatus(s asyncsearcher.AsyncSearchStatus) AsyncSearchSt } return v } + +var typeMappings = []DataType{ + query.DataTypeBytes: DataType_BYTES, + query.DataTypeSeqID: DataType_SEQ_ID, + query.DataTypeDocument: DataType_RAW_DOCUMENT, + query.DataTypeString: DataType_STRING, + query.DataTypeUint32: DataType_UINT32, + query.DataTypeUint64: DataType_UINT64, + query.DataTypeInt32: DataType_INT32, + query.DataTypeInt64: DataType_INT64, + query.DataTypeFloat64: DataType_FLOAT64, + query.DataTypeFloat64Array: DataType_FLOAT64_ARRAY, + query.DataTypeStringArray: DataType_STRING_ARRAY, +} + +var typeMappingsPb = func() []query.DataType { + mappings := make([]query.DataType, len(typeMappings)) + for from, to := range typeMappings { + mappings[to] = query.DataType(from) + } + return mappings +}() + +func (t DataType) ToQueryDataType() (query.DataType, error) { + if int(t) >= len(typeMappingsPb) || t < 0 { + return 0, fmt.Errorf("unknown data type: %d", t) + } + return typeMappingsPb[t], nil +} + +func (t DataType) MustQueryDataType() query.DataType { + v, err := t.ToQueryDataType() + if err != nil { + panic(err) + } + return v +} + +func ToProtoDataType(t query.DataType) (DataType, error) { + if int(t) >= len(typeMappings) { + return 0, fmt.Errorf("unknown data type: %d", t) + } + return typeMappings[t], nil +} + +func MustProtoDataType(t query.DataType) DataType { + v, err := ToProtoDataType(t) + if err != nil { + panic(err) + } + return v +} diff --git a/proxy/search/stream_search.go b/proxy/search/stream_search.go index fc21486f..2f84a3f2 100644 --- a/proxy/search/stream_search.go +++ b/proxy/search/stream_search.go @@ -75,6 +75,11 @@ func (si *Ingestor) StreamSearch( } } + if len(streams) == 0 { + // nothing to read from, return an empty stream + return exec.NewNMergedProducers(nil, query.SeqIDColumn(0), sr.Order), newControlBroadcaster(streams), partialRespErr + } + broadcaster := newControlBroadcaster(streams) producers := make([]query.RecordProducer, 0, len(streams)) for _, s := range streams { @@ -88,12 +93,28 @@ func (si *Ingestor) StreamSearch( offset = 0 } + // TODO: store can change docs' schema in case of filter pipe, need to handle it on proxy. + var expectedSchema *query.Schema + if sr.Agg != nil { + expectedSchema = query.AggsSchema + } + // Shard schemas come off the wire, so validate them before any consumer reads records. + if err := validateShardSchemas(streams, expectedSchema); err != nil { + closeStreams(streams) + return nil, nil, err + } + var mergedStream query.RecordProducer if sr.Agg != nil { mergedStream = exec.NewDistributedAggregator(producers, sr.Agg.Func, sr.Agg.Quantiles) } else { - const seqIdColIdx = 0 - mergedDocsStream := exec.NewNMergedProducers(producers, query.SeqIDColumn(seqIdColIdx), sr.Order) + // streams[0] is safe - streams len is already checked + idCol, err := streams[0].OutSchema().Column[seq.ID](query.DocsIDCol) + if err != nil { + closeStreams(streams) + return nil, nil, err + } + mergedDocsStream := exec.NewNMergedProducers(producers, idCol, sr.Order) mergedStream = exec.NewLimiter(mergedDocsStream, uint32(sr.Size), uint32(offset)) } @@ -267,6 +288,40 @@ func closeStreams(streams []*StreamSearchIterator) { } } +func schemaFromTyping(typing []*storeapi.Typing) (*query.Schema, []query.DataType, error) { + cols := make([]query.ColumnDesc, 0, len(typing)) + types := make([]query.DataType, 0, len(typing)) + for _, t := range typing { + dt, err := t.GetType().ToQueryDataType() + if err != nil { + return nil, nil, err + } + cols = append(cols, query.ColumnDesc{Name: t.GetTitle(), Type: dt}) + types = append(types, dt) + } + schema, err := query.NewSchema(cols...) + if err != nil { + return nil, nil, err + } + return schema, types, nil +} + +func validateShardSchemas(streams []*StreamSearchIterator, expected *query.Schema) error { + if len(streams) == 0 { + return nil + } + base := streams[0].OutSchema() + for _, s := range streams[1:] { + if !s.OutSchema().Equal(base) { + return fmt.Errorf("shard schemas mismatch: %v vs %v", base.Cols(), s.OutSchema().Cols()) + } + } + if expected != nil && !base.Equal(expected) { + return fmt.Errorf("shard schema mismatch: got %v, want %v", base.Cols(), expected.Cols()) + } + return nil +} + // NewStreamSearchIterator reads one message ahead after the header so that a // summary-with-error sent immediately after the header (before any data) is // detected on the open-stream phase and can trigger fail-fast in the @@ -276,7 +331,11 @@ func NewStreamSearchIterator( header *storeapi.ResponseHeader, stream storeapi.StoreApi_StreamSearchClient, ) (*StreamSearchIterator, error) { - it := &StreamSearchIterator{tr: tr, typing: header.Typing, stream: stream} + schema, types, err := schemaFromTyping(header.Typing) + if err != nil { + return nil, fmt.Errorf("bad response header: %w", err) + } + it := &StreamSearchIterator{tr: tr, schema: schema, types: types, stream: stream} msg, err := stream.Recv() if errors.Is(err, io.EOF) { @@ -295,7 +354,8 @@ func NewStreamSearchIterator( type StreamSearchIterator struct { tr *querytracer.Tracer - typing []*storeapi.Typing + schema *query.Schema + types []query.DataType stream storeapi.StoreApi_StreamSearchClient curBatch []*storeapi.Record @@ -336,11 +396,15 @@ func (it *StreamSearchIterator) Next() *query.Record { recordVals := make([]*query.RecordVals, 0, len(record.RawData)) for i, rawData := range record.RawData { - recordVals = append(recordVals, query.NewRecordVals(query.DataType(it.typing[i].Type), rawData)) + recordVals = append(recordVals, query.NewRecordVals(it.types[i], rawData)) } return query.NewRecord(recordVals) } +func (it *StreamSearchIterator) OutSchema() *query.Schema { + return it.schema +} + // push handles a single message received from the store stream. func (it *StreamSearchIterator) push(msg *storeapi.StreamSearchResponse) error { switch v := msg.ResponseType.(type) { diff --git a/proxy/search/stream_search_test.go b/proxy/search/stream_search_test.go new file mode 100644 index 00000000..318a0c27 --- /dev/null +++ b/proxy/search/stream_search_test.go @@ -0,0 +1,87 @@ +package search + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + insaneJSON "github.com/ozontech/insane-json" + + "github.com/ozontech/seq-db/pkg/storeapi" + "github.com/ozontech/seq-db/query" + "github.com/ozontech/seq-db/seq" +) + +func TestSchemaFromTyping(t *testing.T) { + typing := []*storeapi.Typing{ + {Title: "id", Type: storeapi.DataType_SEQ_ID}, + {Title: "data", Type: storeapi.DataType_RAW_DOCUMENT}, + } + + schema, types, err := schemaFromTyping(typing) + require.NoError(t, err) + assert.Equal(t, 2, schema.Len()) + assert.Equal(t, []query.DataType{query.DataTypeSeqID, query.DataTypeDocument}, types) + assert.Equal(t, 0, schema.MustColumn[seq.ID]("id").Idx()) + assert.Equal(t, query.DataTypeDocument, schema.MustColumn[*insaneJSON.Root]("data").DataType()) +} + +func TestSchemaFromTypingDuplicateName(t *testing.T) { + typing := []*storeapi.Typing{ + {Title: "id", Type: storeapi.DataType_SEQ_ID}, + {Title: "id", Type: storeapi.DataType_SEQ_ID}, + } + _, _, err := schemaFromTyping(typing) + assert.Error(t, err) +} + +func TestValidateShardSchemas(t *testing.T) { + newIterator := func(typing []*storeapi.Typing) *StreamSearchIterator { + schema, types, err := schemaFromTyping(typing) + require.NoError(t, err) + return &StreamSearchIterator{schema: schema, types: types} + } + + docsTyping := []*storeapi.Typing{ + {Title: "id", Type: storeapi.DataType_SEQ_ID}, + {Title: "data", Type: storeapi.DataType_RAW_DOCUMENT}, + } + otherTyping := []*storeapi.Typing{ + {Title: "id", Type: storeapi.DataType_SEQ_ID}, + {Title: "payload", Type: storeapi.DataType_RAW_DOCUMENT}, + } + aggsTyping := []*storeapi.Typing{ + {Title: "token", Type: storeapi.DataType_STRING}, + } + + t.Run("no shards", func(t *testing.T) { + assert.NoError(t, validateShardSchemas(nil, nil)) + }) + + t.Run("matching", func(t *testing.T) { + streams := []*StreamSearchIterator{newIterator(docsTyping), newIterator(docsTyping)} + assert.NoError(t, validateShardSchemas(streams, nil)) + }) + + t.Run("mismatch", func(t *testing.T) { + streams := []*StreamSearchIterator{newIterator(docsTyping), newIterator(otherTyping)} + err := validateShardSchemas(streams, nil) + require.Error(t, err) + assert.Contains(t, err.Error(), "shard schemas mismatch") + }) + + t.Run("expected schema mismatch", func(t *testing.T) { + streams := []*StreamSearchIterator{newIterator(docsTyping), newIterator(docsTyping)} + err := validateShardSchemas(streams, query.AggsSchema) + require.Error(t, err) + assert.Contains(t, err.Error(), "shard schema mismatch") + }) + + t.Run("expected schema match", func(t *testing.T) { + streams := []*StreamSearchIterator{newIterator(aggsTyping), newIterator(aggsTyping)} + expected, _, err := schemaFromTyping(aggsTyping) + require.NoError(t, err) + assert.NoError(t, validateShardSchemas(streams, expected)) + }) +} diff --git a/proxyapi/grpc_complex_search.go b/proxyapi/grpc_complex_search.go index de6e21c3..9a9399f4 100644 --- a/proxyapi/grpc_complex_search.go +++ b/proxyapi/grpc_complex_search.go @@ -171,11 +171,11 @@ func (g *grpcV1) useStreamSearch( func readDocuments(storesStream query.RecordProducer) []*seqproxyapi.Document { var docs []*seqproxyapi.Document for r := storesStream.Next(); r != nil; r = storesStream.Next() { - id := r.Vals[0].AsSeqID() + id := docIDCol.Val(r) docs = append(docs, &seqproxyapi.Document{ Id: id.String(), Time: timestamppb.New(id.MID.Time()), - Data: r.Vals[1].RawData(), + Data: docDataCol.RawData(r), }) } return docs @@ -185,13 +185,13 @@ func readAggregations(storesStream query.RecordProducer) []*seqproxyapi.Aggregat buckets := make([]*seqproxyapi.Aggregation_Bucket, 0) for r := storesStream.Next(); r != nil; r = storesStream.Next() { bucket := &seqproxyapi.Aggregation_Bucket{ - Key: r.Vals[0].AsString(), - Value: r.Vals[1].AsFloat64(), + Key: aggKeyCol.Val(r), + Value: aggValueCol.Val(r), } - if ts := r.Vals[2].AsUint64(); ts != consts.DummyMID { + if ts := aggTsCol.Val(r); ts != consts.DummyMID { bucket.Ts = timestamppb.New(seq.MID(ts).Time()) } - if quantiles := r.Vals[3].AsFloat64Array(); len(quantiles) > 0 { + if quantiles := aggQuantilesCol.Val(r); len(quantiles) > 0 { bucket.Quantiles = quantiles } buckets = append(buckets, bucket) diff --git a/proxyapi/grpc_stream_search.go b/proxyapi/grpc_stream_search.go index b6b0ea7d..e42eb39a 100644 --- a/proxyapi/grpc_stream_search.go +++ b/proxyapi/grpc_stream_search.go @@ -6,6 +6,7 @@ import ( "fmt" "io" + insaneJSON "github.com/ozontech/insane-json" "google.golang.org/grpc/codes" "google.golang.org/grpc/status" @@ -286,15 +287,25 @@ func aggsTyping() []*seqproxyapi.Typing { } } +var ( + docIDCol = query.DocsSchema.MustColumn[seq.ID](query.DocsIDCol) + docDataCol = query.DocsSchema.MustColumn[*insaneJSON.Root](query.DocsDataCol) + + aggKeyCol = query.AggResultSchema.MustColumn[string]("token") + aggValueCol = query.AggResultSchema.MustColumn[float64]("value") + aggTsCol = query.AggResultSchema.MustColumn[uint64]("ts") + aggQuantilesCol = query.AggResultSchema.MustColumn[[]float64]("quantiles") +) + // converts *query.Record to *seqproxyapi.Record according to hardcoded schemas from both store and proxy func docToRecord(r *query.Record) *seqproxyapi.Record { - id := r.Vals[0].AsSeqID() + id := docIDCol.Val(r) return &seqproxyapi.Record{ RawData: [][]byte{ []byte(id.String()), // id encoding.Uint64ToBytes(uint64(id.MID)), // time - r.Vals[1].RawData(), // data + docDataCol.RawData(r), // data }, } } @@ -303,10 +314,10 @@ func docToRecord(r *query.Record) *seqproxyapi.Record { func aggToRecord(r *query.Record) *seqproxyapi.Record { return &seqproxyapi.Record{ RawData: [][]byte{ - r.Vals[0].RawData(), // key - r.Vals[1].RawData(), // value - r.Vals[2].RawData(), // ts - r.Vals[3].RawData(), // quantiles + aggKeyCol.RawData(r), // key + aggValueCol.RawData(r), // value + aggTsCol.RawData(r), // ts + aggQuantilesCol.RawData(r), // quantiles }, } } diff --git a/query/column.go b/query/column.go index 7443d541..b49be7d2 100644 --- a/query/column.go +++ b/query/column.go @@ -1,6 +1,8 @@ package query import ( + "cmp" + insaneJSON "github.com/ozontech/insane-json" "github.com/ozontech/seq-db/seq" @@ -10,6 +12,7 @@ type Column[T any] struct { idx int dataType DataType get func(*RecordVals) T + cmp func(T, T) int } func (c Column[T]) Idx() int { @@ -20,6 +23,10 @@ func (c Column[T]) DataType() DataType { return c.dataType } +func (c Column[T]) Cmp() func(T, T) int { + return c.cmp +} + // Val returns value of r's column at c.idx index func (c Column[T]) Val(r *Record) T { return c.get(r.Vals[c.idx]) @@ -35,6 +42,18 @@ func SeqIDColumn(idx int) Column[seq.ID] { idx: idx, dataType: DataTypeSeqID, get: (*RecordVals).AsSeqID, + cmp: cmpSeqID, + } +} + +func cmpSeqID(a, b seq.ID) int { + switch { + case seq.Less(a, b): + return -1 + case seq.Less(b, a): + return 1 + default: + return 0 } } @@ -51,6 +70,7 @@ func StringColumn(idx int) Column[string] { idx: idx, dataType: DataTypeString, get: (*RecordVals).AsString, + cmp: cmp.Compare[string], } } @@ -67,6 +87,7 @@ func Float64Column(idx int) Column[float64] { idx: idx, dataType: DataTypeFloat64, get: (*RecordVals).AsFloat64, + cmp: cmp.Compare[float64], } } @@ -75,6 +96,7 @@ func Uint64Column(idx int) Column[uint64] { idx: idx, dataType: DataTypeUint64, get: (*RecordVals).AsUint64, + cmp: cmp.Compare[uint64], } } @@ -83,6 +105,7 @@ func Int64Column(idx int) Column[int64] { idx: idx, dataType: DataTypeInt64, get: (*RecordVals).AsInt64, + cmp: cmp.Compare[int64], } } @@ -91,6 +114,7 @@ func Uint32Column(idx int) Column[uint32] { idx: idx, dataType: DataTypeUint32, get: (*RecordVals).AsUint32, + cmp: cmp.Compare[uint32], } } @@ -99,6 +123,7 @@ func Int32Column(idx int) Column[int32] { idx: idx, dataType: DataTypeInt32, get: (*RecordVals).AsInt32, + cmp: cmp.Compare[int32], } } diff --git a/query/exec/aggregator.go b/query/exec/aggregator.go index 238b6e7a..3c3250ea 100644 --- a/query/exec/aggregator.go +++ b/query/exec/aggregator.go @@ -20,21 +20,22 @@ const ( ExecutorStateDone ) +// Input and output columns are resolved from existing hardcoded schemas. var ( - aggColTokenIn = query.StringColumn(0) - aggColMinIn = query.Float64Column(1) - aggColMaxIn = query.Float64Column(2) - aggColSumIn = query.Float64Column(3) - aggColTotalIn = query.Uint64Column(4) - aggColTsIn = query.Uint64Column(6) - aggColSamplesIn = query.Float64ArrayColumn(7) - aggColValuesIn = query.StringArrayColumn(8) + aggColTokenIn = query.AggsSchema.MustColumn[string]("token") + aggColMinIn = query.AggsSchema.MustColumn[float64]("min") + aggColMaxIn = query.AggsSchema.MustColumn[float64]("max") + aggColSumIn = query.AggsSchema.MustColumn[float64]("sum") + aggColTotalIn = query.AggsSchema.MustColumn[uint64]("total") + aggColTsIn = query.AggsSchema.MustColumn[uint64]("ts") + aggColSamplesIn = query.AggsSchema.MustColumn[[]float64]("samples") + aggColValuesIn = query.AggsSchema.MustColumn[[]string]("values") ) var ( - aggColTokenOut = query.StringColumn(0) - aggColValueOut = query.Float64Column(1) - aggColTsOut = query.Uint64Column(2) + aggColTokenOut = query.AggResultSchema.MustColumn[string]("token") + aggColValueOut = query.AggResultSchema.MustColumn[float64]("value") + aggColTsOut = query.AggResultSchema.MustColumn[uint64]("ts") ) // aggKey identifies a single timeseries bin: the grouping token plus the @@ -123,6 +124,7 @@ func (a *DistributedAggregator) Next() *query.Record { panic(fmt.Errorf("unimplemented aggregation func")) } + // val order must match query.AggResultSchema a.sortingBuf = append(a.sortingBuf, query.NewRecord([]*query.RecordVals{ query.NewRecordVals(query.DataTypeString, []byte(key.token)), query.NewRecordVals(query.DataTypeFloat64, encoding.Float64ToBytes(value)), diff --git a/query/exec/datasource.go b/query/exec/datasource.go index 9d2e86e6..cd627bbf 100644 --- a/query/exec/datasource.go +++ b/query/exec/datasource.go @@ -288,7 +288,7 @@ func makeAggRecord(bin *storeapi.SearchResponse_Bin, valuesPool []string) *query } return &query.Record{ Vals: []*query.RecordVals{ - query.NewRecordVals(query.DataTypeBytes, []byte(bin.Label)), + query.NewRecordVals(query.DataTypeString, []byte(bin.Label)), query.NewRecordVals(query.DataTypeFloat64, encoding.Float64ToBytes(bin.Hist.Min)), query.NewRecordVals(query.DataTypeFloat64, encoding.Float64ToBytes(bin.Hist.Max)), query.NewRecordVals(query.DataTypeFloat64, encoding.Float64ToBytes(bin.Hist.Sum)), diff --git a/query/exec/extractor.go b/query/exec/extractor.go new file mode 100644 index 00000000..23f0c3bc --- /dev/null +++ b/query/exec/extractor.go @@ -0,0 +1,56 @@ +package exec + +import ( + insaneJSON "github.com/ozontech/insane-json" + + "github.com/ozontech/seq-db/query" +) + +type DocFieldsExtractor struct { + input query.RecordProducer + docCol query.Column[*insaneJSON.Root] + + // extractFields lists scalar JSON fields to extract out of the document column as string values. + extractFields []string + + // roots holds every record whose val has been decoded (spawned an insaneJSON root). + // They are released back to the library pool in Finalize. + roots []*query.Record +} + +func NewDocFieldsExtractor( + input query.RecordProducer, + docCol query.Column[*insaneJSON.Root], + extractFields ...string, +) *DocFieldsExtractor { + return &DocFieldsExtractor{ + input: input, + docCol: docCol, + extractFields: extractFields, + } +} + +func (e *DocFieldsExtractor) Next() *query.Record { + r := e.input.Next() + if r == nil { + return nil + } + + root := e.docCol.Val(r) + for _, field := range e.extractFields { + // now we treat everything as strings, but later will support other types as well + fieldVal := root.Dig(field).AsBytes() + r.Vals = append(r.Vals, query.NewRecordVals(query.DataTypeString, fieldVal)) + } + e.roots = append(e.roots, r) + + return r +} + +func (e *DocFieldsExtractor) Finalize() *query.Summary { + for _, r := range e.roots { + r.Release() + } + e.roots = nil + return e.input.Finalize() +} diff --git a/query/exec/extractor_test.go b/query/exec/extractor_test.go new file mode 100644 index 00000000..501d73b6 --- /dev/null +++ b/query/exec/extractor_test.go @@ -0,0 +1,130 @@ +package exec + +import ( + "fmt" + "testing" + + "github.com/stretchr/testify/assert" + + "github.com/ozontech/seq-db/query" + "github.com/ozontech/seq-db/query/encoding" +) + +func TestDocFieldsExtractorExtractsScalars(t *testing.T) { + input := &testProducer{data: makeExtractorInputRecords([]string{ + `{"service":"svc-1","level":3,"active":true,"absent":null}`, + `{"service":"svc-2","level":7,"active":false}`, + })} + + extractor := NewDocFieldsExtractor( + input, + query.DocColumn(1), + "service", "level", "active", "absent", + ) + + r1 := extractor.Next() + assert.NotNil(t, r1) + // original vals are untouched, extracted fields are appended in order + assert.Equal(t, 6, len(r1.Vals)) + assert.Equal(t, "svc-1", r1.Vals[2].AsString()) + assert.Equal(t, "3", r1.Vals[3].AsString()) + assert.Equal(t, "true", r1.Vals[4].AsString()) + assert.Equal(t, "null", r1.Vals[5].AsString()) + + r2 := extractor.Next() + assert.NotNil(t, r2) + assert.Equal(t, 6, len(r2.Vals)) + assert.Equal(t, "svc-2", r2.Vals[2].AsString()) + assert.Equal(t, "7", r2.Vals[3].AsString()) + assert.Equal(t, "false", r2.Vals[4].AsString()) + assert.Equal(t, "", r2.Vals[5].AsString()) // absent field is missing + + assert.Nil(t, extractor.Next()) +} + +func TestDocFieldsExtractorMissingFieldYieldsEmpty(t *testing.T) { + input := &testProducer{data: makeExtractorInputRecords([]string{ + `{"service":"svc-1"}`, + })} + + extractor := NewDocFieldsExtractor(input, query.DocColumn(1), "nope") + + r := extractor.Next() + assert.NotNil(t, r) + assert.Equal(t, 3, len(r.Vals)) + assert.Equal(t, "", r.Vals[2].AsString()) +} + +func TestDocFieldsExtractorObjectAndArrayFields(t *testing.T) { + // This case documents the scalar-only contract: object and array fields yield an + // empty string, not their JSON representation. + input := &testProducer{data: makeExtractorInputRecords([]string{ + `{"service":"svc-1","obj":{"a":1},"arr":[1,2]}`, + })} + + extractor := NewDocFieldsExtractor(input, query.DocColumn(1), "obj", "arr") + + r := extractor.Next() + assert.NotNil(t, r) + assert.Equal(t, "", r.Vals[2].AsString()) + assert.Equal(t, "", r.Vals[3].AsString()) +} + +func TestDocFieldsExtractorNoFields(t *testing.T) { + docs := []string{`{"service":"svc-1"}`, `{"service":"svc-2"}`} + input := &testProducer{data: makeExtractorInputRecords(docs)} + + extractor := NewDocFieldsExtractor(input, query.DocColumn(1)) + + count := 0 + for r := extractor.Next(); r != nil; r = extractor.Next() { + // only the original two vals, nothing appended + assert.Equal(t, 2, len(r.Vals)) + count++ + } + assert.Equal(t, 2, count) +} + +func TestDocFieldsExtractorFinalize(t *testing.T) { + input := &testProducer{data: makeExtractorInputRecords([]string{ + `{"service":"svc-1"}`, + `{"service":"svc-2"}`, + }), total: 2} + + extractor := NewDocFieldsExtractor(input, query.DocColumn(1), "service") + for r := extractor.Next(); r != nil; r = extractor.Next() { + } + + summary := extractor.Finalize() + assert.Equal(t, uint64(2), summary.Total) + assert.Nil(t, extractor.roots) +} + +func TestDocFieldsExtractorPipeline(t *testing.T) { + // E2E check: a downstream reader uses the extractor's field + // values via a string column appended right after the original vals. + const field = "k8s_pod" + inputData := makeTestInputRecords(10) + input := &testProducer{data: inputData} + + extractor := NewDocFieldsExtractor(input, query.DocColumn(1), field) + podCol := query.StringColumn(2) + + for r := extractor.Next(); r != nil; r = extractor.Next() { + want := fmt.Sprintf("pod-%d", r.Vals[0].AsUint32()) + assert.Equal(t, want, podCol.Val(r)) + } +} + +func makeExtractorInputRecords(docs []string) []*query.Record { + out := make([]*query.Record, 0, len(docs)) + for _, doc := range docs { + out = append(out, &query.Record{ + Vals: []*query.RecordVals{ + query.NewRecordVals(query.DataTypeUint64, encoding.Uint64ToBytes(0)), + query.NewRecordVals(query.DataTypeDocument, []byte(doc)), + }, + }) + } + return out +} diff --git a/query/exec/merger.go b/query/exec/merger.go index 7700dcfc..c2d2ea4d 100644 --- a/query/exec/merger.go +++ b/query/exec/merger.go @@ -1,7 +1,7 @@ package exec import ( - "cmp" + "fmt" "github.com/ozontech/seq-db/query" "github.com/ozontech/seq-db/seq" @@ -14,13 +14,14 @@ type Merger[T any] struct { col query.Column[T] order seq.DocsOrder - cmp func(any, any) int + cmp func(T, T) int // dedup drops records whose sort key repeats the previously emitted one. // It is enabled only for the seq.ID merge: shards may match the same // document, and the merged document stream must contain each seq.ID once. dedup bool - lastVal any + hasLast bool + lastVal T // dups counts records dropped by dedup so Finalize can subtract it from the merged total. dups uint64 @@ -33,12 +34,16 @@ func NewMerger[T any]( col query.Column[T], order seq.DocsOrder, ) *Merger[T] { + if col.Cmp() == nil { + panic(fmt.Sprintf("BUG: %s column cannot be merged", col.DataType())) + } + return &Merger[T]{ left: left, right: right, col: col, order: order, - cmp: createCmpFunc(col.DataType()), + cmp: col.Cmp(), dedup: col.DataType() == query.DataTypeSeqID, curLeft: nil, curRight: nil, @@ -60,12 +65,13 @@ func (m *Merger[T]) Next() *query.Record { return r } val := m.col.Val(r) - if m.lastVal != nil && m.cmp(val, m.lastVal) == 0 { + if m.hasLast && m.cmp(val, m.lastVal) == 0 { // Skip duplicate. m.dups++ continue } m.lastVal = val + m.hasLast = true return r } } @@ -144,37 +150,6 @@ func combineSummaries(left, right *query.Summary) *query.Summary { return summary } -func createCmpFunc(dataType query.DataType) func(any, any) int { - switch dataType { - case query.DataTypeSeqID: - return func(a, b any) int { - v, w := a.(seq.ID), b.(seq.ID) - switch { - case seq.Less(v, w): - return -1 - case seq.Less(w, v): - return 1 - default: - return 0 - } - } - case query.DataTypeUint32: - return func(a, b any) int { return cmp.Compare(a.(uint32), b.(uint32)) } - case query.DataTypeUint64: - return func(a, b any) int { return cmp.Compare(a.(uint64), b.(uint64)) } - case query.DataTypeInt32: - return func(a, b any) int { return cmp.Compare(a.(int32), b.(int32)) } - case query.DataTypeInt64: - return func(a, b any) int { return cmp.Compare(a.(int64), b.(int64)) } - case query.DataTypeFloat64: - return func(a, b any) int { return cmp.Compare(a.(float64), b.(float64)) } - case query.DataTypeString: - return func(a, b any) int { return cmp.Compare(a.(string), b.(string)) } - default: - return func(a, b any) int { return 0 } - } -} - func NewNMergedProducers[T any]( producers []query.RecordProducer, col query.Column[T], diff --git a/query/plan/ops.go b/query/plan/ops.go new file mode 100644 index 00000000..95883150 --- /dev/null +++ b/query/plan/ops.go @@ -0,0 +1,70 @@ +package plan + +import ( + insaneJSON "github.com/ozontech/insane-json" + + "github.com/ozontech/seq-db/parser" + "github.com/ozontech/seq-db/query" + "github.com/ozontech/seq-db/query/exec" +) + +// Op is one operation wrapping the record stream. +// Each op has its own parameters, resolves its input columns from the +// schema of the pipeline built so far, and knows its output schema. +type Op interface { + // Apply wraps the input producer and returns it with the op's output schema. + // Ops that do not change the record shape return the input schema unchanged. + Apply(input query.RecordProducer, in *query.Schema) (query.RecordProducer, *query.Schema) +} + +type ExtractOp struct { + Fields []string + // OpField is the name of the column to perform operation on. + OpField string +} + +func (op *ExtractOp) Apply(input query.RecordProducer, in *query.Schema) (query.RecordProducer, *query.Schema) { + descs := make([]query.ColumnDesc, len(op.Fields)) + for i, f := range op.Fields { + descs[i] = query.ColumnDesc{Name: f, Type: query.DataTypeString} + } + + out := in.Extend(descs...) + docCol := in.MustColumn[*insaneJSON.Root](op.OpField) + return exec.NewDocFieldsExtractor(input, docCol, op.Fields...), out +} + +type FilterOp struct { + Cond parser.FilterCondition + WithTotal bool +} + +func (op *FilterOp) Apply(input query.RecordProducer, in *query.Schema) (query.RecordProducer, *query.Schema) { + col := in.MustColumn[string](op.Cond.Field) + return exec.NewFilter(input, col, exec.NewEq(op.Cond.Value), op.WithTotal), in +} + +type ProjectOp struct { + Fields []string + AllowList bool + // OpField is the name of the column to perform operation on. + OpField string +} + +func (op *ProjectOp) Apply(input query.RecordProducer, in *query.Schema) (query.RecordProducer, *query.Schema) { + docCol := in.MustColumn[*insaneJSON.Root](op.OpField) + return exec.NewDocProjector(input, docCol, &exec.FieldsFilter{ + Fields: op.Fields, + AllowList: op.AllowList, + }), in +} + +type LimitOp struct { + Limit int + Offset int +} + +func (op *LimitOp) Apply(input query.RecordProducer, in *query.Schema) (query.RecordProducer, *query.Schema) { + // set limit=limit+offset and offset=0 to merge stores' results correctly on proxy + return exec.NewLimiter(input, uint32(op.Limit+op.Offset), 0), in +} diff --git a/query/plan/plan.go b/query/plan/plan.go new file mode 100644 index 00000000..bc978548 --- /dev/null +++ b/query/plan/plan.go @@ -0,0 +1,220 @@ +package plan + +import ( + "fmt" + + "github.com/ozontech/seq-db/consts" + "github.com/ozontech/seq-db/frac/processor" + "github.com/ozontech/seq-db/parser" + "github.com/ozontech/seq-db/query" + "github.com/ozontech/seq-db/seq" + "github.com/ozontech/seq-db/util" +) + +type Plan struct { + Scan Scan + Ops []Op + Schema *query.Schema +} + +func (p *Plan) IsAgg() bool { + return len(p.Scan.AggQ) > 0 +} + +// Scan describes the datasource. +type Scan struct { + AST *parser.ASTNode + + From, To seq.MID + Order seq.DocsOrder + WithTotal bool + + AggQ []processor.AggQuery + + OffsetID seq.ID +} + +type BuildParams struct { + SeqQL *parser.SeqQLQuery + + // Input is the record schema the pipeline starts from + Input *query.Schema + // DocField is the title of document field in Input. + // We still keep some hardcoded params. Will ger rid of it later. + DocField string + + From, To seq.MID + OffsetID string + WithTotal bool +} + +// Build converts a parsed SeqQL query into a validated logical plan. +func Build(params BuildParams) (*Plan, error) { + if params.Input == nil { + return nil, fmt.Errorf("input schema is not set") + } + if params.DocField == "" { + return nil, fmt.Errorf("doc column is not set") + } + if _, err := params.Input.Index(params.DocField); err != nil { + return nil, fmt.Errorf("doc column: %w", err) + } + + p := &Plan{ + Scan: Scan{ + AST: params.SeqQL.Root, + From: params.From, + To: params.To, + WithTotal: params.WithTotal, + }, + Schema: params.Input, + } + + var limit, offset int + for _, pipe := range params.SeqQL.Pipes { + switch pipe := pipe.(type) { + case *parser.PipeStats: + aggQ, err := convertStatsAggToAggQuery(pipe.Agg) + if err != nil { + return nil, fmt.Errorf("failed to convert stats aggs: %w", err) + } + p.Scan.AggQ = []processor.AggQuery{aggQ} + p.Schema = query.AggsSchema + case *parser.PipeFilter: + // we need to extract the field so filter can compare based on it. + if _, err := params.Input.Index(pipe.Condition.Field); err == nil { + return nil, fmt.Errorf("cannot filter on reserved column %q", pipe.Condition.Field) + } + p.Ops = append(p.Ops, + &ExtractOp{Fields: []string{pipe.Condition.Field}, OpField: params.DocField}, + &FilterOp{Cond: pipe.Condition, WithTotal: params.WithTotal}, + ) + case *parser.PipeFields: + p.Ops = append(p.Ops, &ProjectOp{Fields: pipe.Fields, AllowList: !pipe.Except, OpField: params.DocField}) + case *parser.PipeSort: + if pipe.Order == "desc" { + p.Scan.Order = seq.DocsOrderDesc + } else { + p.Scan.Order = seq.DocsOrderAsc + } + case *parser.PipeLimit: + limit = pipe.Limit + case *parser.PipeOffset: + offset = pipe.Offset + } + } + + if limit > 0 { + p.Ops = append(p.Ops, &LimitOp{Limit: limit, Offset: offset}) + } + + if err := p.applyOffsetID(params.OffsetID, offset); err != nil { + return nil, err + } + return p, nil +} + +var searchAllTerms = []parser.Term{{ + Kind: parser.TermSymbol, Data: "*", +}} + +func convertStatsAggToAggQuery(statsAgg parser.StatsAgg) (processor.AggQuery, error) { + aggFunc, err := convertStringToAggFunc(statsAgg.Func) + if err != nil { + return processor.AggQuery{}, err + } + + // 'groupBy' is required for Count and Unique. + if statsAgg.GroupBy == "" && (aggFunc == seq.AggFuncCount || aggFunc == seq.AggFuncUnique) { + return processor.AggQuery{}, fmt.Errorf("%w: groupBy is required for %s func", consts.ErrInvalidAggQuery, aggFunc) + } + + // 'field' is required for stat functions like sum, avg, max and min. + if statsAgg.Field == "" && aggFunc != seq.AggFuncCount && aggFunc != seq.AggFuncUnique { + return processor.AggQuery{}, fmt.Errorf("%w: field is required for %s func", consts.ErrInvalidAggQuery, aggFunc) + } + + // Check 'quantiles' is not empty for Quantile func. + if len(statsAgg.Quantiles) == 0 && aggFunc == seq.AggFuncQuantile { + return processor.AggQuery{}, fmt.Errorf("%w: expect an argument for Quantile func", consts.ErrInvalidAggQuery) + } + + var field *parser.Literal + if statsAgg.Field != "" { + field = &parser.Literal{ + Field: statsAgg.Field, + Terms: searchAllTerms, + } + } + + var groupBy *parser.Literal + if statsAgg.GroupBy != "" { + groupBy = &parser.Literal{ + Field: statsAgg.GroupBy, + Terms: searchAllTerms, + } + } + + procAgg := processor.AggQuery{ + Field: field, + GroupBy: groupBy, + Func: aggFunc, + Quantiles: statsAgg.Quantiles, + } + + if statsAgg.Interval != "" { + interval, err := util.ParseDuration(statsAgg.Interval) + if err != nil { + return processor.AggQuery{}, fmt.Errorf("failed to parse interval: %w", err) + } + procAgg.Interval = int64(seq.MIDToMillis(seq.MID(interval.Nanoseconds()))) + } + + return procAgg, nil +} + +func convertStringToAggFunc(funcName string) (seq.AggFunc, error) { + switch funcName { + case "count": + return seq.AggFuncCount, nil + case "sum": + return seq.AggFuncSum, nil + case "min": + return seq.AggFuncMin, nil + case "max": + return seq.AggFuncMax, nil + case "avg": + return seq.AggFuncAvg, nil + case "quantile": + return seq.AggFuncQuantile, nil + case "unique": + return seq.AggFuncUnique, nil + case "unique_count": + return seq.AggFuncUniqueCount, nil + default: + return 0, fmt.Errorf("unknown aggregation function: %s", funcName) + } +} + +func (p *Plan) applyOffsetID(offsetID string, offset int) error { + if offsetID == "" { + return nil + } + if offset != 0 { + return fmt.Errorf(`only one of "offset" and "offset_id" must be provided`) + } + id, err := seq.FromString(offsetID) + if err != nil { + return fmt.Errorf("could not parse offset_id: %s", offsetID) + } + if p.IsAgg() { + return fmt.Errorf("offset_id is not supported for aggregation requests") + } + p.Scan.OffsetID = id + if p.Scan.Order == seq.DocsOrderDesc { + p.Scan.To = id.MID + } else { + p.Scan.From = id.MID + } + return nil +} diff --git a/query/plan/plan_test.go b/query/plan/plan_test.go new file mode 100644 index 00000000..0844a80c --- /dev/null +++ b/query/plan/plan_test.go @@ -0,0 +1,214 @@ +package plan + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/ozontech/seq-db/parser" + "github.com/ozontech/seq-db/query" + "github.com/ozontech/seq-db/seq" +) + +const ( + testFrom seq.MID = 1000 + testTo seq.MID = 2000 +) + +func TestBuildDocsPipes(t *testing.T) { + seqql := &parser.SeqQLQuery{ + Pipes: []parser.Pipe{ + &parser.PipeFilter{Condition: parser.FilterCondition{Field: "service", Value: "api"}}, + &parser.PipeFields{Fields: []string{"service", "level"}}, + &parser.PipeSort{Order: "asc"}, + &parser.PipeLimit{Limit: 10}, + &parser.PipeOffset{Offset: 5}, + }, + } + + p, err := build(seqql, "", true) + require.NoError(t, err) + + assert.False(t, p.IsAgg()) + assert.Equal(t, query.DocsSchema, p.Schema) + assert.Equal(t, seq.DocsOrderAsc, p.Scan.Order) + assert.True(t, p.Scan.WithTotal) + + // the filter needs the field to be extracted ahead of it + require.Len(t, p.Ops, 4) + extractOp, ok := p.Ops[0].(*ExtractOp) + require.True(t, ok) + assert.Equal(t, []string{"service"}, extractOp.Fields) + filterOp, ok := p.Ops[1].(*FilterOp) + require.True(t, ok) + assert.Equal(t, "service", filterOp.Cond.Field) + assert.Equal(t, "api", filterOp.Cond.Value) + assert.True(t, filterOp.WithTotal) + projectOp, ok := p.Ops[2].(*ProjectOp) + require.True(t, ok) + assert.Equal(t, []string{"service", "level"}, projectOp.Fields) + assert.True(t, projectOp.AllowList) + limitOp, ok := p.Ops[3].(*LimitOp) + require.True(t, ok) + assert.Equal(t, 10, limitOp.Limit) + assert.Equal(t, 5, limitOp.Offset) +} + +func TestBuildFieldsExcept(t *testing.T) { + seqql := &parser.SeqQLQuery{ + Pipes: []parser.Pipe{&parser.PipeFields{Fields: []string{"k8s_pod"}, Except: true}}, + } + + p, err := build(seqql, "", false) + require.NoError(t, err) + + require.Len(t, p.Ops, 1) + projectOp, ok := p.Ops[0].(*ProjectOp) + require.True(t, ok) + assert.Equal(t, []string{"k8s_pod"}, projectOp.Fields) + assert.False(t, projectOp.AllowList) +} + +type testProducer struct{} + +func (testProducer) Next() *query.Record { return nil } +func (testProducer) Finalize() *query.Summary { return nil } + +func TestOpsSchemaChange(t *testing.T) { + // the schema flows through the ops: the extract extends it, + // the filter resolves its column from the extension + var input query.RecordProducer = testProducer{} + schema := query.DocsSchema + + extractOp := &ExtractOp{Fields: []string{"service"}, OpField: query.DocsDataCol} + producer, schema := extractOp.Apply(input, schema) + assert.Equal(t, 3, schema.Len()) + assert.Equal(t, 2, schema.MustColumn[string]("service").Idx()) + + filterOp := &FilterOp{Cond: parser.FilterCondition{Field: "service", Value: "api"}} + producer, schema = filterOp.Apply(producer, schema) + assert.Equal(t, 3, schema.Len()) // filter does't change the schema + require.NotNil(t, producer) +} + +func TestBuildDefaultOrder(t *testing.T) { + // no sort pipe: Order stays zero (desc), the scan range is untouched + p, err := build(&parser.SeqQLQuery{}, "", false) + require.NoError(t, err) + assert.Equal(t, seq.DocsOrder(0), p.Scan.Order) + assert.Equal(t, testFrom, p.Scan.From) + assert.Equal(t, testTo, p.Scan.To) +} + +func TestBuildAgg(t *testing.T) { + seqql := &parser.SeqQLQuery{ + Pipes: []parser.Pipe{ + &parser.PipeStats{Agg: parser.StatsAgg{Func: "sum", Field: "level", GroupBy: "service", Interval: "1m"}}, + }, + } + + p, err := build(seqql, "", false) + require.NoError(t, err) + + assert.True(t, p.IsAgg()) + assert.Equal(t, query.AggsSchema, p.Schema) + assert.Empty(t, p.Ops) + require.Len(t, p.Scan.AggQ, 1) + assert.EqualValues(t, seq.AggFuncSum, p.Scan.AggQ[0].Func) +} + +func TestBuildAggValidation(t *testing.T) { + testCases := []struct { + name string + agg parser.StatsAgg + want string + }{ + {"unknown func", parser.StatsAgg{Func: "median"}, "unknown aggregation function"}, + {"count without groupBy", parser.StatsAgg{Func: "count"}, "groupBy is required"}, + {"sum without field", parser.StatsAgg{Func: "sum", GroupBy: "service"}, "field is required"}, + {"quantile without args", parser.StatsAgg{Func: "quantile", Field: "level", GroupBy: "service"}, "expect an argument"}, + } + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + seqql := &parser.SeqQLQuery{ + Pipes: []parser.Pipe{&parser.PipeStats{Agg: tc.agg}}, + } + _, err := build(seqql, "", false) + require.Error(t, err) + assert.Contains(t, err.Error(), tc.want) + }) + } +} + +func TestBuildOffsetID(t *testing.T) { + const offsetID = "0000000000000001-0000000000000001" + id, err := seq.FromString(offsetID) + require.NoError(t, err) + + t.Run("explicit asc narrows from", func(t *testing.T) { + seqql := &parser.SeqQLQuery{Pipes: []parser.Pipe{&parser.PipeSort{Order: "asc"}}} + p, err := build(seqql, offsetID, false) + require.NoError(t, err) + assert.Equal(t, id, p.Scan.OffsetID) + assert.Equal(t, id.MID, p.Scan.From) + assert.Equal(t, testTo, p.Scan.To) + }) + + t.Run("default desc narrows to", func(t *testing.T) { + p, err := build(&parser.SeqQLQuery{}, offsetID, false) + require.NoError(t, err) + assert.Equal(t, testFrom, p.Scan.From) + assert.Equal(t, id.MID, p.Scan.To) + }) + + t.Run("with offset", func(t *testing.T) { + seqql := &parser.SeqQLQuery{ + Pipes: []parser.Pipe{&parser.PipeOffset{Offset: 5}}, + } + _, err := build(seqql, "0000000000000001-0000000000000001", false) + require.Error(t, err) + assert.Contains(t, err.Error(), `only one of "offset" and "offset_id"`) + }) + + t.Run("with agg", func(t *testing.T) { + seqql := &parser.SeqQLQuery{ + Pipes: []parser.Pipe{&parser.PipeStats{Agg: parser.StatsAgg{Func: "count", GroupBy: "service"}}}, + } + _, err := build(seqql, "0000000000000001-0000000000000001", false) + require.Error(t, err) + assert.Contains(t, err.Error(), "offset_id is not supported") + }) + + t.Run("invalid id", func(t *testing.T) { + _, err := build(&parser.SeqQLQuery{}, "not-an-id", false) + require.Error(t, err) + assert.Contains(t, err.Error(), "could not parse offset_id") + }) +} + +func TestBuildInputValidation(t *testing.T) { + t.Run("nil input", func(t *testing.T) { + _, err := Build(BuildParams{SeqQL: &parser.SeqQLQuery{}, DocField: query.DocsDataCol}) + require.Error(t, err) + assert.Contains(t, err.Error(), "input schema is not set") + }) + t.Run("doc column not in schema", func(t *testing.T) { + s := query.MustNewSchema(query.ColumnDesc{Name: "other", Type: query.DataTypeString}) + _, err := Build(BuildParams{SeqQL: &parser.SeqQLQuery{}, Input: s, DocField: query.DocsDataCol}) + require.Error(t, err) + assert.Contains(t, err.Error(), "doc column:") + }) +} + +func build(seqql *parser.SeqQLQuery, offsetID string, withTotal bool) (*Plan, error) { + return Build(BuildParams{ + SeqQL: seqql, + Input: query.DocsSchema, + DocField: query.DocsDataCol, + From: testFrom, + To: testTo, + OffsetID: offsetID, + WithTotal: withTotal, + }) +} diff --git a/query/schema.go b/query/schema.go new file mode 100644 index 00000000..250c6b20 --- /dev/null +++ b/query/schema.go @@ -0,0 +1,161 @@ +package query + +import ( + "cmp" + "fmt" + "slices" +) + +type ColumnDesc struct { + Name string + Type DataType +} + +// Schema is an immutable ordered list of column descriptors. +type Schema struct { + cols []ColumnDesc + byName map[string]int +} + +func NewSchema(cols ...ColumnDesc) (*Schema, error) { + s := &Schema{ + cols: cols, + byName: make(map[string]int, len(cols)), + } + for i, c := range cols { + if _, ok := s.byName[c.Name]; ok { + return nil, fmt.Errorf("duplicate column name: %s", c.Name) + } + s.byName[c.Name] = i + } + return s, nil +} + +func MustNewSchema(cols ...ColumnDesc) *Schema { + s, err := NewSchema(cols...) + if err != nil { + panic(err) + } + return s +} + +// Index returns the index of the named column, or error if the column is absent. +func (s *Schema) Index(name string) (int, error) { + i, ok := s.byName[name] + if !ok { + return 0, fmt.Errorf("schema has no column %q", name) + } + return i, nil +} + +// Column resolves the column by name. It fails when the +// column is absent or its declared type is incompatible with T. +func (s *Schema) Column[T any](name string) (Column[T], error) { + idx, err := s.Index(name) + if err != nil { + return Column[T]{}, err + } + + t := s.cols[idx].Type + get, ok := columnGetters[t].(func(*RecordVals) T) + if !ok { + return Column[T]{}, fmt.Errorf("column %q has type %s, incompatible with %T", name, t, *new(T)) + } + + // cmp stays nil for types without a total order (document, bytes, arrays) + cmpFunc, _ := columnCmps[t].(func(T, T) int) + + return Column[T]{idx: idx, dataType: t, get: get, cmp: cmpFunc}, nil +} + +// MustColumn resolves the column by name. Panics instead of returning an error. +func (s *Schema) MustColumn[T any](name string) Column[T] { + c, err := s.Column[T](name) + if err != nil { + panic(err) + } + return c +} + +// Equal reports whether both schemas list identical (name, type) pairs in the same order. +func (s *Schema) Equal(other *Schema) bool { + if s == nil || other == nil || len(s.cols) != len(other.cols) { + return false + } + for i, c := range s.cols { + if c != other.cols[i] { + return false + } + } + return true +} + +// Extend returns a new schema with the given columns appended. +// The original schema is untouched. Panics on duplicate names. +func (s *Schema) Extend(cols ...ColumnDesc) *Schema { + return MustNewSchema(append(slices.Clone(s.cols), cols...)...) +} + +// Len returns the number of columns. +func (s *Schema) Len() int { + return len(s.cols) +} + +// Cols returns the column descriptors in schema order. +func (s *Schema) Cols() []ColumnDesc { + return s.cols +} + +var ( + columnGetters = map[DataType]any{ + DataTypeSeqID: (*RecordVals).AsSeqID, + DataTypeBytes: (*RecordVals).AsBytes, + DataTypeString: (*RecordVals).AsString, + DataTypeDocument: (*RecordVals).AsDoc, + DataTypeUint32: (*RecordVals).AsUint32, + DataTypeUint64: (*RecordVals).AsUint64, + DataTypeInt32: (*RecordVals).AsInt32, + DataTypeInt64: (*RecordVals).AsInt64, + DataTypeFloat64: (*RecordVals).AsFloat64, + DataTypeFloat64Array: (*RecordVals).AsFloat64Array, + DataTypeStringArray: (*RecordVals).AsStringArray, + } + + columnCmps = map[DataType]any{ + DataTypeSeqID: cmpSeqID, + DataTypeString: cmp.Compare[string], + DataTypeUint32: cmp.Compare[uint32], + DataTypeUint64: cmp.Compare[uint64], + DataTypeInt32: cmp.Compare[int32], + DataTypeInt64: cmp.Compare[int64], + DataTypeFloat64: cmp.Compare[float64], + } +) + +// We still get hardcoded schemas at the edges of the pipeline. +var ( + DocsIDCol = "id" + DocsDataCol = "data" + + DocsSchema = MustNewSchema( + ColumnDesc{Name: DocsIDCol, Type: DataTypeSeqID}, + ColumnDesc{Name: DocsDataCol, Type: DataTypeDocument}, + ) + AggsSchema = MustNewSchema( + ColumnDesc{Name: "token", Type: DataTypeString}, + ColumnDesc{Name: "min", Type: DataTypeFloat64}, + ColumnDesc{Name: "max", Type: DataTypeFloat64}, + ColumnDesc{Name: "sum", Type: DataTypeFloat64}, + ColumnDesc{Name: "total", Type: DataTypeUint64}, + ColumnDesc{Name: "not_exists", Type: DataTypeUint64}, + ColumnDesc{Name: "ts", Type: DataTypeUint64}, + ColumnDesc{Name: "samples", Type: DataTypeFloat64Array}, + ColumnDesc{Name: "values", Type: DataTypeStringArray}, + ) + AggResultSchema = MustNewSchema( + ColumnDesc{Name: "token", Type: DataTypeString}, + ColumnDesc{Name: "value", Type: DataTypeFloat64}, + ColumnDesc{Name: "ts", Type: DataTypeUint64}, + ColumnDesc{Name: "quantiles", Type: DataTypeFloat64Array}, + ) +) diff --git a/query/schema_test.go b/query/schema_test.go new file mode 100644 index 00000000..2b918321 --- /dev/null +++ b/query/schema_test.go @@ -0,0 +1,144 @@ +package query + +import ( + "testing" + + insaneJSON "github.com/ozontech/insane-json" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/ozontech/seq-db/seq" +) + +func TestNewSchemaDuplicateName(t *testing.T) { + _, err := NewSchema( + ColumnDesc{Name: "id", Type: DataTypeSeqID}, + ColumnDesc{Name: "id", Type: DataTypeDocument}, + ) + assert.Error(t, err) + assert.Equal(t, err.Error(), `duplicate column name: id`) +} + +func TestColumnAllTypes(t *testing.T) { + t.Run("seq id", func(t *testing.T) { + s := MustNewSchema(ColumnDesc{Name: "id", Type: DataTypeSeqID}) + c := s.MustColumn[seq.ID]("id") + assert.Equal(t, 0, c.Idx()) + assert.Equal(t, DataTypeSeqID, c.DataType()) + }) + t.Run("bytes", func(t *testing.T) { + s := MustNewSchema(ColumnDesc{Name: "b", Type: DataTypeBytes}) + c := s.MustColumn[[]byte]("b") + assert.Equal(t, DataTypeBytes, c.DataType()) + }) + t.Run("string", func(t *testing.T) { + s := MustNewSchema(ColumnDesc{Name: "s", Type: DataTypeString}) + c := s.MustColumn[string]("s") + assert.Equal(t, DataTypeString, c.DataType()) + }) + t.Run("document", func(t *testing.T) { + s := MustNewSchema(ColumnDesc{Name: "d", Type: DataTypeDocument}) + c := s.MustColumn[*insaneJSON.Root]("d") + assert.Equal(t, DataTypeDocument, c.DataType()) + }) + t.Run("numeric", func(t *testing.T) { + s := MustNewSchema( + ColumnDesc{Name: "u32", Type: DataTypeUint32}, + ColumnDesc{Name: "u64", Type: DataTypeUint64}, + ColumnDesc{Name: "i32", Type: DataTypeInt32}, + ColumnDesc{Name: "i64", Type: DataTypeInt64}, + ColumnDesc{Name: "f64", Type: DataTypeFloat64}, + ColumnDesc{Name: "f64s", Type: DataTypeFloat64Array}, + ColumnDesc{Name: "strs", Type: DataTypeStringArray}, + ) + assert.Equal(t, DataTypeUint32, s.MustColumn[uint32]("u32").DataType()) + assert.Equal(t, DataTypeUint64, s.MustColumn[uint64]("u64").DataType()) + assert.Equal(t, DataTypeInt32, s.MustColumn[int32]("i32").DataType()) + assert.Equal(t, DataTypeInt64, s.MustColumn[int64]("i64").DataType()) + assert.Equal(t, DataTypeFloat64, s.MustColumn[float64]("f64").DataType()) + assert.Equal(t, DataTypeFloat64Array, s.MustColumn[[]float64]("f64s").DataType()) + assert.Equal(t, DataTypeStringArray, s.MustColumn[[]string]("strs").DataType()) + }) + t.Run("cmp", func(t *testing.T) { + s := MustNewSchema( + ColumnDesc{Name: "id", Type: DataTypeSeqID}, + ColumnDesc{Name: "s", Type: DataTypeString}, + ColumnDesc{Name: "f64", Type: DataTypeFloat64}, + ColumnDesc{Name: "f64s", Type: DataTypeFloat64Array}, + ) + require.NotNil(t, s.MustColumn[seq.ID]("id").Cmp()) + require.NotNil(t, s.MustColumn[string]("s").Cmp()) + require.NotNil(t, s.MustColumn[float64]("f64").Cmp()) + // Array columns have no total order: Cmp must be nil, not a panic. + assert.Nil(t, s.MustColumn[[]float64]("f64s").Cmp()) + }) +} + +func TestColumnMissing(t *testing.T) { + s := DocsSchema + _, err := s.Column[seq.ID]("missing") + require.Error(t, err) + assert.Equal(t, err.Error(), `schema has no column "missing"`) +} + +func TestColumnTypeMismatch(t *testing.T) { + _, err := DocsSchema.Column[float64]("id") + require.Error(t, err) + assert.Equal(t, err.Error(), `column "id" has type seq_id, incompatible with float64`) + + _, err = DocsSchema.Column[float64]("data") + require.Error(t, err) + assert.Equal(t, err.Error(), `column "data" has type document, incompatible with float64`) +} + +func TestMustColumnPanics(t *testing.T) { + assert.Panics(t, func() { + DocsSchema.MustColumn[seq.ID]("missing") + }) + assert.Panics(t, func() { + DocsSchema.MustColumn[float64]("id") + }) +} + +func TestSchemaEqual(t *testing.T) { + base := MustNewSchema( + ColumnDesc{Name: "id", Type: DataTypeSeqID}, + ColumnDesc{Name: "data", Type: DataTypeDocument}, + ) + same := MustNewSchema( + ColumnDesc{Name: "id", Type: DataTypeSeqID}, + ColumnDesc{Name: "data", Type: DataTypeDocument}, + ) + reordered := MustNewSchema( + ColumnDesc{Name: "data", Type: DataTypeDocument}, + ColumnDesc{Name: "id", Type: DataTypeSeqID}, + ) + retyped := MustNewSchema( + ColumnDesc{Name: "id", Type: DataTypeSeqID}, + ColumnDesc{Name: "data", Type: DataTypeString}, + ) + + assert.True(t, base.Equal(same)) + assert.False(t, base.Equal(reordered)) + assert.False(t, base.Equal(retyped)) + assert.False(t, base.Equal(DocsSchema.Extend(ColumnDesc{Name: "x", Type: DataTypeString}))) +} + +func TestSchemaExtend(t *testing.T) { + extended := DocsSchema.Extend( + ColumnDesc{Name: "service", Type: DataTypeString}, + ColumnDesc{Name: "level", Type: DataTypeString}, + ) + assert.Equal(t, 4, extended.Len()) + assert.Equal(t, "service", extended.Cols()[2].Name) + assert.Equal(t, 2, extended.MustColumn[string]("service").Idx()) + + assert.Panics(t, func() { + DocsSchema.Extend(ColumnDesc{Name: "data", Type: DataTypeString}) + }) + + // Extend must not mutate the original. + assert.Equal(t, 2, DocsSchema.Len()) + _, err := DocsSchema.Column[string]("service") + assert.Error(t, err) +} diff --git a/storeapi/grpc_stream_search.go b/storeapi/grpc_stream_search.go index 1992ae8d..ac27008b 100644 --- a/storeapi/grpc_stream_search.go +++ b/storeapi/grpc_stream_search.go @@ -20,6 +20,7 @@ import ( "github.com/ozontech/seq-db/pkg/storeapi" "github.com/ozontech/seq-db/query" "github.com/ozontech/seq-db/query/exec" + "github.com/ozontech/seq-db/query/plan" "github.com/ozontech/seq-db/querytracer" "github.com/ozontech/seq-db/seq" "github.com/ozontech/seq-db/tracing" @@ -117,7 +118,7 @@ func (g *GrpcV1) doStreamSearch( parseQueryTr.Done() buildProducerTr := tr.NewChild("build producer") - producer, typing, err := g.buildProducer(ctx, req, tr, seqql) + producer, schema, err := g.buildProducer(ctx, req, tr, seqql) if err != nil { buildProducerTr.Done() return fmt.Errorf("can't build record producer: %w", err) @@ -127,7 +128,7 @@ func (g *GrpcV1) doStreamSearch( err = stream.Send(&storeapi.StreamSearchResponse{ ResponseType: &storeapi.StreamSearchResponse_Header{ Header: &storeapi.ResponseHeader{ - Typing: typing, + Typing: schemaToTyping(schema), }, }, }) @@ -308,196 +309,75 @@ func sendSummary( return nil } +// buildProducer translates the logical plan into the physical one. Returns the plan's output schema. func (g *GrpcV1) buildProducer( ctx context.Context, req *storeapi.StreamSearchQuery, tr *querytracer.Tracer, seqql parser.SeqQLQuery, -) (query.RecordProducer, []*storeapi.Typing, error) { - // The data source is limitless and walks the matched set via cursor pagination; - // the real request limit is applied by a Limiter executor. - searchParams := processor.SearchParams{ - AST: seqql.Root, +) (query.RecordProducer, *query.Schema, error) { + p, err := plan.Build(plan.BuildParams{ + SeqQL: &seqql, + Input: query.DocsSchema, + DocField: query.DocsDataCol, From: seq.MillisToMID(uint64(seq.TimeToMID(req.From.AsTime()))), To: seq.MillisToMID(uint64(seq.TimeToMID(req.To.AsTime()))), + OffsetID: req.OffsetId, WithTotal: req.WithTotal, - } - - typing := docsTyping() - var offset int - var fieldsFilter *exec.FieldsFilter - var docFilter *exec.DocFilter - - for _, pipe := range seqql.Pipes { - switch p := pipe.(type) { - case *parser.PipeLimit: - searchParams.Limit = p.Limit - case *parser.PipeOffset: - offset = p.Offset - case *parser.PipeSort: - order := seq.DocsOrderAsc - if p.Order == "desc" { - order = seq.DocsOrderDesc - } - searchParams.Order = order - case *parser.PipeStats: - aggQ, err := convertStatsAggToAggQuery(p.Agg) - if err != nil { - return nil, nil, fmt.Errorf("failed to convert stats aggs: %w", err) - } - searchParams.AggQ = []processor.AggQuery{aggQ} - typing = aggsTyping() - case *parser.PipeFilter: - docFilter = exec.NewDocFilter(p.Condition.Field, exec.NewEq(p.Condition.Value)) - case *parser.PipeFields: - fieldsFilter = &exec.FieldsFilter{ - Fields: p.Fields, - AllowList: !p.Except, - } - default: - continue - } - } - - if req.OffsetId != "" { - // offset_id pagination - if offset != 0 { - return nil, nil, fmt.Errorf(`only one of "offset" and "offset_id" must be provided`) - } - offsetId, err := seq.FromString(req.OffsetId) - if err != nil { - return nil, nil, fmt.Errorf("could not parse offset_id: %s", req.OffsetId) - } - if len(searchParams.AggQ) > 0 { - return nil, nil, fmt.Errorf("offset_id is not supported for aggregation requests") - } - searchParams.OffsetId = offsetId - // fractions take into an account offset id, but we also limit time range here - // to filter out unneeded fractions - if searchParams.Order == seq.DocsOrderDesc { - searchParams.To = offsetId.MID - } else { - searchParams.From = offsetId.MID - } - } - - const docDataColIdx = 1 - var producer query.RecordProducer - producer = exec.NewSearcherDataSource(ctx, tr, searchParams, g.fracManager, g.searchData.searcher, g.fetchData.docFetcher) - if len(searchParams.AggQ) > 0 { - return producer, typing, nil - } - if docFilter != nil { - producer = exec.NewFilter(producer, query.DocColumn(docDataColIdx), docFilter, req.WithTotal) - } - if fieldsFilter != nil { - producer = exec.NewDocProjector(producer, query.DocColumn(docDataColIdx), fieldsFilter) - } - if searchParams.Limit > 0 { - // set limit=limit+offset and offset=0 to merge stores' results correctly on proxy - producer = exec.NewLimiter(producer, uint32(searchParams.Limit+offset), 0) - } - - return producer, typing, nil -} - -// hardcoded schema -func docsTyping() []*storeapi.Typing { - return []*storeapi.Typing{ - {Title: "id", Type: storeapi.DataType_SEQ_ID}, - {Title: "data", Type: storeapi.DataType_RAW_DOCUMENT}, - } -} - -// hardcoded schema -func aggsTyping() []*storeapi.Typing { - return []*storeapi.Typing{ - {Title: "token", Type: storeapi.DataType_STRING}, - {Title: "min", Type: storeapi.DataType_FLOAT64}, - {Title: "max", Type: storeapi.DataType_FLOAT64}, - {Title: "sum", Type: storeapi.DataType_FLOAT64}, - {Title: "total", Type: storeapi.DataType_UINT64}, - {Title: "not_exists", Type: storeapi.DataType_UINT64}, - {Title: "ts", Type: storeapi.DataType_UINT64}, - {Title: "samples", Type: storeapi.DataType_FLOAT64_ARRAY}, - {Title: "values", Type: storeapi.DataType_STRING_ARRAY}, - } -} - -func convertStatsAggToAggQuery(statsAgg parser.StatsAgg) (processor.AggQuery, error) { - aggFunc, err := convertStringToAggFunc(statsAgg.Func) + }) if err != nil { - return processor.AggQuery{}, err - } - - // 'groupBy' is required for Count and Unique. - if statsAgg.GroupBy == "" && (aggFunc == seq.AggFuncCount || aggFunc == seq.AggFuncUnique) { - return processor.AggQuery{}, fmt.Errorf("%w: groupBy is required for %s func", consts.ErrInvalidAggQuery, aggFunc) + return nil, nil, err } - // 'field' is required for stat functions like sum, avg, max and min. - if statsAgg.Field == "" && aggFunc != seq.AggFuncCount && aggFunc != seq.AggFuncUnique { - return processor.AggQuery{}, fmt.Errorf("%w: field is required for %s func", consts.ErrInvalidAggQuery, aggFunc) - } - - // Check 'quantiles' is not empty for Quantile func. - if len(statsAgg.Quantiles) == 0 && aggFunc == seq.AggFuncQuantile { - return processor.AggQuery{}, fmt.Errorf("%w: expect an argument for Quantile func", consts.ErrInvalidAggQuery) - } - - var field *parser.Literal - if statsAgg.Field != "" { - field = &parser.Literal{ - Field: statsAgg.Field, - Terms: searchAll, + // The data source is limitless and walks the matched set via cursor pagination; + // the real request limit is applied by a Limiter executor. The limit op + // also caps the per-batch scan size. + searchParams := processor.SearchParams{ + AST: p.Scan.AST, + From: p.Scan.From, + To: p.Scan.To, + WithTotal: p.Scan.WithTotal, + Order: p.Scan.Order, + OffsetId: p.Scan.OffsetID, + AggQ: p.Scan.AggQ, + } + for _, op := range p.Ops { + if op, ok := op.(*plan.LimitOp); ok { + searchParams.Limit = op.Limit } } - var groupBy *parser.Literal - if statsAgg.GroupBy != "" { - groupBy = &parser.Literal{ - Field: statsAgg.GroupBy, - Terms: searchAll, - } - } + var producer query.RecordProducer = exec.NewSearcherDataSource( + ctx, + tr, + searchParams, + g.fracManager, + g.searchData.searcher, + g.fetchData.docFetcher, + ) - procAgg := processor.AggQuery{ - Field: field, - GroupBy: groupBy, - Func: aggFunc, - Quantiles: statsAgg.Quantiles, + if p.IsAgg() { + return producer, p.Schema, nil } - if statsAgg.Interval != "" { - interval, err := util.ParseDuration(statsAgg.Interval) - if err != nil { - return processor.AggQuery{}, fmt.Errorf("failed to parse interval: %w", err) - } - procAgg.Interval = int64(seq.MIDToMillis(seq.MID(interval.Nanoseconds()))) + // ops wrap the source, the schema flows through them so an op that changes the record + // shape extends it here and downstream ops resolve their columns from the extension + schema := p.Schema + for _, op := range p.Ops { + producer, schema = op.Apply(producer, schema) } - return procAgg, nil + return producer, schema, nil } -func convertStringToAggFunc(funcName string) (seq.AggFunc, error) { - switch funcName { - case "count": - return seq.AggFuncCount, nil - case "sum": - return seq.AggFuncSum, nil - case "min": - return seq.AggFuncMin, nil - case "max": - return seq.AggFuncMax, nil - case "avg": - return seq.AggFuncAvg, nil - case "quantile": - return seq.AggFuncQuantile, nil - case "unique": - return seq.AggFuncUnique, nil - case "unique_count": - return seq.AggFuncUniqueCount, nil - default: - return 0, fmt.Errorf("unknown aggregation function: %s", funcName) +func schemaToTyping(s *query.Schema) []*storeapi.Typing { + cols := s.Cols() + out := make([]*storeapi.Typing, 0, len(cols)) + for _, c := range cols { + out = append(out, &storeapi.Typing{ + Title: c.Name, + Type: storeapi.MustProtoDataType(c.Type), + }) } + return out } diff --git a/tests/integration_tests/integration_test.go b/tests/integration_tests/integration_test.go index e5e0f3d2..5a757077 100644 --- a/tests/integration_tests/integration_test.go +++ b/tests/integration_tests/integration_test.go @@ -2121,6 +2121,30 @@ func (s *IntegrationTestSuite) TestStreamSearch() { } }) + t.Run("query with filter", func(t *testing.T) { + stream, conn, _, cancel := newStreamSearchClient(t, env) + defer cancel() + defer conn.Close() + + sendStreamSearchQuery(t, stream, `service:a | filter service:a | limit 10`) + docs, summary := collectStreamData(t, stream) + r.Len(docs, 10) + + gotDocs := make([]string, 0, len(docs)) + for _, d := range docs { + gotDocs = append(gotDocs, string(d)) + } + wantDocs := make([]string, 0, len(origDocs)) + for _, d := range origDocs { + wantDocs = append(wantDocs, d) + } + r.Equal(wantDocs[:10], gotDocs, "streamed documents must match the ingested ones") + + r.NotNil(summary) + r.Equal(uint64(totalDocs), summary.GetTotal(), "summary total must match the document count") + r.Equal(seqproxyapi.ErrorCode_ERROR_CODE_NO, summary.GetError().GetCode()) + }) + t.Run("aggregation stream", func(t *testing.T) { stream, conn, _, cancel := newStreamSearchClient(t, env) defer cancel()