From 2b34976c22a541d95d8b520a0ddc918310d5f4cc Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Kuba=20Podg=C3=B3rski?= Date: Tue, 11 Aug 2020 17:53:36 +0200 Subject: [PATCH 01/17] porting sqlmapper from vdsm --- executor.go | 123 ++++- executor_test.go | 80 ++- field.go | 30 ++ go.mod | 1 + go.sum | 2 + proto/interface.go | 77 +++ server/grpc.go | 96 ++-- server/grpc_test.go | 581 ++++++++++++++++++++- sql/column.go | 112 ++++ sql/extract.go | 1177 +++++++++++++++++++++++++++++++++++++++++++ sql/mapper.go | 113 +++++ sql/mapper_test.go | 394 +++++++++++++++ sql/mask.go | 399 +++++++++++++++ sql/model.go | 141 ++++++ sql/query.go | 295 +++++++++++ sql/reduce.go | 500 ++++++++++++++++++ sql/reduce_test.go | 120 +++++ sql/router.go | 162 ++++++ sql/select.go | 792 +++++++++++++++++++++++++++++ 19 files changed, 5134 insertions(+), 61 deletions(-) create mode 100644 sql/column.go create mode 100644 sql/extract.go create mode 100644 sql/mapper.go create mode 100644 sql/mapper_test.go create mode 100644 sql/mask.go create mode 100644 sql/model.go create mode 100644 sql/query.go create mode 100644 sql/reduce.go create mode 100644 sql/reduce_test.go create mode 100644 sql/router.go create mode 100644 sql/select.go diff --git a/executor.go b/executor.go index 67932cf8d..f329e336e 100644 --- a/executor.go +++ b/executor.go @@ -2820,6 +2820,119 @@ type ExtractedTable struct { Columns []ExtractedTableColumn `json:"columns"` } +// ToRows implements the ToRowser interface. +func (t ExtractedTable) ToRows(callback func(*pb.RowResponse) error) error { + if len(t.Columns) == 0 { + return nil + } + + headers := make([]*pb.ColumnInfo, len(t.Fields)+1) + colType := "uint64" + if t.Columns[0].Column.Keyed { + colType = "string" + } + headers[0] = &pb.ColumnInfo{ + Name: "_id", + Datatype: colType, + } + dataHeaders := headers[1:] + for i, f := range t.Fields { + dataHeaders[i] = &pb.ColumnInfo{ + Name: f.Name, + Datatype: f.Type, + } + } + + for _, c := range t.Columns { + cols := make([]*pb.ColumnResponse, len(c.Rows)+1) + if c.Column.Keyed { + cols[0] = &pb.ColumnResponse{ + ColumnVal: &pb.ColumnResponse_StringVal{ + StringVal: c.Column.Key, + }, + } + } else { + cols[0] = &pb.ColumnResponse{ + ColumnVal: &pb.ColumnResponse_Uint64Val{ + Uint64Val: c.Column.ID, + }, + } + } + valCols := cols[1:] + for i, r := range c.Rows { + var col *pb.ColumnResponse + switch r := r.(type) { + case bool: + col = &pb.ColumnResponse{ + ColumnVal: &pb.ColumnResponse_BoolVal{ + BoolVal: r, + }, + } + case int64: + col = &pb.ColumnResponse{ + ColumnVal: &pb.ColumnResponse_Int64Val{ + Int64Val: r, + }, + } + case uint64: + col = &pb.ColumnResponse{ + ColumnVal: &pb.ColumnResponse_Uint64Val{ + Uint64Val: r, + }, + } + case string: + col = &pb.ColumnResponse{ + ColumnVal: &pb.ColumnResponse_StringVal{ + StringVal: r, + }, + } + case []uint64: + col = &pb.ColumnResponse{ + ColumnVal: &pb.ColumnResponse_Uint64ArrayVal{ + Uint64ArrayVal: &pb.Uint64Array{ + Vals: r, + }, + }, + } + case []string: + col = &pb.ColumnResponse{ + ColumnVal: &pb.ColumnResponse_StringArrayVal{ + StringArrayVal: &pb.StringArray{ + Vals: r, + }, + }, + } + case pql.Decimal: + col = &pb.ColumnResponse{ + ColumnVal: &pb.ColumnResponse_DecimalVal{ + DecimalVal: &pb.Decimal{ + Value: r.Value, + Scale: r.Scale, + }, + }, + } + default: + return errors.Errorf("unsupported field value: %v (type: %T)", r, r) + } + valCols[i] = col + } + err := callback(&pb.RowResponse{ + Headers: headers, + Columns: cols, + }) + if err != nil { + return err + } + } + + return nil +} + +// ToTable converts the table to protobuf format. +func (t ExtractedTable) ToTable() (*pb.TableResponse, error) { + return pb.RowsToTable(t, len(t.Columns)) +} + type ExtractedIDColumn struct { ColumnID uint64 Rows [][]uint64 @@ -5019,15 +5132,17 @@ func (e *executor) translateResult(ctx context.Context, index string, idx *Index return nil, ErrFieldNotFound } - typ := field.Type() - + datatype, err := field.Datatype() + if err != nil { + return nil, errors.Wrapf(err, "field %s", v) + } fields[i] = ExtractedTableField{ Name: v, - Type: typ, + Type: datatype, } var mapper fieldMapper - switch typ { + switch typ := field.Type(); typ { case FieldTypeBool: mapper = func(ids []uint64) (interface{}, error) { switch len(ids) { diff --git a/executor_test.go b/executor_test.go index 3bd9c1139..e29b8d93b 100644 --- a/executor_test.go +++ b/executor_test.go @@ -4344,7 +4344,11 @@ func TestExecutor_Execute_Extract(t *testing.T) { c := test.MustRunCluster(t, 3) defer c.Close() - c.CreateField(t, "i", pilosa.IndexOptions{TrackExistence: true}, "set") + set := c.CreateField(t, "i", pilosa.IndexOptions{TrackExistence: true}, "set") + dtSet, err := set.Datatype() + if err != nil { + t.Fatal(err) + } c.ImportBits(t, "i", "set", [][2]uint64{ {0, 1}, {0, 2}, @@ -4355,56 +4359,88 @@ func TestExecutor_Execute_Extract(t *testing.T) { }) c.Query(t, "i", fmt.Sprintf("Clear(%d, set=5)", ShardWidth)) - c.CreateField(t, "i", pilosa.IndexOptions{TrackExistence: true}, "keyset", pilosa.OptFieldKeys()) + keyset := c.CreateField(t, "i", pilosa.IndexOptions{TrackExistence: true}, "keyset", pilosa.OptFieldKeys()) + dtKeyset, err := keyset.Datatype() + if err != nil { + t.Fatal(err) + } c.Query(t, "i", ` Set(0, keyset="h") Set(1, keyset="xyzzy") Set(0, keyset="plugh") `) - c.CreateField(t, "i", pilosa.IndexOptions{TrackExistence: true}, "mutex", pilosa.OptFieldTypeMutex(pilosa.CacheTypeRanked, 5000)) + mutex := c.CreateField(t, "i", pilosa.IndexOptions{TrackExistence: true}, "mutex", pilosa.OptFieldTypeMutex(pilosa.CacheTypeRanked, 5000)) + dtMutex, err := mutex.Datatype() + if err != nil { + t.Fatal(err) + } c.ImportBits(t, "i", "mutex", [][2]uint64{ {0, 1}, {0, 2}, {4, 4 * ShardWidth}, }) - c.CreateField(t, "i", pilosa.IndexOptions{TrackExistence: true}, "keymutex", pilosa.OptFieldKeys(), pilosa.OptFieldTypeMutex(pilosa.CacheTypeRanked, 5000)) + keymutex := c.CreateField(t, "i", pilosa.IndexOptions{TrackExistence: true}, "keymutex", pilosa.OptFieldKeys(), pilosa.OptFieldTypeMutex(pilosa.CacheTypeRanked, 5000)) + dtKeyMutex, err := keymutex.Datatype() + if err != nil { + t.Fatal(err) + } c.Query(t, "i", ` Set(0, keymutex="h") Set(1, keymutex="xyzzy") Set(3, keymutex="plugh") `) - c.CreateField(t, "i", pilosa.IndexOptions{TrackExistence: true}, "time", pilosa.OptFieldTypeTime("YMDH")) + tm := c.CreateField(t, "i", pilosa.IndexOptions{TrackExistence: true}, "time", pilosa.OptFieldTypeTime("YMDH")) + dtTm, err := tm.Datatype() + if err != nil { + t.Fatal(err) + } c.Query(t, "i", ` Set(0, time=1, 2016-01-01T00:00) Set(1, time=2, 2017-01-01T00:00) Set(3, time=3, 2018-01-01T00:00) `) - c.CreateField(t, "i", pilosa.IndexOptions{TrackExistence: true}, "keytime", pilosa.OptFieldKeys(), pilosa.OptFieldTypeTime("YMDH")) + keytm := c.CreateField(t, "i", pilosa.IndexOptions{TrackExistence: true}, "keytime", pilosa.OptFieldKeys(), pilosa.OptFieldTypeTime("YMDH")) + dtKeyTm, err := keytm.Datatype() + if err != nil { + t.Fatal(err) + } c.Query(t, "i", ` Set(0, keytime="h", 2016-01-01T00:00) Set(1, keytime="xyzzy", 2017-01-01T00:00) Set(0, keytime="plugh", 2018-01-01T00:00) `) - c.CreateField(t, "i", pilosa.IndexOptions{TrackExistence: true}, "bsint", pilosa.OptFieldTypeInt(-100, 100)) + bsiInt := c.CreateField(t, "i", pilosa.IndexOptions{TrackExistence: true}, "bsint", pilosa.OptFieldTypeInt(-100, 100)) + dtBsiInt, err := bsiInt.Datatype() + if err != nil { + t.Fatal(err) + } c.Query(t, "i", ` Set(0, bsint=1) Set(1, bsint=-1) Set(3, bsint=2) `) - c.CreateField(t, "i", pilosa.IndexOptions{TrackExistence: true}, "bsidecimal", pilosa.OptFieldTypeDecimal(2)) + bsidecimal := c.CreateField(t, "i", pilosa.IndexOptions{TrackExistence: true}, "bsidecimal", pilosa.OptFieldTypeDecimal(2)) + dtBsiDecimal, err := bsidecimal.Datatype() + if err != nil { + t.Fatal(err) + } c.Query(t, "i", ` Set(0, bsidecimal=0.01) Set(1, bsidecimal=1.00) Set(3, bsidecimal=-1.01) `) - c.CreateField(t, "i", pilosa.IndexOptions{TrackExistence: true}, "bool", pilosa.OptFieldTypeBool()) + boolean := c.CreateField(t, "i", pilosa.IndexOptions{TrackExistence: true}, "bool", pilosa.OptFieldTypeBool()) + dtBoolean, err := boolean.Datatype() + if err != nil { + t.Fatal(err) + } c.Query(t, "i", ` Set(0, bool=true) Set(1, bool=false) @@ -4417,39 +4453,39 @@ func TestExecutor_Execute_Extract(t *testing.T) { Fields: []pilosa.ExtractedTableField{ { Name: "set", - Type: pilosa.FieldTypeSet, + Type: dtSet, }, { Name: "keyset", - Type: pilosa.FieldTypeSet, + Type: dtKeyset, }, { Name: "mutex", - Type: pilosa.FieldTypeMutex, + Type: dtMutex, }, { Name: "keymutex", - Type: pilosa.FieldTypeMutex, + Type: dtKeyMutex, }, { Name: "time", - Type: pilosa.FieldTypeTime, + Type: dtTm, }, { Name: "keytime", - Type: pilosa.FieldTypeTime, + Type: dtKeyTm, }, { Name: "bsint", - Type: pilosa.FieldTypeInt, + Type: dtBsiInt, }, { Name: "bsidecimal", - Type: pilosa.FieldTypeDecimal, + Type: dtBsiDecimal, }, { Name: "bool", - Type: pilosa.FieldTypeBool, + Type: dtBoolean, }, }, Columns: []pilosa.ExtractedTableColumn{ @@ -4574,7 +4610,11 @@ func TestExecutor_Execute_Extract_Keyed(t *testing.T) { c := test.MustRunCluster(t, 3) defer c.Close() - c.CreateField(t, "i", pilosa.IndexOptions{TrackExistence: true, Keys: true}, "set") + set := c.CreateField(t, "i", pilosa.IndexOptions{TrackExistence: true, Keys: true}, "set") + dtSet, err := set.Datatype() + if err != nil { + t.Fatal(err) + } c.Query(t, "i", ` Set("h", set=1) Set("h", set=2) @@ -4589,7 +4629,7 @@ func TestExecutor_Execute_Extract_Keyed(t *testing.T) { Fields: []pilosa.ExtractedTableField{ { Name: "set", - Type: "set", + Type: dtSet, }, }, Columns: []pilosa.ExtractedTableColumn{ diff --git a/field.go b/field.go index c557f0e81..704f05ae9 100644 --- a/field.go +++ b/field.go @@ -1321,6 +1321,36 @@ func (f *Field) ClearBit(tx Tx, rowID, colID uint64) (changed bool, err error) { return changed, nil } +// Datatype returns a useful data type (string, +// uint64, bool, etc.) based on the field type. +func (f *Field) Datatype() (string, error) { + switch t := f.Type(); t { + case "set": + if f.Keys() { + return "[]string", nil + } + return "[]uint64", nil + case "mutex": + if f.Keys() { + return "string", nil + } + return "uint64", nil + case "int": + if f.Keys() { + return "string", nil + } + return "int64", nil + case "decimal": + return "decimal", nil + case "bool": + return "bool", nil + case "time": + return "int64", nil // TODO: this is a placeholder + default: + return "", fmt.Errorf("unimplemented field Datatype: %s", t) + } +} + func groupCompare(a, b string, offset int) (lt, eq bool) { if len(a) > offset { a = a[:offset] diff --git a/go.mod b/go.mod index 70221964e..8e545a410 100644 --- a/go.mod +++ b/go.mod @@ -43,6 +43,7 @@ require ( google.golang.org/grpc v1.28.0 modernc.org/mathutil v1.0.0 modernc.org/strutil v1.0.0 + vitess.io/vitess v3.0.0-rc.3.0.20190602171040-12bfde34629c+incompatible ) go 1.13 diff --git a/go.sum b/go.sum index f97f5d051..ad4e3d287 100644 --- a/go.sum +++ b/go.sum @@ -288,3 +288,5 @@ modernc.org/mathutil v1.0.0 h1:93vKjrJopTPrtTNpZ8XIovER7iCIH1QU7wNbOQXC60I= modernc.org/mathutil v1.0.0/go.mod h1:wU0vUrJsVWBZ4P6e7xtFJEhFSNsfRLJ8H458uRjg03k= modernc.org/strutil v1.0.0 h1:XVFtQwFVwc02Wk+0L/Z/zDDXO81r5Lhe6iMKmGX3KhE= modernc.org/strutil v1.0.0/go.mod h1:lstksw84oURvj9y3tn8lGvRxyRC1S2+g5uuIzNfIOBs= +vitess.io/vitess v3.0.0-rc.3.0.20190602171040-12bfde34629c+incompatible h1:GWnLrAdetgJM0Co5bwwczO49iFZBSInpyGAT77BP9Y0= +vitess.io/vitess v3.0.0-rc.3.0.20190602171040-12bfde34629c+incompatible/go.mod h1:h4qvkyNYTOC0xI+vcidSWoka0gQAZc9ZPHbkHo48gP0= diff --git a/proto/interface.go b/proto/interface.go index 34bf53829..660aa6273 100644 --- a/proto/interface.go +++ b/proto/interface.go @@ -16,6 +16,7 @@ package pilosa import ( "fmt" + "io" "strings" "github.com/pkg/errors" @@ -30,6 +31,36 @@ type StreamClient interface { Recv() (*RowResponse, error) } +// ReadIntoTable reads from a StreamClient and stores the result into a table response. +func ReadIntoTable(cli StreamClient) (*TableResponse, error) { + var headers []*ColumnInfo + rows := []*Row{} + +rx: + for { + row, err := cli.Recv() + switch err { + case nil: + case io.EOF: + break rx + default: + return nil, err + } + + if headers == nil { + headers = row.Headers + } + rows = append(rows, &Row{ + Columns: row.Columns, + }) + } + + return &TableResponse{ + Headers: headers, + Rows: rows, + }, nil +} + // StreamServer is an interface for a stream // which can accept a RowResponse to be later // returned by the stream via Recv(). @@ -84,6 +115,52 @@ func RowsToTable(tr ToRowser, n int) (*TableResponse, error) { }, nil } +// RowBuffer acts as a Sender/Receiver of RowResponses. +// Note that sending a nil value will cause the Recv +// method to return an io.EOF error. +type RowBuffer struct { + ch chan *RowResponse +} + +// NewRowBuffer returns a new instance of RowBuffer. +// sz is the size of the buffer. +func NewRowBuffer(sz int) *RowBuffer { + var chSz int + if sz > 0 { + // Add one to allow for the EOF record. + chSz = sz + 1 + } + return &RowBuffer{ + ch: make(chan *RowResponse, chSz), + } +} + +// Recv returns a RowResponse. When the buffer is empty, +// calling Recv will return an io.EOF error. +func (rb *RowBuffer) Recv() (*RowResponse, error) { + r := <-rb.ch + + // If the StatusError contains a message then return + // with the approprate error. + se := r.GetStatusError() + code := codes.Code(se.GetCode()) + msg := se.GetMessage() + if code != codes.OK { + return nil, status.Error(code, msg) + } else if msg == "EOF" { + return nil, io.EOF + } else if msg != "" { + return nil, errors.New(msg) + } + + return r, nil +} + +func (rb *RowBuffer) Send(rr *RowResponse) error { + rb.ch <- rr + return nil +} + // EOF acts as an io.EOF encoded into a RowResponse. var EOF *RowResponse = &RowResponse{ StatusError: &StatusError{ diff --git a/server/grpc.go b/server/grpc.go index 5fd16692b..3dcc420a1 100644 --- a/server/grpc.go +++ b/server/grpc.go @@ -18,6 +18,7 @@ import ( "context" "crypto/tls" "fmt" + "io" "net" "strings" "time" @@ -25,6 +26,7 @@ import ( "github.com/pilosa/pilosa/v2" "github.com/pilosa/pilosa/v2/logger" pb "github.com/pilosa/pilosa/v2/proto" + "github.com/pilosa/pilosa/v2/sql" "github.com/pilosa/pilosa/v2/stats" "github.com/pkg/errors" "google.golang.org/grpc" @@ -117,9 +119,49 @@ func (h *GRPCHandler) DeleteVDS(ctx context.Context, req *pb.DeleteVDSRequest) ( return &pb.DeleteVDSResponse{}, nil } +func (h *GRPCHandler) execSQL(ctx context.Context, queryStr string) (pb.StreamClient, error) { + mapper := sql.NewMapper() + mapper.Logger = h.logger + query, err := mapper.MapSQL(queryStr) + if err != nil { + return nil, errors.Wrap(err, "failed to map SQL") + } + var results pb.StreamClient + switch query.SQLType { + case sql.SQLTypeSelect: + handler := sql.NewSelectHandler(h.api) + results, err = handler.Handle(ctx, query) + if err != nil { + return nil, errors.Wrap(err, "failed to start SQL query") + } + default: + return nil, status.Errorf(codes.Unimplemented, "query type not supported") + } + return results, nil +} + // QuerySQL handles the SQL request and sends RowResponses to the stream. -func (*GRPCHandler) QuerySQL(req *pb.QuerySQLRequest, stream pb.Pilosa_QuerySQLServer) error { - return status.Errorf(codes.Unimplemented, "method QuerySQL not implemented") +func (h *GRPCHandler) QuerySQL(req *pb.QuerySQLRequest, stream pb.Pilosa_QuerySQLServer) error { + results, err := h.execSQL(stream.Context(), req.Sql) + if err != nil { + return err + } + + for { + row, err := results.Recv() + switch err { + case nil: + case io.EOF: + return nil + default: + return errors.Wrap(err, "failed to load next row") + } + + err = stream.Send(row) + if err != nil { + return errors.Wrap(err, "failed to send row") + } + } } // QuerySQLUnary is a unary-response (non-streaming) version of QuerySQL, returning a TableResponse. @@ -134,8 +176,12 @@ func (*GRPCHandler) QuerySQL(req *pb.QuerySQLRequest, stream pb.Pilosa_QuerySQLS // Futures, which are used by python-molecula to perform multiple queries // concurrently. There is additional discussion and historical context here: // https://github.com/molecula/pilosa/pull/644 -func (*GRPCHandler) QuerySQLUnary(ctx context.Context, req *pb.QuerySQLRequest) (*pb.TableResponse, error) { - return nil, status.Errorf(codes.Unimplemented, "method QuerySQLUnary not implemented") +func (h *GRPCHandler) QuerySQLUnary(ctx context.Context, req *pb.QuerySQLRequest) (*pb.TableResponse, error) { + results, err := h.execSQL(ctx, req.Sql) + if err != nil { + return nil, err + } + return pb.ReadIntoTable(results) } // QueryPQL handles the PQL request and sends RowResponses to the stream. @@ -298,36 +344,6 @@ func ToRowserWrapper(result interface{}) (pb.ToRowser, error) { return toRowser, nil } -// fieldDataType returns a useful data type (string, -// uint64, bool, etc.) based on the Pilosa field type. -func fieldDataType(f *pilosa.Field) string { - switch f.Type() { - case "set": - if f.Keys() { - return "[]string" - } - return "[]uint64" - case "mutex": - if f.Keys() { - return "string" - } - return "uint64" - case "int": - if f.Keys() { - return "string" - } - return "int64" - case "decimal": - return "decimal" - case "bool": - return "bool" - case "time": - return "int64" // TODO: this is a placeholder - default: - panic(fmt.Sprintf("unimplemented fieldDataType: %s", f.Type())) - } -} - // Inspect handles the inspect request and sends an InspectResponse to the stream. func (h *GRPCHandler) Inspect(req *pb.InspectRequest, stream pb.Pilosa_InspectServer) error { const defaultLimit = 100000 @@ -422,7 +438,11 @@ func (h *GRPCHandler) Inspect(req *pb.InspectRequest, stream pb.Pilosa_InspectSe {Name: "_id", Datatype: "uint64"}, } for _, field := range fields { - ci = append(ci, &pb.ColumnInfo{Name: field.Name(), Datatype: fieldDataType(field)}) + fdt, err := field.Datatype() + if err != nil { + return errors.Wrapf(err, "field %s", field.Name()) + } + ci = append(ci, &pb.ColumnInfo{Name: field.Name(), Datatype: fdt}) } // If Columns is empty, then get the _exists list (via All()), @@ -666,7 +686,11 @@ func (h *GRPCHandler) Inspect(req *pb.InspectRequest, stream pb.Pilosa_InspectSe {Name: "_id", Datatype: "string"}, } for _, field := range fields { - ci = append(ci, &pb.ColumnInfo{Name: field.Name(), Datatype: fieldDataType(field)}) + fdt, err := field.Datatype() + if err != nil { + return errors.Wrapf(err, "field %s", field.Name()) + } + ci = append(ci, &pb.ColumnInfo{Name: field.Name(), Datatype: fdt}) } // If Columns is empty, then get the _exists list (via All()), diff --git a/server/grpc_test.go b/server/grpc_test.go index 2009bbfc6..e143c9168 100644 --- a/server/grpc_test.go +++ b/server/grpc_test.go @@ -16,9 +16,13 @@ package server_test import ( "context" + "fmt" + "reflect" + "strconv" "testing" "github.com/pilosa/pilosa/v2" + "github.com/pilosa/pilosa/v2/pql" pb "github.com/pilosa/pilosa/v2/proto" "github.com/pilosa/pilosa/v2/server" "github.com/pilosa/pilosa/v2/test" @@ -338,7 +342,6 @@ func TestQueryPQLUnary(t *testing.T) { i := m.MustCreateIndex(t, "i", pilosa.IndexOptions{}) m.MustCreateField(t, i.Name(), "f", pilosa.OptFieldKeys()) - ctx := context.Background() gh := server.NewGRPCHandler(m.API) @@ -361,3 +364,579 @@ func TestQueryPQLUnary(t *testing.T) { t.Fatalf("expected error: InvalidArgument, got: %v", err) } } + +type ( + tableResponse struct { + headers []columnInfo + rows []row + } + columnInfo struct { + name string + datatype string + } + row struct { + columns []columnResponse + } + columnResponse interface{} +) + +func TestQuerySQLUnary(t *testing.T) { + + ctx := context.Background() + gh, tearDownFunc := setUpTestQuerySQLUnary(ctx, t) + defer tearDownFunc() + + tests := []struct { + sql string + exp tableResponse + eq func(tableResponse, tableResponse) error + }{ + { + // Extract(Limit(All(), limit=100, offset=0),Rows(age)) + sql: "select age from grouper", + exp: tableResponse{ + headers: []columnInfo{ + {"age", "int64"}, + }, + rows: []row{ + {[]columnResponse{int64(27)}}, + {[]columnResponse{int64(16)}}, + {[]columnResponse{int64(19)}}, + {[]columnResponse{int64(27)}}, + {[]columnResponse{int64(16)}}, + {[]columnResponse{int64(34)}}, + {[]columnResponse{int64(27)}}, + {[]columnResponse{int64(16)}}, + {[]columnResponse{int64(16)}}, + {[]columnResponse{int64(31)}}, + }, + }, + eq: equal, + }, + { + // Extract(Limit(ConstRow(columns=[2]), limit=100, offset=0),Rows(age),Rows(color),Rows(height),Rows(score)) + sql: "select * from grouper where _id=2", + exp: tableResponse{ + headers: []columnInfo{ + {"_id", "uint64"}, + {"age", "int64"}, + {"color", "[]string"}, + {"height", "int64"}, + {"score", "int64"}, + }, + rows: []row{ + {[]columnResponse{uint64(2), int64(16), []string{"blue"}, int64(30), int64(-8)}}, + }, + }, + eq: equal, + }, + // join + { + // Count(Intersect(All(),Distinct(Row(grouperid!=null),index='joiner',field='grouperid'))) + sql: "select count(*) from grouper g INNER JOIN joiner j ON g._id = j.grouperid", + exp: tableResponse{ + headers: []columnInfo{ + {"count(*)", "uint64"}, + }, + rows: []row{ + {[]columnResponse{uint64(8)}}, + }, + }, + eq: equal, + }, + { + // Intersect(All(),Distinct(Row(grouperid!=null),index='joiner',field='grouperid')) + sql: "select _id from grouper g INNER JOIN joiner j ON g._id = j.grouperid", + exp: tableResponse{ + headers: []columnInfo{{"_id", "uint64"}}, + rows: []row{ + {[]columnResponse{uint64(1)}}, + {[]columnResponse{uint64(2)}}, + {[]columnResponse{uint64(3)}}, + {[]columnResponse{uint64(5)}}, + {[]columnResponse{uint64(6)}}, + {[]columnResponse{uint64(7)}}, + {[]columnResponse{uint64(8)}}, + {[]columnResponse{uint64(9)}}, + }, + }, + eq: equalUnordered, + }, + { + // Intersect(Row(color='red'),Distinct(Row(grouperid!=null),index='joiner',field='grouperid')) + sql: "select _id from grouper g INNER JOIN joiner j ON g._id = j.grouperid where g.color = 'red'", + exp: tableResponse{ + headers: []columnInfo{{"_id", "uint64"}}, + rows: []row{ + {[]columnResponse{uint64(3)}}, + {[]columnResponse{uint64(8)}}, + {[]columnResponse{uint64(9)}}, + }, + }, + eq: equalUnordered, + }, + { + // Intersect(Row(color='red'),Distinct(Row(jointype=2),index='joiner',field='grouperid')) + sql: "select _id from grouper g INNER JOIN joiner j ON g._id = j.grouperid where g.color = 'red' and j.jointype = 2", + exp: tableResponse{ + headers: []columnInfo{{"_id", "uint64"}}, + rows: []row{ + {[]columnResponse{uint64(3)}}, + {[]columnResponse{uint64(8)}}, + {[]columnResponse{uint64(9)}}, + }, + }, + eq: equalUnordered, + }, + // order by + { + // Distinct(Row(score!=null),index='grouper',field='score') + sql: "select distinct score from grouper order by score asc", + exp: tableResponse{ + headers: []columnInfo{{"score", "int64"}}, + rows: []row{ + {[]columnResponse{int64(-13)}}, + {[]columnResponse{int64(-10)}}, + {[]columnResponse{int64(-8)}}, + {[]columnResponse{int64(-2)}}, + {[]columnResponse{int64(0)}}, + {[]columnResponse{int64(6)}}, + {[]columnResponse{int64(80)}}, + {[]columnResponse{int64(100)}}, + }, + }, + eq: equal, + }, + { + // Distinct(Row(score!=null),index='grouper',field='score') + sql: "select distinct score from grouper order by score desc", + exp: tableResponse{ + headers: []columnInfo{{"score", "int64"}}, + rows: []row{ + {[]columnResponse{int64(100)}}, + {[]columnResponse{int64(80)}}, + {[]columnResponse{int64(6)}}, + {[]columnResponse{int64(0)}}, + {[]columnResponse{int64(-2)}}, + {[]columnResponse{int64(-8)}}, + {[]columnResponse{int64(-10)}}, + {[]columnResponse{int64(-13)}}, + }, + }, + eq: equal, + }, + { + // Distinct(Row(score!=null),index='grouper',field='score') + sql: "select distinct score from grouper order by score asc limit 5", + exp: tableResponse{ + headers: []columnInfo{{"score", "int64"}}, + rows: []row{ + {[]columnResponse{int64(-13)}}, + {[]columnResponse{int64(-10)}}, + {[]columnResponse{int64(-8)}}, + {[]columnResponse{int64(-2)}}, + {[]columnResponse{int64(0)}}, + }, + }, + eq: equal, + }, + + { + // Distinct(Row(score!=null),index='grouper',field='score') + sql: "select distinct score from grouper order by score desc limit 5", + exp: tableResponse{ + headers: []columnInfo{{"score", "int64"}}, + rows: []row{ + {[]columnResponse{int64(100)}}, + {[]columnResponse{int64(80)}}, + {[]columnResponse{int64(6)}}, + {[]columnResponse{int64(0)}}, + {[]columnResponse{int64(-2)}}, + }, + }, + eq: equal, + }, + + // distinct + { + // Distinct(Row(score!=null),index='grouper',field='score') + sql: "select distinct score from grouper", + exp: tableResponse{ + headers: []columnInfo{{"score", "int64"}}, + rows: []row{ + {[]columnResponse{int64(-13)}}, + {[]columnResponse{int64(-10)}}, + {[]columnResponse{int64(-8)}}, + {[]columnResponse{int64(-2)}}, + {[]columnResponse{int64(0)}}, + {[]columnResponse{int64(6)}}, + {[]columnResponse{int64(80)}}, + {[]columnResponse{int64(100)}}, + }, + }, + eq: equalUnordered, + }, + { + + // Distinct(Row(height!=null),index='grouper',field='height') + sql: "select distinct height from grouper", + exp: tableResponse{ + headers: []columnInfo{{"height", "int64"}}, + rows: []row{ + {[]columnResponse{int64(20)}}, + {[]columnResponse{int64(30)}}, + {[]columnResponse{int64(40)}}, + {[]columnResponse{int64(50)}}, + {[]columnResponse{int64(60)}}, + {[]columnResponse{int64(70)}}, + {[]columnResponse{int64(80)}}, + {[]columnResponse{int64(90)}}, + {[]columnResponse{int64(100)}}, + {[]columnResponse{int64(110)}}, + }, + }, + eq: equalUnordered, + }, + + // groupby + { + // GroupBy(Rows(field='age'),limit=100) + sql: "select age as yrs, count(*) as cnt from grouper group by age", + exp: tableResponse{ + headers: []columnInfo{ + {"yrs", "int64"}, + {"cnt", "uint64"}, + }, + rows: []row{ + {[]columnResponse{int64(16), uint64(4)}}, + {[]columnResponse{int64(19), uint64(1)}}, + {[]columnResponse{int64(27), uint64(3)}}, + {[]columnResponse{int64(31), uint64(1)}}, + {[]columnResponse{int64(34), uint64(1)}}, + }, + }, + eq: equalUnordered, + }, + { + // GroupBy(Rows(field='age'),Rows(field='color'),limit=100) + sql: "select age, color, count(*) from grouper group by age, color", + exp: tableResponse{ + headers: []columnInfo{ + {"age", "int64"}, + {"color", "string"}, + {"count(*)", "uint64"}, + }, + rows: []row{ + {[]columnResponse{int64(16), "blue", uint64(2)}}, + {[]columnResponse{int64(16), "red", uint64(2)}}, + {[]columnResponse{int64(19), "red", uint64(1)}}, + {[]columnResponse{int64(27), "blue", uint64(2)}}, + {[]columnResponse{int64(27), "green", uint64(1)}}, + {[]columnResponse{int64(31), "red", uint64(1)}}, + {[]columnResponse{int64(34), "blue", uint64(1)}}, + }, + }, + eq: equalUnordered, + }, + { + // GroupBy(Rows(field='age'),Rows(field='color'),limit=100,filter=Row(age=27),aggregate=Sum(field='height')) + sql: "select age, color, sum(height) from grouper where age = 27 group by age, color", + exp: tableResponse{ + headers: []columnInfo{ + {"age", "int64"}, + {"color", "string"}, + {"sum(height)", "int64"}, + }, + rows: []row{ + {[]columnResponse{int64(27), "blue", int64(100)}}, + {[]columnResponse{int64(27), "green", int64(50)}}, + }, + }, + eq: equalUnordered, + }, + { + // GroupBy(Rows(field='age'),limit=100,having=Condition(count>1)) + sql: "select age, count(*) from grouper group by age having count > 1", + exp: tableResponse{ + headers: []columnInfo{ + {"age", "int64"}, + {"count(*)", "uint64"}, + }, + rows: []row{ + {[]columnResponse{int64(16), uint64(4)}}, + {[]columnResponse{int64(27), uint64(3)}}, + }, + }, + eq: equalUnordered, + }, + { + // GroupBy(Rows(field='age'),limit=100,having=Condition(1<=count<=3)) + sql: "select age, count(*) from grouper group by age having count between 1 and 3", + exp: tableResponse{ + headers: []columnInfo{ + {"age", "int64"}, + {"count(*)", "uint64"}, + }, + rows: []row{ + {[]columnResponse{int64(19), uint64(1)}}, + {[]columnResponse{int64(27), uint64(3)}}, + {[]columnResponse{int64(31), uint64(1)}}, + {[]columnResponse{int64(34), uint64(1)}}, + }, + }, + eq: equalUnordered, + }, + + { + // GroupBy(Rows(field='age'),limit=3) + sql: "select age, count(*) as cnt from grouper group by age order by cnt desc, age desc limit 3", + exp: tableResponse{ + headers: []columnInfo{ + {"age", "int64"}, + {"cnt", "uint64"}, + }, + rows: []row{ + {[]columnResponse{int64(16), uint64(4)}}, + {[]columnResponse{int64(27), uint64(3)}}, + {[]columnResponse{int64(19), uint64(1)}}, + }, + }, + eq: equal, + }, + } + + for i, test := range tests { + t.Run("test-"+strconv.Itoa(i), func(t *testing.T) { + resp, err := gh.QuerySQLUnary(ctx, &pb.QuerySQLRequest{Sql: test.sql}) + if err != nil { + t.Fatalf("sql: %s, error: %v", test.sql, err) + } else { + tr := toTableResponse(resp) + if err := test.eq(test.exp, tr); err != nil { + t.Fatalf("sql: %s, error: %+v", test.sql, err) + } + } + }) + } +} + +func setUpTestQuerySQLUnary(ctx context.Context, t *testing.T) (gh *server.GRPCHandler, tearDownFunc func()) { + t.Helper() + + m := test.RunCommand(t) + gh = server.NewGRPCHandler(m.API) + + // grouper + grouper := m.MustCreateIndex(t, "grouper", pilosa.IndexOptions{Keys: false, TrackExistence: true}) + m.MustCreateField(t, grouper.Name(), "color", pilosa.OptFieldKeys()) + for id, color := range map[int]string{ + 1: "blue", + 2: "blue", + 5: "blue", + 6: "blue", + 7: "blue", + 3: "red", + 8: "red", + 9: "red", + 10: "red", + 4: "green", + } { + if _, err := gh.QueryPQLUnary(ctx, &pb.QueryPQLRequest{ + Index: grouper.Name(), + Pql: fmt.Sprintf(`Set(%d, color="%s")`, id, color), + }); err != nil { + t.Fatal(err) + } + } + m.MustCreateField(t, grouper.Name(), "score", pilosa.OptFieldTypeInt(-1000, 1000)) + for id, score := range map[int]int{ + 1: -10, + 2: -8, + 3: 6, + 4: 0, + 5: -2, + 6: 100, + 7: 0, + 8: -13, + 9: 80, + 10: -2, + } { + if _, err := gh.QueryPQLUnary(ctx, &pb.QueryPQLRequest{ + Index: grouper.Name(), + Pql: fmt.Sprintf(`Set(%d, score=%d)`, id, score), + }); err != nil { + t.Fatal(err) + } + } + m.MustCreateField(t, grouper.Name(), "age", pilosa.OptFieldTypeInt(0, 100)) + for id, age := range map[int]int{ + 2: 16, + 5: 16, + 8: 16, + 9: 16, + 3: 19, + 1: 27, + 4: 27, + 7: 27, + 10: 31, + 6: 34, + } { + if _, err := gh.QueryPQLUnary(ctx, &pb.QueryPQLRequest{ + Index: grouper.Name(), + Pql: fmt.Sprintf(`Set(%d, age=%d)`, id, age), + }); err != nil { + t.Fatal(err) + } + } + + m.MustCreateField(t, grouper.Name(), "height", pilosa.OptFieldTypeInt(0, 1000)) + for id, height := range map[int]int{ + 1: 20, + 2: 30, + 3: 40, + 4: 50, + 5: 60, + 6: 70, + 7: 80, + 8: 90, + 9: 100, + 10: 110, + } { + if _, err := gh.QueryPQLUnary(ctx, &pb.QueryPQLRequest{ + Index: grouper.Name(), + Pql: fmt.Sprintf(`Set(%d, height=%d)`, id, height), + }); err != nil { + t.Fatal(err) + } + } + + // joiner + joiner := m.MustCreateIndex(t, "joiner", pilosa.IndexOptions{TrackExistence: true}) + m.MustCreateField(t, joiner.Name(), "grouperid", pilosa.OptFieldTypeInt(0, 1000), pilosa.OptFieldForeignIndex(grouper.Name())) + m.MustCreateField(t, joiner.Name(), "jointype", pilosa.OptFieldTypeInt(-1000, 1000)) + for id, grouperid := range map[int]int{ + 1: 1, + 2: 2, + 3: 5, + 4: 6, + 5: 7, + 6: 3, + 7: 8, + 8: 9, + 9: 1, + 10: 2, + } { + if _, err := gh.QueryPQLUnary(ctx, &pb.QueryPQLRequest{ + Index: joiner.Name(), + Pql: fmt.Sprintf(`Set(%d, grouperid=%d)`, id, grouperid), + }); err != nil { + t.Fatal(err) + } + } + for id, jointype := range map[int]int{ + 1: 1, + 2: 1, + 3: 1, + 4: 1, + 5: 1, + 6: 2, + 7: 2, + 8: 2, + 9: 3, + 10: 3, + } { + if _, err := gh.QueryPQLUnary(ctx, &pb.QueryPQLRequest{ + Index: joiner.Name(), + Pql: fmt.Sprintf(`Set(%d, jointype=%d)`, id, jointype), + }); err != nil { + t.Fatal(err) + } + } + + return gh, func() { + if err := m.API.DeleteIndex(ctx, joiner.Name()); err != nil { + panic(err) + } + if err := m.API.DeleteIndex(ctx, grouper.Name()); err != nil { + panic(err) + } + if err := m.Close(); err != nil { + panic(err) + } + } +} + +func toTableResponse(resp *pb.TableResponse) tableResponse { + tr := tableResponse{ + headers: make([]columnInfo, len(resp.Headers)), + rows: make([]row, len(resp.Rows)), + } + + for i, h := range resp.Headers { + tr.headers[i] = columnInfo{ + name: h.Name, + datatype: h.Datatype, + } + } + + for i, r := range resp.Rows { + tr.rows[i].columns = make([]columnResponse, len(r.Columns)) + for j, c := range r.Columns { + + switch v := c.GetColumnVal().(type) { + case *pb.ColumnResponse_StringVal: + tr.rows[i].columns[j] = v.StringVal + case *pb.ColumnResponse_Uint64Val: + tr.rows[i].columns[j] = v.Uint64Val + case *pb.ColumnResponse_Int64Val: + tr.rows[i].columns[j] = v.Int64Val + case *pb.ColumnResponse_BoolVal: + tr.rows[i].columns[j] = v.BoolVal + case *pb.ColumnResponse_BlobVal: + tr.rows[i].columns[j] = v.BlobVal + case *pb.ColumnResponse_Uint64ArrayVal: + tr.rows[i].columns[j] = v.Uint64ArrayVal.Vals + case *pb.ColumnResponse_StringArrayVal: + tr.rows[i].columns[j] = v.StringArrayVal.Vals + case *pb.ColumnResponse_Float64Val: + tr.rows[i].columns[j] = v.Float64Val + case *pb.ColumnResponse_DecimalVal: + tr.rows[i].columns[j] = pql.NewDecimal(v.DecimalVal.Value, v.DecimalVal.Scale) + default: + tr.rows[i].columns[j] = nil + } + } + } + + return tr +} + +func equal(exp tableResponse, got tableResponse) error { + if !reflect.DeepEqual(exp, got) { + return fmt.Errorf("got: %+v %[1]T, but expected: %+v", got, exp) + } + return nil +} + +func equalUnordered(exp tableResponse, got tableResponse) error { + if len(exp.headers) != len(got.headers) || !reflect.DeepEqual(exp.headers, got.headers) { + return fmt.Errorf("header does not match: got %+v, but expected %+v", got.headers, exp.headers) + } + + if len(exp.rows) != len(got.rows) { + return fmt.Errorf("rows count does not match: got %+v, but expected %+v", len(got.rows), len(exp.rows)) + } + for _, er := range exp.rows { + for j, gr := range got.rows { + if reflect.DeepEqual(er.columns, gr.columns) { + got.rows[j] = got.rows[len(got.rows)-1] + got.rows = got.rows[:len(got.rows)-1] + break + } + } + } + if len(got.rows) > 0 { + return fmt.Errorf("got incorrect rows: %+v", got.rows) + } + return nil +} diff --git a/sql/column.go b/sql/column.go new file mode 100644 index 000000000..1bf8f8885 --- /dev/null +++ b/sql/column.go @@ -0,0 +1,112 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package sql + +import ( + "strings" + "time" +) + +const ColID = "_id" + +type FuncName string + +const ( + FuncCount FuncName = "count" + FuncMin FuncName = "min" + FuncMax FuncName = "max" + FuncSum FuncName = "sum" + FuncAvg FuncName = "avg" +) + +// Column is an interface which supports mapping to +// source headers, and aliasing column names. +// Alias() should always return a value; +// either a unique alias, or the same value +// return by Name(), but never an empty string. +type Column interface { + Source() string + Name() string + Alias() string +} + +type BasicColumn struct { + source string + name string + alias string +} + +func NewBasicColumn(s, n, a string) *BasicColumn { + return &BasicColumn{ + source: s, + name: n, + alias: a, + } +} + +func (b *BasicColumn) Source() string { + return b.source +} +func (b *BasicColumn) Name() string { + return b.name +} +func (b *BasicColumn) Alias() string { + if b.alias != "" { + return b.alias + } + return b.name +} + +type StarColumn struct{} + +func NewStarColumn() *StarColumn { + return &StarColumn{} +} + +func (s *StarColumn) Source() string { + return "" +} +func (s *StarColumn) Name() string { + return "" +} +func (s *StarColumn) Alias() string { + return "" +} + +func ConvertToTime(text string) (time.Time, bool) { + timeFormats := []string{ + "2006-01-02 15:04:05", + "2006-01-02 15:04", + "2006-01-02", + } + for _, timeFormat := range timeFormats { + if t, err := time.Parse(timeFormat, text); err == nil { + return t, true + } + } + return time.Now(), false +} + +func ExtractFieldName(columName string) (fieldName string, isSpecial bool) { + isSpecial = false + fieldName = columName + if strings.HasPrefix(columName, "_") { + isSpecial = true + if strings.HasSuffix(columName, "_time") { + fieldName = columName[1 : len(columName)-5] + } + } + return +} diff --git a/sql/extract.go b/sql/extract.go new file mode 100644 index 000000000..8ea1c00cf --- /dev/null +++ b/sql/extract.go @@ -0,0 +1,1177 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package sql + +import ( + "fmt" + "reflect" + "strconv" + "strings" + + "github.com/pilosa/pilosa/v2" + "github.com/pilosa/pilosa/v2/pql" + "github.com/pkg/errors" + "vitess.io/vitess/go/vt/sqlparser" +) + +// parseColumn is a column parsed from a sql query. Its qualifier +// value should map to either the name or alias of a parseTable. +type parseColumn struct { + name string + qualifier string +} + +// parseTable is a table parsed from a sql query. It includes +// its name and alias, along with a boolean indicating whether +// the table is the primary side of a join statement (i.e. it +// refers to the Pilosa column _id). +type parseTable struct { + name string + alias string + primary bool // indicates the side of the join representing the column _id + column *parseColumn // contains the related column from the ON clause + index *pilosa.Index // the pilosa index related to this table +} + +// parseTables is a slice of parseTable parsed from a single sql query. +type parseTables []*parseTable + +// byName returns the parseTable from the slice which matches on name. +// If there is no match it returns nil. +func (j parseTables) byName(n string) *parseTable { + for i := range j { + if j[i].name == n { + return j[i] + } + } + return nil +} + +// byName returns the parseTable from the slice which matches on alias. +// If there is no match it returns nil. +func (j parseTables) byAlias(a string) *parseTable { + for i := range j { + if j[i].alias == a { + return j[i] + } + } + return nil +} + +// primary returns the primary parseTable from the slice. +// In order for this to be useful, it is assumed that +// joinTables contains exactly two joinTable pointers +// (a primary and a secondary). +func (j parseTables) primary() *parseTable { + for i := range j { + if j[i].primary { + return j[i] + } + } + return nil +} + +// secondary returns the secondary parseTable from the slice. +func (j parseTables) secondary() *parseTable { + for i := range j { + if !j[i].primary { + return j[i] + } + } + return nil +} + +// tableWhere represents a parseTable from a sql query along +// with the portion of the where clause that relates to +// that table. For example, if a sql query had: +// from tbl1, tbl2 +// where tbl1.field1=1 and tbl2.field2=2 +// then each table would have a separate tableWhere object +// with the where made up of only the field with matching qualifier. +type tableWhere struct { + table *parseTable + where string +} + +// tableWheres is a slice of tableWhere. +type tableWheres []*tableWhere + +// extractParseTable returns a parseTable for the sqlparser.TableExpr. +func extractParseTable(tableExpr sqlparser.TableExpr) (*parseTable, error) { + switch tbl := tableExpr.(type) { + case *sqlparser.AliasedTableExpr: + tableName := tbl.Expr.(sqlparser.TableName).ToViewName().Name.String() + alias := tbl.As.String() + if alias == "" { + alias = tableName + } + return &parseTable{ + name: tableName, + alias: alias, + }, nil + } + + return nil, errors.New("unsupported table expression") +} + +func extractSelectFields(index *pilosa.Index, stmt *sqlparser.Select) ([]Column, selectFeatures, error) { + columns := []Column{} + features := selectFeatures{} + for _, item := range stmt.SelectExprs { + switch expr := item.(type) { + case *sqlparser.AliasedExpr: + var column Column + var alias string = expr.As.String() + switch colExpr := expr.Expr.(type) { + case *sqlparser.ColName: + fieldName := colExpr.Name.String() + + if fieldName == ColID { + if index.Options().Keys { + column = NewKeyIndexColumn(index, alias) + } else { + column = NewIDIndexColumn(index, alias) + } + } else { + column = NewFieldColumn(index.Field(fieldName), alias) + } + case *sqlparser.FuncExpr: + funcName := FuncName(strings.ToLower(colExpr.Name.String())) + var field *pilosa.Field + + if len(colExpr.Exprs) != 1 { + return nil, features, errors.New("function should have a single argument") + } + + switch expr := colExpr.Exprs[0].(type) { + case *sqlparser.AliasedExpr: + if colExpr, ok := expr.Expr.(*sqlparser.ColName); ok { + fieldName := colExpr.Name.String() + field = index.Field(fieldName) + } else { + return nil, features, errors.New("table name is required") + } + case *sqlparser.StarExpr: + // We don't currently track this; it either has a field or doesn't. + default: + return nil, features, errors.New("table name is required") + } + + switch funcName { + case FuncCount, FuncMin, FuncMax, FuncSum, FuncAvg: + column = NewFuncColumn(funcName, field, alias) + default: + return nil, features, fmt.Errorf("unknown function: %s", funcName) + } + features.funcs = append(features.funcs, selectFunc{ + funcName: funcName, + field: field, + }) + default: + return nil, features, errors.New("table name is required") + } + columns = append(columns, column) + case *sqlparser.StarExpr: + columns = append(columns, NewStarColumn()) + default: + return nil, features, errors.New("only column names or * are supported in select") + } + } + return columns, features, nil +} + +func extractIndexName(stmt *sqlparser.Select) (string, error) { + if len(stmt.From) != 1 { + return "", errors.New("selecting from multiple tables is not supported") + } + + fromExpr := stmt.From[0] + switch from := fromExpr.(type) { + case *sqlparser.AliasedTableExpr: + indexName := from.Expr.(sqlparser.TableName).ToViewName().Name.String() + return indexName, nil + } + + return "", errors.New("unsupported from clause") +} + +// extractParseColumn returns a parseColumn for the sqlparser.ColName. +func extractParseColumn(col *sqlparser.ColName) (*parseColumn, error) { + colName := col.Name.String() + qualifier := col.Qualifier.ToViewName().Name.String() + + return &parseColumn{ + name: colName, + qualifier: qualifier, + }, nil +} + +func extractWhere(index *pilosa.Index, expr sqlparser.Expr) (string, error) { + switch e := expr.(type) { + case *sqlparser.ComparisonExpr: + parseCol, op, val, err := extractComparison(e) + if err != nil { + return "", err + } + + if parseCol.name == "_id" { + switch op { + case "=": + return ConstRow(val), nil + case "in": + switch valExpr := val.(type) { + case []interface{}: + return ConstRow(valExpr...), nil + } + } + } + + field := index.Field(parseCol.name) + if field == nil { + return "", pilosa.ErrFieldNotFound + } + + switch field.Type() { + case pilosa.FieldTypeInt, pilosa.FieldTypeDecimal: + switch op { + case "=": + return Equals(field.Name(), val), nil + case "<": + return LT(field.Name(), val), nil + case "<=": + return LTE(field.Name(), val), nil + case ">": + return GT(field.Name(), val), nil + case ">=": + return GTE(field.Name(), val), nil + case "<>", "!=": + return NotEquals(field.Name(), val), nil + } + default: + switch op { + case "=": + return Row(field.Name(), val) + case "in": + var qs []string + switch valExpr := val.(type) { + case []interface{}: + for _, v := range valExpr { + q, err := Row(field.Name(), v) + if err != nil { + return "", err + } + qs = append(qs, q) + } + return Union(qs...), nil + default: + return "", fmt.Errorf("in operator expects `[]interface{}` but got: %T", valExpr) + } + case "like": + sval, ok := val.(string) + if !ok { + return "", fmt.Errorf("like operator expects `string` but got: %T", val) + } + return Like(field.Name(), sval), nil + } + } + case *sqlparser.AndExpr: + pql, err := extractWhereDateRange(index, e.Left, e.Right) + if err == nil { + return pql, err + } + left, err := extractWhere(index, e.Left) + if err != nil { + return "", err + } + right, err := extractWhere(index, e.Right) + if err != nil { + return "", err + } + return Intersect(left, right), nil + case *sqlparser.OrExpr: + left, err := extractWhere(index, e.Left) + if err != nil { + return "", err + } + right, err := extractWhere(index, e.Right) + if err != nil { + return "", err + } + return Union(left, right), nil + case *sqlparser.NotExpr: + expr, err := extractWhere(index, e.Expr) + if err != nil { + return "", err + } + return Not(expr), nil + case *sqlparser.ParenExpr: + expr, err := extractWhere(index, e.Expr) + if err != nil { + return "", err + } + return expr, nil + case *sqlparser.RangeCond: + if e.Operator != "between" { + return "", errors.New("only between is supported") + } + left, ok := e.Left.(*sqlparser.ColName) + if !ok { + return "", errors.New("left operand must be a column name") + } + columnName := left.Name.String() + fieldName, isSpecial := ExtractFieldName(columnName) + if isSpecial { + return "", errors.New("special fields are not allowed here") + } + field := index.Field(fieldName) + + switch field.Type() { + case pilosa.FieldTypeInt: + fromNum, err := extractInt(e.From) + if err != nil { + return "", err + } + toNum, err := extractInt(e.To) + if err != nil { + return "", err + } + return Between(field.Name(), fromNum, toNum), nil + case pilosa.FieldTypeDecimal: + fromNum, err := extractFloat(e.From) + if err != nil { + return "", err + } + toNum, err := extractFloat(e.To) + if err != nil { + return "", err + } + return Between(field.Name(), fromNum, toNum), nil + default: + return "", errors.New("only int and float64 fields are supported") + } + case *sqlparser.IsExpr: + left, ok := e.Expr.(*sqlparser.ColName) + if !ok { + return "", errors.New("left operand must be a column name") + } + field := index.Field(left.Name.String()) + if field.Type() == pilosa.FieldTypeInt { + if e.Operator == "is not null" { + return NotNull(field.Name()), nil + } + return "", errors.New("only `is not null` is supported for int fields") + } + return "", errors.New("`is` expression is supported only for int fields") + } + return "", errors.New("cannot extract where") +} + +func extractWhereDateRange(index *pilosa.Index, leftExpr sqlparser.Expr, rightExpr sqlparser.Expr) (string, error) { + var fieldName string + var fromExpr sqlparser.Expr + var toExpr sqlparser.Expr + var isSpecial bool + var idOrKey interface{} + + // Where part is expected to be in one of the following forms: + // where FIELD=VALUE and FIELD between DATETIME_FORMAT and DATETIME_FORMAT + // Or: + // where FIELD between DATETIME_FORMAT and DATETIME_FORMAT and FIELD=VALUE + + extract := func(compExpr *sqlparser.ComparisonExpr, rangeCond *sqlparser.RangeCond) error { + parseCol, op, val, err := extractComparison(compExpr) + if err != nil { + return err + } + idOrKey = val + if op != "=" { + // only = operator can exist here. + return errors.New("only = operator can exist here") + } + fieldName, isSpecial = ExtractFieldName(parseCol.name) + if isSpecial { + // column name cannot be special here. + return errors.New("column name cannot be special here") + } + if rangeCond.Operator != "between" { + // only between is accepted at this point. + return errors.New("only between is accepted at this point") + } + condFieldName, _ := ExtractFieldName(rangeCond.Left.(*sqlparser.ColName).Name.String()) + if condFieldName != fieldName { + // the field names in both sides should be the same, otherwise reject. + return errors.New("the field names in both sides should be the same, otherwise reject") + } + fromExpr = rangeCond.From + toExpr = rangeCond.To + return nil + } + + if left, ok := leftExpr.(*sqlparser.ComparisonExpr); ok { + // if FIELD=VALUE part is on the left, and range part is on the right + if right, ok := rightExpr.(*sqlparser.RangeCond); ok { + if err := extract(left, right); err != nil { + return "", err + } + } else { + return "", errors.New("right is not a range cond") + } + } else if right, ok := rightExpr.(*sqlparser.ComparisonExpr); ok { + // if FIELD=VALUE part is on the right, and left part is on the left + if left, ok := leftExpr.(*sqlparser.RangeCond); ok { + if err := extract(right, left); err != nil { + return "", err + } + } else { + return "", errors.New("left is not a range cond") + } + } + + fromStr, err := extractStr(fromExpr) + if err != nil { + return "", err + } + toStr, err := extractStr(toExpr) + if err != nil { + return "", err + } + // try to convert `from` and `to` to time + fromTime, ok := ConvertToTime(fromStr) + if !ok { + return "", errors.New("from operand must be in the correct time format") + } + toTime, ok := ConvertToTime(toStr) + if !ok { + return "", errors.New("to operand must be in the correct time format") + } + + field := index.Field(fieldName) + return RowRange(field.Name(), idOrKey, fromTime, toTime) +} + +func extractVal(e sqlparser.Expr) (interface{}, error) { + val, ok := e.(*sqlparser.SQLVal) + if !ok { + return nil, errors.New("expression must be a value") + } + switch val.Type { + case sqlparser.StrVal: + return string(val.Val), nil + case sqlparser.IntVal: + value, err := strconv.Atoi(string(val.Val)) + if err != nil { + return nil, err + } + return value, nil + case sqlparser.FloatVal: + value, err := strconv.ParseFloat(string(val.Val), 64) + if err != nil { + return nil, err + } + return value, nil + default: + return nil, fmt.Errorf("unknown type: %d", val.Type) + } +} + +func extractInt(e sqlparser.Expr) (int, error) { + val, err := extractVal(e) + if err != nil { + return 0, err + } + num, ok := val.(int) + if !ok { + return 0, errors.New("value must be an integer") + } + return num, nil +} + +func extractFloat(e sqlparser.Expr) (float64, error) { + val, err := extractVal(e) + if err != nil { + return 0, err + } + + var num float64 + switch v := val.(type) { + case int: + num = float64(v) + case float64: + num = v + default: + return 0, errors.New("value must be convertable to a float64") + } + return num, nil +} + +func extractStr(e sqlparser.Expr) (string, error) { + val, err := extractVal(e) + if err != nil { + return "", err + } + s, ok := val.(string) + if !ok { + return "", errors.New("value must be a string") + } + return s, nil +} + +func extractTuple(tuple sqlparser.ValTuple) ([]interface{}, error) { + result := []interface{}{} + for _, item := range tuple { + if val, ok := item.(*sqlparser.SQLVal); ok { + v, err := extractVal(val) + if err != nil { + return nil, err + } + result = append(result, v) + } else { + return nil, errors.New("tuple should contain integers or strings") + } + } + return result, nil +} + +func extractComparison(expr *sqlparser.ComparisonExpr) (col *parseColumn, op string, value interface{}, err error) { + op = expr.Operator + if colExpr, ok := expr.Left.(*sqlparser.ColName); ok { + col, err = extractParseColumn(colExpr) + if err != nil { + return + } + if op == "in" { + switch valExpr := expr.Right.(type) { + case sqlparser.ValTuple: + value, err = extractTuple(valExpr) + + default: + err = fmt.Errorf("in operator excepts only a tuple or a query, received `%s`", + reflect.TypeOf(valExpr).String()) + return + } + } else { + value, err = extractVal(expr.Right) + } + if err != nil { + return + } + } else { + if colExpr, ok := expr.Right.(*sqlparser.ColName); ok { + col, err = extractParseColumn(colExpr) + if err != nil { + return + } + if op == "in" { + value, err = extractTuple(expr.Left.(sqlparser.ValTuple)) + } else { + value, err = extractVal(expr.Left) + } + if err != nil { + return + } + } else { + err = errors.New("either left or right operand should be a column name") + } + } + return +} + +func extractLimitOffset(stmt *sqlparser.Select) (uint, uint, error) { + if stmt.Limit == nil { + return 100, 0, nil + } + var offset, limit uint + if offsetExpr, ok := stmt.Limit.Offset.(*sqlparser.SQLVal); ok { + val, err := extractVal(offsetExpr) + if err != nil { + return 0, 0, err + } + if offsetVal, ok := val.(int); ok { + offset = uint(offsetVal) + } else { + return 0, 0, errors.New("offset must be an integer") + } + } + if limitExpr, ok := stmt.Limit.Rowcount.(*sqlparser.SQLVal); ok { + val, err := extractVal(limitExpr) + if err != nil { + return 0, 0, err + } + if limitVal, ok := val.(int); ok { + limit = uint(limitVal) + } else { + return 0, 0, errors.New("limit must be an integer") + } + } + return limit, offset, nil +} + +// extractOrderBy returns the order by fields and directions (asc/desc) +// as separate string slices. +func extractOrderBy(stmt *sqlparser.Select) ([]string, []string, error) { + var flds []string + var dirs []string + + for _, item := range stmt.OrderBy { + switch colExpr := item.Expr.(type) { + case *sqlparser.ColName: + colName := colExpr.Name.String() + flds = append(flds, colName) + dirs = append(dirs, item.Direction) + } + } + + return flds, dirs, nil +} + +func extractTableNames(stmt *sqlparser.Select) ([]string, error) { + if len(stmt.From) == 0 { + return []string{}, nil + } + + switch from := stmt.From[0].(type) { + case *sqlparser.AliasedTableExpr: + tableName := from.Expr.(sqlparser.TableName).ToViewName().Name.String() + return []string{tableName}, nil + case *sqlparser.JoinTableExpr: + ret := []string{} + switch left := from.LeftExpr.(type) { + case *sqlparser.AliasedTableExpr: + leftTableName := left.Expr.(sqlparser.TableName).ToViewName().Name.String() + ret = append(ret, leftTableName) + } + switch right := from.RightExpr.(type) { + case *sqlparser.AliasedTableExpr: + rightTableName := right.Expr.(sqlparser.TableName).ToViewName().Name.String() + ret = append(ret, rightTableName) + } + return ret, nil + } + + return []string{}, nil +} + +func extractGroupByFieldNames(stmt sqlparser.GroupBy) ([]string, error) { + fields := make([]string, len(stmt)) + for i, item := range stmt { + col, ok := item.(*sqlparser.ColName) + if !ok { + return nil, errors.New("group by accepts columns") + } + fields[i] = col.Name.String() + } + return fields, nil +} + +func extractHavingClause(stmt *sqlparser.Where) (*HavingClause, error) { + if stmt == nil { + return nil, nil + } + if stmt.Type != "having" { + return nil, fmt.Errorf("invalid having type: %s", stmt.Type) + } + + hc := &HavingClause{} + + switch having := stmt.Expr.(type) { + case *sqlparser.RangeCond: + if having.Operator != "between" { + return nil, errors.New("only between is supported") + } + hc.Subj = having.Left.(*sqlparser.ColName).Name.String() + hc.Cond.Op = pql.BETWEEN + fromPred, err := extractInt(having.From) + if err != nil { + return nil, err + } + toPred, err := extractInt(having.To) + if err != nil { + return nil, err + } + vals := make([]interface{}, 2) + switch hc.Subj { + case "count": + vals[0] = uint64(fromPred) + vals[1] = uint64(toPred) + case "sum": + vals[0] = int64(fromPred) + vals[1] = int64(toPred) + } + hc.Cond.Value = vals + return hc, nil + case *sqlparser.AndExpr: + left := having.Left.(*sqlparser.ComparisonExpr) + right := having.Right.(*sqlparser.ComparisonExpr) + leftName := left.Left.(*sqlparser.ColName).Name.String() + rightName := right.Left.(*sqlparser.ColName).Name.String() + if leftName != rightName { + return nil, fmt.Errorf("having comparitors do not match: %s, %s", leftName, rightName) + } + hc.Subj = leftName + + leftOp := extractComparisonOp(left) + rightOp := extractComparisonOp(right) + + leftPred, err := extractInt(left.Right) + if err != nil { + return nil, err + } + rightPred, err := extractInt(right.Right) + if err != nil { + return nil, err + } + + intVals := make([]int, 2) + if leftOp == pql.GT && rightOp == pql.LT { + intVals[0] = leftPred + intVals[1] = rightPred + hc.Cond.Op = pql.BTWN_LT_LT + } else if leftOp == pql.GT && rightOp == pql.LTE { + intVals[0] = leftPred + intVals[1] = rightPred + hc.Cond.Op = pql.BTWN_LT_LTE + } else if leftOp == pql.GTE && rightOp == pql.LT { + intVals[0] = leftPred + intVals[1] = rightPred + hc.Cond.Op = pql.BTWN_LTE_LT + } else if leftOp == pql.GTE && rightOp == pql.LTE { + intVals[0] = leftPred + intVals[1] = rightPred + hc.Cond.Op = pql.BETWEEN + } else if leftOp == pql.LT && rightOp == pql.GT { + intVals[0] = rightPred + intVals[1] = leftPred + hc.Cond.Op = pql.BTWN_LT_LT + } else if leftOp == pql.LT && rightOp == pql.GTE { + intVals[0] = rightPred + intVals[1] = leftPred + hc.Cond.Op = pql.BTWN_LTE_LT + } else if leftOp == pql.LTE && rightOp == pql.GT { + intVals[0] = rightPred + intVals[1] = leftPred + hc.Cond.Op = pql.BTWN_LT_LTE + } else if leftOp == pql.LTE && rightOp == pql.GTE { + intVals[0] = rightPred + intVals[1] = leftPred + hc.Cond.Op = pql.BETWEEN + } + + vals := make([]interface{}, 2) + switch hc.Subj { + case "count": + vals[0] = uint64(intVals[0]) + vals[1] = uint64(intVals[1]) + case "sum": + vals[0] = int64(intVals[0]) + vals[1] = int64(intVals[1]) + } + hc.Cond.Value = vals + + return hc, nil + case *sqlparser.ComparisonExpr: + hc.Subj = having.Left.(*sqlparser.ColName).Name.String() + hc.Cond.Op = extractComparisonOp(having) + switch hc.Subj { + case "count": + pred, err := extractInt(having.Right) + if err != nil { + return nil, err + } + hc.Cond.Value = uint64(pred) + case "sum": + pred, err := extractInt(having.Right) + if err != nil { + return nil, err + } + hc.Cond.Value = int64(pred) + } + return hc, nil + } + + return nil, errors.New("unsupported having clause") +} + +func extractComparisonOp(expr *sqlparser.ComparisonExpr) pql.Token { + switch expr.Operator { + case "==": + return pql.EQ + case "!=": + return pql.NEQ + case "<": + return pql.LT + case "<=": + return pql.LTE + case ">": + return pql.GT + case ">=": + return pql.GTE + } + return pql.ILLEGAL +} + +// extractJoinTables returns a slice of parseTable containing two +// items, the primary and secondary join tables. This function does +// not extract join tables of the form: +// from tbl1, tbl2 +// The from clause must be of the form: +// from tbl1 INNER JOIN tbl2 ON ... +// +func extractJoinTables(stmt *sqlparser.Select) (parseTables, error) { + if len(stmt.From) != 1 { + return nil, errors.New("selecting from multiple tables is not supported") + } + + tbls := make([]*parseTable, 2) + + from, ok := stmt.From[0].(*sqlparser.JoinTableExpr) + if !ok { + return nil, errors.New("unsupported join clause") + } + + leftTable, err := extractParseTable(from.LeftExpr) + if err != nil { + return nil, errors.Wrap(err, "extracting left join table") + } + rightTable, err := extractParseTable(from.RightExpr) + if err != nil { + return nil, errors.Wrap(err, "extracting right join table") + } + + // It is not important which table goes in which tbls position; + // the primary/secondary table will be determined later. + tbls[0] = leftTable + tbls[1] = rightTable + + // Get the ON condition and determine which table is primary. + switch onCond := from.Condition.On.(type) { + case *sqlparser.ComparisonExpr: + if onCond.Operator != "=" { + return nil, errors.Errorf("unsupported on condition comparison type: %s", onCond.Operator) + } + // get left ColName + left, ok := onCond.Left.(*sqlparser.ColName) + if !ok { + return nil, errors.New("left join operand must be a column name") + } + leftJoinCol, err := extractParseColumn(left) + if err != nil { + return nil, errors.Wrap(err, "extracting left join column") + } + // get right ColName + right, ok := onCond.Right.(*sqlparser.ColName) + if !ok { + return nil, errors.New("right join operand must be a column name") + } + rightJoinCol, err := extractParseColumn(right) + if err != nil { + return nil, errors.Wrap(err, "extracting right join column") + } + + // The primary column is set as the column referencing the "_id" field. + var primaryColumn *parseColumn + if leftJoinCol.name == ColID && rightJoinCol.name != ColID { + primaryColumn = leftJoinCol + } else if leftJoinCol.name != ColID && rightJoinCol.name == ColID { + primaryColumn = rightJoinCol + } else { + return nil, errors.Errorf("exactly one join column must be %s, have: %s, %s", ColID, leftJoinCol.name, rightJoinCol.name) + } + + // populate the joinTable column and primary fields + for _, jc := range []*parseColumn{leftJoinCol, rightJoinCol} { + var found bool + for i := range tbls { + if tbls[i].alias == jc.qualifier { + tbls[i].column = jc + if jc == primaryColumn { + tbls[i].primary = true + } + found = true + } + } + if !found { + return nil, errors.Errorf("no tables match qualifier: %s", jc.qualifier) + } + } + + default: + return nil, errors.Errorf("unsupported on condition type: %T", onCond) + } + + return tbls, nil +} + +// extractWheres returns the slice of tableWhere for the sql query. +func extractWheres(indexes []*pilosa.Index, tbls parseTables, expr sqlparser.Expr) (tableWheres, error) { + wheres := make([]*tableWhere, 0) + + // Set the index associated with each parseTable. + // TODO: may be able to move this to parseTables creation? + for _, idx := range indexes { + tbl := tbls.byName(idx.Name()) + if tbl == nil { + return nil, errors.Errorf("index not in parseTables: %s", idx.Name()) + } + tbl.index = idx + } + + // make a map of tbls alias to slice index. + m := make(map[string]int) + for i, tbl := range tbls { + m[tbl.alias] = i + } + + switch e := expr.(type) { + case *sqlparser.ComparisonExpr: + pCol, op, val, err := extractComparison(e) + if err != nil { + return nil, err + } + + pTable := tbls.byAlias(pCol.qualifier) + if pTable == nil { + return nil, errors.Errorf("no index for qaulifier: %s", pCol.qualifier) + } else if pTable.index == nil { + return nil, errors.Errorf("parse table has no index: %s", pTable.name) + } + + field := pTable.index.Field(pCol.name) + + tw := &tableWhere{ + table: pTable, + } + + if field.Type() == pilosa.FieldTypeInt { + num, ok := val.(int) + if !ok { + return nil, errors.New("right operand must be a number") + } + switch op { + case "=": + tw.where = Equals(field.Name(), num) + case "<": + tw.where = LT(field.Name(), num) + case "<=": + tw.where = LTE(field.Name(), num) + case ">": + tw.where = GT(field.Name(), num) + case ">=": + tw.where = GTE(field.Name(), num) + case "<>": + fallthrough + case "!=": + tw.where = NotEquals(field.Name(), num) + } + return append(wheres, tw), nil + } + if op == "=" { + if tw.where, err = Row(field.Name(), val); err != nil { + return nil, err + } + return append(wheres, tw), nil + } + + if op == "in" { + var qs []string + switch valExpr := val.(type) { + case []interface{}: + for _, v := range valExpr { + r, e := Row(field.Name(), v) + if e != nil { + return nil, errors.Wrap(err, "extracting where statements") + } + qs = append(qs, r) + } + tw.where = Union(qs...) + return append(wheres, tw), nil + + default: + return nil, fmt.Errorf("in operator expects `[]interface{}` but got: %T", valExpr) + } + } + + case *sqlparser.AndExpr: + left, err := extractWheres(indexes, tbls, e.Left) + if err != nil { + return nil, err + } + right, err := extractWheres(indexes, tbls, e.Right) + if err != nil { + return nil, err + } + + // The following logic is used to build the where portion of the query + // related to each table. The goal is to return one or two tableWhere + // objects (either 0 or 1 for each table in the join). + if len(left) == 1 && len(right) == 1 && left[0].table == right[0].table { + // if left(1) and right(1) are from the same alias, + // then intersect them into left and return left(1) + left[0].where = Intersect(left[0].where, right[0].where) + return left, nil + } else if len(left) == 1 && len(right) == 1 && left[0].table != right[0].table { + // if left(1) and right(1) are NOT from the same alias, + // then return final(2) + return []*tableWhere{left[0], right[0]}, nil + } else if len(left) == 1 && len(right) == 2 { + // if left(1) and right(2) + // then intersect the 1's and return final(2) + if left[0].table == right[0].table { + left[0].where = Intersect(left[0].where, right[0].where) + return []*tableWhere{left[0], right[1]}, nil + } else if left[0].table == right[1].table { + left[0].where = Intersect(left[0].where, right[1].where) + return []*tableWhere{left[0], right[0]}, nil + } + return nil, errors.Errorf("no matching table on right: %s", left[0].table.name) + } else if len(left) == 1 && len(right) == 2 { + // if left(2) and right(1), + // then intersect the 1's and return final(2) + if right[0].table == left[0].table { + right[0].where = Intersect(right[0].where, left[0].where) + return []*tableWhere{left[0], right[1]}, nil + } else if right[0].table == left[1].table { + right[0].where = Intersect(right[0].where, left[1].where) + return []*tableWhere{left[0], right[0]}, nil + } + return nil, errors.Errorf("no matching table on left: %s", right[0].table.name) + } else if len(left) == 2 && len(right) == 2 { + // if left(2) and right(2) + // then intsect both and return final(2) + if left[0].table == right[0].table && left[1].table == right[1].table { + left[0].where = Intersect(left[0].where, right[0].where) + left[1].where = Intersect(left[1].where, right[1].where) + return left, nil + } else if left[0].table == right[1].table && left[1].table == right[0].table { + left[0].where = Intersect(left[0].where, right[1].where) + left[1].where = Intersect(left[1].where, right[0].where) + return left, nil + } + return nil, errors.Errorf("non-matching tables: %s/%s, %s/%s", left[0].table.name, left[1].table.name, right[0].table.name, right[1].table.name) + } + return nil, errors.Errorf("invalid table count; expected 1 or 2, but got: %d, %d", len(left), len(right)) + case *sqlparser.OrExpr: + left, err := extractWheres(indexes, tbls, e.Left) + if err != nil { + return nil, err + } + right, err := extractWheres(indexes, tbls, e.Right) + if err != nil { + return nil, err + } + + if len(left) == 1 && len(right) == 1 && left[0].table == right[0].table { + // if left(1) and right(1) are from the same alias, + // then union them and return final(1) + left[0].where = Union(left[0].where, right[0].where) + return left, nil + } + return nil, errors.Errorf("invalid table count; expected 1/1, but got: %d/%d", len(left), len(right)) + case *sqlparser.NotExpr: + expr, err := extractWheres(indexes, tbls, e.Expr) + if err != nil { + return nil, err + } + if len(expr) == 1 { + expr[0].where = Not(expr[0].where) + return expr, nil + } + return nil, errors.Errorf("not support a single expression. got: %d", len(expr)) + case *sqlparser.ParenExpr: + expr, err := extractWheres(indexes, tbls, e.Expr) + if err != nil { + return nil, err + } + if len(expr) == 1 { + return expr, nil + } + return nil, errors.Errorf("not support a single expression. got: %d", len(expr)) + case *sqlparser.RangeCond: + if e.Operator != "between" { + return nil, errors.New("only between is supported") + } + left, ok := e.Left.(*sqlparser.ColName) + if !ok { + return nil, errors.New("left operand must be a column name") + } + + pCol, err := extractParseColumn(left) + if err != nil { + return nil, errors.Wrap(err, "extracting parse column") + } + + pTable := tbls.byAlias(pCol.qualifier) + if pTable == nil { + return nil, errors.Errorf("no index for qaulifier: %s", pCol.qualifier) + } else if pTable.index == nil { + return nil, errors.Errorf("parse table has no index: %s", pTable.name) + } + + fieldName, isSpecial := ExtractFieldName(pCol.name) + if isSpecial { + return nil, errors.New("special fields are not allowed here") + } + + tw := &tableWhere{ + table: pTable, + } + + field := pTable.index.Field(fieldName) + if field.Type() == pilosa.FieldTypeInt { + fromNum, err := extractInt(e.From) + if err != nil { + return nil, err + } + toNum, err := extractInt(e.To) + if err != nil { + return nil, err + } + tw.where = Between(field.Name(), fromNum, toNum) + return append(wheres, tw), nil + } + return nil, errors.New("only int fields are supported") + case *sqlparser.IsExpr: + left, ok := e.Expr.(*sqlparser.ColName) + if !ok { + return nil, errors.New("left operand must be a column name") + } + + pCol, err := extractParseColumn(left) + if err != nil { + return nil, errors.Wrap(err, "extracting parse column") + } + + pTable := tbls.byAlias(pCol.qualifier) + if pTable == nil { + return nil, errors.Errorf("no index for qaulifier: %s", pCol.qualifier) + } else if pTable.index == nil { + return nil, errors.Errorf("parse table has no index: %s", pTable.name) + } + + tw := &tableWhere{ + table: pTable, + } + + field := pTable.index.Field(pCol.name) + if field.Type() == pilosa.FieldTypeInt { + if e.Operator == "is not null" { + tw.where = NotNull(field.Name()) + return append(wheres, tw), nil + } + return nil, errors.New("only `is not null` is supported for int fields") + } + return nil, errors.New("`is` expression is supported only for int fields") + } + return nil, errors.New("cannot extract where") +} diff --git a/sql/mapper.go b/sql/mapper.go new file mode 100644 index 000000000..a43dba06d --- /dev/null +++ b/sql/mapper.go @@ -0,0 +1,113 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package sql + +import ( + "strings" + + "github.com/pilosa/pilosa/v2/logger" + "github.com/pkg/errors" + "vitess.io/vitess/go/vt/sqlparser" +) + +const ( + SQLTypeSelect = "select" + SQLTypeShow = "show" + SQLTypeEmpty = "" +) + +// System errors. +var ( + ErrMultipleSQLStatements = errors.New("statement contains multiple sql queries") +) + +type Attributes map[string]interface{} + +type MappedSQL struct { + SQLType string + Statement sqlparser.Statement + Mask QueryMask + Tables []string +} + +// Mapper is responsible for mapping a SQL query to structure representation +type Mapper struct { + Logger logger.Logger +} + +func NewMapper() *Mapper { + return &Mapper{ + Logger: logger.NopLogger, + } +} + +// Parse parses SQL query +func (m *Mapper) Parse(sql string) (sqlparser.Statement, QueryMask, error) { + parsed, err := sqlparser.Parse(sql) + if err != nil { + return nil, QueryMask{}, errors.Wrap(err, "parsing sql") + } + + qm := GenerateMask(parsed) + + return parsed, qm, nil +} + +// MapSQL converts a sql string into a MappedSQL object, +// which includes the parsed query and the query mask, +// among other information about the query. +func (m *Mapper) MapSQL(sql string) (*MappedSQL, error) { + // In the case where `sql` contains more than one query—since + // we don't support multiple return sets—we're going to just + // ignore everything and return a specific error type. This + // will allow the caller to handle it as needed (i.e. it can + // return the error, or return an empty result set). + if parts := strings.Split(sql, ";"); len(parts) > 1 { + var partCount int + for _, part := range parts { + if trimmed := strings.TrimSpace(part); trimmed != "" && trimmed != "\x00" { + partCount++ + } + } + if partCount != 1 { + return nil, ErrMultipleSQLStatements + } + } + + stmt, qm, err := m.Parse(sql) + if err != nil { + return nil, errors.Wrap(err, "parsing sql") + } + + var sqlType string + var tableNames []string + switch slct := stmt.(type) { + case *sqlparser.Select: + sqlType = SQLTypeSelect + tableNames, err = extractTableNames(slct) + if err != nil { + return nil, errors.Wrap(err, "extracting table names") + } + case *sqlparser.Show: + sqlType = SQLTypeShow + } + + return &MappedSQL{ + SQLType: sqlType, + Statement: stmt, + Mask: qm, + Tables: tableNames, + }, nil +} diff --git a/sql/mapper_test.go b/sql/mapper_test.go new file mode 100644 index 000000000..216ffda4f --- /dev/null +++ b/sql/mapper_test.go @@ -0,0 +1,394 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package sql + +import ( + "fmt" + "testing" + + "vitess.io/vitess/go/vt/sqlparser" +) + +func TestParse(t *testing.T) { + tests := []struct { + sql string + expMask QueryMask + }{ + { + sql: "select _id from tbl", + expMask: QueryMask{ + SelectMask: SelectPartID, + FromMask: FromPartTable, + }, + }, + { + sql: "select * from tbl", + expMask: QueryMask{ + SelectMask: SelectPartStar, + FromMask: FromPartTable, + }, + }, + { + sql: "select fld from tbl", + expMask: QueryMask{ + SelectMask: SelectPartField, + FromMask: FromPartTable, + }, + }, + { + sql: "select fld1, fld2 from tbl", + expMask: QueryMask{ + SelectMask: SelectPartFields, + FromMask: FromPartTable, + }, + }, + { + sql: "select count(*) from tbl", + expMask: QueryMask{ + SelectMask: SelectPartCountStar, + FromMask: FromPartTable, + }, + }, + { + sql: "select min(fld) from tbl", + expMask: QueryMask{ + SelectMask: SelectPartMinField, + FromMask: FromPartTable, + }, + }, + { + sql: "select max(fld) from tbl", + expMask: QueryMask{ + SelectMask: SelectPartMaxField, + FromMask: FromPartTable, + }, + }, + { + sql: "select sum(fld) from tbl", + expMask: QueryMask{ + SelectMask: SelectPartSumField, + FromMask: FromPartTable, + }, + }, + { + sql: "select avg(fld) from tbl", + expMask: QueryMask{ + SelectMask: SelectPartAvgField, + FromMask: FromPartTable, + }, + }, + { + sql: "select _id, count(*) from tbl", + expMask: QueryMask{ + SelectMask: SelectPartID | SelectPartCountStar, + FromMask: FromPartTable, + }, + }, + { + sql: "select _id from tbl1, tbl2", + expMask: QueryMask{ + SelectMask: SelectPartID, + FromMask: FromPartTables, + }, + }, + { + sql: "select _id from tbl1 INNER JOIN tbl2", + expMask: QueryMask{ + SelectMask: SelectPartID, + FromMask: FromPartJoin, + }, + }, + { + sql: "select * from tbl where _id = 1", + expMask: QueryMask{ + SelectMask: SelectPartStar, + FromMask: FromPartTable, + WhereMask: WherePartIDCondition, + }, + }, + { + sql: "select * from tbl where fld = 1", + expMask: QueryMask{ + SelectMask: SelectPartStar, + FromMask: FromPartTable, + WhereMask: WherePartFieldCondition, + }, + }, + { + sql: "select * from tbl group by fld", + expMask: QueryMask{ + SelectMask: SelectPartStar, + FromMask: FromPartTable, + GroupByMask: GroupByPartField, + }, + }, + { + sql: "select * from tbl group by fld1, fld2", + expMask: QueryMask{ + SelectMask: SelectPartStar, + FromMask: FromPartTable, + GroupByMask: GroupByPartFields, + }, + }, + { + sql: "select fld, sum(fld) from tbl group by fld", + expMask: QueryMask{ + SelectMask: SelectPartField | SelectPartSumField, + FromMask: FromPartTable, + GroupByMask: GroupByPartField, + }, + }, + { + sql: "select fld, sum(fld) from tbl group by fld having sum > 10", + expMask: QueryMask{ + SelectMask: SelectPartField | SelectPartSumField, + FromMask: FromPartTable, + GroupByMask: GroupByPartField, + HavingMask: HavingPartCondition, + }, + }, + { + sql: "select fld from tbl order by fld", + expMask: QueryMask{ + SelectMask: SelectPartField, + FromMask: FromPartTable, + OrderByMask: OrderByPartField, + }, + }, + { + sql: "select fld from tbl order by fld1, fld2", + expMask: QueryMask{ + SelectMask: SelectPartField, + FromMask: FromPartTable, + OrderByMask: OrderByPartFields, + }, + }, + { + sql: "select fld from tbl limit 10", + expMask: QueryMask{ + SelectMask: SelectPartField, + FromMask: FromPartTable, + LimitMask: LimitPartLimit, + }, + }, + { + sql: "select fld from tbl limit 10, 5", + expMask: QueryMask{ + SelectMask: SelectPartField, + FromMask: FromPartTable, + LimitMask: LimitPartLimit | LimitPartOffset, + }, + }, + { + sql: "select distinct fld from tbl", + expMask: QueryMask{ + SelectMask: SelectPartDistinct | SelectPartField, + FromMask: FromPartTable, + }, + }, + { + sql: "select count(*) from tbl where fld = 1", + expMask: QueryMask{ + SelectMask: SelectPartCountStar, + FromMask: FromPartTable, + WhereMask: WherePartFieldCondition, + }, + }, + { + sql: "select count(*) from tbl where fld1 = 1 and fld2 = 2", + expMask: QueryMask{ + SelectMask: SelectPartCountStar, + FromMask: FromPartTable, + WhereMask: WherePartMultiFieldCondition, + }, + }, + { + sql: "select count(*) from tbl where fld1 = 1 or fld2 = 2", + expMask: QueryMask{ + SelectMask: SelectPartCountStar, + FromMask: FromPartTable, + WhereMask: WherePartMultiFieldCondition, + }, + }, + { + sql: "select _id from tbl where not fld = 1 limit 10", + expMask: QueryMask{ + SelectMask: SelectPartID, + FromMask: FromPartTable, + WhereMask: WherePartFieldCondition, + LimitMask: LimitPartLimit, + }, + }, + { + sql: "select _id from tbl where fld between 1 and 3", + expMask: QueryMask{ + SelectMask: SelectPartID, + FromMask: FromPartTable, + WhereMask: WherePartFieldCondition, + }, + }, + { + sql: "select _id from tbl where fld1 between 1 and 3 and fld2 = 2", + expMask: QueryMask{ + SelectMask: SelectPartID, + FromMask: FromPartTable, + WhereMask: WherePartMultiFieldCondition, + }, + }, + { + sql: "select count(*) from tbl where fld is not null", + expMask: QueryMask{ + SelectMask: SelectPartCountStar, + FromMask: FromPartTable, + WhereMask: WherePartFieldCondition, + }, + }, + { + sql: "select fld, count(*) from tbl group by fld", + expMask: QueryMask{ + SelectMask: SelectPartField | SelectPartCountStar, + FromMask: FromPartTable, + GroupByMask: GroupByPartField, + }, + }, + { + sql: "select fld1, fld2, count(*) from grouper group by fld1, fld2", + expMask: QueryMask{ + SelectMask: SelectPartFields | SelectPartCountStar, + FromMask: FromPartTable, + GroupByMask: GroupByPartFields, + }, + }, + { + sql: "select fld1, fld2, count(*) from tbl where fld1 = 1 group by fld1, fld2 limit 1", + expMask: QueryMask{ + SelectMask: SelectPartFields | SelectPartCountStar, + FromMask: FromPartTable, + WhereMask: WherePartFieldCondition, + GroupByMask: GroupByPartFields, + LimitMask: LimitPartLimit, + }, + }, + { + sql: "select count(distinct fld) from tbl", + expMask: QueryMask{ + SelectMask: SelectPartCountDistinctField, + FromMask: FromPartTable, + }, + }, + { + sql: "select count(*) from tbl1 INNER JOIN tbl2 ON tbl1._id = tbl2.bsi", + expMask: QueryMask{ + SelectMask: SelectPartCountStar, + FromMask: FromPartJoin, + }, + }, + { + sql: "select count(*) from tbl1 INNER JOIN tbl2 ON tbl1._id = tbl2.bsi where tbl1.fld1 = 1 and tbl2.fld2 = 2", + expMask: QueryMask{ + SelectMask: SelectPartCountStar, + FromMask: FromPartJoin, + WhereMask: WherePartMultiFieldCondition, + }, + }, + } + + mapper := NewMapper() + for i, test := range tests { + t.Run(fmt.Sprintf("test-%d", i), func(t *testing.T) { + _, mask, err := mapper.Parse(test.sql) + if err != nil { + t.Fatal(err) + } + + if mask != test.expMask { + t.Fatalf("expected mask: %v, but got: %v", test.expMask, mask) + } + }) + } +} + +// This thing passes even if join is not implemented. +// TODO: make this actually a test +func TestSelectJoin(t *testing.T) { + tests := []struct { + sql string + }{ + { + // Count(Intersect(All(),Distinct(Row(grouperid!=null),index='joiner',field='grouperid'))) + sql: "select count(*) from grouper g INNER JOIN joiner j ON g._id = j.grouperid", + }, + { + // Intersect(All(),Distinct(Row(grouperid!=null),index='joiner',field='grouperid')) + sql: "select _id from grouper g INNER JOIN joiner j ON g._id = j.grouperid", + }, + { + // Intersect(Row(color='red'),Distinct(Row(grouperid!=null),index='joiner',field='grouperid')) + sql: "select _id from grouper g INNER JOIN joiner j ON g._id = j.grouperid where g.color = 'red'", + }, + { + // Intersect(All(),Distinct(Row(grouperid!=null),index='joiner',field='grouperid')) + sql: "select _id from grouper g INNER JOIN joiner j ON g._id = j.grouperid where g.color = 'red' and j.jointype = 2", + }, + } + + mapper := NewMapper() + for i, test := range tests { + t.Run(fmt.Sprintf("test-%d", i), func(t *testing.T) { + if m, err := mapper.MapSQL(test.sql); err != nil { + t.Fatal(err) + } else { + if stmt, ok := m.Statement.(*sqlparser.Select); !ok { + t.Fatalf("%s: expected Statement: sqlparser.Select, got %T", test.sql, m.Statement) + } else { + t.Logf("%+v\n", stmt) + } + } + }) + } +} + +func TestOrderBy(t *testing.T) { + tests := []struct { + sql string + }{ + { + sql: "select distinct score from grouper order by score asc", + }, + { + sql: "select distinct score from grouper order by score desc", + }, + { + sql: "select distinct score from grouper order by score asc limit 5", + }, + { + sql: "select distinct score from grouper order by score desc limit 5", + }, + } + mapper := NewMapper() + for i, test := range tests { + t.Run(fmt.Sprintf("test-%d", i), func(t *testing.T) { + if m, err := mapper.MapSQL(test.sql); err != nil { + t.Fatal(err) + } else { + if stmt, ok := m.Statement.(*sqlparser.Select); !ok { + t.Fatalf("%s: expected Statement: sqlparser.Select, got %T", test.sql, m.Statement) + } else { + t.Logf("%+v\n", stmt) + } + } + }) + } +} diff --git a/sql/mask.go b/sql/mask.go new file mode 100644 index 000000000..e1f859dff --- /dev/null +++ b/sql/mask.go @@ -0,0 +1,399 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package sql + +import ( + "strings" + + "vitess.io/vitess/go/vt/sqlparser" +) + +type selectPart int + +const ( + SelectPartDistinct selectPart = 1 << iota + SelectPartID + SelectPartStar + SelectPartField + SelectPartFields + SelectPartCountStar + SelectPartCountField + SelectPartCountDistinctField + SelectPartMinField + SelectPartMaxField + SelectPartSumField + SelectPartAvgField +) + +type fromPart int + +const ( + FromPartTable fromPart = 1 << iota + FromPartTables + FromPartJoin +) + +type wherePart int + +const ( + WherePartIDCondition wherePart = 1 << iota + WherePartFieldCondition + WherePartMultiFieldCondition +) + +type groupByPart int + +const ( + GroupByPartField groupByPart = 1 << iota + GroupByPartFields +) + +type havingPart int + +const ( + HavingPartCondition havingPart = 1 << iota +) + +type orderByPart int + +const ( + OrderByPartField orderByPart = 1 << iota + OrderByPartFields +) + +type limitPart int + +const ( + LimitPartLimit limitPart = 1 << iota + LimitPartOffset +) + +type QueryMask struct { + SelectMask selectPart + FromMask fromPart + WhereMask wherePart + GroupByMask groupByPart + HavingMask havingPart + OrderByMask orderByPart + LimitMask limitPart +} + +func NewQueryMask(sp selectPart, fp fromPart, wp wherePart, gp groupByPart, hp havingPart) QueryMask { + qm := QueryMask{} + qm.orSelect(sp) + qm.orFrom(fp) + qm.orWhere(wp) + qm.orGroupBy(gp) + qm.orHaving(hp) + return qm +} + +// ApplyFilter returns true if m passes the filter f. +// Note: only certain query parts are included; namely, +// the orderBy and limit masks are not applied to the +// filter. +func (m *QueryMask) ApplyFilter(f QueryMask) bool { + if m.SelectMask&f.SelectMask != m.SelectMask { + return false + } + if m.FromMask&f.FromMask != m.FromMask { + return false + } + if m.WhereMask&f.WhereMask != m.WhereMask { + return false + } + if m.GroupByMask&f.GroupByMask != m.GroupByMask { + return false + } + if m.HavingMask&f.HavingMask != m.HavingMask { + return false + } + return true +} + +// orSelect applies the bitwise-OR operation to the selectMask. +func (m *QueryMask) orSelect(o selectPart) { + m.SelectMask |= o +} + +// orFrom applies the bitwise-OR operation to the fromMask. +func (m *QueryMask) orFrom(o fromPart) { + m.FromMask |= o +} + +// orWhere applies the bitwise-OR operation to the whereMask. +func (m *QueryMask) orWhere(parts ...wherePart) { + for _, o := range parts { + m.WhereMask |= o + } +} + +// orGroupBy applies the bitwise-OR operation to the groupByMask. +func (m *QueryMask) orGroupBy(o groupByPart) { + m.GroupByMask |= o +} + +// orHaving applies the bitwise-OR operation to the havingMask. +func (m *QueryMask) orHaving(o havingPart) { + m.HavingMask |= o +} + +// orOrderBy applies the bitwise-OR operation to the orderByMask. +func (m *QueryMask) orOrderBy(o orderByPart) { + m.OrderByMask |= o +} + +// orLimit applies the bitwise-OR operation to the limitMask. +func (m *QueryMask) orLimit(o limitPart) { + m.LimitMask |= o +} + +// HasSelect returns true if the mask contains a supported select clause. +func (m *QueryMask) HasSelect() bool { + return m.SelectMask > 0 +} + +// HasSelectPart returns true if the mask contains the provided select part. +func (m *QueryMask) HasSelectPart(p selectPart) bool { + return (m.SelectMask & p) > 0 +} + +// HasFrom returns true if the mask contains a supported from clause. +func (m *QueryMask) HasFrom() bool { + return m.FromMask > 0 +} + +// HasWhere returns true if the mask contains a supported where clause. +func (m *QueryMask) HasWhere() bool { + return m.WhereMask > 0 +} + +// HasGroupBy returns true if the mask contains a supported group by clause. +func (m *QueryMask) HasGroupBy() bool { + return m.GroupByMask > 0 +} + +// HasHaving returns true if the mask contains a supported having clause. +func (m *QueryMask) HasHaving() bool { + return m.HavingMask > 0 +} + +// HasOrderBy returns true if the mask contains a supported order by clause. +func (m *QueryMask) HasOrderBy() bool { + return m.OrderByMask > 0 +} + +// HasLimit returns true if the mask contains a supported limit clause. +func (m *QueryMask) HasLimit() bool { + return m.LimitMask > 0 +} + +////////////////////////////////////////////////////////////////////////// + +func MustGenerateMask(sql string) QueryMask { + parsed, err := sqlparser.Parse(sql) + if err != nil { + return QueryMask{} + } + return GenerateMask(parsed) +} + +func GenerateMask(parsed sqlparser.Statement) QueryMask { + qm := QueryMask{} + + switch stmt := parsed.(type) { + case *sqlparser.Select: + if strings.ToLower(strings.TrimSpace(stmt.Distinct)) == "distinct" { + qm.orSelect(SelectPartDistinct) + } + // select parts + var fldCount int + for _, item := range stmt.SelectExprs { + switch expr := item.(type) { + case *sqlparser.AliasedExpr: + switch colExpr := expr.Expr.(type) { + case *sqlparser.ColName: + name := colExpr.Name.String() + if name == ColID { + qm.orSelect(SelectPartID) + } else { + fldCount++ + } + case *sqlparser.FuncExpr: + funcName := FuncName(strings.ToLower(colExpr.Name.String())) + switch len(colExpr.Exprs) { + case 1: + var isStar bool + var isField bool + var isDistinctField bool + switch exp := colExpr.Exprs[0].(type) { + case *sqlparser.AliasedExpr: + switch exp.Expr.(type) { + case *sqlparser.ColName: + isField = true + isDistinctField = colExpr.Distinct + } + case *sqlparser.StarExpr: + isStar = true + } + + switch funcName { + case FuncCount: + if isStar { + qm.orSelect(SelectPartCountStar) + } else if isDistinctField { + qm.orSelect(SelectPartCountDistinctField) + } else if isField { + qm.orSelect(SelectPartCountField) + } + case FuncMin: + if isField { + qm.orSelect(SelectPartMinField) + } + case FuncMax: + if isField { + qm.orSelect(SelectPartMaxField) + } + case FuncSum: + if isField { + qm.orSelect(SelectPartSumField) + } + case FuncAvg: + if isField { + qm.orSelect(SelectPartAvgField) + } + } + } + } + case *sqlparser.StarExpr: + qm.orSelect(SelectPartStar) + } + } + if fldCount == 1 { + qm.orSelect(SelectPartField) + } else if fldCount > 1 { + qm.orSelect(SelectPartFields) + } + + // from parts + switch len(stmt.From) { + case 1: + switch stmt.From[0].(type) { + case *sqlparser.AliasedTableExpr: + qm.orFrom(FromPartTable) + case *sqlparser.JoinTableExpr: + qm.orFrom(FromPartJoin) + } + case 2: + switch stmt.From[0].(type) { + case *sqlparser.AliasedTableExpr: + switch stmt.From[1].(type) { + case *sqlparser.AliasedTableExpr: + qm.orFrom(FromPartTables) + } + } + } + + // where parts + where := stmt.Where + if where != nil { + switch where.Type { + case "where": + qm.orWhere(generateWhereMask(where.Expr)...) + } + } + + // group by parts + var groupByFieldCount int + for _, item := range stmt.GroupBy { + switch item.(type) { + case *sqlparser.ColName: + groupByFieldCount++ + } + } + if groupByFieldCount == 1 { + qm.orGroupBy(GroupByPartField) + } else if groupByFieldCount > 1 { + qm.orGroupBy(GroupByPartFields) + } + + // having parts + if stmt.Having != nil { + qm.orHaving(HavingPartCondition) + } + + // order by parts + var orderByFieldCount int + for _, item := range stmt.OrderBy { + switch item.Expr.(type) { + case *sqlparser.ColName: + orderByFieldCount++ + } + } + if orderByFieldCount == 1 { + qm.orOrderBy(OrderByPartField) + } else if orderByFieldCount > 1 { + qm.orOrderBy(OrderByPartFields) + } + + // limit parts + if stmt.Limit != nil { + switch stmt.Limit.Rowcount.(type) { + case *sqlparser.SQLVal: + qm.orLimit(LimitPartLimit) + } + switch stmt.Limit.Offset.(type) { + case *sqlparser.SQLVal: + qm.orLimit(LimitPartOffset) + } + } + } + return qm +} + +// TODO: add more recursion within the comparison operators (left/right parts) +func generateWhereMask(e sqlparser.Expr) []wherePart { + var wp []wherePart + + switch expr := e.(type) { + case *sqlparser.ComparisonExpr: + var leftName string + if colExpr, ok := expr.Left.(*sqlparser.ColName); ok { + leftName = colExpr.Name.String() + } + switch leftName { + case ColID: + wp = append(wp, WherePartIDCondition) + case "": + // + default: + wp = append(wp, WherePartFieldCondition) + } + case *sqlparser.RangeCond: + wp = append(wp, WherePartFieldCondition) + case *sqlparser.AndExpr, *sqlparser.OrExpr: + // TODO: we need to recursively ensure that the left/right + // sides of these and/or expressions are field-op-val, and + // that none of the fields are "_id" + // TODO: could we use extractComparison or something like it? + wp = append(wp, WherePartMultiFieldCondition) + case *sqlparser.NotExpr: + wp = append(wp, generateWhereMask(expr.Expr)...) + case *sqlparser.IsExpr: + wp = append(wp, WherePartFieldCondition) + } + + return wp +} diff --git a/sql/model.go b/sql/model.go new file mode 100644 index 000000000..76ba80498 --- /dev/null +++ b/sql/model.go @@ -0,0 +1,141 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package sql + +import ( + "fmt" + + "github.com/pilosa/pilosa/v2" + "github.com/pkg/errors" +) + +var ( + ErrUnsupportedQuery = errors.New("unsupported query") +) + +// TODO: what the difference between this and IDIndexColumn? +type KeyIndexColumn struct { + Index *pilosa.Index + text string + alias string +} + +func NewKeyIndexColumn(index *pilosa.Index, alias string) *KeyIndexColumn { + return &KeyIndexColumn{ + Index: index, + text: ColID, + alias: alias, + } +} + +func (i *KeyIndexColumn) Source() string { + return "" // TODO +} +func (i *KeyIndexColumn) Name() string { + return i.text +} +func (i *KeyIndexColumn) Alias() string { + if i.alias != "" { + return i.alias + } + return i.Name() +} + +type IDIndexColumn struct { + Index *pilosa.Index + text string + alias string +} + +func NewIDIndexColumn(index *pilosa.Index, alias string) *IDIndexColumn { + return &IDIndexColumn{ + Index: index, + text: ColID, + alias: alias, + } +} + +func (i *IDIndexColumn) Source() string { + return "" // TODO +} +func (i *IDIndexColumn) Name() string { + return i.text +} +func (i *IDIndexColumn) Alias() string { + if i.alias != "" { + return i.alias + } + return i.Name() +} + +type FieldColumn struct { + Field *pilosa.Field + text string + alias string +} + +func NewFieldColumn(field *pilosa.Field, alias string) *FieldColumn { + return &FieldColumn{ + Field: field, + text: field.Name(), + alias: alias, + } +} + +func (f *FieldColumn) Source() string { + return "" +} +func (f *FieldColumn) Name() string { + return f.text +} +func (f *FieldColumn) Alias() string { + if f.alias != "" { + return f.alias + } + return f.Name() +} + +type FuncColumn struct { + Field *pilosa.Field + FuncName FuncName + alias string +} + +func NewFuncColumn(funcName FuncName, field *pilosa.Field, alias string) *FuncColumn { + return &FuncColumn{ + Field: field, + FuncName: funcName, + alias: alias, + } +} + +func (f *FuncColumn) Source() string { + return string(f.FuncName) +} + +func (f *FuncColumn) Name() string { + fieldName := "*" + if f.Field != nil { + fieldName = f.Field.Name() + } + return fmt.Sprintf("%s(%s)", f.FuncName, fieldName) +} + +func (f *FuncColumn) Alias() string { + if f.alias != "" { + return f.alias + } + return f.Name() +} diff --git a/sql/query.go b/sql/query.go new file mode 100644 index 000000000..9996110a9 --- /dev/null +++ b/sql/query.go @@ -0,0 +1,295 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package sql + +import ( + "encoding/json" + "fmt" + "strconv" + "strings" + "time" + + "github.com/pkg/errors" +) + +const timeFormat = "2006-01-02T15:04" + +// LT creates a less than query. +func LT(fieldName string, value interface{}) string { + return fmt.Sprintf("Row(%s<%s)", fieldName, intOrFloat(value)) +} + +// LTE creates a less than or equal query. +func LTE(fieldName string, value interface{}) string { + return fmt.Sprintf("Row(%s<=%s)", fieldName, intOrFloat(value)) +} + +// GT creates a greater than query. +func GT(fieldName string, value interface{}) string { + return fmt.Sprintf("Row(%s>%s)", fieldName, intOrFloat(value)) +} + +// GTE creates a greater than or equal query. +func GTE(fieldName string, value interface{}) string { + return fmt.Sprintf("Row(%s>=%s)", fieldName, intOrFloat(value)) +} + +// Equals creates an equals query. +func Equals(fieldName string, value interface{}) string { + return fmt.Sprintf("Row(%s=%s)", fieldName, intOrFloat(value)) +} + +// NotEquals creates a not equals query. +func NotEquals(fieldName string, value interface{}) string { + return fmt.Sprintf("Row(%s!=%s)", fieldName, intOrFloat(value)) +} + +// NotNull creates a not equal to null query. +func NotNull(fieldName string) string { + return fmt.Sprintf("Row(%s!=null)", fieldName) +} + +// Row query +func Row(fieldName string, rowIDOrKey interface{}) (string, error) { + rowStr, err := formatIDKeyBool(rowIDOrKey) + if err != nil { + return "", err + } + text := fmt.Sprintf("Row(%s=%s)", fieldName, rowStr) + return text, nil +} + +// RowRange is a Row query with from,to times +func RowRange(fieldName string, rowIDOrKey interface{}, start time.Time, end time.Time) (string, error) { + rowStr, err := formatIDKeyBool(rowIDOrKey) + if err != nil { + return "", err + } + text := fmt.Sprintf("Row(%s=%s,from='%s',to='%s')", fieldName, rowStr, start.Format(timeFormat), end.Format(timeFormat)) + return text, nil +} + +// Union query - see rowOperation +func Union(rows ...string) string { + return rowOperation("Union", rows...) +} + +// Intersect query - see rowOperation +func Intersect(rows ...string) string { + return rowOperation("Intersect", rows...) +} + +// Not query +func Not(rows ...string) string { + return rowOperation("Not", rows...) +} + +// Like creates a Rows query filtered by a pattern. +// An underscore ('_') can be used as a placeholder for a single UTF-8 codepoint or a percent sign ('%') can be used as a placeholder for 0 or more codepoints. +// All other codepoints in the pattern are matched exactly. +func Like(fieldName string, pattern string) string { + pattern = strings.ReplaceAll(pattern, `\`, `\\`) + pattern = strings.ReplaceAll(pattern, `'`, `\'`) + return fmt.Sprintf("UnionRows(Rows(field='%s',like='%s'))", fieldName, pattern) +} + +// Between creates a between query. +func Between(fieldName string, a interface{}, b interface{}) string { + return fmt.Sprintf("Row(%s >< [%s,%s])", fieldName, intOrFloat(a), intOrFloat(b)) +} + +// Distinct creates a Distinct query. +func Distinct(indexName, fieldName string) string { + return fmt.Sprintf("Distinct(Row(%s!=null),index='%s',field='%s')", fieldName, indexName, fieldName) +} + +// RowDistinct creates a Distinct query with the given row filter. +func RowDistinct(indexName, fieldName string, row string) string { + return fmt.Sprintf("Distinct(%s,index='%s',field='%s')", row, indexName, fieldName) +} + +// Rows creates a Rows query with defaults +func Rows(fieldName string) string { + return fmt.Sprintf("Rows(field='%s')", fieldName) +} + +// RowsLimit creates a Rows query with the given limit +func RowsLimit(fieldName string, limit int64) (string, error) { + if limit < 0 { + return "", errors.New("rows limit must be non-negative") + } + text := fmt.Sprintf("Rows(field='%s',limit=%d)", fieldName, limit) + return text, nil +} + +// All creates an All query. +// Returns the set columns with existence true. +func All() string { + return "All()" +} + +// Count creates a Count query. +// Returns the number of set columns in the ROW_CALL passed in. +func Count(rowCall string) string { + return fmt.Sprintf("Count(%s)", rowCall) +} + +// Sum creates a sum query. +func Sum(fieldName string, row string) string { + return valQuery(fieldName, "Sum", row) +} + +// Min creates a min query. +func Min(fieldName string, row string) string { + return valQuery(fieldName, "Min", row) +} + +// Max creates a max query. +func Max(fieldName string, row string) string { + return valQuery(fieldName, "Max", row) +} + +// TopN creates a TopN query with the given item count. +// Returns the id and count of the top n rows (by count of columns) in the field. +func TopN(fieldName string, n uint64) string { + return fmt.Sprintf("TopN(%s,n=%d)", fieldName, n) +} + +// RowTopN creates a TopN query with the given item count and row. +// This variant supports customizing the row query. +func RowTopN(fieldName string, n uint64, row string) string { + return fmt.Sprintf("TopN(%s,%s,n=%d)", fieldName, row, n) +} + +// GroupByBase creates a GroupBy query with the given functional options. +func GroupByBase(rows []string, limit int64, filter, aggregate, having string) (string, error) { + if len(rows) == 0 { + return "", errors.New("there should be at least one rows query") + } + if limit < 0 { + return "", errors.New("limit must be non-negative") + } + + // rows + text := fmt.Sprintf("GroupBy(%s", strings.Join(rows, ",")) + + // limit + if limit > 0 { + text += fmt.Sprintf(",limit=%d", limit) + } + + // filter + if filter != "" { + text += fmt.Sprintf(",filter=%s", filter) + } + + // aggregate + if aggregate != "" { + text += fmt.Sprintf(",aggregate=%s", aggregate) + } + + // having + if having != "" { + text += fmt.Sprintf(",having=%s", having) + } + + text += ")" + return text, nil +} + +// Limit creates a limit query. +func Limit(row string, limit uint, offset uint) string { + return fmt.Sprintf("Limit(%s, limit=%d, offset=%d)", row, limit, offset) +} + +// Offset creates a limit query but only with an offset. +func Offset(row string, offset uint) string { + return fmt.Sprintf("Limit(%s, offset=%d)", row, offset) +} + +// ConstRow creates a query value that uses a list of columns in place of a Row query. +func ConstRow(ids ...interface{}) string { + if ids == nil { + ids = []interface{}{} + } + data, _ := json.Marshal(ids) + return fmt.Sprintf("ConstRow(columns=%s)", data) +} + +// Extract creates an Extract query. +// It accepts a bitmap query to select columns and a list of fields to select rows. +func Extract(rowCall string, fields ...string) string { + var rowsCall string + for _, r := range fields { + rowsCall += "," + fmt.Sprintf("Rows(%s)", r) + } + + return fmt.Sprintf("Extract(%s%s)", rowCall, rowsCall) +} + +func valQuery(fieldName string, op string, row string) string { + if row != "" { + row += "," + } + return fmt.Sprintf("%s(%sfield='%s')", op, row, fieldName) +} + +func rowOperation(name string, rows ...string) string { + return fmt.Sprintf("%s(%s)", name, strings.Join(rows, ",")) +} + +func formatIDKeyBool(idKeyBool interface{}) (string, error) { + if b, ok := idKeyBool.(bool); ok { + return strconv.FormatBool(b), nil + } + if flt, ok := idKeyBool.(float64); ok { + return fmt.Sprintf("%f", flt), nil + } + return formatIDKey(idKeyBool) +} + +func formatIDKey(idKey interface{}) (string, error) { + switch v := idKey.(type) { + case uint: + return strconv.FormatUint(uint64(v), 10), nil + case uint32: + return strconv.FormatUint(uint64(v), 10), nil + case uint64: + return strconv.FormatUint(v, 10), nil + case int: + return strconv.FormatInt(int64(v), 10), nil + case int32: + return strconv.FormatInt(int64(v), 10), nil + case int64: + return strconv.FormatInt(v, 10), nil + case string: + v = strings.ReplaceAll(v, `\`, `\\`) + return fmt.Sprintf(`'%s'`, strings.ReplaceAll(v, `'`, `\'`)), nil + default: + return "", errors.Errorf("id/key is not a string or integer type: %#v", idKey) + } +} + +func intOrFloat(value interface{}) string { + switch value.(type) { + case float64, float32: + // In order to test expected values, we set the precision + // to 8. TODO: It's likely we'll need to address this + // at some point. + return fmt.Sprintf("%.8f", value) + default: + return fmt.Sprintf("%d", value) + } +} diff --git a/sql/reduce.go b/sql/reduce.go new file mode 100644 index 000000000..be0e98695 --- /dev/null +++ b/sql/reduce.go @@ -0,0 +1,500 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package sql + +import ( + "io" + "sort" + + "github.com/pilosa/pilosa/v2/pql" + pproto "github.com/pilosa/pilosa/v2/proto" + "github.com/pkg/errors" + "google.golang.org/grpc/codes" +) + +// DataType contants describe the possible values +// for the Datatype value in the RowResponse header. +const ( + DataTypeDecimal = "decimal" + DataTypeFloat64 = "float64" + DataTypeInt64 = "int64" + DataTypeString = "string" + DataTypeUint64Array = "[]uint64" +) + +type Reducer interface { + Reduce(pproto.StreamClient, pproto.StreamServer) error +} + +// LimitReducer limits the number of messages passed through. +type LimitReducer struct { + limit uint + offset uint +} + +// NewLimitReducer returns a new instance of LimitReducer. +func NewLimitReducer(limit, offset uint) *LimitReducer { + return &LimitReducer{ + limit: limit, + offset: offset, + } +} + +// Reduce applies the limit reducer to the client stream and sends the results +// to the server stream. +func (l *LimitReducer) Reduce(c pproto.StreamClient, s pproto.StreamServer) error { + offsetCountdown := l.offset + + // in the case of an offset, since we'll be skipping the first record + // which contains the headers, we need to pull the headers, save them, + // and apply them to the first record that we actually send through. + var headers []*pproto.ColumnInfo + + for i := uint(0); i < l.limit+l.offset || l.limit == 0; i++ { + r, err := c.Recv() + if err == io.EOF { + break + } else if err != nil { + return s.Send(pproto.ErrorWrap(err, "receiving on client stream")) + } + if offsetCountdown > 0 { + if headers == nil { + headers = r.Headers + } + offsetCountdown-- + continue + } + if headers != nil { + r.Headers = headers + headers = nil + } + if err := s.Send(r); err != nil { + return s.Send(pproto.ErrorWrap(err, "sending on server stream")) + } + } + return s.Send(pproto.EOF) +} + +// OrderByReducer orders the results based on the provide conditions. +// It also takes limit and offset to reduce the amount of items +// needing to be held in memory for sorting. +type OrderByReducer struct { + fields []string + isDescending []bool // direction[asc: false, desc: true] + limit uint + offset uint +} + +// NewOrderByReducer returns a new instance of OrderByReducer. +func NewOrderByReducer(fields, dirs []string, limit, offset uint) *OrderByReducer { + descendings := make([]bool, len(fields)) + for i := range dirs { + if dirs[i] == "desc" { + descendings[i] = true + } + } + return &OrderByReducer{ + fields: fields, + isDescending: descendings, + limit: limit, + offset: offset, + } +} + +// Reduce applies the order by reducer to the client stream and sends the results +// to the server stream. +func (o *OrderByReducer) Reduce(c pproto.StreamClient, s pproto.StreamServer) error { + // hold is a slice of row responses, to be sent to the output + // stream sorted by the sort conditions. + var hold []*pproto.RowResponse + + // sortColNames contains the names of the columns to + // sort on. + sortColNames := o.fields + + // sortColIdxs contains the positions of the sort columns + // in the result set. + sortColIdxs := make([]int, len(sortColNames)) + + // sortColTypes contains the data types of the columns + // to be sorted. (ex: "uint64", "string", etc.). This + // is used to determine how to convert it to a typed + // field for sorting. + sortColTypes := make([]string, len(sortColNames)) + + // holdHeaders is used to stash the headers (from + // the first row) so they can be applied later + // to what will eventually be the first row after + // sorting has occurred. + var holdHeaders []*pproto.ColumnInfo + + ii := 0 + for { + rr, err := c.Recv() + if err != nil { + if err == io.EOF { + break + } + return s.Send(pproto.ErrorWrap(err, "receiving row response")) + } + + // On the first row, get the sort column information + // from the headers. Also, stash the headers for + // later in the `holdHeaders` var. + if ii == 0 { + holdHeaders = rr.Headers + for i, rrHdr := range rr.Headers { + hdrName := rrHdr.GetName() + hdrType := rrHdr.GetDatatype() + for j := range sortColNames { + if sortColNames[j] == hdrName { + sortColIdxs[j] = i + sortColTypes[j] = hdrType + } + } + } + // Clear the headers in case this record is + // no longer first (we re-apply the headers + // to the first outgoing record later). + rr.Headers = nil + } + + // Put each row in the hold. + hold = append(hold, rr) + + ii++ + + // TODO: in the case where limit is provided and the number of possible + // rows is large, it might be more efficient to periodically sort/trim + // the hold so it doesn't become too large. For example, it could + // be constrained to size (limit + offset + buffer), where buffer is + // an amount that the hold can grow before being trimmed. + } + + // Sort the hold. + sorter, err := pproto.NewRowResponseSorter( + sortColIdxs, + o.isDescending, + sortColTypes, + hold, + ) + if err != nil { + return s.Send(pproto.ErrorWrap(err, "creating row response sorter")) + } + sort.Sort(sorter) + + var rowsToConsider uint = uint(len(hold)) + var offsetCountdown uint + if o.limit > 0 { + offsetCountdown = o.offset + if o.limit+o.offset < rowsToConsider { + rowsToConsider = o.limit + o.offset + } + } + + // Loop over hold and send each row response. + // Apply the header to the first row that is sent. + var headerApplied bool + for i := uint(0); i < rowsToConsider; i++ { + if offsetCountdown > 0 { + offsetCountdown-- + continue + } + // Re-apply the headers to the first record. + if !headerApplied { + hold[i].Headers = holdHeaders + headerApplied = true + } + err := s.Send(hold[i]) + if err != nil { + return s.Send(pproto.ErrorWrap(err, "sending hold row")) + } + } + return s.Send(pproto.EOF) +} + +// ValCountFuncReducer converts a ValCount result to the proper +// result for Func. +type ValCountFuncReducer struct { + fn FuncName +} + +// NewValCountFuncReducer returns a new instance of ValCountFuncReducer. +func NewValCountFuncReducer(fn FuncName) *ValCountFuncReducer { + return &ValCountFuncReducer{ + fn: fn, + } +} + +// Reduce modifies the stream according to the function. +func (v *ValCountFuncReducer) Reduce(c pproto.StreamClient, s pproto.StreamServer) error { + r, err := c.Recv() + if err != nil { + if err == io.EOF { + return s.Send(pproto.EOF) + } + return s.Send(pproto.Error(err)) + } + + // Get the index of the column with header of "value". + var idxVal int = -1 + var idxCnt int = -1 + headers := r.GetHeaders() + for i, hdr := range headers { + switch hdr.GetName() { + case "value": + idxVal = i + case "count": + idxCnt = i + } + } + + var sourceDataType string + var returnDataType string + + sourceDataType = headers[idxVal].GetDatatype() + returnDataType = sourceDataType + switch v.fn { + case FuncAvg: + returnDataType = DataTypeFloat64 + } + + rr := pproto.RowResponse{ + Headers: []*pproto.ColumnInfo{ + {Name: string(v.fn), Datatype: returnDataType}, + }, + Columns: make([]*pproto.ColumnResponse, 1), + } + + cols := r.GetColumns() + if len(cols) == 0 { + return s.Send(pproto.ErrorCode( + errors.New("empty column set"), + codes.Unknown, + )) + } + + if idxVal == -1 { + return s.Send(pproto.ErrorCode( + errors.New("result set has no column: value"), + codes.Unknown, + )) + } + if idxCnt == -1 { + return s.Send(pproto.ErrorCode( + errors.New("result set has no column: count"), + codes.Unknown, + )) + } + + switch v.fn { + case FuncAvg: + var avg float64 + if sourceDataType == DataTypeDecimal { + val := cols[idxVal].GetDecimalVal() + dec := pql.NewDecimal(val.Value, val.Scale) + cnt := cols[idxCnt].GetInt64Val() + avg = dec.Float64() / float64(cnt) + } else { + val := cols[idxVal].GetInt64Val() + cnt := cols[idxCnt].GetInt64Val() + avg = float64(val) / float64(cnt) + } + rr.Columns[0] = &pproto.ColumnResponse{ColumnVal: &pproto.ColumnResponse_Float64Val{Float64Val: avg}} + default: + if sourceDataType == DataTypeDecimal { + val := cols[idxVal].GetDecimalVal() + rr.Columns[0] = &pproto.ColumnResponse{ColumnVal: &pproto.ColumnResponse_DecimalVal{DecimalVal: &pproto.Decimal{Value: val.Value, Scale: val.Scale}}} + } else { + val := cols[idxVal].GetInt64Val() + rr.Columns[0] = &pproto.ColumnResponse{ColumnVal: &pproto.ColumnResponse_Int64Val{Int64Val: val}} + } + } + + if err := s.Send(&rr); err != nil { + return errors.Wrap(err, "sending row response") + } + return s.Send(pproto.EOF) +} + +// CountIDReducer returns a stream of _id's as a count. +type CountIDReducer struct{} + +// Reduce counts the stream of IDs and returns a single record. +func (r *CountIDReducer) Reduce(c pproto.StreamClient, s pproto.StreamServer) error { + var cnt uint64 + + for { + _, err := c.Recv() + if err != nil { + if err == io.EOF { + break + } + return s.Send(pproto.ErrorWrap(err, "receiving on client stream")) + } + cnt++ + } + + rr := pproto.RowResponse{ + Headers: []*pproto.ColumnInfo{ + {Name: string(FuncCount), Datatype: "uint64"}, + }, + Columns: []*pproto.ColumnResponse{ + &pproto.ColumnResponse{ColumnVal: &pproto.ColumnResponse_Uint64Val{Uint64Val: cnt}}, + }, + } + + if err := s.Send(&rr); err != nil { + return errors.Wrap(err, "sending row response") + } + return s.Send(pproto.EOF) +} + +// AssignHeadersReducer overwrites the headers on the first record +// according to field names and aliases from sql. It also reorders +// the columns in the result stream to match the sql select clause. +type AssignHeadersReducer struct { + cols []Column +} + +// NewAssignHeadersReducer returns a new instance of AssignHeadersReducer. +func NewAssignHeadersReducer(cols []Column) *AssignHeadersReducer { + return &AssignHeadersReducer{ + cols: cols, + } +} + +// Reduce modifies the stream. +func (r *AssignHeadersReducer) Reduce(c pproto.StreamClient, s pproto.StreamServer) error { + var placement []uint + var labels []string + + var cnt int + for { + rr, err := c.Recv() + if err != nil { + if err == io.EOF { + break + } + return s.Send(pproto.ErrorWrap(err, "receiving on client stream")) + } + + // If the placement slice is [0-n] where n == len(Headers) + // then we don't need to alter rr on records after cnt == 0. + // If we don't apply aliases, we don't have to alter Headers + // either, but that may not be worth messing with. + + if cnt == 0 { + placement, labels, err = headerAssignment(r.cols, rr.Headers) + if err != nil { + return s.Send(pproto.ErrorWrap(err, "getting header assignment")) + } + + // mod is the modified RowResponse object that gets populated + // according to placement and labels, then sent. + mod := &pproto.RowResponse{ + Headers: make([]*pproto.ColumnInfo, len(placement)), + Columns: make([]*pproto.ColumnResponse, len(placement)), + } + + // For now, we assume that the column count in each RowResponse + // is consistent (i.e. we can validate one time, here, on the + // first row, and not every time, in the `else` statement below). + if len(placement) > len(rr.Columns) { + return s.Send(pproto.ErrorCode( + errors.New("mismatched header placement and column count"), + codes.Unknown, + )) + } + + for i := 0; i < len(placement); i++ { + mod.Headers[i] = rr.Headers[placement[i]] + mod.Headers[i].Name = labels[i] + mod.Columns[i] = rr.Columns[placement[i]] + } + if err := s.Send(mod); err != nil { + return errors.Wrap(err, "sending mod") + } + } else { + // mod is the modified RowResponse object that gets populated + // according to placement and labels, then sent. + mod := &pproto.RowResponse{ + Columns: make([]*pproto.ColumnResponse, len(placement)), + } + for i := 0; i < len(placement); i++ { + mod.Columns[i] = rr.Columns[placement[i]] + } + if err := s.Send(mod); err != nil { + return errors.Wrap(err, "sending mod") + } + } + cnt++ + } + + return s.Send(pproto.EOF) +} + +var ( + ErrIncompleteHeaders = errors.New("incomplete header assignment") + ErrFieldNotInHeaders = errors.New("field not found in source header") +) + +func headerAssignment(cols []Column, hdrs []*pproto.ColumnInfo) ([]uint, []string, error) { + // If any of the columns are "*" (i.e. type StarColumn), + // then ignore everything else and just use all result + // headers. + var hasStar bool + for _, col := range cols { + if _, ok := col.(*StarColumn); ok { + hasStar = true + break + } + } + if hasStar { + placement := make([]uint, len(hdrs)) + labels := make([]string, len(hdrs)) + for i, hdr := range hdrs { + placement[i] = uint(i) + labels[i] = hdr.Name + } + return placement, labels, nil + } + + if len(cols) > len(hdrs) { + return nil, nil, ErrIncompleteHeaders + } + placement := make([]uint, len(cols)) + labels := make([]string, len(cols)) + + // Make a map of the RowResponse headers. + hdrMap := make(map[string]uint) + for i, hdr := range hdrs { + hdrMap[hdr.Name] = uint(i) + } + + // Lookup each column in the hdrMap and determine the desired placement. + for i, col := range cols { + if srcHdrIdx, ok := hdrMap[col.Source()]; ok { + placement[i] = srcHdrIdx + labels[i] = col.Alias() + } else if nameHdrIdx, ok := hdrMap[col.Name()]; ok { + placement[i] = nameHdrIdx + labels[i] = col.Alias() + } else { + return nil, nil, errors.Wrapf(ErrFieldNotInHeaders, "field: %s", col.Name()) + } + } + return placement, labels, nil +} diff --git a/sql/reduce_test.go b/sql/reduce_test.go new file mode 100644 index 000000000..27f645c99 --- /dev/null +++ b/sql/reduce_test.go @@ -0,0 +1,120 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package sql + +import ( + "fmt" + "reflect" + "testing" + + pproto "github.com/pilosa/pilosa/v2/proto" + "github.com/pkg/errors" +) + +func TestHeaderAssignment(t *testing.T) { + abcdHdrs := []*pproto.ColumnInfo{ + {Name: "a"}, {Name: "b"}, {Name: "c"}, {Name: "d"}, + } + tests := []struct { + cols []Column + hdrs []*pproto.ColumnInfo + expPlacement []uint + expLabels []string + expErr error + }{ + { + cols: []Column{ + NewBasicColumn("a", "namea", ""), + }, + hdrs: abcdHdrs, + expPlacement: []uint{0}, + expLabels: []string{"namea"}, + }, + { + cols: []Column{ + NewBasicColumn("", "a", ""), + }, + hdrs: abcdHdrs, + expPlacement: []uint{0}, + expLabels: []string{"a"}, + }, + { + cols: []Column{ + NewBasicColumn("", "a", "aliasa"), + }, + hdrs: abcdHdrs, + expPlacement: []uint{0}, + expLabels: []string{"aliasa"}, + }, + { + cols: []Column{ + NewBasicColumn("", "a", "aliasa"), + NewBasicColumn("c", "namec", ""), + }, + hdrs: abcdHdrs, + expPlacement: []uint{0, 2}, + expLabels: []string{"aliasa", "namec"}, + }, + { + cols: []Column{ + NewBasicColumn("d", "c", "aliasd"), + NewBasicColumn("b", "nameb", ""), + }, + hdrs: abcdHdrs, + expPlacement: []uint{3, 1}, + expLabels: []string{"aliasd", "nameb"}, + }, + // Errors + { + cols: []Column{ + NewBasicColumn("", "x", ""), + }, + hdrs: abcdHdrs, + expErr: ErrFieldNotInHeaders, + }, + { + cols: []Column{ + NewBasicColumn("", "a", ""), + NewBasicColumn("", "b", ""), + NewBasicColumn("", "c", ""), + NewBasicColumn("", "d", ""), + NewBasicColumn("", "e", ""), + }, + hdrs: abcdHdrs, + expErr: ErrIncompleteHeaders, + }, + } + for i, test := range tests { + t.Run(fmt.Sprintf("test-%d", i), func(t *testing.T) { + placement, labels, err := headerAssignment(test.cols, test.hdrs) + + if test.expErr == nil { + if err != nil { + t.Fatal(err) + } + } else { + if test.expErr != errors.Cause(err) { + t.Fatalf("expected error: %v, but got: %v", test.expErr, err) + } + } + + if !reflect.DeepEqual(placement, test.expPlacement) { + t.Fatalf("expected placement: %v, but got: %v", test.expPlacement, placement) + } else if !reflect.DeepEqual(labels, test.expLabels) { + t.Fatalf("expected labels: %v, but got: %v", test.expLabels, labels) + } + }) + } +} diff --git a/sql/router.go b/sql/router.go new file mode 100644 index 000000000..f64daa338 --- /dev/null +++ b/sql/router.go @@ -0,0 +1,162 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package sql + +type router struct { + direct map[QueryMask]handler + filters []maskFilter +} + +type maskFilter struct { + optional QueryMask + required []QueryMask + handler handler +} + +func newRouter() *router { + selectRouter := &router{ + direct: make(map[QueryMask]handler), + } + + selectRouter.addFilter( + NewQueryMask( + SelectPartID|SelectPartStar|SelectPartField|SelectPartFields, + FromPartTable, + WherePartFieldCondition|WherePartMultiFieldCondition|WherePartIDCondition, + 0, + 0, + ), + []QueryMask{}, + handlerSelectFieldsFromTableWhere{}, + ) + //// + selectRouter.addRoute("select distinct fld from tbl", handlerSelectDistinctFromTable{}) + //// + selectRouter.addFilter( + NewQueryMask( + SelectPartCountStar|SelectPartCountField|SelectPartCountDistinctField, + FromPartTable, + WherePartFieldCondition|WherePartMultiFieldCondition|WherePartIDCondition, + 0, + 0, + ), + []QueryMask{}, + handlerSelectCountFromTableWhere{}, + ) + //// + selectRouter.addFilter( + NewQueryMask( + SelectPartMinField|SelectPartMaxField|SelectPartSumField|SelectPartAvgField, + FromPartTable, + WherePartFieldCondition|WherePartMultiFieldCondition|WherePartIDCondition, + 0, + 0, + ), + []QueryMask{}, + handlerSelectFuncFromTableWhere{}, + ) + //// + groupByOptional := NewQueryMask( + SelectPartField|SelectPartFields|SelectPartCountStar|SelectPartSumField, + FromPartTable, + WherePartFieldCondition, // TODO: this can probably handle fields as well + GroupByPartField|GroupByPartFields, + HavingPartCondition, + ) + selectRouter.addFilter( + groupByOptional, + []QueryMask{NewQueryMask(0, 0, 0, GroupByPartField, 0)}, + handlerSelectGroupBy{}, + ) + selectRouter.addFilter( + groupByOptional, + []QueryMask{NewQueryMask(0, 0, 0, GroupByPartFields, 0)}, + handlerSelectGroupBy{}, + ) + selectRouter.addRoute("select fld, count(fld) from tbl group by fld", handlerSelectGroupBy{}) + selectRouter.addRoute("select fld1, count(fld1) from tbl where fld2=1 group by fld1", handlerSelectGroupBy{}) + + selectRouter.addRoute("select count(*) from tbl1 INNER JOIN tbl2 ON tbl1._id = tbl2.bsi", handlerSelectJoin{}) + selectRouter.addRoute("select _id from tbl1 INNER JOIN tbl2 ON tbl1._id = tbl2.bsi", handlerSelectJoin{}) + selectRouter.addRoute("select _id from tbl1 INNER JOIN tbl2 ON tbl1._id = tbl2.bsi where fld1=1", handlerSelectJoin{}) + selectRouter.addRoute("select _id from tbl1 INNER JOIN tbl2 ON tbl1._id = tbl2.bsi where fld1=1 and fld2=2", handlerSelectJoin{}) + + return selectRouter +} + +func (r *router) addRoute(sql string, handler handler) { + r.direct[MustGenerateMask(sql)] = handler +} + +func (r *router) addFilter(opt QueryMask, req []QueryMask, handler handler) { + mf := maskFilter{ + optional: opt, + required: req, + handler: handler, + } + r.filters = append(r.filters, mf) +} + +func (r *router) handler(qm QueryMask) handler { + // First, check for a direct mapping. + // Zero out the orderBy and limit mask, because + // those are not specific to the query processing. + zm := QueryMask{ + SelectMask: qm.SelectMask, + FromMask: qm.FromMask, + WhereMask: qm.WhereMask, + GroupByMask: qm.GroupByMask, + HavingMask: qm.HavingMask, + } + if h, ok := r.direct[zm]; ok { + return h + } + for _, mf := range r.filters { + if applyMaskFilter(&qm, mf) { + return mf.handler + } + } + return nil +} + +// applyMaskFilter returns true if m passes the filter mf. +// Note: only certain query parts are included; namely, +// the orderBy and limit masks are not applied to the +// filter. A mask can satisfy any part of the optional +// filter to pass through, but it MUST satisfy all parts +// of the required filter. +func applyMaskFilter(m *QueryMask, mf maskFilter) bool { + if !m.ApplyFilter(mf.optional) { + return false + } + for _, req := range mf.required { + if m.SelectMask&req.SelectMask != req.SelectMask { + return false + } + if m.FromMask&req.FromMask != req.FromMask { + return false + } + if m.WhereMask&req.WhereMask != req.WhereMask { + return false + } + if m.GroupByMask&req.GroupByMask != req.GroupByMask { + return false + } + if m.HavingMask&req.HavingMask != req.HavingMask { + return false + } + } + return true +} diff --git a/sql/select.go b/sql/select.go new file mode 100644 index 000000000..1c11cdfca --- /dev/null +++ b/sql/select.go @@ -0,0 +1,792 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package sql + +import ( + "context" + "fmt" + "strings" + + "github.com/pilosa/pilosa/v2" + "github.com/pilosa/pilosa/v2/pql" + pproto "github.com/pilosa/pilosa/v2/proto" + "github.com/pkg/errors" + "vitess.io/vitess/go/vt/sqlparser" +) + +// SelectHandler executes SQL select statements +type SelectHandler struct { + api *pilosa.API + router *router +} + +// NewSelectHandler constructor +func NewSelectHandler(api *pilosa.API) *SelectHandler { + return &SelectHandler{ + api: api, + router: newRouter(), + } +} + +// Handle executes mapped SQL +func (s *SelectHandler) Handle(ctx context.Context, mapped *MappedSQL) (pproto.StreamClient, error) { + mr, err := s.mapSelect(ctx, mapped.Statement.(*sqlparser.Select), mapped.Mask) + if err != nil { + return nil, errors.Wrap(err, "mapping select") + } + return s.execMappingResult(ctx, mr) +} + +func (s *SelectHandler) mapSelect(ctx context.Context, selectStmt *sqlparser.Select, qm QueryMask) (*MappingResult, error) { + // Get the handler for this query mask. + handler := s.router.handler(qm) + if handler == nil { + return nil, ErrUnsupportedQuery + } + indexFunc := func(indexName string) *pilosa.Index { + idx, err := s.api.Index(ctx, indexName) + if err != nil { + return nil + } + return idx + } + + mr, err := handler.Apply(selectStmt, qm, indexFunc) + if err != nil { + return nil, errors.Wrap(err, "handling") + } + return mr, nil +} + +func (s *SelectHandler) execMappingResult(ctx context.Context, mr *MappingResult) (stream pproto.StreamClient, err error) { + if mr.Query == "" { + return nil, errors.New("no pql query created") + } + + fmt.Println("PQL:", mr.Query) + resp, err := s.api.Query(ctx, &pilosa.QueryRequest{Index: mr.IndexName, Query: mr.Query}) + if err != nil { + return nil, errors.Wrap(err, "doing pql query") + } + res := resp.Results[0] + + // TODO: synchronize this properly somehow. + // It would probbably help to get rid of the streaming too. + respRows := pproto.NewRowBuffer(0) + switch res := res.(type) { + case pproto.ToRowser: + go func() { + if err := res.ToRows(respRows.Send); err != nil { + respRows.Send(pproto.Error(err)) //nolint:errcheck + } else { + _ = respRows.Send(pproto.EOF) //nolint:errcheck + } + }() + case []pilosa.GroupCount: + go func() { + if err := pilosa.GroupCounts(res).ToRows(respRows.Send); err != nil { + respRows.Send(pproto.Error(err)) //nolint:errcheck + } else { + respRows.Send(pproto.EOF) //nolint:errcheck + } + }() + case uint64: + go func() { + respRows.Send(&pproto.RowResponse{ //nolint:errcheck + Headers: []*pproto.ColumnInfo{ + { + Name: "count", + Datatype: "uint64", + }, + }, + Columns: []*pproto.ColumnResponse{ + { + ColumnVal: &pproto.ColumnResponse_Uint64Val{ + Uint64Val: res, + }, + }, + }, + }) + respRows.Send(pproto.EOF) //nolint:errcheck + }() + case bool: + go func() { + respRows.Send(&pproto.RowResponse{ //nolint:errcheck + Headers: []*pproto.ColumnInfo{ + { + Name: "result", + Datatype: "bool", + }, + }, + Columns: []*pproto.ColumnResponse{ + { + ColumnVal: &pproto.ColumnResponse_BoolVal{ + BoolVal: res, + }, + }, + }, + }) + respRows.Send(pproto.EOF) //nolint:errcheck + }() + default: + return nil, fmt.Errorf("unsupported result type %T", res) + } + + // Apply Reducers + result := respRows + for _, red := range mr.Reducers { + out := pproto.NewRowBuffer(0) + + // Run Reducers asyncronously. + // TODO: stop swallowing this error. + // TODO: does this need an EOF as input? + go red.Reduce(result, out) //nolint:errcheck + + result = out + } + + return result, nil +} + +type MappingResult struct { + IndexName string + ColumnIDs []uint64 + ColumnKeys []string + FieldFilters []string + Limit uint64 + Offset uint64 + Query string + Header []Column + Reducers []Reducer +} + +func (mr *MappingResult) addReducer(r Reducer) { + mr.Reducers = append(mr.Reducers, r) +} + +type SelectProperties struct { + Index *pilosa.Index + Fields []Column + Features selectFeatures + WherePQL string + WhereIDs []uint64 + WhereKeys []string + Offset uint + Limit uint + GroupByFieldNames []string + Having *HavingClause +} + +type selectFunc struct { + funcName FuncName + field *pilosa.Field +} + +type selectFeatures struct { + HasRowAttrs bool + HasColAttrs bool + funcs []selectFunc +} + +type HavingClause struct { + Subj string + Cond pql.Condition +} + +type handler interface { + Apply(*sqlparser.Select, QueryMask, func(string) *pilosa.Index) (*MappingResult, error) +} + +// handlerSelectFieldsFromTable: Inspect() +type handlerSelectFieldsFromTableWhere struct{} + +func (h handlerSelectFieldsFromTableWhere) Apply(stmt *sqlparser.Select, qm QueryMask, indexFunc func(string) *pilosa.Index) (*MappingResult, error) { + indexName, err := extractIndexName(stmt) + if err != nil { + return nil, errors.Wrapf(err, "extracting index name") + } + index := indexFunc(indexName) + if index == nil { + return nil, errors.WithMessage(pilosa.ErrIndexNotFound, indexName) + } + + var whereQuery string + if qm.HasWhere() { + whereQuery, err = extractWhere(index, stmt.Where.Expr) + if err != nil { + return nil, err + } + } else { + whereQuery = "All()" + } + + selectFields, _, err := extractSelectFields(index, stmt) + if err != nil { + return nil, errors.Wrap(err, "extracting select fields") + } + + var fields []string + for _, fld := range selectFields { + if _, ok := fld.(*StarColumn); ok { + pflds := index.Fields() + fields = []string{"_id"} + for _, f := range pflds { + name := f.Name() + if strings.HasPrefix(name, "_") { + continue + } + fields = append(fields, name) + } + break + } + fields = append(fields, fld.Name()) + } + for i, fld := range fields { + if fld == "_id" && i != 0 { + return nil, errors.New("_id can only be the first field in a select") + } + } + + limit, offset, err := extractLimitOffset(stmt) + if err != nil { + return nil, errors.Wrap(err, "extracting limit") + } + + orderByFlds, orderByDirs, err := extractOrderBy(stmt) + if err != nil { + return nil, errors.Wrap(err, "extracting order by") + } + + mr := &MappingResult{ + IndexName: indexName, + //FieldFilters: fields, + Header: selectFields, + } + + // TODO: assign headers + mr.addReducer(NewAssignHeadersReducer(selectFields)) + + // TODO: If both order and limit/offset are required, then + // we can't supply limit/offset to the InspectRequest; we + // have to get all records, which we don't want to do on + // a large data set. We need to come up with a better + // way to handle that situation. + switch { + case qm.HasOrderBy(): + mr.addReducer(NewOrderByReducer(orderByFlds, orderByDirs, limit, offset)) + case limit != 0: + whereQuery = Limit(whereQuery, limit, offset) + case offset != 0: + whereQuery = Offset(whereQuery, offset) + } + + if len(fields) > 0 && fields[0] == "_id" { + fields = fields[1:] + } + mr.Query = Extract(whereQuery, fields...) + + return mr, nil +} + +// handlerSelectDistinctFromTable: Rows, Rows(limit): select distinct fld from tbl +type handlerSelectDistinctFromTable struct{} + +func (h handlerSelectDistinctFromTable) Apply(stmt *sqlparser.Select, qm QueryMask, indexFunc func(string) *pilosa.Index) (*MappingResult, error) { + indexName, err := extractIndexName(stmt) + if err != nil { + return nil, errors.Wrapf(err, "extracting index name") + } + index := indexFunc(indexName) + if index == nil { + return nil, errors.WithMessage(pilosa.ErrIndexNotFound, indexName) + } + + selectFields, _, err := extractSelectFields(index, stmt) + if err != nil { + return nil, errors.Wrap(err, "extracting select fields") + } + + fieldCol, ok := selectFields[0].(*FieldColumn) + if !ok { + return nil, errors.New("distinct requires a valid field column") + } + + limit, offset, err := extractLimitOffset(stmt) + if err != nil { + return nil, errors.Wrap(err, "extracting limit") + } + + orderByFlds, orderByDirs, err := extractOrderBy(stmt) + if err != nil { + return nil, errors.Wrap(err, "extracting order by") + } + + // Determine the type of the field needing distinct. + // If the pilosa field is type int, handle it as a Distinct() query. + // Otherwise, use Rows() + // TODO: ensure this works for all field types (bool, time, etc). + var qo string + if fieldCol.Field.Type() == pilosa.FieldTypeInt { + qo = Distinct(fieldCol.Field.Index(), fieldCol.Field.Name()) + } else { + if !qm.HasOrderBy() && limit > 0 { + if qo, err = RowsLimit(fieldCol.Field.Name(), int64(limit)); err != nil { + return nil, errors.Wrap(err, "creating Rows query") + } + } else { + qo = Rows(fieldCol.Field.Name()) + } + } + + mr := &MappingResult{ + IndexName: indexName, + Header: selectFields, + Query: qo, + } + + mr.addReducer(NewAssignHeadersReducer(selectFields)) + if qm.HasOrderBy() { + mr.addReducer(NewOrderByReducer(orderByFlds, orderByDirs, limit, offset)) + } else { + mr.addReducer(NewLimitReducer(limit, offset)) + } + + return mr, nil +} + +// handlerSelectCountFromTableWhere: Count() +type handlerSelectCountFromTableWhere struct{} + +func (h handlerSelectCountFromTableWhere) Apply(stmt *sqlparser.Select, qm QueryMask, indexFunc func(string) *pilosa.Index) (*MappingResult, error) { + var qo string + var reducers []Reducer + + indexName, err := extractIndexName(stmt) + if err != nil { + return nil, errors.Wrapf(err, "extracting index name") + } + index := indexFunc(indexName) + if index == nil { + return nil, errors.WithMessage(pilosa.ErrIndexNotFound, indexName) + } + + selectFields, features, err := extractSelectFields(index, stmt) + if err != nil { + return nil, errors.Wrap(err, "extracting select fields") + } + + var wherePQL string + if stmt.Where != nil { + wherePQL, err = extractWhere(index, stmt.Where.Expr) + if err != nil { + return nil, err + } + } else { + wherePQL = All() + } + + funcs := features.funcs + if len(funcs) != 1 { + return nil, errors.New("handler does not support multiple functions") + } else if funcs[0].funcName != FuncCount { + return nil, errors.Errorf("handler expected func: %s", FuncCount) + } + + if funcs[0].field == nil { + qo = Count(wherePQL) + } else { + // TODO: add the Distinct (for Int fields) here (like we do in handlerSelectDistinctFromTable) + qo = Rows(funcs[0].field.Name()) + reducers = append(reducers, &CountIDReducer{}) + } + mr := &MappingResult{ + IndexName: indexName, + Header: selectFields, + Query: qo, + Reducers: reducers, + } + + mr.addReducer(NewAssignHeadersReducer(selectFields)) + // NOTE: limit and order by don't make sense in this handler + // because it just returns a single row. + + return mr, nil +} + +// handlerSelectFuncFromTableWhere: min(), max(), sum(), avg() +type handlerSelectFuncFromTableWhere struct{} + +func (h handlerSelectFuncFromTableWhere) Apply(stmt *sqlparser.Select, qm QueryMask, indexFunc func(string) *pilosa.Index) (*MappingResult, error) { + var qo string + + indexName, err := extractIndexName(stmt) + if err != nil { + return nil, errors.Wrapf(err, "extracting index name") + } + index := indexFunc(indexName) + if index == nil { + return nil, errors.WithMessage(pilosa.ErrIndexNotFound, indexName) + } + + selectFields, features, err := extractSelectFields(index, stmt) + if err != nil { + return nil, errors.Wrap(err, "extracting select fields") + } + + funcs := features.funcs + if len(funcs) != 1 { + return nil, errors.New("handler does not support multiple functions") + } + + funcField := funcs[0].field + if funcField == nil { + return nil, errors.New("function contains no field") + } + + var wherePQL string + if qm.HasWhere() { + wherePQL, err = extractWhere(index, stmt.Where.Expr) + if err != nil { + return nil, err + } + } + + limit, offset, err := extractLimitOffset(stmt) + if err != nil { + return nil, errors.Wrap(err, "extracting limit offset") + } + + orderByFlds, orderByDirs, err := extractOrderBy(stmt) + if err != nil { + return nil, errors.Wrap(err, "extracting order by") + } + + switch funcs[0].funcName { + case FuncMin: + qo = Min(funcField.Name(), wherePQL) + case FuncMax: + qo = Max(funcField.Name(), wherePQL) + case FuncAvg: + fallthrough + case FuncSum: + qo = Sum(funcField.Name(), wherePQL) + } + + mr := &MappingResult{ + IndexName: indexName, + Header: selectFields, + Query: qo, + } + + mr.addReducer(NewValCountFuncReducer(funcs[0].funcName)) + mr.addReducer(NewAssignHeadersReducer(selectFields)) + if qm.HasOrderBy() { + mr.addReducer(NewOrderByReducer(orderByFlds, orderByDirs, limit, offset)) + } else { + mr.addReducer(NewLimitReducer(limit, offset)) + } + + return mr, nil +} + +// handlerSelectGroupBy: GroupBy +type handlerSelectGroupBy struct{} + +func (h handlerSelectGroupBy) Apply(stmt *sqlparser.Select, qm QueryMask, indexFunc func(string) *pilosa.Index) (*MappingResult, error) { + var qo string + + indexName, err := extractIndexName(stmt) + if err != nil { + return nil, errors.Wrapf(err, "extracting index name") + } + index := indexFunc(indexName) + if index == nil { + return nil, errors.WithMessage(pilosa.ErrIndexNotFound, indexName) + } + + selectFields, features, err := extractSelectFields(index, stmt) + if err != nil { + return nil, errors.Wrap(err, "extracting select fields") + } + + orderByFlds, orderByDirs, err := extractOrderBy(stmt) + if err != nil { + return nil, errors.Wrap(err, "extracting order by") + } + + // If the query can be supported by TopN, + // i.e. if it's of the form: + // select fld, count(fld) as cnt from tbl group by fld order by cnt desc limit 1 + // select fld, count(fld) as cnt from tbl where fld2=1 group by fld order by cnt desc limit 1 + // then redirect it to handlerSelectIDCountFromTable. + // Otherwise, handle it as a normal GroupBy query. + // TODO: this level of inspection on the query needs to be built into + // an official query planner. The existing solution, which uses very + // broad masks to route the query to specific handlers, doesn't do + // this kind of finer-grain inspection of, for example, the order by + // fields themselves. + if func() bool { + if !qm.HasLimit() { + return false + } + if len(orderByFlds) != 1 { + return false + } + if orderByDirs[0] != "desc" { + return false + } + if qm == MustGenerateMask("select fld, count(fld) from tbl group by fld order by cnt limit 1") || + qm == MustGenerateMask("select fld, count(fld) from tbl where fld=1 group by fld order by cnt limit 1") { + // Check that the order-by field is the count field. + for i := range selectFields { + if s, ok := selectFields[i].(*FuncColumn); !ok { + continue + } else if s.FuncName == FuncCount && s.Alias() == orderByFlds[0] { + return true + } + } + } + return false + }() { + return handlerSelectIDCountFromTable{}.Apply(stmt, qm, indexFunc) + } + + groupByFieldNames, err := extractGroupByFieldNames(stmt.GroupBy) + if err != nil { + return nil, errors.Wrap(err, "extracting group by fields") + } + + having, err := extractHavingClause(stmt.Having) + if err != nil { + return nil, errors.Wrap(err, "extracting having clause") + } + + rowsQueries := []string{} + for _, fieldName := range groupByFieldNames { + field := index.Field(fieldName) + rowsQueries = append(rowsQueries, Rows(field.Name())) + } + + var wherePQL string + if stmt.Where != nil { + wherePQL, err = extractWhere(index, stmt.Where.Expr) + if err != nil { + return nil, errors.Wrap(err, "extracting where") + } + } + + limit, offset, err := extractLimitOffset(stmt) + if err != nil { + return nil, errors.Wrap(err, "extracting limit offset") + } + + // Group by queries can any combination of count() and sum() + // in the select fields. + var idxSum int = -1 + funcs := features.funcs + for i := range funcs { + switch funcs[i].funcName { + case FuncSum: + idxSum = i + } + } + + var sumQuery string + if idxSum >= 0 { + sumQuery = Sum(funcs[idxSum].field.Name(), "") + } + + var havingQuery string + if having != nil { + havingQuery = fmt.Sprintf("Condition(%s)", having.Cond.StringWithSubj(having.Subj)) + } + + qo, err = GroupByBase(rowsQueries, int64(limit+offset), wherePQL, sumQuery, havingQuery) + if err != nil { + return nil, err + } + + mr := &MappingResult{ + IndexName: indexName, + Header: selectFields, + Query: qo, + } + + mr.addReducer(NewAssignHeadersReducer(selectFields)) + if qm.HasOrderBy() { + mr.addReducer(NewOrderByReducer(orderByFlds, orderByDirs, limit, offset)) + } else { + mr.addReducer(NewLimitReducer(limit, offset)) + } + + return mr, nil +} + +// handlerSelectIDCountFromTable: TopN +type handlerSelectIDCountFromTable struct{} + +func (f handlerSelectIDCountFromTable) Apply(stmt *sqlparser.Select, qm QueryMask, indexFunc func(string) *pilosa.Index) (*MappingResult, error) { + var qo string + + indexName, err := extractIndexName(stmt) + if err != nil { + return nil, errors.Wrapf(err, "extracting index name") + } + index := indexFunc(indexName) + if index == nil { + return nil, errors.WithMessage(pilosa.ErrIndexNotFound, indexName) + } + + selectFields, features, err := extractSelectFields(index, stmt) + if err != nil { + return nil, errors.Wrap(err, "extracting select fields") + } + + var wherePQL string + if stmt.Where != nil { + wherePQL, err = extractWhere(index, stmt.Where.Expr) + if err != nil { + return nil, errors.Wrap(err, "extracting where") + } + } + + limit, offset, err := extractLimitOffset(stmt) + if err != nil { + return nil, errors.Wrap(err, "extracting limit offset") + } + + funcs := features.funcs + if len(funcs) != 1 { + return nil, errors.New("handler does not support multiple functions") + } else if funcs[0].funcName != FuncCount { + return nil, errors.Errorf("handler expected func: %s", FuncCount) + } + + if wherePQL == "" { + qo = TopN(funcs[0].field.Name(), uint64(limit+offset)) + } else { + qo = RowTopN(funcs[0].field.Name(), uint64(limit+offset), wherePQL) + } + + mr := &MappingResult{ + IndexName: indexName, + Header: selectFields, + Query: qo, + } + + mr.addReducer(NewAssignHeadersReducer(selectFields)) + mr.addReducer(NewLimitReducer(limit, offset)) + // TODO: order by is not implemented on this method because order desc + // is handled in pilosa TopN. In order to support asc here, we would + // have to return the entire TopN cache. Instead, we should consider + // supported something like this in Pilosa itself. + + return mr, nil +} + +// handlerSelectJoin: Join/Distinct() +type handlerSelectJoin struct{} + +func (h handlerSelectJoin) Apply(stmt *sqlparser.Select, qm QueryMask, indexFunc func(string) *pilosa.Index) (*MappingResult, error) { + var qo string + + pts, err := extractJoinTables(stmt) + if err != nil { + return nil, errors.Wrap(err, "extracting join tables") + } + + primary := pts.primary() + secondary := pts.secondary() + + primaryIndexName := primary.name + primaryIndex := indexFunc(primaryIndexName) + primaryField := primaryIndex.Field(primary.column.name) + + secondaryIndexName := secondary.name + secondaryIndex := indexFunc(secondaryIndexName) + secondaryField := secondaryIndex.Field(secondary.column.name) + + var wheres tableWheres + if qm.HasWhere() { + indexes := []*pilosa.Index{primaryIndex, secondaryIndex} + wheres, err = extractWheres(indexes, pts, stmt.Where.Expr) + if err != nil { + return nil, err + } + } + + var primaryWhere string + var secondaryWhere string + for i, w := range wheres { + switch w.table.index { + case primaryIndex: + primaryWhere = wheres[i].where + case secondaryIndex: + secondaryWhere = wheres[i].where + } + } + + // Build the Distinct() portion of the query on the secondary. + var distinctQry string + if secondaryWhere == "" { + distinctQry = Distinct(secondaryField.Index(), secondaryField.Name()) + } else { + distinctQry = RowDistinct(secondaryField.Index(), secondaryField.Name(), secondaryWhere) + } + + var rowQry string + if primaryWhere == "" { + rowQry = Intersect(All(), distinctQry) + } else { + _ = primaryField + rowQry = Intersect(primaryWhere, distinctQry) + } + + selectFields, _, err := extractSelectFields(primaryIndex, stmt) + if err != nil { + return nil, errors.Wrap(err, "extracting select fields") + } + + orderByFlds, orderByDirs, err := extractOrderBy(stmt) + if err != nil { + return nil, errors.Wrap(err, "extracting order by") + } + + if qm.HasSelectPart(SelectPartCountStar) { + qo = Count(rowQry) + } else { + qo = rowQry + } + + limit, offset, err := extractLimitOffset(stmt) + if err != nil { + return nil, errors.Wrap(err, "extracting limit") + } + + mr := &MappingResult{ + IndexName: primaryIndex.Name(), + Header: selectFields, + Query: qo, + } + + mr.addReducer(NewAssignHeadersReducer(selectFields)) + if qm.HasOrderBy() { + mr.addReducer(NewOrderByReducer(orderByFlds, orderByDirs, limit, offset)) + } else { + mr.addReducer(NewLimitReducer(limit, offset)) + } + + return mr, nil +} From 5849794a1b4d6ec961319e63ebaedae0d2371d1a Mon Sep 17 00:00:00 2001 From: Ben Johnson Date: Fri, 14 Aug 2020 08:58:59 -0600 Subject: [PATCH 02/17] Add RBF WAL write cache --- rbf/wal.go | 29 +++++++++++++++++++++++++---- 1 file changed, 25 insertions(+), 4 deletions(-) diff --git a/rbf/wal.go b/rbf/wal.go index 0f8461e16..fae0dc713 100644 --- a/rbf/wal.go +++ b/rbf/wal.go @@ -29,6 +29,7 @@ type WALSegment struct { path string // path to file w *os.File // write handle data []byte // read-only mmap data + buf []byte // write buffer pageN int // number of written pages } @@ -121,6 +122,12 @@ func (s *WALSegment) Close() error { // CloseForWrite closes the write handle, if initialized. func (s *WALSegment) CloseForWrite() error { + // Ensure write buffer is flushed out. + if err := s.Sync(); err != nil { + return err + } + + // Close underlying file writer. if s.w != nil { if err := s.w.Close(); err != nil { return err @@ -138,6 +145,15 @@ func (s *WALSegment) ReadWALPage(walID int64) ([]byte, error) { } offset := (walID - s.minWALID) * PageSize + + // If offset is within write buffer, return from write buffer. + writeBufferOffset := int64((s.pageN * PageSize) - len(s.buf)) + if offset >= writeBufferOffset { + buf := s.buf[offset-writeBufferOffset:] + return buf[:PageSize:PageSize], nil + } + + // Otherwise return from on-disk mmap. return s.data[offset : offset+PageSize], nil } @@ -161,10 +177,8 @@ func (s *WALSegment) WriteWALPage(page []byte, isMeta bool) (walID int64, err er // TODO: Write meta page checksum } - // Write page at position & increment page count. - if _, err := s.w.WriteAt(page, int64(s.pageN*PageSize)); err != nil { - return 0, fmt.Errorf("wal segment write: %w", err) - } + // Append write to write buffer & increment page count. + s.buf = append(s.buf, page...) s.pageN++ return walID, nil @@ -175,6 +189,13 @@ func (s *WALSegment) Sync() error { if s.w == nil { return nil } + + // Flush buffer to disk. + if _, err := s.w.WriteAt(s.buf, int64((s.pageN*PageSize)-len(s.buf))); err != nil { + return fmt.Errorf("wal segment write: %w", err) + } + s.buf = nil + return s.w.Sync() } From 51504f4fe4aa3809202aaebe1ec177d4fab84316 Mon Sep 17 00:00:00 2001 From: Ben Johnson Date: Tue, 18 Aug 2020 08:41:30 -0600 Subject: [PATCH 03/17] Add WAL write cache mutex; update name; add benchmarks --- rbf/wal.go | 48 ++++++++++++++++++++++++++++++++---------------- rbf/wal_test.go | 44 ++++++++++++++++++++++++++++++++++++++++++-- 2 files changed, 74 insertions(+), 18 deletions(-) diff --git a/rbf/wal.go b/rbf/wal.go index fae0dc713..b0fa9fbd0 100644 --- a/rbf/wal.go +++ b/rbf/wal.go @@ -18,6 +18,7 @@ import ( "fmt" "os" "path/filepath" + "sync" "syscall" "github.com/pilosa/pilosa/v2/syswrap" @@ -25,12 +26,13 @@ import ( // WALSegment represents a single file in the WAL. type WALSegment struct { - minWALID int64 // base WALID; calculated from path - path string // path to file - w *os.File // write handle - data []byte // read-only mmap data - buf []byte // write buffer - pageN int // number of written pages + mu sync.RWMutex + minWALID int64 // base WALID; calculated from path + path string // path to file + w *os.File // write handle + data []byte // read-only mmap data + writeCache []byte // write buffer + pageN int // number of written pages } // NewWALSegment returns a new instance of WALSegment for a given path. @@ -139,6 +141,9 @@ func (s *WALSegment) CloseForWrite() error { // ReadWALPage reads a single page at the given WAL ID. func (s *WALSegment) ReadWALPage(walID int64) ([]byte, error) { + s.mu.RLock() + defer s.mu.RUnlock() + // Ensure requested ID is contained in this file. if walID < s.minWALID || walID > s.minWALID+int64(s.pageN) { return nil, fmt.Errorf("wal segment page read out of range: id=%d base=%d pageN=%d", walID, s.minWALID, s.pageN) @@ -147,9 +152,9 @@ func (s *WALSegment) ReadWALPage(walID int64) ([]byte, error) { offset := (walID - s.minWALID) * PageSize // If offset is within write buffer, return from write buffer. - writeBufferOffset := int64((s.pageN * PageSize) - len(s.buf)) + writeBufferOffset := int64((s.pageN * PageSize) - len(s.writeCache)) if offset >= writeBufferOffset { - buf := s.buf[offset-writeBufferOffset:] + buf := s.writeCache[offset-writeBufferOffset:] return buf[:PageSize:PageSize], nil } @@ -161,6 +166,9 @@ func (s *WALSegment) ReadWALPage(walID int64) ([]byte, error) { func (s *WALSegment) WriteWALPage(page []byte, isMeta bool) (walID int64, err error) { assert(len(page) == PageSize, "invalid page size: %d", len(page)) + s.mu.Lock() + defer s.mu.Unlock() + // Initialize write file handle if not yet initialized. if s.w == nil { if s.w, err = os.OpenFile(s.path, os.O_WRONLY, 0666); err != nil { @@ -178,24 +186,32 @@ func (s *WALSegment) WriteWALPage(page []byte, isMeta bool) (walID int64, err er } // Append write to write buffer & increment page count. - s.buf = append(s.buf, page...) + s.writeCache = append(s.writeCache, page...) s.pageN++ return walID, nil } -// Sync flushes all changes to disk. +// Flush flushes the write buffer to the OS cache. +func (s *WALSegment) Flush() error { + s.mu.Lock() + defer s.mu.Unlock() + + if _, err := s.w.WriteAt(s.writeCache, int64((s.pageN*PageSize)-len(s.writeCache))); err != nil { + return fmt.Errorf("wal segment write: %w", err) + } + s.writeCache = nil + return nil +} + +// Sync flushes the write buffer and invokes a file sync to flush data to disk. func (s *WALSegment) Sync() error { if s.w == nil { return nil } - - // Flush buffer to disk. - if _, err := s.w.WriteAt(s.buf, int64((s.pageN*PageSize)-len(s.buf))); err != nil { - return fmt.Errorf("wal segment write: %w", err) + if err := s.Flush(); err != nil { + return err } - s.buf = nil - return s.w.Sync() } diff --git a/rbf/wal_test.go b/rbf/wal_test.go index 4511578c8..b41e7e792 100644 --- a/rbf/wal_test.go +++ b/rbf/wal_test.go @@ -16,7 +16,7 @@ package rbf_test import ( "bytes" - "encoding/hex" + "encoding/hex" "io/ioutil" "math/rand" "os" @@ -26,7 +26,6 @@ import ( "github.com/pilosa/pilosa/v2/rbf" ) - func TestWALSegment_Open(t *testing.T) { t.Run("OK", func(t *testing.T) { s := MustOpenWALSegment(t, 10) @@ -108,6 +107,47 @@ func TestParseWALSegmentPath(t *testing.T) { }) } +func BenchmarkWALSegment_WriteWALPage(b *testing.B) { + b.Run("8KB", func(b *testing.B) { benchmarkWALSegment_WriteWALPage(b, 8*(1<<10)) }) + b.Run("16KB", func(b *testing.B) { benchmarkWALSegment_WriteWALPage(b, 16*(1<<10)) }) + b.Run("64KB", func(b *testing.B) { benchmarkWALSegment_WriteWALPage(b, 64*(1<<10)) }) + b.Run("256KB", func(b *testing.B) { benchmarkWALSegment_WriteWALPage(b, 256*(1<<10)) }) + b.Run("1MB", func(b *testing.B) { benchmarkWALSegment_WriteWALPage(b, (1 << 20)) }) + b.Run("10MB", func(b *testing.B) { benchmarkWALSegment_WriteWALPage(b, 10*(1<<20)) }) +} + +func benchmarkWALSegment_WriteWALPage(b *testing.B, flushSize int) { + page := make([]byte, rbf.PageSize) + + for i := 0; i < b.N; i++ { + func() { + s := MustOpenWALSegment(b, 0) + defer MustCloseWALSegment(b, s) + + // Fill the segment but stop after each flush interval to flush the write buffer. + for j := 0; j < rbf.MaxWALSegmentFileSize; j += rbf.PageSize { + if _, err := s.WriteWALPage(page, false); err != nil { + b.Fatal(err) + } + + // Flush write buffer. + if j != 0 && j%flushSize == 0 { + if err := s.Flush(); err != nil { + b.Fatal(err) + } + } + } + + // Fsync to disk at the end. + if err := s.Sync(); err != nil { + b.Fatal(err) + } + }() + } + + b.SetBytes(rbf.MaxWALSegmentFileSize) +} + // MustOpenWALSegment opens a WAL segment in a temporary path. Fails on error. func MustOpenWALSegment(tb testing.TB, walID int64) *rbf.WALSegment { tb.Helper() From 9dbb82cf3e72ce8c9c93d1bbe3d891db82f5cfc7 Mon Sep 17 00:00:00 2001 From: Ben Johnson Date: Wed, 19 Aug 2020 08:42:53 -0600 Subject: [PATCH 04/17] WAL mutex fixes --- rbf/wal.go | 47 +++++++++++++++++++++++++++++++++++++++++------ 1 file changed, 41 insertions(+), 6 deletions(-) diff --git a/rbf/wal.go b/rbf/wal.go index b0fa9fbd0..aee4b07a3 100644 --- a/rbf/wal.go +++ b/rbf/wal.go @@ -46,20 +46,37 @@ func NewWALSegment(path string) *WALSegment { func (s *WALSegment) Path() string { return s.path } // MinWALID returns the initial WAL ID of the segment. Only available after Open(). -func (s *WALSegment) MinWALID() int64 { return s.minWALID } +func (s *WALSegment) MinWALID() int64 { + s.mu.RLock() + defer s.mu.RUnlock() + return s.minWALID +} // MaxWALID returns the maximum WAL ID of the segment. Only available after Open(). func (s *WALSegment) MaxWALID() int64 { + s.mu.RLock() + defer s.mu.RUnlock() return s.minWALID + int64(s.pageN) - 1 } // PageN returns the number of pages in the segment. -func (s *WALSegment) PageN() int { return s.pageN } +func (s *WALSegment) PageN() int { + s.mu.RLock() + defer s.mu.RUnlock() + return s.pageN +} // Size returns the current size of the segment, in bytes. -func (s *WALSegment) Size() int64 { return int64(s.pageN) * PageSize } +func (s *WALSegment) Size() int64 { + s.mu.RLock() + defer s.mu.RUnlock() + return int64(s.pageN) * PageSize +} func (s *WALSegment) Open() (err error) { + s.mu.Lock() + defer s.mu.Unlock() + // Extract base WAL ID and validate path. if s.minWALID, err = ParseWALSegmentPath(s.path); err != nil { return err @@ -110,7 +127,10 @@ func (s *WALSegment) Open() (err error) { // Close closes the write handle and the read-only mmap. func (s *WALSegment) Close() error { - if err := s.CloseForWrite(); err != nil { + s.mu.Lock() + defer s.mu.Unlock() + + if err := s.closeForWrite(); err != nil { return err } if s.data != nil { @@ -124,8 +144,14 @@ func (s *WALSegment) Close() error { // CloseForWrite closes the write handle, if initialized. func (s *WALSegment) CloseForWrite() error { + s.mu.Lock() + defer s.mu.Unlock() + return s.closeForWrite() +} + +func (s *WALSegment) closeForWrite() error { // Ensure write buffer is flushed out. - if err := s.Sync(); err != nil { + if err := s.sync(); err != nil { return err } @@ -196,7 +222,10 @@ func (s *WALSegment) WriteWALPage(page []byte, isMeta bool) (walID int64, err er func (s *WALSegment) Flush() error { s.mu.Lock() defer s.mu.Unlock() + return s.flush() +} +func (s *WALSegment) flush() error { if _, err := s.w.WriteAt(s.writeCache, int64((s.pageN*PageSize)-len(s.writeCache))); err != nil { return fmt.Errorf("wal segment write: %w", err) } @@ -206,10 +235,16 @@ func (s *WALSegment) Flush() error { // Sync flushes the write buffer and invokes a file sync to flush data to disk. func (s *WALSegment) Sync() error { + s.mu.Lock() + defer s.mu.Unlock() + return s.sync() +} + +func (s *WALSegment) sync() error { if s.w == nil { return nil } - if err := s.Flush(); err != nil { + if err := s.flush(); err != nil { return err } return s.w.Sync() From 0a96043c23e2a3d5c620f0846c92f9181e4c7354 Mon Sep 17 00:00:00 2001 From: Alan Bernstein Date: Mon, 17 Aug 2020 22:16:47 -0500 Subject: [PATCH 05/17] Use _buffer instead of buffer --- pql/parser.go | 4 +- pql/parser_test.go | 19 ++++++++++ pql/pql.peg | 8 ++-- pql/pql.peg.go | 93 +++++++++++++++++----------------------------- pql/pqlpeg_test.go | 8 +--- 5 files changed, 60 insertions(+), 72 deletions(-) diff --git a/pql/parser.go b/pql/parser.go index 8a5907dbf..e09b431ef 100644 --- a/pql/parser.go +++ b/pql/parser.go @@ -61,9 +61,7 @@ func (p *parser) Parse() (*Query, error) { p.PQL = PQL{ Buffer: string(buf), } - if err := p.Init(); err != nil { - return nil, errors.Wrap(err, "initializing") - } + p.Init() err = p.PQL.Parse() if err != nil { return nil, errors.Wrap(err, "parsing") diff --git a/pql/parser_test.go b/pql/parser_test.go index 8e5cf8cc4..94d97d1a1 100644 --- a/pql/parser_test.go +++ b/pql/parser_test.go @@ -194,6 +194,25 @@ func TestParser_Parse(t *testing.T) { } }) + // Parse unicode keys + t.Run("UnicodeKey", func(t *testing.T) { + // s := `���t` + s := `Æ` + q, err := pql.ParseString(`Row(unicode="` + s + `")`) + if err != nil { + t.Fatal(err) + } else if !reflect.DeepEqual(q.Calls[0], + &pql.Call{ + Name: "Row", + Args: map[string]interface{}{ + "unicode": s, + }, + }, + ) { + t.Fatalf("uexpected call: %#v", q.Calls[0]) + } + }) + } func TestUnquote(t *testing.T) { diff --git a/pql/pql.peg b/pql/pql.peg index 7e6c039bd..dcb10d26d 100644 --- a/pql/pql.peg +++ b/pql/pql.peg @@ -62,10 +62,10 @@ itema <- ( 'null' &(comma / sp close) { p.addVal(nil) } / 'false' &(comma / sp close) { p.addVal(false) } / timestampfmt { p.addVal(buffer[begin:end]) } ) -itemb <- ( < IDENT > { p.startCall(buffer[begin:end]) } open allargs comma? close { p.addVal(p.endCall()) } - / < ([[A-Z]] / [0-9] / '-' / '_' / ':')+ > { p.addVal(buffer[begin:end]) } - / < '"' doublequotedstring '"' > { p.addVal(buffer[begin:end]) } - / < '\'' singlequotedstring '\'' > { p.addVal(buffer[begin:end]) } +itemb <- ( < IDENT > { p.startCall(string(_buffer[begin:end])) } open allargs comma? close { p.addVal(p.endCall()) } + / < ([[A-Z]] / [0-9] / '-' / '_' / ':')+ > { p.addVal(string(_buffer[begin:end])) } + / < '"' doublequotedstring '"' > { p.addVal(string(_buffer[begin:end])) } + / < '\'' singlequotedstring '\'' > { p.addVal(string(_buffer[begin:end])) } ) float <- ( < '-'? [0-9]+ ('.'[0-9]*)? > { p.addNumVal(buffer[begin:end], true) } / < '-'? '.'[0-9]+ > { p.addNumVal(buffer[begin:end], true) } diff --git a/pql/pql.peg.go b/pql/pql.peg.go index 3cba1b702..e16002005 100644 --- a/pql/pql.peg.go +++ b/pql/pql.peg.go @@ -1,11 +1,10 @@ package pql -// Code generated by peg -inline pql.peg DO NOT EDIT. +//go:generate peg -inline pql.peg import ( "fmt" - "io" - "os" + "math" "sort" "strconv" ) @@ -245,19 +244,19 @@ type node32 struct { up, next *node32 } -func (node *node32) print(w io.Writer, pretty bool, buffer string) { +func (node *node32) print(pretty bool, buffer string) { var print func(node *node32, depth int) print = func(node *node32, depth int) { for node != nil { for c := 0; c < depth; c++ { - fmt.Fprintf(w, " ") + fmt.Printf(" ") } rule := rul3s[node.pegRule] quote := strconv.Quote(string(([]rune(buffer)[node.begin:node.end]))) if !pretty { - fmt.Fprintf(w, "%v %v\n", rule, quote) + fmt.Printf("%v %v\n", rule, quote) } else { - fmt.Fprintf(w, "\x1B[34m%v\x1B[m %v\n", rule, quote) + fmt.Printf("\x1B[34m%v\x1B[m %v\n", rule, quote) } if node.up != nil { print(node.up, depth+1) @@ -268,12 +267,12 @@ func (node *node32) print(w io.Writer, pretty bool, buffer string) { print(node, 0) } -func (node *node32) Print(w io.Writer, buffer string) { - node.print(w, false, buffer) +func (node *node32) Print(buffer string) { + node.print(false, buffer) } -func (node *node32) PrettyPrint(w io.Writer, buffer string) { - node.print(w, true, buffer) +func (node *node32) PrettyPrint(buffer string) { + node.print(true, buffer) } type tokens32 struct { @@ -316,24 +315,24 @@ func (t *tokens32) AST() *node32 { } func (t *tokens32) PrintSyntaxTree(buffer string) { - t.AST().Print(os.Stdout, buffer) -} - -func (t *tokens32) WriteSyntaxTree(w io.Writer, buffer string) { - t.AST().Print(w, buffer) + t.AST().Print(buffer) } func (t *tokens32) PrettyPrintSyntaxTree(buffer string) { - t.AST().PrettyPrint(os.Stdout, buffer) + t.AST().PrettyPrint(buffer) } func (t *tokens32) Add(rule pegRule, begin, end, index uint32) { - tree, i := t.tree, int(index) - if i >= len(tree) { - t.tree = append(tree, token32{pegRule: rule, begin: begin, end: end}) - return + if tree := t.tree; int(index) >= len(tree) { + expanded := make([]token32, 2*len(tree)) + copy(expanded, tree) + t.tree = expanded + } + t.tree[index] = token32{ + pegRule: rule, + begin: begin, + end: end, } - tree[i] = token32{pegRule: rule, begin: begin, end: end} } func (t *tokens32) Tokens() []token32 { @@ -397,7 +396,7 @@ type parseError struct { } func (e *parseError) Error() string { - tokens, err := []token32{e.max}, "\n" + tokens, error := []token32{e.max}, "\n" positions, p := make([]int, 2*len(tokens)), 0 for _, token := range tokens { positions[p], p = int(token.begin), p+1 @@ -410,14 +409,14 @@ func (e *parseError) Error() string { } for _, token := range tokens { begin, end := int(token.begin), int(token.end) - err += fmt.Sprintf(format, + error += fmt.Sprintf(format, rul3s[token.pegRule], translations[begin].line, translations[begin].symbol, translations[end].line, translations[end].symbol, strconv.Quote(string(e.p.buffer[begin:end]))) } - return err + return error } func (p *PQL) PrintSyntaxTree() { @@ -428,10 +427,6 @@ func (p *PQL) PrintSyntaxTree() { } } -func (p *PQL) WriteSyntaxTree(w io.Writer) { - p.tokens32.WriteSyntaxTree(w, p.Buffer) -} - func (p *PQL) Execute() { buffer, _buffer, text, begin, end := p.Buffer, p.buffer, "", 0, 0 for _, token := range p.Tokens() { @@ -530,15 +525,15 @@ func (p *PQL) Execute() { case ruleAction43: p.addVal(buffer[begin:end]) case ruleAction44: - p.startCall(buffer[begin:end]) + p.startCall(string(_buffer[begin:end])) case ruleAction45: p.addVal(p.endCall()) case ruleAction46: - p.addVal(buffer[begin:end]) + p.addVal(string(_buffer[begin:end])) case ruleAction47: - p.addVal(buffer[begin:end]) + p.addVal(string(_buffer[begin:end])) case ruleAction48: - p.addVal(buffer[begin:end]) + p.addVal(string(_buffer[begin:end])) case ruleAction49: p.addNumVal(buffer[begin:end], true) case ruleAction50: @@ -571,31 +566,12 @@ func (p *PQL) Execute() { _, _, _, _, _ = buffer, _buffer, text, begin, end } -func Pretty(pretty bool) func(*PQL) error { - return func(p *PQL) error { - p.Pretty = pretty - return nil - } -} - -func Size(size int) func(*PQL) error { - return func(p *PQL) error { - p.tokens32 = tokens32{tree: make([]token32, 0, size)} - return nil - } -} -func (p *PQL) Init(options ...func(*PQL) error) error { +func (p *PQL) Init() { var ( max token32 position, tokenIndex uint32 buffer []rune ) - for _, option := range options { - err := option(p) - if err != nil { - return err - } - } p.reset = func() { max = token32{} position, tokenIndex = 0, 0 @@ -609,7 +585,7 @@ func (p *PQL) Init(options ...func(*PQL) error) error { p.reset() _rules := p.rules - tree := p.tokens32 + tree := tokens32{tree: make([]token32, math.MaxInt16)} p.parse = func(rule ...int) error { r := 1 if len(rule) > 0 { @@ -3631,15 +3607,15 @@ func (p *PQL) Init(options ...func(*PQL) error) error { nil, /* 86 Action43 <- <{ p.addVal(buffer[begin:end]) }> */ nil, - /* 87 Action44 <- <{ p.startCall(buffer[begin:end]) }> */ + /* 87 Action44 <- <{ p.startCall(string(_buffer[begin:end])) }> */ nil, /* 88 Action45 <- <{ p.addVal(p.endCall()) }> */ nil, - /* 89 Action46 <- <{ p.addVal(buffer[begin:end]) }> */ + /* 89 Action46 <- <{ p.addVal(string(_buffer[begin:end])) }> */ nil, - /* 90 Action47 <- <{ p.addVal(buffer[begin:end]) }> */ + /* 90 Action47 <- <{ p.addVal(string(_buffer[begin:end])) }> */ nil, - /* 91 Action48 <- <{ p.addVal(buffer[begin:end]) }> */ + /* 91 Action48 <- <{ p.addVal(string(_buffer[begin:end])) }> */ nil, /* 92 Action49 <- <{ p.addNumVal(buffer[begin:end], true) }> */ nil, @@ -3669,5 +3645,4 @@ func (p *PQL) Init(options ...func(*PQL) error) error { nil, } p.rules = _rules - return nil } diff --git a/pql/pqlpeg_test.go b/pql/pqlpeg_test.go index 4ebf1dcc9..e6aba444d 100644 --- a/pql/pqlpeg_test.go +++ b/pql/pqlpeg_test.go @@ -24,9 +24,7 @@ import ( func TestPEG(t *testing.T) { p := PQL{Buffer: ` SetBit(Union(Zitmap(row==4), Intersect(Qitmap(blah>4), Ritmap(field="http://zoo9.com=\\'hello' and \"hello\"")), Hitmap(row=ag-bee)), a="4z", b=5) Count(Union(Witmap(row=5.73, frame=.10), Row(zztop><[2, 9]))) TopN(blah, fields=["hello", "goodbye", "zero"])`[1:]} - if err := p.Init(); err != nil { - t.Fatalf("initialization error: %v", err) - } + p.Init() err := p.Parse() if err != nil { t.Fatalf("parse error: %v", err) @@ -34,9 +32,7 @@ SetBit(Union(Zitmap(row==4), Intersect(Qitmap(blah>4), Ritmap(field="http://zoo9 p.Execute() p = PQL{Buffer: `SetRowAttrs(attr="http://zoo9.com=\\'hello' "and \"hello\"")`} - if err := p.Init(); err != nil { - t.Fatalf("initialization error: %v", err) - } + p.Init() err = p.Parse() if err == nil { t.Fatalf("should have been an error because of the interior unescaped double quote") From 6fc5465aaa0f8fb8f08d5dfa8a764c0689c5c339 Mon Sep 17 00:00:00 2001 From: Alan Bernstein Date: Wed, 19 Aug 2020 16:19:43 -0500 Subject: [PATCH 06/17] Use rune slice in all cases, add tests --- pql/parser_test.go | 20 ---------- pql/pql.peg | 46 +++++++++++----------- pql/pql.peg.go | 98 +++++++++++++++++++++++----------------------- pql/pqlpeg_test.go | 53 +++++++++++++++++++++++++ 4 files changed, 125 insertions(+), 92 deletions(-) diff --git a/pql/parser_test.go b/pql/parser_test.go index 94d97d1a1..d2ee45748 100644 --- a/pql/parser_test.go +++ b/pql/parser_test.go @@ -193,26 +193,6 @@ func TestParser_Parse(t *testing.T) { t.Fatalf("unexpected call: %#v", q.Calls[0]) } }) - - // Parse unicode keys - t.Run("UnicodeKey", func(t *testing.T) { - // s := `���t` - s := `Æ` - q, err := pql.ParseString(`Row(unicode="` + s + `")`) - if err != nil { - t.Fatal(err) - } else if !reflect.DeepEqual(q.Calls[0], - &pql.Call{ - Name: "Row", - Args: map[string]interface{}{ - "unicode": s, - }, - }, - ) { - t.Fatalf("uexpected call: %#v", q.Calls[0]) - } - }) - } func TestUnquote(t *testing.T) { diff --git a/pql/pql.peg b/pql/pql.peg index dcb10d26d..b58e36723 100644 --- a/pql/pql.peg +++ b/pql/pql.peg @@ -14,8 +14,8 @@ Call <- 'Set' {p.startCall("Set")} open col comma dargs (comma timestamp)? clos / 'Store' {p.startCall("Store")} open Call comma darg close {p.endCall()} / 'TopN' {p.startCall("TopN")} open posfield (comma allargs)? close {p.endCall()} / 'Rows' {p.startCall("Rows")} open posfield (comma allargs)? close {p.endCall()} - / 'Range' {p.startCall("Range")} open field sp '=' sp fvalue comma 'from='? {p.addField("from")} timestampfmt {p.addVal(buffer[begin:end])} comma 'to='? sp {p.addField("to")} timestampfmt {p.addVal(buffer[begin:end])} close {p.endCall()} - / < IDENT > { p.startCall(buffer[begin:end] ) } open allargs comma? close { p.endCall() } + / 'Range' {p.startCall("Range")} open field sp '=' sp fvalue comma 'from='? {p.addField("from")} timestampfmt {p.addVal(text)} comma 'to='? sp {p.addField("to")} timestampfmt {p.addVal(text)} close {p.endCall()} + / < IDENT > { p.startCall(text ) } open allargs comma? close { p.endCall() } allargs <- Call (comma Call)* (comma dargs)? / dargs / sp fargs <- farg (comma fargs)? sp farg <- ( field sp '=' sp fvalue @@ -37,9 +37,9 @@ COND <- ( '><' { p.addBTWN() } ) conditional <- {p.startConditional()} condint condLT condfield condLT condint {p.endConditional()} -condint <- < '-'? [0-9]* '.' [0-9]+ / '0' / '-'? [1-9] [0-9]* > sp {p.condAdd(buffer[begin:end])} -condLT <- <('<=' / '<')> sp {p.condAdd(buffer[begin:end])} -condfield <- sp {p.condAdd(buffer[begin:end])} +condint <- < '-'? [0-9]* '.' [0-9]+ / '0' / '-'? [1-9] [0-9]* > sp {p.condAdd(text)} +condLT <- <('<=' / '<')> sp {p.condAdd(text)} +condfield <- sp {p.condAdd(text)} dvalue <- ( ditem / lbrack { p.startList() } dlist rbrack { p.endList() } @@ -60,35 +60,35 @@ fitem <- ( itema itema <- ( 'null' &(comma / sp close) { p.addVal(nil) } / 'true' &(comma / sp close) { p.addVal(true) } / 'false' &(comma / sp close) { p.addVal(false) } - / timestampfmt { p.addVal(buffer[begin:end]) } + / timestampfmt { p.addVal(text) } ) -itemb <- ( < IDENT > { p.startCall(string(_buffer[begin:end])) } open allargs comma? close { p.addVal(p.endCall()) } - / < ([[A-Z]] / [0-9] / '-' / '_' / ':')+ > { p.addVal(string(_buffer[begin:end])) } - / < '"' doublequotedstring '"' > { p.addVal(string(_buffer[begin:end])) } - / < '\'' singlequotedstring '\'' > { p.addVal(string(_buffer[begin:end])) } +itemb <- ( < IDENT > { p.startCall(text) } open allargs comma? close { p.addVal(p.endCall()) } + / < ([[A-Z]] / [0-9] / '-' / '_' / ':')+ > { p.addVal(text) } + / < '"' doublequotedstring '"' > { p.addVal(text) } + / < '\'' singlequotedstring '\'' > { p.addVal(text) } ) -float <- ( < '-'? [0-9]+ ('.'[0-9]*)? > { p.addNumVal(buffer[begin:end], true) } - / < '-'? '.'[0-9]+ > { p.addNumVal(buffer[begin:end], true) } +float <- ( < '-'? [0-9]+ ('.'[0-9]*)? > { p.addNumVal(text, true) } + / < '-'? '.'[0-9]+ > { p.addNumVal(text, true) } ) -decimal <- ( < '-'? [0-9]+ ('.'[0-9]*)? > { p.addNumVal(buffer[begin:end], false) } - / < '-'? '.'[0-9]+ > { p.addNumVal(buffer[begin:end], false) } +decimal <- ( < '-'? [0-9]+ ('.'[0-9]*)? > { p.addNumVal(text, false) } + / < '-'? '.'[0-9]+ > { p.addNumVal(text, false) } ) doublequotedstring <- ( '\\"' / '\\\\' / '\\n' / '\\t' / [^"\\] )* singlequotedstring <- ( '\\\'' / '\\\\' / '\\n' / '\\t' / [^'\\] )* fieldExpr <- ( [[A-Z]] / '_' ) ( [[A-Z]] / [0-9] / '_' / '-' )* -field <- { p.addField(buffer[begin:end]) } +field <- { p.addField(text) } reserved <- ('_row' / '_col' / '_start' / '_end' / '_timestamp' / '_field') -posfield <- { p.addPosStr("_field", buffer[begin:end]) } +posfield <- { p.addPosStr("_field", text) } uint <- [1-9] [0-9]* / '0' -col <- ( {p.addPosNum("_col", buffer[begin:end])} - / < '\'' singlequotedstring '\'' > {p.addPosStr("_col", buffer[begin:end])} - / < '"' doublequotedstring '"' > {p.addPosStr("_col", buffer[begin:end])} +col <- ( {p.addPosNum("_col", text)} + / < '\'' singlequotedstring '\'' > {p.addPosStr("_col", text)} + / < '"' doublequotedstring '"' > {p.addPosStr("_col", text)} ) -row <- ( {p.addPosNum("_row", buffer[begin:end])} - / < '\'' singlequotedstring '\'' > {p.addPosStr("_row", buffer[begin:end])} - / < '"' doublequotedstring '"' > {p.addPosStr("_row", buffer[begin:end])} +row <- ( {p.addPosNum("_row", text)} + / < '\'' singlequotedstring '\'' > {p.addPosStr("_row", text)} + / < '"' doublequotedstring '"' > {p.addPosStr("_row", text)} ) open <- '(' sp @@ -102,4 +102,4 @@ IDENT <- [[A-Z]] ([[A-Z]] / [0-9])* timestampbasicfmt <- [0-9][0-9][0-9][0-9]'-'[01][0-9]'-'[0-3][0-9]'T'[0-9][0-9]':'[0-9][0-9] timestampfmt <- '"' '"' / '\'' '\'' / -timestamp <- {p.addPosStr("_timestamp", buffer[begin:end])} +timestamp <- {p.addPosStr("_timestamp", text)} diff --git a/pql/pql.peg.go b/pql/pql.peg.go index e16002005..3c2a70940 100644 --- a/pql/pql.peg.go +++ b/pql/pql.peg.go @@ -1,6 +1,6 @@ package pql -//go:generate peg -inline pql.peg +// Code generated by peg -inline pql.peg DO NOT EDIT. import ( "fmt" @@ -473,15 +473,15 @@ func (p *PQL) Execute() { case ruleAction17: p.addField("from") case ruleAction18: - p.addVal(buffer[begin:end]) + p.addVal(text) case ruleAction19: p.addField("to") case ruleAction20: - p.addVal(buffer[begin:end]) + p.addVal(text) case ruleAction21: p.endCall() case ruleAction22: - p.startCall(buffer[begin:end]) + p.startCall(text) case ruleAction23: p.endCall() case ruleAction24: @@ -503,11 +503,11 @@ func (p *PQL) Execute() { case ruleAction32: p.endConditional() case ruleAction33: - p.condAdd(buffer[begin:end]) + p.condAdd(text) case ruleAction34: - p.condAdd(buffer[begin:end]) + p.condAdd(text) case ruleAction35: - p.condAdd(buffer[begin:end]) + p.condAdd(text) case ruleAction36: p.startList() case ruleAction37: @@ -523,43 +523,43 @@ func (p *PQL) Execute() { case ruleAction42: p.addVal(false) case ruleAction43: - p.addVal(buffer[begin:end]) + p.addVal(text) case ruleAction44: - p.startCall(string(_buffer[begin:end])) + p.startCall(text) case ruleAction45: p.addVal(p.endCall()) case ruleAction46: - p.addVal(string(_buffer[begin:end])) + p.addVal(text) case ruleAction47: - p.addVal(string(_buffer[begin:end])) + p.addVal(text) case ruleAction48: - p.addVal(string(_buffer[begin:end])) + p.addVal(text) case ruleAction49: - p.addNumVal(buffer[begin:end], true) + p.addNumVal(text, true) case ruleAction50: - p.addNumVal(buffer[begin:end], true) + p.addNumVal(text, true) case ruleAction51: - p.addNumVal(buffer[begin:end], false) + p.addNumVal(text, false) case ruleAction52: - p.addNumVal(buffer[begin:end], false) + p.addNumVal(text, false) case ruleAction53: - p.addField(buffer[begin:end]) + p.addField(text) case ruleAction54: - p.addPosStr("_field", buffer[begin:end]) + p.addPosStr("_field", text) case ruleAction55: - p.addPosNum("_col", buffer[begin:end]) + p.addPosNum("_col", text) case ruleAction56: - p.addPosStr("_col", buffer[begin:end]) + p.addPosStr("_col", text) case ruleAction57: - p.addPosStr("_col", buffer[begin:end]) + p.addPosStr("_col", text) case ruleAction58: - p.addPosNum("_row", buffer[begin:end]) + p.addPosNum("_row", text) case ruleAction59: - p.addPosStr("_row", buffer[begin:end]) + p.addPosStr("_row", text) case ruleAction60: - p.addPosStr("_row", buffer[begin:end]) + p.addPosStr("_row", text) case ruleAction61: - p.addPosStr("_timestamp", buffer[begin:end]) + p.addPosStr("_timestamp", text) } } @@ -3554,16 +3554,16 @@ func (p *PQL) Init() { nil, /* 59 Action17 <- <{p.addField("from")}> */ nil, - /* 60 Action18 <- <{p.addVal(buffer[begin:end])}> */ + /* 60 Action18 <- <{p.addVal(text)}> */ nil, /* 61 Action19 <- <{p.addField("to")}> */ nil, - /* 62 Action20 <- <{p.addVal(buffer[begin:end])}> */ + /* 62 Action20 <- <{p.addVal(text)}> */ nil, /* 63 Action21 <- <{p.endCall()}> */ nil, nil, - /* 65 Action22 <- <{ p.startCall(buffer[begin:end] ) }> */ + /* 65 Action22 <- <{ p.startCall(text ) }> */ nil, /* 66 Action23 <- <{ p.endCall() }> */ nil, @@ -3585,11 +3585,11 @@ func (p *PQL) Init() { nil, /* 75 Action32 <- <{p.endConditional()}> */ nil, - /* 76 Action33 <- <{p.condAdd(buffer[begin:end])}> */ + /* 76 Action33 <- <{p.condAdd(text)}> */ nil, - /* 77 Action34 <- <{p.condAdd(buffer[begin:end])}> */ + /* 77 Action34 <- <{p.condAdd(text)}> */ nil, - /* 78 Action35 <- <{p.condAdd(buffer[begin:end])}> */ + /* 78 Action35 <- <{p.condAdd(text)}> */ nil, /* 79 Action36 <- <{ p.startList() }> */ nil, @@ -3605,43 +3605,43 @@ func (p *PQL) Init() { nil, /* 85 Action42 <- <{ p.addVal(false) }> */ nil, - /* 86 Action43 <- <{ p.addVal(buffer[begin:end]) }> */ + /* 86 Action43 <- <{ p.addVal(text) }> */ nil, - /* 87 Action44 <- <{ p.startCall(string(_buffer[begin:end])) }> */ + /* 87 Action44 <- <{ p.startCall(text) }> */ nil, /* 88 Action45 <- <{ p.addVal(p.endCall()) }> */ nil, - /* 89 Action46 <- <{ p.addVal(string(_buffer[begin:end])) }> */ + /* 89 Action46 <- <{ p.addVal(text) }> */ nil, - /* 90 Action47 <- <{ p.addVal(string(_buffer[begin:end])) }> */ + /* 90 Action47 <- <{ p.addVal(text) }> */ nil, - /* 91 Action48 <- <{ p.addVal(string(_buffer[begin:end])) }> */ + /* 91 Action48 <- <{ p.addVal(text) }> */ nil, - /* 92 Action49 <- <{ p.addNumVal(buffer[begin:end], true) }> */ + /* 92 Action49 <- <{ p.addNumVal(text, true) }> */ nil, - /* 93 Action50 <- <{ p.addNumVal(buffer[begin:end], true) }> */ + /* 93 Action50 <- <{ p.addNumVal(text, true) }> */ nil, - /* 94 Action51 <- <{ p.addNumVal(buffer[begin:end], false) }> */ + /* 94 Action51 <- <{ p.addNumVal(text, false) }> */ nil, - /* 95 Action52 <- <{ p.addNumVal(buffer[begin:end], false) }> */ + /* 95 Action52 <- <{ p.addNumVal(text, false) }> */ nil, - /* 96 Action53 <- <{ p.addField(buffer[begin:end]) }> */ + /* 96 Action53 <- <{ p.addField(text) }> */ nil, - /* 97 Action54 <- <{ p.addPosStr("_field", buffer[begin:end]) }> */ + /* 97 Action54 <- <{ p.addPosStr("_field", text) }> */ nil, - /* 98 Action55 <- <{p.addPosNum("_col", buffer[begin:end])}> */ + /* 98 Action55 <- <{p.addPosNum("_col", text)}> */ nil, - /* 99 Action56 <- <{p.addPosStr("_col", buffer[begin:end])}> */ + /* 99 Action56 <- <{p.addPosStr("_col", text)}> */ nil, - /* 100 Action57 <- <{p.addPosStr("_col", buffer[begin:end])}> */ + /* 100 Action57 <- <{p.addPosStr("_col", text)}> */ nil, - /* 101 Action58 <- <{p.addPosNum("_row", buffer[begin:end])}> */ + /* 101 Action58 <- <{p.addPosNum("_row", text)}> */ nil, - /* 102 Action59 <- <{p.addPosStr("_row", buffer[begin:end])}> */ + /* 102 Action59 <- <{p.addPosStr("_row", text)}> */ nil, - /* 103 Action60 <- <{p.addPosStr("_row", buffer[begin:end])}> */ + /* 103 Action60 <- <{p.addPosStr("_row", text)}> */ nil, - /* 104 Action61 <- <{p.addPosStr("_timestamp", buffer[begin:end])}> */ + /* 104 Action61 <- <{p.addPosStr("_timestamp", text)}> */ nil, } p.rules = _rules diff --git a/pql/pqlpeg_test.go b/pql/pqlpeg_test.go index e6aba444d..9339cad3a 100644 --- a/pql/pqlpeg_test.go +++ b/pql/pqlpeg_test.go @@ -375,6 +375,48 @@ func TestPQLDeepEquality(t *testing.T) { "_timestamp": "2010-07-08T14:44", }, }}, + { + name: "SetWithUnicode", + call: `Set(0, unicode="Æ�漢д ☮♬ ♞🜻💣")`, + exp: &Call{ + Name: "Set", + Args: map[string]interface{}{ + "_col": int64(0), + "unicode": `Æ�漢д ☮♬ ♞🜻💣`, + }, + }}, + { + name: "RowWithUnicode", + call: `Row(unicode="Æ�漢д ☮♬ ♞🜻💣")`, + exp: &Call{ + Name: "Row", + Args: map[string]interface{}{ + "unicode": `Æ�漢д ☮♬ ♞🜻💣`, + }, + }}, + { + name: "RowsWithUnicode", + call: `Rows(job, previous="💣")`, + exp: &Call{ + Name: "Rows", + Args: map[string]interface{}{ + "_field": "job", + "previous": `💣`, + }, + }}, + { + name: "TopNWithUnicode", + call: `TopN(stargazer, Row(unicode="Æ�漢д ☮♬ ♞🜻💣"), a="∑")`, + exp: &Call{ + Name: "TopN", + Args: map[string]interface{}{ + "_field": "stargazer", + "a": "∑", + }, + Children: []*Call{ + {Name: "Row", Args: map[string]interface{}{"unicode": "Æ�漢д ☮♬ ♞🜻💣"}}, + }, + }}, { name: "SetRowAttrs", call: "SetRowAttrs(myfield, 9, z=4)", @@ -409,6 +451,17 @@ func TestPQLDeepEquality(t *testing.T) { }, }}, { + name: "SetRowAttrsWithUnicodeValues", + call: `SetRowAttrs(myfield, "∫", z="∀", a="∑")`, + exp: &Call{ + Name: "SetRowAttrs", + Args: map[string]interface{}{ + "z": "∀", + "a": "∑", + "_field": "myfield", + "_row": "∫", + }, + }}, { name: "SetColumnAttrs", call: "SetColumnAttrs(9, z=4)", exp: &Call{ From 128e02046ab7b3a02b15a321a801bbe9628b14f9 Mon Sep 17 00:00:00 2001 From: Nia Weiss Date: Tue, 11 Aug 2020 08:24:13 -0400 Subject: [PATCH 07/17] add a postgres endpoint to pilosa --- ctl/server.go | 3 + go.mod | 1 + go.sum | 2 + pg/io.go | 249 ++++++++++++++++++++++ pg/message/io.go | 148 +++++++++++++ pg/message/message.go | 344 ++++++++++++++++++++++++++++++ pg/pgtest/handler.go | 31 +++ pg/pgtest/memnet.go | 71 +++++++ pg/pgtest/server.go | 91 ++++++++ pg/pgtest/tls.go | 101 +++++++++ pg/protocol.go | 471 ++++++++++++++++++++++++++++++++++++++++++ pg/query.go | 130 ++++++++++++ pg/server.go | 136 ++++++++++++ pg/server_test.go | 218 +++++++++++++++++++ pg/type.go | 54 +++++ server/config.go | 26 +++ server/pg.go | 399 +++++++++++++++++++++++++++++++++++ server/server.go | 25 +++ 18 files changed, 2500 insertions(+) create mode 100644 pg/io.go create mode 100644 pg/message/io.go create mode 100644 pg/message/message.go create mode 100644 pg/pgtest/handler.go create mode 100644 pg/pgtest/memnet.go create mode 100644 pg/pgtest/server.go create mode 100644 pg/pgtest/tls.go create mode 100644 pg/protocol.go create mode 100644 pg/query.go create mode 100644 pg/server.go create mode 100644 pg/server_test.go create mode 100644 pg/type.go create mode 100644 server/pg.go diff --git a/ctl/server.go b/ctl/server.go index 6dc4eb0a6..378f951a3 100644 --- a/ctl/server.go +++ b/ctl/server.go @@ -88,4 +88,7 @@ func BuildServerFlags(cmd *cobra.Command, srv *server.Command) { // Transactional storage engine flags.StringVarP(&srv.Config.Txsrc, "tx", "", "", "transaction/storage to use: one of roaring, rbf, badger, rbf_roaring, roaring_rbf, badger_roaring, roaring_badger, badger_rbf, or rbf_badger (default roaring)") + + // Postgres endpoint + flags.StringVar(&srv.Config.Postgres.Addr, "postgres.addr", "", "address to which to bind a postgres endpoint") } diff --git a/go.mod b/go.mod index 70221964e..577efd92f 100644 --- a/go.mod +++ b/go.mod @@ -20,6 +20,7 @@ require ( github.com/gorilla/handlers v1.3.0 github.com/gorilla/mux v1.7.0 github.com/hashicorp/memberlist v0.1.3 + github.com/lib/pq v1.8.0 github.com/opentracing/opentracing-go v1.1.0 github.com/pelletier/go-toml v1.2.0 github.com/pkg/errors v0.8.1 diff --git a/go.sum b/go.sum index f97f5d051..10c1da24b 100644 --- a/go.sum +++ b/go.sum @@ -117,6 +117,8 @@ github.com/kr/pretty v0.1.0/go.mod h1:dAy3ld7l9f0ibDNOQOHHMYYIIbhfbHSm3C4ZsoJORN github.com/kr/pty v1.1.1/go.mod h1:pFQYn66WHrOpPYNljwOMqo10TkYh1fy3cYio2l3bCsQ= github.com/kr/text v0.1.0 h1:45sCR5RtlFHMR4UwH9sdQ5TC8v0qDQCHnXt+kaKSTVE= github.com/kr/text v0.1.0/go.mod h1:4Jbv+DJW3UT/LiOwJeYQe1efqtUx/iVham/4vfdArNI= +github.com/lib/pq v1.8.0 h1:9xohqzkUwzR4Ga4ivdTcawVS89YSDVxXMa3xJX3cGzg= +github.com/lib/pq v1.8.0/go.mod h1:AlVN5x4E4T544tWzH6hKfbfQvm3HdbOxrmggDNAPY9o= github.com/magiconair/properties v1.8.0 h1:LLgXmsheXeRoUOBOjtwPQCWIYqM/LU1ayDtDePerRcY= github.com/magiconair/properties v1.8.0/go.mod h1:PppfXfuXeibc/6YijjN8zIbojt8czPbwD3XqdrwzmxQ= github.com/matttproud/golang_protobuf_extensions v1.0.1 h1:4hp9jkHxhMHkqkrB3Ix0jegS5sx/RkqARlsWZ6pIwiU= diff --git a/pg/io.go b/pg/io.go new file mode 100644 index 000000000..80ae8cd24 --- /dev/null +++ b/pg/io.go @@ -0,0 +1,249 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package pg + +import ( + "net" + "sync" + "sync/atomic" + "time" + "unsafe" + + "github.com/pkg/errors" +) + +// timeoutWriter wraps a connection and implements io.Writer with a timeout for each write. +type timeoutWriter struct { + conn net.Conn + timeout time.Duration +} + +func (w *timeoutWriter) Write(data []byte) (int, error) { + err := w.conn.SetWriteDeadline(time.Now().Add(w.timeout)) + if err != nil { + return 0, err + } + return w.conn.Write(data) +} + +// errPreempted is an error used to indicate preemption of an idle connection. +var errPreempted = errors.New("preempted during idle") + +// idleState is an atomic state value used to track a preemptible connection. +type idleState uint32 + +const ( + idleStateActive idleState = iota + idleStateIdle + idleStatePreempted + idleStatePendingPreemption +) + +func (s *idleState) load() idleState { + return idleState(atomic.LoadUint32((*uint32)(unsafe.Pointer(s)))) +} + +func (s *idleState) cas(old, new idleState) bool { + return atomic.CompareAndSwapUint32((*uint32)(unsafe.Pointer(s)), uint32(old), uint32(new)) +} + +// idleReader is an io.Reader implementation on a preemptible network connection. +// The connection has 2 modes: "idle" and "active". +// While in idle mode, the connection has no timeout but can be preempted. +// While in active mode, the connection may have a read timeout but cannot be immediately preempted. +// When a read completes in idle mode, the connection returns to active mode. +// If the connection is preempted in active mode, the preemption will be deferred until the connection returns to idle mode. +// This also provides a read timeout. +type idleReader struct { + conn net.Conn + timeout time.Duration + state idleState + preemptMu sync.Mutex +} + +// setIdle pushes the reader into idle mode. +// If a preemption is pending, it will be delivered on the next call to Read. +func (r *idleReader) setIdle() error { + // Clear the read deadline. + err := r.conn.SetReadDeadline(time.Time{}) + if err != nil { + return errors.Wrap(err, "failed to clear deadline") + } + + for { + // Transition to idle mode. + state := r.state.load() + var target idleState + switch state { + case idleStateActive: + // active -> idle + target = idleStateIdle + case idleStatePendingPreemption: + // pending preemption -> preempted + // Switching to idle mode activates the preemption. + target = idleStatePreempted + default: + panic("inconsistent state") + } + + if r.state.cas(state, target) { + return nil + } + } +} + +// preempt the reader. +// If the reader is not currently idle, the preemption will be delivered next time the connection enters idle mode. +// This does not wait until the preemption error is delivered. +func (r *idleReader) preempt() error { + r.preemptMu.Lock() + defer r.preemptMu.Unlock() + for { + state := r.state.load() + var target idleState + switch state { + case idleStateActive: + // active -> pending preemption + target = idleStatePendingPreemption + case idleStateIdle: + // idle -> preempted + target = idleStatePreempted + case idleStatePendingPreemption, idleStatePreempted: + // A preemption has already been delivered. + return nil + default: + panic("inconsistent state") + } + + ok := r.state.cas(state, target) + if ok && target == idleStatePreempted { + // We have entered preemption mode. + // Preempt the current read on the connection. + return r.conn.SetReadDeadline(time.Now()) + } + } +} + +// Read from the connection. +func (r *idleReader) Read(data []byte) (int, error) { + var needsDeadlineReset bool + state := r.state.load() + switch state { + case idleStateActive, idleStatePendingPreemption: + // Connection is active. + // There is no need to worry about preemption. + if r.timeout != 0 { + // Apply a read timeout. + err := r.conn.SetReadDeadline(time.Now().Add(r.timeout)) + if err != nil { + return 0, err + } + } + return r.conn.Read(data) + + case idleStateIdle: + // Read, and handle preemption. + n, err := r.conn.Read(data) + if err != nil { + // Check if the error was caused by preemption. + state = r.state.load() + switch { + case state == idleStatePreempted && n != 0: + // Some data was read before the preemption was delivered. + // Re-activate and discard the error. + + // Synchronize against the preempter. + // This is necessary to ensure that the cancellation deadline is cleared. + r.preemptMu.Lock() + defer r.preemptMu.Unlock() + + // Re-activate the connection. + // No CAS loop is necessary since we are synchronized against preempters. + r.state = idleStatePendingPreemption + + // The deadline may need to reset since the preempter may have changed it. + needsDeadlineReset = true + + case state == idleStatePreempted: + // The read was preempted. + return 0, errPreempted + + case state != idleStateIdle: + // No other states make sense here. + panic("inconsistent state") + + default: + // No preemption was involved. + // It is just a regular network error. + return n, err + } + } else { + // The read went through. + + // Exit from idle mode. + // Ideally, transition to active mode. + // However, a preemption may trigger while this is running. + for state == idleStateIdle { + if r.state.cas(idleStateIdle, idleStateActive) { + state = idleStateActive + break + } + + state = r.state.load() + } + + switch state { + case idleStateActive: + // The connection was reactivated normally. + + case idleStatePreempted: + // The connection was preempted after the read completed. + // Defer the preemption and complete successfully. + + // Synchronize against the preempter. + // This is necessary to ensure that the cancellation deadline is cleared. + r.preemptMu.Lock() + defer r.preemptMu.Unlock() + + // Re-activate the connection. + // No CAS loop is necessary since we are synchronized against preempters. + r.state = idleStatePendingPreemption + + // The deadline may need to reset since the preempter may have changed it. + needsDeadlineReset = true + + default: + panic("inconsistent state") + } + } + + if needsDeadlineReset && r.timeout == 0 { + // Clear the deadline. + err := r.conn.SetReadDeadline(time.Time{}) + if err != nil { + return n, err + } + } + + return n, nil + + case idleStatePreempted: + // The connection is preempted. + return 0, errPreempted + + default: + panic("inconsistent state") + } +} diff --git a/pg/message/io.go b/pg/message/io.go new file mode 100644 index 000000000..115048229 --- /dev/null +++ b/pg/message/io.go @@ -0,0 +1,148 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package message + +import ( + "bufio" + "encoding/binary" + "errors" + "fmt" + "io" +) + +// Reader reads messages. +type Reader interface { + ReadMessage() (Message, error) +} + +// Writer writes messages. +type Writer interface { + WriteMessage(Message) error + Flush() error +} + +// WireReader reads messages in Postgres wire protocol format. +type WireReader struct { + buf []byte + r *bufio.Reader + scratch [4]byte +} + +// ReadMessage reads a single message off of the wire. +// The returned message is only valid until the next read call, as the data buffer may be re-used. +func (r *WireReader) ReadMessage() (Message, error) { + t, err := r.r.ReadByte() + if err != nil { + return Message{}, err + } + + _, err = r.r.Read(r.scratch[:]) + if err != nil { + return Message{}, err + } + + len := binary.BigEndian.Uint32(r.scratch[:4]) + if len < 4 { + return Message{}, errors.New("invalid message length") + } + len -= 4 + + if cap(r.buf) < int(len) { + r.buf = make([]byte, len) + } else { + r.buf = r.buf[:len] + } + _, err = io.ReadFull(r.r, r.buf) + if err != nil { + return Message{}, err + } + + return Message{ + Type: Type(t), + Data: r.buf, + }, nil +} + +var _ Reader = (*WireReader)(nil) + +// NewWireReader returns a message reader that reads postgres wire protocol format. +func NewWireReader(r *bufio.Reader) *WireReader { + return &WireReader{r: r} +} + +// ErrMessageTooBig is an error indicating that a message is to big to be sent or recieved. +var ErrMessageTooBig = errors.New("message is too big") + +// WireWriter writes messages in Postgres wire protocol. +type WireWriter struct { + w *bufio.Writer + scratch [4]byte +} + +// WriteMessage writes a message onto the wire. +func (w *WireWriter) WriteMessage(message Message) error { + if uint(len(message.Data))+4 >= 1<<31 { + return ErrMessageTooBig + } + + err := w.w.WriteByte(byte(message.Type)) + if err != nil { + return err + } + + binary.BigEndian.PutUint32(w.scratch[:], uint32(len(message.Data))+4) + _, err = w.w.Write(w.scratch[:]) + if err != nil { + return err + } + + _, err = w.w.Write(message.Data) + return err +} + +// Flush writes any buffered data to the underlying stream. +func (w *WireWriter) Flush() error { + return w.w.Flush() +} + +var _ Writer = (*WireWriter)(nil) + +// NewWireWriter returns a message writer that writes in postgres wire protocol format. +func NewWireWriter(w *bufio.Writer) *WireWriter { + return &WireWriter{w: w} +} + +type DumpWriter struct { + Writer + Out io.Writer +} + +func (w *DumpWriter) WriteMessage(msg Message) error { + fmt.Fprintf(w.Out, "write type %s: %x %q\n", string(msg.Type), msg.Data, string(msg.Data)) + return w.Writer.WriteMessage(msg) +} + +type DumpReader struct { + Reader + Out io.Writer +} + +func (r *DumpReader) ReadMessage() (Message, error) { + msg, err := r.Reader.ReadMessage() + if err == nil { + fmt.Fprintf(r.Out, "read type %s: %x %q\n", string(msg.Type), msg.Data, string(msg.Data)) + } + return msg, err +} diff --git a/pg/message/message.go b/pg/message/message.go new file mode 100644 index 000000000..b7de6683f --- /dev/null +++ b/pg/message/message.go @@ -0,0 +1,344 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package message + +import ( + "bytes" + "encoding/binary" +) + +// Type is a byte indicating the type of a Postgres message. +type Type byte + +const ( + // TypeAuthentication is a message used to transfer authentication info. + TypeAuthentication Type = 'R' + + // TypeReadyForQuery is a message used to indicate that the server is ready for another query. + TypeReadyForQuery Type = 'Z' + + // TypeCommandComplete is a message used to indicate that a query has completed. + TypeCommandComplete Type = 'C' + + // TypeError is an error message. + TypeError Type = 'E' + + // TypeRowDescription is a message indicating the column types of the result rows from a query. + TypeRowDescription Type = 'T' + + // TypeDataRow is a message with the contents of a single row. + TypeDataRow Type = 'D' + + // TypeTermination is a message indicating a request to terminate a connection. + TypeTermination Type = 'X' + + // TypeNegotiateProtocolVersion is a message used when a client attempts to connect with a newer minor version than the server supports. + TypeNegotiateProtocolVersion Type = 'v' + + // TypeSimpleQuery is a simple query request. + TypeSimpleQuery Type = 'Q' +) + +// AuthenticationOK is a message indicating that authentication has completed. +var AuthenticationOK = Message{ + Type: TypeAuthentication, + Data: []byte{0, 0, 0, 0}, +} + +// Message is a Postgres message value. +type Message struct { + Type Type + Data []byte +} + +// TransactionStatus is the current transaction state. +type TransactionStatus byte + +const ( + // TransactionStatusIdle indicates that there is no active transaction. + TransactionStatusIdle TransactionStatus = 'I' + + // TransactionStatusActive indicates that the connection currently has an active transaction. + TransactionStatusActive TransactionStatus = 'T' + + // TransactionStatusFailed indicates that the connection currently has a failed transaction. + TransactionStatusFailed TransactionStatus = 'E' +) + +// Encoder encodes messages. +type Encoder struct { + buf bytes.Buffer + scratch [4]byte +} + +func (e *Encoder) i16(i int16) error { + binary.BigEndian.PutUint16(e.scratch[:2], uint16(i)) + _, err := e.buf.Write(e.scratch[:2]) + return err +} + +func (e *Encoder) i32(i int32) error { + binary.BigEndian.PutUint32(e.scratch[:], uint32(i)) + _, err := e.buf.Write(e.scratch[:]) + return err +} + +// ReadyForQuery encodes a "ready for query" message. +func (e *Encoder) ReadyForQuery(status TransactionStatus) (Message, error) { + e.buf.Reset() + + err := e.buf.WriteByte(byte(status)) + if err != nil { + return Message{}, err + } + + return Message{ + Type: TypeReadyForQuery, + Data: e.buf.Bytes(), + }, nil +} + +// CommandComplete encodes a command completion message. +func (e *Encoder) CommandComplete(tag string) (Message, error) { + e.buf.Reset() + + _, err := e.buf.WriteString(tag) + if err != nil { + return Message{}, err + } + + err = e.buf.WriteByte(0) + if err != nil { + return Message{}, err + } + + return Message{ + Type: TypeCommandComplete, + Data: e.buf.Bytes(), + }, nil +} + +// NoticeFieldType indicates the type of a notice/error field. +// https://www.postgresql.org/docs/9.3/protocol-error-fields.html +type NoticeFieldType byte + +const ( + // NoticeFieldSeverity indicates the severity of a notice/error. + NoticeFieldSeverity NoticeFieldType = 'S' + + // NoticeFieldMessage is a short human-readable error/notice message. + NoticeFieldMessage NoticeFieldType = 'M' + + // NoticeFieldDetail is an optional extended description of the error. + NoticeFieldDetail NoticeFieldType = 'D' + + // NoticeFieldHint is a suggestion of how to address the issue. + NoticeFieldHint NoticeFieldType = 'H' +) + +// NoticeField is a field in an error or notice. +type NoticeField struct { + Type NoticeFieldType + Data string +} + +func (e *Encoder) messageOrNotice(fields ...NoticeField) error { + for _, f := range fields { + err := e.buf.WriteByte(byte(f.Type)) + if err != nil { + return err + } + + _, err = e.buf.WriteString(f.Data) + if err != nil { + return err + } + + err = e.buf.WriteByte(0) + if err != nil { + return err + } + } + + return e.buf.WriteByte(0) +} + +// Error encodes a Postgres error message. +func (e *Encoder) Error(fields ...NoticeField) (Message, error) { + e.buf.Reset() + err := e.messageOrNotice(fields...) + if err != nil { + return Message{}, err + } + return Message{ + Type: TypeError, + Data: e.buf.Bytes(), + }, nil +} + +// GoError creates a simple Postgres error message from a Go error value. +func (e *Encoder) GoError(err error) (Message, error) { + return e.Error( + NoticeField{ + Type: NoticeFieldSeverity, + Data: "ERROR", + }, + NoticeField{ + Type: NoticeFieldMessage, + Data: err.Error(), + }, + ) +} + +// ColumnDescription is a description of a data column. +type ColumnDescription struct { + Name string + TableID int32 //either a table/col id or 0 + FieldID int16 //either a table/col id or 0 + TypeID int32 //field type + TypeLen int16 //size in bytes of field + TypeModifier int32 //type modifer? + Mode int16 //0=text 1=binary +} + +// RowDescription describes the response rows from a query. +func (e *Encoder) RowDescription(cols ...ColumnDescription) (Message, error) { + if len(cols) >= 1<<15 { + return Message{}, ErrMessageTooBig + } + + e.buf.Reset() + + err := e.i16(int16(len(cols))) + if err != nil { + return Message{}, nil + } + + for _, col := range cols { + _, err := e.buf.WriteString(col.Name) + if err != nil { + return Message{}, err + } + err = e.buf.WriteByte(0) + if err != nil { + return Message{}, err + } + + err = e.i32(col.TableID) + if err != nil { + return Message{}, err + } + + err = e.i16(col.FieldID) + if err != nil { + return Message{}, err + } + + err = e.i32(col.TypeID) + if err != nil { + return Message{}, err + } + + err = e.i16(col.TypeLen) + if err != nil { + return Message{}, err + } + + err = e.i32(col.TypeModifier) + if err != nil { + return Message{}, err + } + + err = e.i16(col.Mode) + if err != nil { + return Message{}, err + } + } + + return Message{ + Type: TypeRowDescription, + Data: e.buf.Bytes(), + }, nil +} + +// TextRow encodes a data row in textual format. +func (e *Encoder) TextRow(row ...string) (Message, error) { + if len(row) >= 1<<15 { + return Message{}, ErrMessageTooBig + } + + e.buf.Reset() + + err := e.i16(int16(len(row))) + if err != nil { + return Message{}, err + } + + for _, val := range row { + if uint(len(val)) >= 1<<31 { + return Message{}, ErrMessageTooBig + } + + err = e.i32(int32(len(val))) + if err != nil { + return Message{}, err + } + + _, err = e.buf.WriteString(val) + if err != nil { + return Message{}, err + } + } + + return Message{ + Type: TypeDataRow, + Data: e.buf.Bytes(), + }, nil +} + +// NegotiateProtocolVersion encodes a protocol negotiation packet. +func (e *Encoder) NegotiateProtocolVersion(maxMinor int32, unrecognizedOptions ...string) (Message, error) { + if uint64(len(unrecognizedOptions)) >= 1<<31 { + return Message{}, ErrMessageTooBig + } + + e.buf.Reset() + + err := e.i32(maxMinor) + if err != nil { + return Message{}, err + } + + err = e.i32(int32(len(unrecognizedOptions))) + if err != nil { + return Message{}, err + } + for _, opt := range unrecognizedOptions { + _, err = e.buf.WriteString(opt) + if err != nil { + return Message{}, err + } + + err = e.buf.WriteByte(0) + if err != nil { + return Message{}, err + } + } + + return Message{ + Type: TypeNegotiateProtocolVersion, + Data: e.buf.Bytes(), + }, nil +} diff --git a/pg/pgtest/handler.go b/pg/pgtest/handler.go new file mode 100644 index 000000000..e6cb22e12 --- /dev/null +++ b/pg/pgtest/handler.go @@ -0,0 +1,31 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package pgtest + +import ( + "context" + + "github.com/pilosa/pilosa/v2/pg" +) + +// HandlerFunc implements a postgres query handler with a function. +type HandlerFunc func(context.Context, pg.QueryResultWriter, pg.Query) error + +// HandleQuery calls the user's query handler function. +func (h HandlerFunc) HandleQuery(ctx context.Context, w pg.QueryResultWriter, q pg.Query) error { + return h(ctx, w, q) +} + +var _ pg.QueryHandler = HandlerFunc(nil) diff --git a/pg/pgtest/memnet.go b/pg/pgtest/memnet.go new file mode 100644 index 000000000..0ba85565c --- /dev/null +++ b/pg/pgtest/memnet.go @@ -0,0 +1,71 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package pgtest + +import ( + "errors" + "net" + "sync" +) + +// errListenerClosed is an error returned when the listener is closed. +var errListenerClosed = errors.New("listener closed") + +type inMemoryListener struct { + ch chan net.Conn + closed chan struct{} + once sync.Once +} + +func (l *inMemoryListener) Accept() (net.Conn, error) { + select { + case <-l.closed: + return nil, errListenerClosed + default: + } + select { + case conn := <-l.ch: + return conn, nil + case <-l.closed: + return nil, errListenerClosed + } +} + +func (l *inMemoryListener) Close() error { + l.once.Do(func() { close(l.closed) }) + + return nil +} + +type memAddr struct{} + +func (a memAddr) Network() string { return "memory" } +func (a memAddr) String() string { return "memory" } + +func (l *inMemoryListener) Addr() net.Addr { + return memAddr{} +} + +func (l *inMemoryListener) Dial() (net.Conn, error) { + serverConn, clientConn := net.Pipe() + select { + case l.ch <- serverConn: + return clientConn, nil + case <-l.closed: + serverConn.Close() + clientConn.Close() + return nil, errListenerClosed + } +} diff --git a/pg/pgtest/server.go b/pg/pgtest/server.go new file mode 100644 index 000000000..e3a48e244 --- /dev/null +++ b/pg/pgtest/server.go @@ -0,0 +1,91 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package pgtest + +import ( + "context" + "net" + "testing" + + "github.com/pilosa/pilosa/v2/pg" + "github.com/pkg/errors" + "golang.org/x/sync/errgroup" +) + +// ShutdownFunc is a function to use to shut down a test fixture. +// This function will send a shutdown signal and then wait for completion. +type ShutdownFunc func() error + +// Finish invokes the shutdown function and fails the test if an error occurs. +func (f ShutdownFunc) Finish(tb testing.TB, name string) { + err := f() + if err != nil { + tb.Errorf("failed to shut down %s: %v", name, err) + } +} + +// ServeTCP creates a TCP listener and serves postgres wire protocol on it. +func ServeTCP(addr string, server *pg.Server) (net.Addr, ShutdownFunc, error) { + listener, err := net.Listen("tcp", addr) + if err != nil { + return nil, nil, errors.Wrap(err, "listening on TCP") + } + + laddr := listener.Addr() + + ctx, cancel := context.WithCancel(context.Background()) + var eg errgroup.Group + eg.Go(func() error { return server.Serve(ctx, listener) }) + + return laddr, + func() error { + cancel() + return eg.Wait() + }, + nil +} + +// ServeTLS sets up TLS on the server and invokes ServeTCP. +func ServeTLS(addr string, server *pg.Server) (net.Addr, ShutdownFunc, error) { + err := SetupTLS(server) + if err != nil { + return nil, nil, errors.Wrap(err, "server TLS setup failed") + } + + return ServeTCP(addr, server) +} + +// ConnectFunc is a function to connect to a server. +type ConnectFunc func() (net.Conn, error) + +// ServeMem serves postgres on in-memory connections. +// TLS does not work here, as it relies on the OS to buffer and discard data. +func ServeMem(server *pg.Server) (ConnectFunc, ShutdownFunc, error) { + listener := &inMemoryListener{ + ch: make(chan net.Conn), + closed: make(chan struct{}), + } + + ctx, cancel := context.WithCancel(context.Background()) + var eg errgroup.Group + eg.Go(func() error { return server.Serve(ctx, listener) }) + + return listener.Dial, + func() error { + cancel() + return eg.Wait() + }, + nil +} diff --git a/pg/pgtest/tls.go b/pg/pgtest/tls.go new file mode 100644 index 000000000..53e942e80 --- /dev/null +++ b/pg/pgtest/tls.go @@ -0,0 +1,101 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package pgtest + +import ( + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/tls" + "crypto/x509" + "crypto/x509/pkix" + "encoding/pem" + "math/big" + "time" + + "github.com/pilosa/pilosa/v2/pg" + "github.com/pkg/errors" +) + +// SetupTLS generates a TLS certificate and installs it into the server. +// TODO: have the client properly trust this (generate a CA to install instead of using self-signed). +func SetupTLS(server *pg.Server) error { + // Generate an ecdsa key for the cert. + key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + if err != nil { + return errors.Wrap(err, "generating TLS key") + } + + // Generate a random 128-bit serial number. + serialNumber, err := rand.Int(rand.Reader, new(big.Int).Lsh(big.NewInt(1), 128)) + if err != nil { + return errors.Wrap(err, "generating serial number") + } + + // Make the certificate valid starting now. + now := time.Now() + + // Create a certificate template. + template := x509.Certificate{ + SerialNumber: serialNumber, + Subject: pkix.Name{ + Organization: []string{"Molecula"}, + }, + NotBefore: now, + NotAfter: now.Add(time.Hour), + + KeyUsage: x509.KeyUsageDigitalSignature, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}, + BasicConstraintsValid: true, + } + + // Generate a self-signed x509 cert from the template and the key. + certData, err := x509.CreateCertificate(rand.Reader, &template, &template, key.Public(), key) + if err != nil { + return errors.Wrap(err, "encoding certificate x509") + } + + // Encode the cert to PEM so that the TLS package can load it. + certPEM := pem.EncodeToMemory(&pem.Block{ + Type: "CERTIFICATE", + Bytes: certData, + }) + + // Encode the private key to x509. + keyData, err := x509.MarshalPKCS8PrivateKey(key) + if err != nil { + return errors.Wrap(err, "encoding key x509") + } + + // Encode the key to PEM so that the TLS package can load it. + keyPEM := pem.EncodeToMemory(&pem.Block{ + Type: "PRIVATE KEY", + Bytes: keyData, + }) + + // Load the certificate and key from their PEM encodings. + cert, err := tls.X509KeyPair(certPEM, keyPEM) + if err != nil { + return errors.Wrap(err, "loading TLS key pair") + } + + // Install the certificate into the server. + if server.TLSConfig == nil { + server.TLSConfig = &tls.Config{} + } + server.TLSConfig.Certificates = append(server.TLSConfig.Certificates, cert) + + return nil +} diff --git a/pg/protocol.go b/pg/protocol.go new file mode 100644 index 000000000..57dc2168f --- /dev/null +++ b/pg/protocol.go @@ -0,0 +1,471 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package pg + +import ( + "bufio" + "bytes" + "context" + "crypto/tls" + "encoding/binary" + "encoding/hex" + "fmt" + "io" + "io/ioutil" + "net" + "strings" + "sync" + "time" + + "github.com/pilosa/pilosa/v2/pg/message" + "github.com/pkg/errors" +) + +// Protocol is a Postgres protocol version. +type Protocol uint32 + +const ( + // ProtocolPostgres30 is version 3.0 of the Postgres wire protocol. + ProtocolPostgres30 Protocol = (3 << 16) | 0 + + // ProtocolCancel is the protocol used for query cancellation. + ProtocolCancel Protocol = (1234 << 16) | 5678 + + // ProtocolSSL is the protocol used for SSL upgrades. + ProtocolSSL Protocol = (1234 << 16) | 5679 + + // ProtocolSupported is the main protocol version supported by this package. + ProtocolSupported Protocol = ProtocolPostgres30 +) + +// Major returns the major revision of the protocol. +func (p Protocol) Major() uint16 { + return uint16(p >> 16) +} + +// Minor returns the minor revision of the protocol. +func (p Protocol) Minor() uint16 { + return uint16(p) +} + +func (p Protocol) String() string { + switch p { + case ProtocolCancel: + return "cancel" + case ProtocolSSL: + return "SSL" + } + + return fmt.Sprintf("v%d.%d", p.Major(), p.Minor()) +} + +// handle reads the startup packet and dispatches an appropriate protocol handler for the connection. +func (s *Server) handle(ctx context.Context, conn net.Conn) (err error) { + defer func() { + cerr := conn.Close() + if cerr != nil && err == nil { + err = errors.Wrap(cerr, "closing connection") + } + }() + + if tcpconn, ok := conn.(*net.TCPConn); ok { + // Postgres does not have any real mechanism for confirming that a connection is still alive. + // Without this, a connection that breaks while idle would live indefinitely. + // With a TCP keepalive, this should return an error after approximately 2 hours (depending on OS configuration). + err := tcpconn.SetKeepAlive(true) + if err != nil { + return errors.Wrap(err, "enabling TCP keepalive") + } + } + + var startupDeadline time.Time + if s.StartupTimeout > 0 { + // Set deadline for processing the startup. + startupDeadline = time.Now().Add(s.StartupTimeout) + err = conn.SetDeadline(startupDeadline) + if err != nil { + return errors.Wrap(err, "setting deadline on protocol startup") + } + } + +startup: + // Read startup packet. + var buf [4]byte + _, err = io.ReadFull(conn, buf[:]) + if err != nil { + return errors.Wrap(err, "reading startup message length") + } + size := binary.BigEndian.Uint32(buf[:]) + if size < 4 { + return errors.Errorf("invalid startup packet length: %d bytes", size) + } + maxLen := s.MaxStartupSize + if maxLen == 0 { + maxLen = 1024 * 1024 + } + if size > maxLen { + return errors.Errorf("oversized startup frame of %d bytes (max: %d bytes)", size, maxLen) + } + data := make([]byte, size-4) + _, err = io.ReadFull(conn, data) + if err != nil { + return errors.Wrap(err, "reading startup packet") + } + + // Extract protocol ID. + if len(data) < 4 { + return errors.Errorf("startup packet is too small for protocol ID: %d bytes", len(data)) + } + proto := Protocol(binary.BigEndian.Uint32(data)) + data = data[4:] + + // Handle special protocols. + switch proto { + case ProtocolCancel: + // TODO: send an actual postgres error message. + return errors.New("cancellation protocol not yet supported") + + case ProtocolSSL: + if s.TLSConfig != nil { + // Upgrade the connection to TLS and renegotiate on the tunneled connection. + _, err = conn.Write([]byte{'S'}) + if err != nil { + return errors.Wrap(err, "sending SSL support confirmation") + } + conn = tls.Server(conn, s.TLSConfig) + if s.StartupTimeout > 0 { + err := conn.SetDeadline(startupDeadline) + if err != nil { + return errors.Wrap(err, "transferring startup deadline to TLS connection") + } + } + goto startup + } + + // Inform the client that SSL is not available and try again. + s.Logger.Debugf("client at %s requested a secure postgres connection but TLS is not configured", conn.RemoteAddr()) + _, err = conn.Write([]byte{'N'}) + if err != nil { + return errors.Wrap(err, "sending SSL unsupported notification") + } + goto startup + } + + // Handle regular postgres. + return s.handleStandard(ctx, proto, conn, data) +} + +// parseParams parses a parameter list from a startup packet. +func parseParams(data []byte) (map[string]string, error) { + params := make(map[string]string) + for { + idx := bytes.IndexByte(data, 0) + switch idx { + case 0: + return params, nil + case -1: + return nil, errors.New("malformed startup parameter list") + } + + key := string(data[:idx]) + data = data[idx+1:] + + idx = bytes.IndexByte(data, 0) + if idx == -1 { + return nil, errors.New("malformed startup parameter list") + } + val := string(data[:idx]) + data = data[idx+1:] + + params[key] = val + } +} + +// handleStandard handles a connection in the standard postgres wire protocol. +// The client is responsible for closing the connection when this finishes. +func (s *Server) handleStandard(ctx context.Context, proto Protocol, conn net.Conn, data []byte) error { + // Wait for helper goroutines to finish. + var wg sync.WaitGroup + defer wg.Wait() + + // Set up context. + ctx, cancel := context.WithCancel(ctx) + defer cancel() + + // Check the major version. + if proto.Major() != ProtocolSupported.Major() { + return errors.Errorf("unsupported protocol %v", proto) + } + + // Parse the parameters bundled in the startup packet. + params, err := parseParams(data) + if err != nil { + return errors.Wrap(err, "parsing parameters") + } + + if user, ok := params["user"]; ok { + // Log the connection. + s.Logger.Debugf("new postgres connection from user %q at %v", user, conn.RemoteAddr()) + } else { + // We do not use this much yet, but the wire protocol says that it is required. + return errors.New("missing username") + } + + // Set up message input and output. + + // Set up a reader that will preempt the connection when the context is canceled. + ir := idleReader{ + conn: conn, + timeout: s.ReadTimeout, + } + wg.Add(1) + go func() { + defer wg.Done() + + <-ctx.Done() + + ir.preempt() //nolint:errcheck + }() + + // Clear the startup deadline. + err = conn.SetDeadline(time.Time{}) + if err != nil { + return err + } + + // Set up a message reader with buffering. + rbuf := bufio.NewReader(&ir) + r := message.NewWireReader(rbuf) + + // Set up a writer on the connection. + var ww io.Writer = conn + if s.WriteTimeout != 0 { + // Apply the write timeout. + ww = &timeoutWriter{ + conn: conn, + timeout: s.WriteTimeout, + } + } + + // Set up a message writer with buffering. + w := message.NewWireWriter(bufio.NewWriter(ww)) + + var encoder message.Encoder + if proto.Minor() > ProtocolSupported.Minor() { + // Negotiate the version down. + s.Logger.Debugf("client requested unsupported protocol version %v; attempting to downgrade to %v", proto, ProtocolSupported) + msg, err := encoder.NegotiateProtocolVersion(int32(ProtocolSupported.Minor())) + if err != nil { + return errors.Wrap(err, "negotiating version") + } + err = w.WriteMessage(msg) + if err != nil { + return errors.Wrap(err, "negotiating version") + } + } + + // TODO: real auth + err = w.WriteMessage(message.AuthenticationOK) + if err != nil { + return errors.Wrap(err, "sending authentication confirmation") + } + + var queryReady bool + for { + if !queryReady { + // Indicate that we are ready for a query. + // TODO: provide a valid transaction state. + msg, err := encoder.ReadyForQuery(message.TransactionStatusActive) + if err != nil { + return errors.Wrap(err, "sending query ready status") + } + err = w.WriteMessage(msg) + if err != nil { + return errors.Wrap(err, "sending query ready status") + } + + // Flush the write buffer so that the client can respond. + err = w.Flush() + if err != nil { + return errors.Wrap(err, "flushing status") + } + + if rbuf.Buffered() == 0 { + // Put the connection into idle mode. + err = ir.setIdle() + if err != nil { + return errors.Wrap(err, "setting idle mode") + } + } else { + // If the client follows the spec, then it should not have sent anything more. + // However, it seems that no clients completely follow the spec, so we shouldn't rely on anything that isn't entirely straightforward. + s.Logger.Debugf("postgres client sent additional data without waiting for completion") + } + } + + // Read the next packet. + msg, err := r.ReadMessage() + if err != nil { + if err == errPreempted { + // The server is shutting down. + return errors.Wrap(s.handleShutdown( + conn, w, &encoder, + + message.NoticeField{ + Type: message.NoticeFieldSeverity, + Data: "ERROR", + }, + message.NoticeField{ + Type: message.NoticeFieldMessage, + Data: "server shutting down", + }, + message.NoticeField{ + Type: message.NoticeFieldHint, + Data: "This is normal. This message is sent when a server is shutting down and terminating its connections.", + }, + ), "processing connection shutdown") + } + + return err + } + + switch msg.Type { + case message.TypeTermination: + // We are done. + return w.Flush() + + case message.TypeSimpleQuery: + // Execute a simple query. + + queryReady = false + + // Parse the query message (a null-terminated string). + query := SimpleQuery(strings.TrimSuffix(string(msg.Data), "\x00")) + + // Set up a result writer. + // SELECT is used as a default tag, which seems to be handled decently by clients. + // The encoder is intentionally not used because its buffer may be huge. + qwriter := &queryResultWriter{ + w: w, + te: s.TypeEngine, + tag: "SELECT", + } + + // Dispatch the query handler. + // TODO: cancellation (requires crazy internode logic and fake process IDs) + // This is not the connection context, since we want the request to finish safely before connection shutdown. + qerr := s.QueryHandler.HandleQuery(context.Background(), qwriter, query) + if qerr != nil { + // There was an error in processing the query. + // Send the error back to the client and keep going. + s.Logger.Debugf("failed to execute query %q: %v", query, qerr) + msg, err = encoder.GoError(qerr) + if err != nil { + return errors.Wrap(err, "failed to send query error to client") + } + err = w.WriteMessage(msg) + if err != nil { + return errors.Wrap(err, "failed to send query error to client") + } + } else { + // The query completed normally. + // Notify the client of completion. + msg, err = encoder.CommandComplete(qwriter.tag) + if err != nil { + return errors.Wrap(err, "sending command completion notification") + } + err = w.WriteMessage(msg) + if err != nil { + return errors.Wrap(err, "sending command completion notification") + } + } + + // The data will be flushed after we write back the "ready for query" state. + + default: + // The message is not supported yet. + // Send an error. + s.Logger.Printf("unrecognized postgres packet %v", msg) + msg, err = encoder.Error( + message.NoticeField{ + Type: message.NoticeFieldSeverity, + Data: "ERROR", + }, + message.NoticeField{ + Type: message.NoticeFieldMessage, + Data: fmt.Sprintf("unrecognized message type %q", msg.Type), + }, + message.NoticeField{ + Type: message.NoticeFieldDetail, + Data: "message body:" + hex.Dump(msg.Data), + }, + ) + if err != nil { + return errors.Wrap(err, "sending unrecognized message error") + } + err = w.WriteMessage(msg) + if err != nil { + return errors.Wrap(err, "sending unrecognized message error") + } + err = w.Flush() + if err != nil { + return errors.Wrap(err, "sending unrecognized message error") + } + } + } +} + +func (s *Server) handleShutdown(conn net.Conn, w message.Writer, encoder *message.Encoder, notice ...message.NoticeField) error { + var wg sync.WaitGroup + defer wg.Wait() + + // Try to send a message to the client before closing the connection. + msg, err := encoder.Error(notice...) + if err != nil { + return errors.Wrap(err, "generating shutdown notification") + } + + if s.WriteTimeout == 0 { + // The client is likely to not listen for incoming messages. + // Force a write timeout to ensure that this terminates. + err := conn.SetWriteDeadline(time.Now().Add(time.Second)) + if err != nil { + return errors.Wrap(err, "setting shutdown write deadline") + } + } + + // The client may be waiting on a write, so we need to drain the incoming data stream. + err = conn.SetReadDeadline(time.Time{}) + if err != nil { + return errors.Wrap(err, "clearing read deadline for shutdown") + } + defer conn.SetReadDeadline(time.Now()) //nolint:errcheck + wg.Add(1) + go func() { + defer wg.Done() + + io.Copy(ioutil.Discard, conn) //nolint:errcheck + }() + + // Attempt to send the shutdown notification. + // This will fail under many scenarios, as the client is not necessarily reading. + err = w.WriteMessage(msg) + if err != nil { + return nil + } + w.Flush() + + return nil +} diff --git a/pg/query.go b/pg/query.go new file mode 100644 index 000000000..8b95b4df8 --- /dev/null +++ b/pg/query.go @@ -0,0 +1,130 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package pg + +import ( + "context" + "fmt" + + "github.com/pilosa/pilosa/v2/pg/message" + "github.com/pkg/errors" +) + +// Query is an interface to be implemented by queries. +type Query interface { + fmt.Stringer +} + +// SimpleQuery is a query sent as only a string. +// It has no parameters. +type SimpleQuery string + +func (q SimpleQuery) String() string { + return string(q) +} + +// ColumnInfo contains metadata about a column. +type ColumnInfo struct { + Name string + Type Type + TableID int32 + FieldID int16 +} + +// QueryResultWriter is used to write the results of a query back over the connection. +type QueryResultWriter interface { + // WriteHeader sets the column header information. + WriteHeader(...ColumnInfo) error + + // WriteRowText sends a row of data in textual format. + WriteRowText(...string) error + + // Tag assigns a tag to the query. + // This should be called before the query is completed. + Tag(tag string) +} + +// QueryHandler handles a query. +type QueryHandler interface { + // HandleQuery executes a query and writes the results back. + HandleQuery(context.Context, QueryResultWriter, Query) error +} + +// queryResultWriter implements QueryResultWrtiter over postgres wire protocol. +// The underlying message writer must be flushed by the caller once the query has finished. +type queryResultWriter struct { + w message.Writer + te TypeEngine + enc message.Encoder + width int + wroteHeaders bool + tag string +} + +func (w *queryResultWriter) WriteHeader(info ...ColumnInfo) error { + if w.wroteHeaders { + return errors.New("double-write of query headers") + } + + // Translate column information into a row description message. + desc := make([]message.ColumnDescription, len(info)) + for i, c := range info { + t, err := w.te.TranslateType(c.Type) + if err != nil { + return errors.Wrap(err, "translating column type") + } + t.Name = c.Name + t.TableID = c.TableID + t.FieldID = c.FieldID + desc[i] = t + } + + // Encode the row description. + msg, err := w.enc.RowDescription(desc...) + if err != nil { + return errors.Wrap(err, "encoding query header") + } + + w.wroteHeaders = true + w.width = len(desc) + + // Write the row description. + return w.w.WriteMessage(msg) +} + +func (w *queryResultWriter) WriteRowText(text ...string) error { + // Check preconditions of the call. + switch { + case !w.wroteHeaders: + return errors.New("writing rows without headers") + case len(text) != w.width: + return errors.Errorf("expected %d columns but found %d", w.width, len(text)) + } + + // Encode the row data as text into a DataRow message. + msg, err := w.enc.TextRow(text...) + if err != nil { + return err + } + + // Write the data row over the network. + return w.w.WriteMessage(msg) +} + +func (w *queryResultWriter) Tag(tag string) { + w.tag = tag +} + +var _ QueryResultWriter = (*queryResultWriter)(nil) diff --git a/pg/server.go b/pg/server.go new file mode 100644 index 000000000..55d374005 --- /dev/null +++ b/pg/server.go @@ -0,0 +1,136 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package pg + +import ( + "context" + "crypto/tls" + "net" + "sync" + "time" + + "github.com/pilosa/pilosa/v2/logger" +) + +// Server is a postgres wire protocol server. +type Server struct { + // QueryHandler will be used to serve query requests. + QueryHandler QueryHandler + + // TypeEngine is the type engine to use to result columns in query requests. + TypeEngine TypeEngine + + // TLSConfig is the TLS configuration to use to serve postgres TLS connections. + TLSConfig *tls.Config + + // StartupTimeout is the timeout to use for connection startup. + // If a connection fails to set up a protocol before this completes, it will be terminated. + StartupTimeout time.Duration + + // ReadTimeout is the timeout to apply for active reads (reads during the lifetime of a command). + // This timeout does not apply to an idling connection. + ReadTimeout time.Duration + + // WriteTimeout is the timeout to apply to network writes. + WriteTimeout time.Duration + + // MaxStartupSize is the maximum size of the startup packet (in bytes). + // This defaults to 2^31-1 bytes, which is the maximum size allowed by the protocol. + MaxStartupSize uint32 + + // ConnectionLimit is the maximum number of connections to allow at once. + ConnectionLimit uint16 + + // Logger is the logger to use for error conditions and state changes. + Logger logger.Logger +} + +// ServeConn serves a single connection. +func (s *Server) ServeConn(ctx context.Context, conn net.Conn) error { + return s.handle(ctx, conn) +} + +// Serve accepts postgres connections from a listener and processes them. +// If the context is cancelled, this will stop accepting requests and wait until all connections have terminated. +// No error will be returned if terminated by context cancellation. +// This will close the connection for the caller. +func (s *Server) Serve(ctx context.Context, l net.Listener) (err error) { + // Ignore errors triggered by a shutdown. + // Also propogate any error from terminating the listener. + var cerr error + defer func(ctx context.Context) { + if ctx.Err() == context.Canceled { + err = cerr + } + }(ctx) + + // Wait for the listener to be closed and all connections to shut down. + var wg sync.WaitGroup + defer wg.Wait() + + // Wrap the context to propogate a shutdown to the listeners and connection handlers. + ctx, cancel := context.WithCancel(ctx) + defer cancel() + + // Start a goroutine to shut down the listener when the context is canceled. + wg.Add(1) + go func() { + defer wg.Done() + + <-ctx.Done() + cerr = l.Close() + }() + + // Set up a semaphore for the connection limit. + var limit chan struct{} + done := ctx.Done() + if s.ConnectionLimit != 0 { + limit = make(chan struct{}, s.ConnectionLimit) + } + + for { + if limit != nil { + // Wait for connection limit. + if len(limit) == cap(limit) { + s.Logger.Printf("postgres connection limit reached") + } + select { + case limit <- struct{}{}: + case <-done: + return nil + } + } + + // Accept a connection. + conn, err := l.Accept() + if err != nil { + return err + } + + // Handle the connection in another goroutine. + wg.Add(1) + go func() { + defer wg.Done() + if limit != nil { + // Restore connection limit when done. + defer func() { <-limit }() + } + err := s.handle(ctx, conn) + if err != nil { + s.Logger.Printf("postgres connection terminated with error: %v", err) + } + }() + } +} diff --git a/pg/server_test.go b/pg/server_test.go new file mode 100644 index 000000000..4903fdb60 --- /dev/null +++ b/pg/server_test.go @@ -0,0 +1,218 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package pg_test + +import ( + "context" + "fmt" + "net" + "os/exec" + "strconv" + "testing" + "time" + + "github.com/lib/pq" + "github.com/pilosa/pilosa/v2/logger" + "github.com/pilosa/pilosa/v2/pg" + "github.com/pilosa/pilosa/v2/pg/pgtest" +) + +// TestStartupTimeout tests that an incoming connection that does nothing times out and gets closed. +func TestStartupTimeout(t *testing.T) { + t.Parallel() + + connect, shutdown, err := pgtest.ServeMem(&pg.Server{ + StartupTimeout: time.Millisecond, + Logger: logger.NewLogfLogger(t), + }) + if err != nil { + t.Fatalf("starting in-memory postgres server: %v", err) + } + defer shutdown.Finish(t, "in-memory postgres server") + + conn, err := connect() + if err != nil { + t.Fatalf("failed to acquire connection: %v", err) + } + defer conn.Close() + + // The server isn't sending anything, so this should block until the connection dies. + conn.Read(make([]byte, 1024)) //nolint:errcheck +} + +// TestStartupInvalidLength tests that sending an HTTP GET request does not cause the server to allocate 1.2 GiB of memory. +func TestStartupInvalidLength(t *testing.T) { + t.Parallel() + + res := testing.Benchmark(func(b *testing.B) { + connect, shutdown, err := pgtest.ServeMem(&pg.Server{ + MaxStartupSize: 1024, + Logger: logger.NewLogfLogger(t), + }) + if err != nil { + t.Fatalf("starting in-memory postgres server: %v", err) + } + defer shutdown.Finish(t, "in-memory postgres server") + + b.ReportAllocs() + + b.ResetTimer() + + for i := 0; i < b.N; i++ { + conn, err := connect() + if err != nil { + t.Fatalf("failed to acquire connection: %v", err) + } + + _, err = conn.Write([]byte("GET ")) + if err != nil { + t.Fatalf("failed to write invalid length: %v", err) + } + + // The server isn't sending anything, so this should block until the connection dies. + conn.Read(make([]byte, 1024)) //nolint:errcheck + + err = conn.Close() + if err != nil { + t.Fatalf("failed to close connection: %v", err) + } + } + }) + bpo := res.AllocedBytesPerOp() + t.Logf("allocated %d bytes per op", bpo) + if bpo > 1024*1024 { + t.Errorf("allocated too much memory: %d bytes/connection", bpo) + } +} + +// TestPQConnect tests connecting the Go SQL driver `pq` to this postgres server. +func TestPQConnect(t *testing.T) { + t.Parallel() + + server := &pg.Server{ + StartupTimeout: time.Second, + Logger: logger.NewLogfLogger(t), + } + addr, shutdown, err := pgtest.ServeTCP(":0", server) + if err != nil { + t.Fatalf("starting postgres server: %v", err) + } + defer shutdown.Finish(t, "postgres server") + + tcpAddr := addr.(*net.TCPAddr) + + connector, err := pq.NewConnector(fmt.Sprintf("user=molecula dbname=pilosa sslmode=disable host=%s port=%d", tcpAddr.IP, tcpAddr.Port)) + if err != nil { + t.Fatalf("failed to create connector: %v", err) + } + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + conn, err := connector.Connect(ctx) + if err != nil { + t.Fatalf("failed to connect to postgres: %v", err) + } + defer pgtest.ShutdownFunc(conn.Close).Finish(t, "postgres TLS conn") +} + +// TestPQConnectSSL tests connecting the Go SQL driver `pq` to this postgres server, with SSL enabled. +func TestPQConnectSSL(t *testing.T) { + t.Parallel() + + server := &pg.Server{ + StartupTimeout: time.Second, + Logger: logger.NewLogfLogger(t), + } + addr, shutdown, err := pgtest.ServeTLS(":0", server) + if err != nil { + t.Fatalf("starting postgres server: %v", err) + } + defer shutdown.Finish(t, "postgres TLS server") + + tcpAddr := addr.(*net.TCPAddr) + + connector, err := pq.NewConnector(fmt.Sprintf("user=molecula dbname=pilosa sslmode=require host=%s port=%d", tcpAddr.IP, tcpAddr.Port)) + if err != nil { + t.Fatalf("failed to create connector: %v", err) + } + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + conn, err := connector.Connect(ctx) + if err != nil { + t.Fatalf("failed to connect to postgres: %v", err) + } + defer pgtest.ShutdownFunc(conn.Close).Finish(t, "postgres TLS conn") +} + +// TestPSQLQuery tests sending a query from the `psql` command line tool. +func TestPSQLQuery(t *testing.T) { + // Check if psql is present. + // Skip this test if it is not. + _, err := exec.LookPath("psql") + if err != nil { + if err, ok := err.(*exec.Error); ok { + if err.Err == exec.ErrNotFound { + t.Skip("psql is not available") + } + } + t.Fatalf("searching for psql: %v", err) + } + + server := &pg.Server{ + QueryHandler: pgtest.HandlerFunc(func(ctx context.Context, w pg.QueryResultWriter, q pg.Query) error { + err := w.WriteHeader(pg.ColumnInfo{ + Name: "field", + Type: pg.TypeCharoid, + }) + if err != nil { + return err + } + + err = w.WriteRowText("h") + if err != nil { + return err + } + + err = w.WriteRowText("xyzzy") + if err != nil { + return err + } + + return nil + }), + TypeEngine: pg.PrimitiveTypeEngine{}, + StartupTimeout: time.Second, + Logger: logger.NewLogfLogger(t), + } + addr, shutdown, err := pgtest.ServeTCP(":0", server) + if err != nil { + t.Fatalf("starting postgres server: %v", err) + } + defer shutdown.Finish(t, "postgres server") + + tcpAddr := addr.(*net.TCPAddr) + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + cmd := exec.CommandContext(ctx, "psql", "-h", tcpAddr.IP.String(), "-p", strconv.Itoa(tcpAddr.Port), "-c", "test query") + data, err := cmd.CombinedOutput() + if err != nil { + t.Fatalf("psql failed: %v", string(data)) + } +} diff --git a/pg/type.go b/pg/type.go new file mode 100644 index 000000000..7a89cc947 --- /dev/null +++ b/pg/type.go @@ -0,0 +1,54 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package pg + +import "github.com/pilosa/pilosa/v2/pg/message" + +// Type represents a postgres type. +type Type struct { + // I am not entirely sure what should be in here long term. + // For now, I am just going to leave it like this. + id int32 +} + +// TypeCharoid is a postgres type for text. +var TypeCharoid = Type{id: 18} + +// TypeData is a type containing raw postgres wire type information. +type TypeData struct { + TypeID int32 + TypeLen int16 + TypeModifier int32 +} + +// TypeEngine is a system for managing types. +// This is necessary for compound types like arrays which need ID generation. +type TypeEngine interface { + // TranslateType populates a column description with type information. + TranslateType(Type) (message.ColumnDescription, error) +} + +// PrimitiveTypeEngine is a simple type engine that only works on primitive types. +type PrimitiveTypeEngine struct{} + +// TranslateType translates a type to a column description. +func (pte PrimitiveTypeEngine) TranslateType(t Type) (message.ColumnDescription, error) { + return message.ColumnDescription{ + TypeID: t.id, // just charoid for now; as far as I can tell most implementations do not really use this + TypeLen: -1, // vdsm had 4. . . but the spec says this should be negative + TypeModifier: -1, + Mode: 0, // send as text + }, nil +} diff --git a/server/config.go b/server/config.go index 2da11a365..7437d65ed 100644 --- a/server/config.go +++ b/server/config.go @@ -167,6 +167,25 @@ type Config struct { MutexFraction int `toml:"mutex-fraction"` } `toml:"profile"` + Postgres struct { + // Addr is the address to which to bind a postgres endpoint. + // If this is empty, no endpoint will be created. + Addr string `toml:"addr"` + // TLS configuration for postgres connections. + TLS TLSConfig `toml:"tls"` + + StartupTimeout toml.Duration `toml:"startup-timeout"` + ReadTimeout toml.Duration `toml:"read-timeout"` + WriteTimeout toml.Duration `toml:"write-timout"` + + MaxStartupSize uint32 `toml:"max-startup-size"` + + // ConnectionLimit is the maximum number of postgres connections to allow simultaneously. + // Setting this to 0 disables the limit. + // This mostly exists because other DBs seem to have it. + ConnectionLimit uint16 `toml:"max-connections"` + } `toml:"postgres"` + // Txsrc determines which Tx implementation the holder/Index will use; one // of the available transactional-storage engines. Choices are listed // in the string constants below. Should be one of @@ -235,6 +254,13 @@ func NewConfig() *Config { c.Profile.BlockRate = 10000000 // 1 sample per 10 ms c.Profile.MutexFraction = 100 // 1% sampling + // Postgres config (off by default). + c.Postgres.MaxStartupSize = 8 * 1024 * 1024 + c.Postgres.StartupTimeout = toml.Duration(5 * time.Second) + c.Postgres.ReadTimeout = toml.Duration(10 * time.Second) + c.Postgres.WriteTimeout = toml.Duration(10 * time.Second) + // we don't really need a connection limit + return c } diff --git a/server/pg.go b/server/pg.go new file mode 100644 index 000000000..98117c374 --- /dev/null +++ b/server/pg.go @@ -0,0 +1,399 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package server + +import ( + "context" + "crypto/tls" + "encoding/json" + "fmt" + "net" + "strconv" + "strings" + "time" + + "github.com/pilosa/pilosa/v2" + "github.com/pilosa/pilosa/v2/logger" + "github.com/pilosa/pilosa/v2/pg" + pb "github.com/pilosa/pilosa/v2/proto" + + "github.com/pkg/errors" + "golang.org/x/sync/errgroup" +) + +// PostgresServer provides a postgres endpoint on pilosa. +type PostgresServer struct { + api *pilosa.API + logger logger.Logger + eg errgroup.Group + s pg.Server + stop context.CancelFunc +} + +// NewPostgresServer creates a postgres server. +func NewPostgresServer(api *pilosa.API, logger logger.Logger, tls *tls.Config) *PostgresServer { + return &PostgresServer{ + api: api, + logger: logger, + s: pg.Server{ + QueryHandler: &queryDecodeHandler{ + child: &pilosaQueryHandler{ + api: api, + }, + }, + TypeEngine: pg.PrimitiveTypeEngine{}, + StartupTimeout: 5 * time.Second, + ReadTimeout: 10 * time.Second, + WriteTimeout: 10 * time.Second, + MaxStartupSize: 8 * 1024 * 1024, + Logger: logger, + TLSConfig: tls, + }, + } +} + +// Start a postgres endpoint at the specified address. +func (s *PostgresServer) Start(addr string) error { + l, err := net.Listen("tcp", addr) + if err != nil { + return errors.Wrap(err, "creating listener") + } + + s.logger.Printf("serving postgres wire protocol on %s", l.Addr()) + + ctx, cancel := context.WithCancel(context.Background()) + s.stop = cancel + + s.eg.Go(func() error { return s.s.Serve(ctx, l) }) + + return nil +} + +func (s *PostgresServer) Close() error { + if s == nil { + return nil + } + + if s.stop == nil { + return nil + } + s.stop() + + s.logger.Printf("waiting for postgres connections to shut down") + + return s.eg.Wait() +} + +type pgPQLQuery struct { + index string + query string +} + +func (q pgPQLQuery) String() string { + return fmt.Sprintf("[%s]%s", q.index, q.query) +} + +func pgDecodePQL(str string) (q pg.Query, err error) { + defer func() { + err = errors.Wrap(err, "not a valid PQL-over-postgres query") + }() + + if !strings.HasPrefix(str, "[") { + return nil, errors.New("missing index specification") + } + idx := strings.IndexRune(str, ']') + if idx == -1 { + return nil, errors.New("unclosed bracket in index specification") + } + + return pgPQLQuery{ + index: str[1:idx], + query: str[idx+1:], + }, nil +} + +type pilosaQueryHandler struct { + api *pilosa.API +} + +func pgWriteRow(w pg.QueryResultWriter, row *pilosa.Row) error { + err := w.WriteHeader(pg.ColumnInfo{ + Name: "_id", + Type: pg.TypeCharoid, + }) + if err != nil { + return errors.Wrap(err, "writing result header") + } + + if row.Keys != nil { + for _, k := range row.Keys { + err = w.WriteRowText(k) + if err != nil { + return errors.Wrap(err, "writing key") + } + } + } else { + for _, col := range row.Columns() { + err = w.WriteRowText(strconv.FormatUint(col, 10)) + if err != nil { + return errors.Wrap(err, "writing column ID") + } + } + } + + return nil +} + +func pgWriteRows(w pg.QueryResultWriter, rows pilosa.RowIdentifiers) error { + err := w.WriteHeader(pg.ColumnInfo{ + Name: rows.Field(), + Type: pg.TypeCharoid, + }) + if err != nil { + return errors.Wrap(err, "writing result header") + } + + if rows.Keys != nil { + for _, k := range rows.Keys { + err = w.WriteRowText(k) + if err != nil { + return errors.Wrap(err, "writing key") + } + } + } else { + for _, row := range rows.Rows { + err = w.WriteRowText(strconv.FormatUint(row, 10)) + if err != nil { + return errors.Wrap(err, "writing row ID") + } + } + } + + return nil +} + +func pgFormatVal(val interface{}) string { + switch val := val.(type) { + case bool: + return strconv.FormatBool(val) + case int64: + return strconv.FormatInt(val, 10) + case uint64: + return strconv.FormatUint(val, 10) + case string: + return val + default: + data, _ := json.Marshal(val) + return string(data) + } +} + +func pgWriteExtractedTable(w pg.QueryResultWriter, tbl pilosa.ExtractedTable) error { + headers := make([]pg.ColumnInfo, len(tbl.Fields)+1) + headers[0] = pg.ColumnInfo{ + Name: "_id", + Type: pg.TypeCharoid, + } + dataHeaders := headers[1:] + for i, f := range tbl.Fields { + dataHeaders[i] = pg.ColumnInfo{ + Name: f.Name, + Type: pg.TypeCharoid, + } + } + err := w.WriteHeader(headers...) + if err != nil { + return errors.Wrap(err, "writing result header") + } + + vals := make([]string, len(headers)) + dataVals := vals[1:] + for _, col := range tbl.Columns { + if col.Column.Keyed { + vals[0] = col.Column.Key + } else { + vals[0] = strconv.FormatUint(col.Column.ID, 10) + } + for i, v := range col.Rows { + dataVals[i] = pgFormatVal(v) + } + err = w.WriteRowText(vals...) + if err != nil { + return errors.Wrap(err, "writing result row") + } + } + + return nil +} + +func pgWriteGroupCount(w pg.QueryResultWriter, counts []pilosa.GroupCount) error { + if len(counts) == 0 { + // Not enough information is available to construct the header. + // This is a significant flaw in the data type. + return nil + } + + headers := make([]pg.ColumnInfo, len(counts[0].Group)+2) + for i, g := range counts[0].Group { + headers[i] = pg.ColumnInfo{ + Name: g.Field, + Type: pg.TypeCharoid, + } + } + headers[len(headers)-2] = pg.ColumnInfo{ + Name: "count", + Type: pg.TypeCharoid, + } + headers[len(headers)-1] = pg.ColumnInfo{ + Name: "sum", + Type: pg.TypeCharoid, + } + err := w.WriteHeader(headers...) + if err != nil { + return errors.Wrap(err, "writing result header") + } + + vals := make([]string, len(headers)) + for _, gc := range counts { + for j, g := range gc.Group { + var v string + switch { + case g.Value != nil: + v = strconv.FormatInt(*g.Value, 10) + case g.RowKey != "": + v = g.RowKey + default: + v = strconv.FormatUint(g.RowID, 10) + } + vals[j] = v + } + vals[len(vals)-2] = strconv.FormatUint(gc.Count, 10) + vals[len(vals)-1] = strconv.FormatInt(gc.Sum, 10) + + err := w.WriteRowText(vals...) + if err != nil { + return errors.Wrap(err, "writing group count result") + } + } + + return nil +} + +func pgWriteRowser(w pg.QueryResultWriter, result pb.ToRowser) error { + var data []string + return result.ToRows(func(row *pb.RowResponse) error { + if data == nil { + headers := make([]pg.ColumnInfo, len(row.Columns)) + for i, h := range row.Headers { + headers[i] = pg.ColumnInfo{ + Name: h.Name, + Type: pg.TypeCharoid, + } + } + err := w.WriteHeader(headers...) + if err != nil { + return errors.Wrap(err, "writing headers") + } + + data = make([]string, len(headers)) + } + + for i, col := range row.Columns { + var v string + switch col := col.ColumnVal.(type) { + case *pb.ColumnResponse_BoolVal: + v = strconv.FormatBool(col.BoolVal) + case *pb.ColumnResponse_DecimalVal: + v = col.DecimalVal.String() + case *pb.ColumnResponse_Float64Val: + v = strconv.FormatFloat(col.Float64Val, 'g', -1, 64) + case *pb.ColumnResponse_Int64Val: + v = strconv.FormatInt(col.Int64Val, 10) + case *pb.ColumnResponse_Uint64Val: + v = strconv.FormatUint(col.Uint64Val, 10) + case *pb.ColumnResponse_StringVal: + v = col.StringVal + case *pb.ColumnResponse_StringArrayVal: + data, _ := json.Marshal(col.StringArrayVal.Vals) + v = string(data) + case *pb.ColumnResponse_Uint64ArrayVal: + data, _ := json.Marshal(col.Uint64ArrayVal.Vals) + v = string(data) + default: + return errors.Errorf("unable to process value of type %T", col) + } + + data[i] = v + } + + return w.WriteRowText(data...) + }) +} + +func pgWriteResult(w pg.QueryResultWriter, result interface{}) error { + switch result := result.(type) { + case *pilosa.Row: + return pgWriteRow(w, result) + case pilosa.RowIdentifiers: + return pgWriteRows(w, result) + case pilosa.ExtractedTable: + return pgWriteExtractedTable(w, result) + case []pilosa.GroupCount: + return pgWriteGroupCount(w, result) + case pb.ToRowser: // we should avoid protobuf where we can... + return pgWriteRowser(w, result) + default: + return errors.Errorf("result type %T not yet supported", result) + } +} + +func (pqh *pilosaQueryHandler) HandleQuery(ctx context.Context, w pg.QueryResultWriter, q pg.Query) error { + switch q := q.(type) { + case pgPQLQuery: + resp, err := pqh.api.Query(ctx, &pilosa.QueryRequest{ + Index: q.index, + Query: q.query, + }) + if err != nil { + return errors.Wrap(err, "executing query") + } + if len(resp.Results) != 1 { + return errors.Errorf("expected 1 query result but found %d", len(resp.Results)) + } + return errors.Wrap(pgWriteResult(w, resp.Results[0]), "writing query result") + + default: + return errors.Errorf("query type %T not yet supported (query: %s)", q, q) + } +} + +type queryDecodeHandler struct { + child pg.QueryHandler +} + +func (qdh *queryDecodeHandler) HandleQuery(ctx context.Context, w pg.QueryResultWriter, q pg.Query) error { + switch qv := q.(type) { + case pg.SimpleQuery: + if strings.HasPrefix(string(qv), "[") { + pqlQuery, err := pgDecodePQL(strings.TrimSuffix(string(qv), ";")) + if err != nil { + return errors.Wrap(err, "decoding query") + } + q = pqlQuery + } + } + + return qdh.child.HandleQuery(ctx, w, q) +} diff --git a/server/server.go b/server/server.go index 2b1cea4bb..5578cf35f 100644 --- a/server/server.go +++ b/server/server.go @@ -87,6 +87,7 @@ type Command struct { listenURI *pilosa.URI tlsConfig *tls.Config closeTimeout time.Duration + pgserver *PostgresServer serverOptions []pilosa.ServerOption } @@ -171,6 +172,29 @@ func (m *Command) Start() (err error) { } }() + // Initialize postgres. + m.pgserver = nil + if m.Config.Postgres.Addr != "" { + var tlsConf *tls.Config + if m.Config.Postgres.TLS.CertificatePath != "" { + conf, err := GetTLSConfig(&m.Config.Postgres.TLS, m.logger.Logger()) + if err != nil { + return errors.Wrap(err, "settuing up postgres TLS") + } + tlsConf = conf + } + m.pgserver = NewPostgresServer(m.API, m.logger, tlsConf) + m.pgserver.s.StartupTimeout = time.Duration(m.Config.Postgres.StartupTimeout) + m.pgserver.s.ReadTimeout = time.Duration(m.Config.Postgres.ReadTimeout) + m.pgserver.s.WriteTimeout = time.Duration(m.Config.Postgres.WriteTimeout) + m.pgserver.s.MaxStartupSize = m.Config.Postgres.MaxStartupSize + m.pgserver.s.ConnectionLimit = m.Config.Postgres.ConnectionLimit + err := m.pgserver.Start(m.Config.Postgres.Addr) + if err != nil { + return errors.Wrap(err, "starting postgres") + } + } + close(m.Started) return nil } @@ -515,6 +539,7 @@ func (m *Command) Close() error { eg.Go(m.Handler.Close) eg.Go(m.Server.Close) eg.Go(m.API.Close) + eg.Go(m.pgserver.Close) if m.gossipMemberSet != nil { eg.Go(m.gossipMemberSet.Close) } From 44061923c70caa9a600d61a4d0b36c76772c0e79 Mon Sep 17 00:00:00 2001 From: Nia Date: Thu, 20 Aug 2020 14:14:14 -0400 Subject: [PATCH 08/17] Update pg/message/io.go Co-authored-by: Travis Turner --- pg/message/io.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pg/message/io.go b/pg/message/io.go index 115048229..706b8e999 100644 --- a/pg/message/io.go +++ b/pg/message/io.go @@ -82,7 +82,7 @@ func NewWireReader(r *bufio.Reader) *WireReader { return &WireReader{r: r} } -// ErrMessageTooBig is an error indicating that a message is to big to be sent or recieved. +// ErrMessageTooBig is an error indicating that a message is too big to be sent or received. var ErrMessageTooBig = errors.New("message is too big") // WireWriter writes messages in Postgres wire protocol. From 14676c07133861e081d5e6f2b1bb0397a21523a3 Mon Sep 17 00:00:00 2001 From: Nia Weiss Date: Thu, 20 Aug 2020 14:17:50 -0400 Subject: [PATCH 09/17] remove postgres debugging types --- pg/message/io.go | 24 ------------------------ 1 file changed, 24 deletions(-) diff --git a/pg/message/io.go b/pg/message/io.go index 706b8e999..e09150b9d 100644 --- a/pg/message/io.go +++ b/pg/message/io.go @@ -18,7 +18,6 @@ import ( "bufio" "encoding/binary" "errors" - "fmt" "io" ) @@ -123,26 +122,3 @@ var _ Writer = (*WireWriter)(nil) func NewWireWriter(w *bufio.Writer) *WireWriter { return &WireWriter{w: w} } - -type DumpWriter struct { - Writer - Out io.Writer -} - -func (w *DumpWriter) WriteMessage(msg Message) error { - fmt.Fprintf(w.Out, "write type %s: %x %q\n", string(msg.Type), msg.Data, string(msg.Data)) - return w.Writer.WriteMessage(msg) -} - -type DumpReader struct { - Reader - Out io.Writer -} - -func (r *DumpReader) ReadMessage() (Message, error) { - msg, err := r.Reader.ReadMessage() - if err == nil { - fmt.Fprintf(r.Out, "read type %s: %x %q\n", string(msg.Type), msg.Data, string(msg.Data)) - } - return msg, err -} From df5dd1557e183607639e111d715dabd6460607ac Mon Sep 17 00:00:00 2001 From: Jason Aten Date: Tue, 18 Aug 2020 07:33:09 -0500 Subject: [PATCH 10/17] AtomicRecord allows the client to request atomic updates. - Atomic record contains multiple ImportRequest and ImportValueRequest, plus ability to Clear individual requests. - adds http handlers for importing AtomicRecord. --- Makefile | 12 + api.go | 118 +++++-- badger_test.go | 4 +- encoding/proto/proto.go | 44 +++ fragment.go | 2 +- fragment_internal_test.go | 18 +- handler.go | 24 +- http/handler.go | 57 ++++ internal/public.pb.go | 641 +++++++++++++++++++++++++++++++------- internal/public.proto | 9 + lmdb_test.go | 7 +- pg/protocol.go | 2 +- row.go | 2 +- tx_test.go | 245 +++++++++++++++ 14 files changed, 1031 insertions(+), 154 deletions(-) create mode 100644 tx_test.go diff --git a/Makefile b/Makefile index 4ca3b029e..5794843ce 100644 --- a/Makefile +++ b/Makefile @@ -234,6 +234,18 @@ bg-rbf: @echo " log.bg-rbf green: \c"; cat log.bg-rbf | grep PASS |wc -l @echo " log.bg-rbf red: \c"; cat log.bg-rbf | grep '\-\-\- FAIL' |wc -l +rbf-lm: + mv log.rbf-lm log.rbf-lm.prev || true + PILOSA_TXSRC=rbf_badger go test -v -tags='$(BUILD_TAGS)' $(TESTFLAGS) $(NOCHECKPTR) 2>&1 | tee log.rbf-lm + @echo " log.rbf-lm green: \c"; cat log.rbf-lm | grep PASS |wc -l + @echo " log.rbf-lm red: \c"; cat log.rbf-lm | grep '\-\-\- FAIL' |wc -l + +lm-rbf: + mv log.lm-rbf log.lm-rbf.prev || true + PILOSA_TXSRC=badger_rbf go test -v -tags='$(BUILD_TAGS)' $(TESTFLAGS) $(NOCHECKPTR) 2>&1 | tee log.lm-rbf + @echo " log.lm-rbf green: \c"; cat log.lm-rbf | grep PASS |wc -l + @echo " log.lm-rbf red: \c"; cat log.lm-rbf | grep '\-\-\- FAIL' |wc -l + lm-rr: mv log.lm-rr log.lm-rr.prev || true PILOSA_TXSRC=lmdb_roaring go test -v -tags='$(BUILD_TAGS)' $(TESTFLAGS) $(NOCHECKPTR) 2>&1 | tee log.lm-rr diff --git a/api.go b/api.go index 287dbc66c..d8cdef042 100644 --- a/api.go +++ b/api.go @@ -1022,6 +1022,9 @@ type ImportOptions struct { Clear bool IgnoreKeyCheck bool Presorted bool + + // test Tx atomicity if > 0 + SimPowerLossAfter int } // ImportOption is a functional option type for API.Import. @@ -1052,8 +1055,71 @@ func OptImportOptionsPresorted(b bool) ImportOption { } } -// Import bulk imports data into a particular index,field,shard. +var ErrAborted = fmt.Errorf("error: update was aborted") + +func (api *API) ImportAtomicRecord(ctx context.Context, req *AtomicRecord, opts ...ImportOption) error { + simPowerLoss := false + lossAfter := -1 + var opt ImportOptions + for _, setter := range opts { + if setter != nil { + err := setter(&opt) + if err != nil { + return errors.Wrap(err, "ImportAtomicRecord ImportOptions") + } + } + } + if opt.SimPowerLossAfter > 0 { + simPowerLoss = true + lossAfter = opt.SimPowerLossAfter + } + + idx, err := api.Index(ctx, req.Index) + if err != nil { + return errors.Wrap(err, "getting index") + } + + // the whole point is to run this part of the import atomically. + // So make a Tx. + tx := idx.Txf.NewTx(Txo{Write: writable, Index: idx}) + defer tx.Rollback() + + tot := 0 + + // BSIs (Values) + for _, ivr := range req.Ivr { + tot++ + if simPowerLoss && tot > lossAfter { + return ErrAborted + } + opts0 := append(opts, OptImportOptionsClear(ivr.Clear)) + err := api.ImportValueWithTx(ctx, tx, ivr, opts0...) + if err != nil { + return errors.Wrap(err, "ImportAtomicRecord ImportValueWithTx") + } + } + + // other bits, non-BSI + for _, ir := range req.Ir { + tot++ + if simPowerLoss && tot > lossAfter { + return ErrAborted + } + opts0 := append(opts, OptImportOptionsClear(ir.Clear)) + err := api.ImportWithTx(ctx, tx, ir, opts0...) + if err != nil { + return errors.Wrap(err, "ImportAtomicRecord ImportWithTx") + } + } + return tx.Commit() +} + func (api *API) Import(ctx context.Context, req *ImportRequest, opts ...ImportOption) error { + return api.ImportWithTx(ctx, nil, req, opts...) +} + +// Import bulk imports data into a particular index,field,shard. +func (api *API) ImportWithTx(ctx context.Context, tx Tx, req *ImportRequest, opts ...ImportOption) error { span, _ := tracing.StartSpanFromContext(ctx, "API.Import") defer span.Finish() @@ -1156,39 +1222,43 @@ func (api *API) Import(ctx context.Context, req *ImportRequest, opts ...ImportOp timestamps[i] = &t } + isLocalTx := false + if tx == nil { + isLocalTx = true + tx = index.Txf.NewTx(Txo{Write: true, Index: index}) + defer tx.Rollback() + } + // Import columnIDs into existence field. if !options.Clear { - if err := func() error { - tx := index.Txf.NewTx(Txo{Write: true, Index: index}) - defer tx.Rollback() - - if err := importExistenceColumns(tx, index, req.ColumnIDs); err != nil { - api.server.logger.Printf("import existence error: index=%s, field=%s, shard=%d, columns=%d, err=%s", req.Index, req.Field, req.Shard, len(req.ColumnIDs), err) - return err - } - return tx.Commit() - }(); err != nil { + if err := importExistenceColumns(tx, index, req.ColumnIDs); err != nil { + api.server.logger.Printf("import existence error: index=%s, field=%s, shard=%d, columns=%d, err=%s", req.Index, req.Field, req.Shard, len(req.ColumnIDs), err) + return err + } + if err != nil { return errors.Wrap(err, "importing existence columns") } } - tx := index.Txf.NewTx(Txo{Write: true, Index: index}) - defer tx.Rollback() - // Import into fragment. err = field.Import(tx, req.RowIDs, req.ColumnIDs, timestamps, opts...) if err != nil { api.server.logger.Printf("import error: index=%s, field=%s, shard=%d, columns=%d, err=%s", req.Index, req.Field, req.Shard, len(req.ColumnIDs), err) - } else { - err = tx.Commit() + return errors.Wrap(err, "importing") } - return errors.Wrap(err, "importing") + if isLocalTx { + err = tx.Commit() + } + return errors.Wrap(err, "committing") +} +func (api *API) ImportValue(ctx context.Context, req *ImportValueRequest, opts ...ImportOption) error { + return api.ImportValueWithTx(ctx, nil, req, opts...) } // ImportValue bulk imports values into a particular field. -func (api *API) ImportValue(ctx context.Context, req *ImportValueRequest, opts ...ImportOption) error { +func (api *API) ImportValueWithTx(ctx context.Context, tx Tx, req *ImportValueRequest, opts ...ImportOption) error { span, _ := tracing.StartSpanFromContext(ctx, "API.ImportValue") defer span.Finish() @@ -1198,7 +1268,7 @@ func (api *API) ImportValue(ctx context.Context, req *ImportValueRequest, opts . index, field, err := api.indexField(req.Index, req.Field, req.Shard) if err != nil { - return errors.Wrap(err, "getting index and field") + return errors.Wrap(err, fmt.Sprintf("getting index '%v' and field '%v'; shard=%v", req.Index, req.Field, req.Shard)) } if err := req.ValidateWithTimestamp(index.CreatedAt(), field.CreatedAt()); err != nil { @@ -1258,11 +1328,15 @@ func (api *API) ImportValue(ctx context.Context, req *ImportValueRequest, opts . sort.Sort(req) } + isLocalTx := false // if we're importing into a specific shard if req.Shard != math.MaxUint64 { // Obtain transaction. - tx := index.Txf.NewTx(Txo{Write: true, Index: index}) - defer tx.Rollback() + if tx == nil { + isLocalTx = true + tx = index.Txf.NewTx(Txo{Write: true, Index: index}) + defer tx.Rollback() + } // Check that column IDs match the stated shard. if s1, s2 := req.ColumnIDs[0]/ShardWidth, req.ColumnIDs[len(req.ColumnIDs)-1]/ShardWidth; s1 != s2 && s2 != req.Shard { @@ -1293,7 +1367,7 @@ func (api *API) ImportValue(ctx context.Context, req *ImportValueRequest, opts . api.server.logger.Printf("import error: index=%s, field=%s, shard=%d, columns=%d, err=%s", req.Index, req.Field, req.Shard, len(req.ColumnIDs), err) } } - if err == nil { + if err == nil && isLocalTx { err = tx.Commit() } return errors.Wrap(err, "importing value") diff --git a/badger_test.go b/badger_test.go index 6c271bb0d..c1627f1db 100644 --- a/badger_test.go +++ b/badger_test.go @@ -1275,6 +1275,7 @@ func getTestBitmapAsRawRoaring(bitsToSet ...uint64) []byte { return buf.Bytes() } +/* func TestBadger_AutoCommit(t *testing.T) { // setup @@ -1330,6 +1331,7 @@ func TestBadger_BigWritesAvoidTxnTooLargeWithAutoCommit(t *testing.T) { err := tx.Commit() panicOn(err) } +*/ func TestBadger_DeleteIndex(t *testing.T) { @@ -1394,7 +1396,7 @@ func TestBadger_DeleteIndex(t *testing.T) { } func TestBadger_DeleteIndex_over100k(t *testing.T) { - + t.Skip("test big and long running, skip") // setup dbwrap, clean := mustOpenEmptyBadgerWrapper("TestBadger_DeleteIndex_over100k") defer clean() diff --git a/encoding/proto/proto.go b/encoding/proto/proto.go index 29cb2ecf7..a1453cd62 100644 --- a/encoding/proto/proto.go +++ b/encoding/proto/proto.go @@ -304,6 +304,14 @@ func (s Serializer) Unmarshal(buf []byte, m pilosa.Message) error { } decodeTransactionMessage(msg, mt) return nil + case *pilosa.AtomicRecord: + msg := &internal.AtomicRecord{} + err := proto.Unmarshal(buf, msg) + if err != nil { + return errors.Wrap(err, "unmarshaling AtomicRecord") + } + s.decodeAtomicRecord(msg, mt) + return nil default: panic(fmt.Sprintf("unhandled pilosa.Message of type %T: %#v", mt, m)) } @@ -375,6 +383,8 @@ func (s Serializer) encodeToProto(m pilosa.Message) proto.Message { return s.encodeTranslateIDsResponse(mt) case *pilosa.TransactionMessage: return s.encodeTransactionMessage(mt) + case *pilosa.AtomicRecord: + return s.encodeAtomicRecord(mt) } return nil } @@ -413,6 +423,7 @@ func (s Serializer) encodeImportRequest(m *pilosa.ImportRequest) *internal.Impor RowKeys: m.RowKeys, ColumnKeys: m.ColumnKeys, Timestamps: m.Timestamps, + Clear: m.Clear, } } @@ -428,6 +439,7 @@ func (s Serializer) encodeImportValueRequest(m *pilosa.ImportValueRequest) *inte Values: m.Values, FloatValues: m.FloatValues, StringValues: m.StringValues, + Clear: m.Clear, } } @@ -873,6 +885,20 @@ func (s Serializer) encodeTransactionMessage(msg *pilosa.TransactionMessage) *in } } +func (s Serializer) encodeAtomicRecord(msg *pilosa.AtomicRecord) *internal.AtomicRecord { + ar := &internal.AtomicRecord{ + Index: msg.Index, + Shard: msg.Shard, + } + for _, ivr := range msg.Ivr { + ar.Ivr = append(ar.Ivr, s.encodeImportValueRequest(ivr)) + } + for _, ir := range msg.Ir { + ar.Ir = append(ar.Ir, s.encodeImportRequest(ir)) + } + return ar +} + func (s Serializer) encodeTransaction(trns *pilosa.Transaction) *internal.Transaction { if trns == nil { return nil @@ -1179,6 +1205,7 @@ func (s Serializer) decodeImportRequest(pb *internal.ImportRequest, m *pilosa.Im m.Timestamps = pb.Timestamps m.IndexCreatedAt = pb.IndexCreatedAt m.FieldCreatedAt = pb.FieldCreatedAt + m.Clear = pb.Clear } func (s Serializer) decodeImportValueRequest(pb *internal.ImportValueRequest, m *pilosa.ImportValueRequest) { @@ -1192,6 +1219,7 @@ func (s Serializer) decodeImportValueRequest(pb *internal.ImportValueRequest, m m.StringValues = pb.StringValues m.IndexCreatedAt = pb.IndexCreatedAt m.FieldCreatedAt = pb.FieldCreatedAt + m.Clear = pb.Clear } func (s Serializer) decodeImportRoaringRequest(pb *internal.ImportRoaringRequest, m *pilosa.ImportRoaringRequest) { @@ -1295,6 +1323,22 @@ func decodeTransactionMessage(pb *internal.TransactionMessage, m *pilosa.Transac decodeTransaction(pb.Transaction, m.Transaction) } +func (s Serializer) decodeAtomicRecord(pb *internal.AtomicRecord, m *pilosa.AtomicRecord) { + m.Index = pb.Index + m.Shard = pb.Shard + m.Ivr = make([]*pilosa.ImportValueRequest, len(pb.Ivr)) + m.Ir = make([]*pilosa.ImportRequest, len(pb.Ir)) + + for i, ivr := range pb.Ivr { + m.Ivr[i] = &pilosa.ImportValueRequest{} + s.decodeImportValueRequest(ivr, m.Ivr[i]) + } + for i, ir := range pb.Ir { + m.Ir[i] = &pilosa.ImportRequest{} + s.decodeImportRequest(ir, m.Ir[i]) + } +} + func decodeTransaction(pb *internal.Transaction, trns *pilosa.Transaction) { trns.ID = pb.ID trns.Active = pb.Active diff --git a/fragment.go b/fragment.go index 523aac338..09a4bde3e 100644 --- a/fragment.go +++ b/fragment.go @@ -622,7 +622,7 @@ func (f *fragment) rowFromStorage(tx Tx, rowID uint64) (*Row, error) { // setBit sets a bit for a given column & row within the fragment. // This updates both the on-disk storage and the in-cache bitmap. func (f *fragment) setBit(tx Tx, rowID, columnID uint64) (changed bool, err error) { - f.mu.Lock() + f.mu.Lock() // controls access to the file. defer f.mu.Unlock() var wp *io.Writer if f.storage != nil { diff --git a/fragment_internal_test.go b/fragment_internal_test.go index a41cee1a9..8ef6d5909 100644 --- a/fragment_internal_test.go +++ b/fragment_internal_test.go @@ -5323,19 +5323,11 @@ func check(t *testing.T, tx Tx, f *fragment, exp map[uint64]map[uint64]struct{}) func TestImportValueConcurrent(t *testing.T) { f, idx := mustOpenBSIFragment("i", "f", viewBSIGroupPrefix+"foo", 0) - types := idx.Txf.TxTypes() - for _, ty := range types { - switch ty { - case roaringTxn: - t.Skip(fmt.Sprintf("skipping TestImportValueConcurrent under " + - "blueGreenTx because the lack of transactional consistency " + - "from Roaring-per-file will create false comparison " + - "failures.")) - case lmdbTxn: - t.Skip(fmt.Sprintf("skipping TestImportValueConcurrent under " + - "lmdb since only a single writer is allowed at once.")) - } - } + + // produces false positives under blue_green because of races + // between commits, and is probematic under single writer + // backends like rbf and lmdb. Marking as roaring-only. + roaringOnlyTest(t) // Since eg.Go gets called multiple times below, each // time needs its own Tx. So close the default one and diff --git a/handler.go b/handler.go index bb5aa245d..26e1bd186 100644 --- a/handler.go +++ b/handler.go @@ -113,6 +113,7 @@ var NopHandler Handler = nopHandler{} // ImportValueRequest describes the import request structure // for a value (BSI) import. +// Note: no RowIDs here. have to convert BSI Values into RowIDs internally. type ImportValueRequest struct { Index string IndexCreatedAt int64 @@ -121,11 +122,25 @@ type ImportValueRequest struct { // if Shard is MaxUint64 (an impossible shard value), this // indicates that the column IDs may come from multiple shards. Shard uint64 - ColumnIDs []uint64 + ColumnIDs []uint64 // e.g. weather stationID ColumnKeys []string - Values []int64 + Values []int64 // e.g. temperature, humidity, barometric pressure FloatValues []float64 StringValues []string + Clear bool // only works for ImportAtomicRecord() at the moment. +} + +// AtomicRecord applies all its Ivr and Ivr atomically, in a Tx. +// The top level Shard has to agree with Ivr[i].Shard and the Iv[i].Shard +// for all i included (in Ivr and Ir). The same goes for the top level Index: all records +// have to be writes to the same Index. These requirements are checked. +// +type AtomicRecord struct { + Index string + Shard uint64 + + Ivr []*ImportValueRequest // BSI values + Ir []*ImportRequest // other field types, e.g. single bit } func (ivr *ImportValueRequest) Len() int { return len(ivr.ColumnIDs) } @@ -176,7 +191,7 @@ func (ivr *ImportValueRequest) ValidateWithTimestamp(indexCreatedAt, fieldCreate } // ImportColumnAttrsRequest describes the import request structure -// for a ColumnAttr import +// for a ColumnAttr import. type ImportColumnAttrsRequest struct { AttrKey string ColumnIDs []uint64 @@ -187,7 +202,7 @@ type ImportColumnAttrsRequest struct { } // ImportRequest describes the import request structure -// for an import. +// for an import. BSIs use the ImportValueRequest instead. type ImportRequest struct { Index string IndexCreatedAt int64 @@ -199,6 +214,7 @@ type ImportRequest struct { RowKeys []string ColumnKeys []string Timestamps []int64 + Clear bool // only works for ImportAtomicRecord() at the moment. } // ValidateWithTimestamp ensures that the payload of the request is valid. diff --git a/http/handler.go b/http/handler.go index 890a4a31d..f0514bda1 100644 --- a/http/handler.go +++ b/http/handler.go @@ -204,6 +204,7 @@ func (h *Handler) populateValidators() { h.validators["PostField"] = queryValidationSpecRequired() h.validators["DeleteField"] = queryValidationSpecRequired() h.validators["PostImport"] = queryValidationSpecRequired().Optional("clear", "ignoreKeyCheck") + h.validators["PostImportAtomicRecord"] = queryValidationSpecRequired().Optional("simPowerLossAfter") h.validators["PostImportRoaring"] = queryValidationSpecRequired().Optional("remote", "clear") h.validators["PostQuery"] = queryValidationSpecRequired().Optional("shards", "columnAttrs", "excludeRowAttrs", "excludeColumns", "profile") h.validators["GetInfo"] = queryValidationSpecRequired() @@ -343,6 +344,7 @@ func newRouter(handler *Handler) *mux.Router { router.Handle("/debug/vars", expvar.Handler()).Methods("GET") router.Handle("/metrics", promhttp.Handler()) router.HandleFunc("/export", handler.handleGetExport).Methods("GET").Name("GetExport") + router.HandleFunc("/import-atomic-record", handler.handlePostImportAtomicRecord).Methods("POST").Name("PostImportAtomicRecord") router.HandleFunc("/index", handler.handleGetIndexes).Methods("GET").Name("GetIndexes") router.HandleFunc("/index", handler.handlePostIndex).Methods("POST").Name("PostIndex") router.HandleFunc("/index/", handler.handlePostIndex).Methods("POST").Name("PostIndex") @@ -2020,6 +2022,61 @@ func GetHTTPClient(t *tls.Config) *http.Client { return &http.Client{Transport: transport} } +// handlePostImportAtomicRecord handles /import-atomic-record requests +func (h *Handler) handlePostImportAtomicRecord(w http.ResponseWriter, r *http.Request) { + + // Verify that request is only communicating over protobufs. + if error, code := validateProtobufHeader(r); error != "" { + http.Error(w, error, code) + return + } + + // Read entire body. + body, err := readBody(r) + if err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } + // Unmarshal request based on field type. + + q := r.URL.Query() + sLoss := q.Get("simPowerLossAfter") + loss := 0 + if sLoss != "" { + l, err := strconv.ParseInt(sLoss, 10, 64) + loss = int(l) + if err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + } + } + opt := func(o *pilosa.ImportOptions) error { + o.SimPowerLossAfter = loss + return nil + } + + req := &pilosa.AtomicRecord{} + if err := proto.DefaultSerializer.Unmarshal(body, req); err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } + + if err := h.api.ImportAtomicRecord(r.Context(), req, opt); err != nil { + switch errors.Cause(err) { + case pilosa.ErrClusterDoesNotOwnShard, pilosa.ErrPreconditionFailed: + http.Error(w, err.Error(), http.StatusPreconditionFailed) + default: + http.Error(w, err.Error(), http.StatusInternalServerError) + } + return + } + + // Write response. + _, err = w.Write(importOk) + if err != nil { + h.logger.Printf("writing import response: %v", err) + } +} + // handlePostImport handles /import requests. func (h *Handler) handlePostImport(w http.ResponseWriter, r *http.Request) { // Verify that request is only communicating over protobufs. diff --git a/internal/public.pb.go b/internal/public.pb.go index 0cf303b0b..8125f04ab 100644 --- a/internal/public.pb.go +++ b/internal/public.pb.go @@ -1751,6 +1751,7 @@ type ImportRequest struct { Timestamps []int64 `protobuf:"varint,6,rep,packed,name=Timestamps,proto3" json:"Timestamps,omitempty"` IndexCreatedAt int64 `protobuf:"varint,9,opt,name=IndexCreatedAt,proto3" json:"IndexCreatedAt,omitempty"` FieldCreatedAt int64 `protobuf:"varint,10,opt,name=FieldCreatedAt,proto3" json:"FieldCreatedAt,omitempty"` + Clear bool `protobuf:"varint,11,opt,name=Clear,proto3" json:"Clear,omitempty"` XXX_NoUnkeyedLiteral struct{} `json:"-"` XXX_unrecognized []byte `json:"-"` XXX_sizecache int32 `json:"-"` @@ -1859,6 +1860,13 @@ func (m *ImportRequest) GetFieldCreatedAt() int64 { return 0 } +func (m *ImportRequest) GetClear() bool { + if m != nil { + return m.Clear + } + return false +} + type ImportValueRequest struct { Index string `protobuf:"bytes,1,opt,name=Index,proto3" json:"Index,omitempty"` Field string `protobuf:"bytes,2,opt,name=Field,proto3" json:"Field,omitempty"` @@ -1870,6 +1878,7 @@ type ImportValueRequest struct { StringValues []string `protobuf:"bytes,9,rep,name=StringValues,proto3" json:"StringValues,omitempty"` IndexCreatedAt int64 `protobuf:"varint,10,opt,name=IndexCreatedAt,proto3" json:"IndexCreatedAt,omitempty"` FieldCreatedAt int64 `protobuf:"varint,11,opt,name=FieldCreatedAt,proto3" json:"FieldCreatedAt,omitempty"` + Clear bool `protobuf:"varint,12,opt,name=Clear,proto3" json:"Clear,omitempty"` XXX_NoUnkeyedLiteral struct{} `json:"-"` XXX_unrecognized []byte `json:"-"` XXX_sizecache int32 `json:"-"` @@ -1978,6 +1987,84 @@ func (m *ImportValueRequest) GetFieldCreatedAt() int64 { return 0 } +func (m *ImportValueRequest) GetClear() bool { + if m != nil { + return m.Clear + } + return false +} + +type AtomicRecord struct { + Index string `protobuf:"bytes,1,opt,name=Index,proto3" json:"Index,omitempty"` + Shard uint64 `protobuf:"varint,2,opt,name=Shard,proto3" json:"Shard,omitempty"` + Ivr []*ImportValueRequest `protobuf:"bytes,3,rep,name=Ivr,proto3" json:"Ivr,omitempty"` + Ir []*ImportRequest `protobuf:"bytes,4,rep,name=Ir,proto3" json:"Ir,omitempty"` + XXX_NoUnkeyedLiteral struct{} `json:"-"` + XXX_unrecognized []byte `json:"-"` + XXX_sizecache int32 `json:"-"` +} + +func (m *AtomicRecord) Reset() { *m = AtomicRecord{} } +func (m *AtomicRecord) String() string { return proto.CompactTextString(m) } +func (*AtomicRecord) ProtoMessage() {} +func (*AtomicRecord) Descriptor() ([]byte, []int) { + return fileDescriptor_413a91106d7bcce8, []int{27} +} +func (m *AtomicRecord) XXX_Unmarshal(b []byte) error { + return m.Unmarshal(b) +} +func (m *AtomicRecord) XXX_Marshal(b []byte, deterministic bool) ([]byte, error) { + if deterministic { + return xxx_messageInfo_AtomicRecord.Marshal(b, m, deterministic) + } else { + b = b[:cap(b)] + n, err := m.MarshalToSizedBuffer(b) + if err != nil { + return nil, err + } + return b[:n], nil + } +} +func (m *AtomicRecord) XXX_Merge(src proto.Message) { + xxx_messageInfo_AtomicRecord.Merge(m, src) +} +func (m *AtomicRecord) XXX_Size() int { + return m.Size() +} +func (m *AtomicRecord) XXX_DiscardUnknown() { + xxx_messageInfo_AtomicRecord.DiscardUnknown(m) +} + +var xxx_messageInfo_AtomicRecord proto.InternalMessageInfo + +func (m *AtomicRecord) GetIndex() string { + if m != nil { + return m.Index + } + return "" +} + +func (m *AtomicRecord) GetShard() uint64 { + if m != nil { + return m.Shard + } + return 0 +} + +func (m *AtomicRecord) GetIvr() []*ImportValueRequest { + if m != nil { + return m.Ivr + } + return nil +} + +func (m *AtomicRecord) GetIr() []*ImportRequest { + if m != nil { + return m.Ir + } + return nil +} + type TranslateKeysRequest struct { Index string `protobuf:"bytes,1,opt,name=Index,proto3" json:"Index,omitempty"` Field string `protobuf:"bytes,2,opt,name=Field,proto3" json:"Field,omitempty"` @@ -1991,7 +2078,7 @@ func (m *TranslateKeysRequest) Reset() { *m = TranslateKeysRequest{} } func (m *TranslateKeysRequest) String() string { return proto.CompactTextString(m) } func (*TranslateKeysRequest) ProtoMessage() {} func (*TranslateKeysRequest) Descriptor() ([]byte, []int) { - return fileDescriptor_413a91106d7bcce8, []int{27} + return fileDescriptor_413a91106d7bcce8, []int{28} } func (m *TranslateKeysRequest) XXX_Unmarshal(b []byte) error { return m.Unmarshal(b) @@ -2052,7 +2139,7 @@ func (m *TranslateKeysResponse) Reset() { *m = TranslateKeysResponse{} } func (m *TranslateKeysResponse) String() string { return proto.CompactTextString(m) } func (*TranslateKeysResponse) ProtoMessage() {} func (*TranslateKeysResponse) Descriptor() ([]byte, []int) { - return fileDescriptor_413a91106d7bcce8, []int{28} + return fileDescriptor_413a91106d7bcce8, []int{29} } func (m *TranslateKeysResponse) XXX_Unmarshal(b []byte) error { return m.Unmarshal(b) @@ -2101,7 +2188,7 @@ func (m *TranslateIDsRequest) Reset() { *m = TranslateIDsRequest{} } func (m *TranslateIDsRequest) String() string { return proto.CompactTextString(m) } func (*TranslateIDsRequest) ProtoMessage() {} func (*TranslateIDsRequest) Descriptor() ([]byte, []int) { - return fileDescriptor_413a91106d7bcce8, []int{29} + return fileDescriptor_413a91106d7bcce8, []int{30} } func (m *TranslateIDsRequest) XXX_Unmarshal(b []byte) error { return m.Unmarshal(b) @@ -2162,7 +2249,7 @@ func (m *TranslateIDsResponse) Reset() { *m = TranslateIDsResponse{} } func (m *TranslateIDsResponse) String() string { return proto.CompactTextString(m) } func (*TranslateIDsResponse) ProtoMessage() {} func (*TranslateIDsResponse) Descriptor() ([]byte, []int) { - return fileDescriptor_413a91106d7bcce8, []int{30} + return fileDescriptor_413a91106d7bcce8, []int{31} } func (m *TranslateIDsResponse) XXX_Unmarshal(b []byte) error { return m.Unmarshal(b) @@ -2210,7 +2297,7 @@ func (m *ImportRoaringRequestView) Reset() { *m = ImportRoaringRequestVi func (m *ImportRoaringRequestView) String() string { return proto.CompactTextString(m) } func (*ImportRoaringRequestView) ProtoMessage() {} func (*ImportRoaringRequestView) Descriptor() ([]byte, []int) { - return fileDescriptor_413a91106d7bcce8, []int{31} + return fileDescriptor_413a91106d7bcce8, []int{32} } func (m *ImportRoaringRequestView) XXX_Unmarshal(b []byte) error { return m.Unmarshal(b) @@ -2269,7 +2356,7 @@ func (m *ImportRoaringRequest) Reset() { *m = ImportRoaringRequest{} } func (m *ImportRoaringRequest) String() string { return proto.CompactTextString(m) } func (*ImportRoaringRequest) ProtoMessage() {} func (*ImportRoaringRequest) Descriptor() ([]byte, []int) { - return fileDescriptor_413a91106d7bcce8, []int{32} + return fileDescriptor_413a91106d7bcce8, []int{33} } func (m *ImportRoaringRequest) XXX_Unmarshal(b []byte) error { return m.Unmarshal(b) @@ -2356,7 +2443,7 @@ func (m *ImportColumnAttrsRequest) Reset() { *m = ImportColumnAttrsReque func (m *ImportColumnAttrsRequest) String() string { return proto.CompactTextString(m) } func (*ImportColumnAttrsRequest) ProtoMessage() {} func (*ImportColumnAttrsRequest) Descriptor() ([]byte, []int) { - return fileDescriptor_413a91106d7bcce8, []int{33} + return fileDescriptor_413a91106d7bcce8, []int{34} } func (m *ImportColumnAttrsRequest) XXX_Unmarshal(b []byte) error { return m.Unmarshal(b) @@ -2455,6 +2542,7 @@ func init() { proto.RegisterType((*QueryResult)(nil), "internal.QueryResult") proto.RegisterType((*ImportRequest)(nil), "internal.ImportRequest") proto.RegisterType((*ImportValueRequest)(nil), "internal.ImportValueRequest") + proto.RegisterType((*AtomicRecord)(nil), "internal.AtomicRecord") proto.RegisterType((*TranslateKeysRequest)(nil), "internal.TranslateKeysRequest") proto.RegisterType((*TranslateKeysResponse)(nil), "internal.TranslateKeysResponse") proto.RegisterType((*TranslateIDsRequest)(nil), "internal.TranslateIDsRequest") @@ -2467,105 +2555,109 @@ func init() { func init() { proto.RegisterFile("public.proto", fileDescriptor_413a91106d7bcce8) } var fileDescriptor_413a91106d7bcce8 = []byte{ - // 1555 bytes of a gzipped FileDescriptorProto - 0x1f, 0x8b, 0x08, 0x00, 0x00, 0x00, 0x00, 0x00, 0x02, 0xff, 0xac, 0x58, 0x4f, 0x73, 0xdb, 0x44, - 0x14, 0x8f, 0x2c, 0x39, 0xb6, 0x9f, 0x93, 0x34, 0xdd, 0xa6, 0x45, 0x53, 0xda, 0xd4, 0xa3, 0x09, - 0x60, 0x38, 0xa4, 0x93, 0xd2, 0x76, 0x7a, 0x01, 0xda, 0xd4, 0x29, 0xd1, 0x94, 0x84, 0xb2, 0xce, - 0x84, 0x1b, 0x33, 0x8a, 0xbd, 0xa4, 0x1a, 0x64, 0xcb, 0xc8, 0x32, 0x4e, 0x2e, 0xcc, 0xf0, 0x19, - 0xb8, 0xf0, 0x11, 0xb8, 0xf2, 0x15, 0x38, 0x71, 0xe4, 0xce, 0x85, 0x29, 0x7c, 0x8b, 0x5e, 0x98, - 0xf7, 0x56, 0xab, 0x5d, 0x29, 0x4a, 0x9a, 0xe9, 0x70, 0xdb, 0xf7, 0x67, 0xdf, 0xbe, 0xf7, 0x7b, - 0x6f, 0xdf, 0x3e, 0x09, 0x96, 0x26, 0xb3, 0xa3, 0x28, 0x1c, 0x6c, 0x4e, 0x92, 0x38, 0x8d, 0x59, - 0x33, 0x1c, 0xa7, 0x22, 0x19, 0x07, 0x91, 0x37, 0x05, 0x9b, 0xc7, 0x73, 0xe6, 0x42, 0xe3, 0x69, - 0x1c, 0xcd, 0x46, 0xe3, 0xa9, 0x6b, 0x75, 0xec, 0xae, 0xc3, 0x15, 0xc9, 0x18, 0x38, 0xcf, 0xc5, - 0xe9, 0xd4, 0xb5, 0x3b, 0x76, 0xb7, 0xc5, 0x69, 0xcd, 0x36, 0xa0, 0xfe, 0x24, 0x4d, 0x93, 0xa9, - 0x5b, 0xeb, 0xd8, 0xdd, 0xf6, 0xbd, 0x95, 0x4d, 0x65, 0x6e, 0x13, 0xd9, 0x5c, 0x0a, 0xd1, 0x26, - 0x8f, 0x83, 0x24, 0x1c, 0x1f, 0xbb, 0x4e, 0xc7, 0xea, 0x2e, 0x71, 0x45, 0x7a, 0x7b, 0xd0, 0xea, - 0x87, 0xc7, 0x63, 0x31, 0xc4, 0xa3, 0xef, 0x80, 0xfd, 0x22, 0xc6, 0x63, 0xad, 0x6e, 0xfb, 0xde, - 0xb2, 0x36, 0xc5, 0xe3, 0x39, 0x47, 0x09, 0x2a, 0xec, 0x8b, 0x63, 0xb7, 0x56, 0xa9, 0xb0, 0x2f, - 0x8e, 0xbd, 0x47, 0xb0, 0xc2, 0xe3, 0xb9, 0x3f, 0x14, 0xe3, 0x34, 0xfc, 0x36, 0x14, 0x09, 0x39, - 0xcd, 0xe3, 0xb9, 0x8a, 0x85, 0xd6, 0x79, 0x20, 0x35, 0x1d, 0x88, 0x77, 0x13, 0x16, 0xfd, 0xde, - 0x17, 0xe1, 0x34, 0x65, 0xab, 0x60, 0xfb, 0x3d, 0xb5, 0x01, 0x97, 0x9e, 0x0f, 0x57, 0x77, 0x4e, - 0xd2, 0x24, 0x18, 0xa4, 0x62, 0xe8, 0xf7, 0x24, 0x1c, 0x6c, 0x05, 0x6a, 0x7e, 0x8f, 0x7c, 0x75, - 0x78, 0xcd, 0xef, 0xb1, 0x0d, 0x70, 0x0e, 0x83, 0x48, 0x01, 0xb1, 0xaa, 0x9d, 0x93, 0x66, 0x39, - 0x49, 0xbd, 0xa3, 0x82, 0xa9, 0xbd, 0x20, 0x4d, 0xc2, 0x13, 0x76, 0x03, 0x16, 0x9f, 0x85, 0x22, - 0x1a, 0xca, 0x43, 0x5b, 0x3c, 0xa3, 0xd8, 0x03, 0x9d, 0x0a, 0x69, 0xf5, 0x5d, 0x6d, 0xf5, 0x8c, - 0x43, 0x79, 0x9e, 0xbc, 0xdb, 0xd0, 0x78, 0x2e, 0x4e, 0x29, 0x16, 0x15, 0xa9, 0x65, 0x44, 0xfa, - 0x97, 0x05, 0xd7, 0xf2, 0xdd, 0x07, 0xc1, 0x51, 0x24, 0x0e, 0x83, 0x68, 0x26, 0xd8, 0x86, 0x8a, - 0xdb, 0xaa, 0xf2, 0x7f, 0x77, 0x81, 0xb0, 0x60, 0x1f, 0xe4, 0xd8, 0xa1, 0xda, 0x55, 0xad, 0x96, - 0x1d, 0xb9, 0xbb, 0x90, 0x55, 0xc6, 0x2d, 0x68, 0x6e, 0xf7, 0x7d, 0x32, 0xed, 0xda, 0x1d, 0xab, - 0x6b, 0xef, 0x2e, 0xf0, 0x9c, 0xc3, 0x6e, 0x42, 0x63, 0x6f, 0x96, 0x8a, 0x13, 0xbf, 0x47, 0x15, - 0xe1, 0xec, 0x2e, 0x70, 0xc5, 0xc0, 0x9d, 0xb4, 0x7c, 0x2e, 0x4e, 0xdd, 0x7a, 0xc7, 0xea, 0xb6, - 0x70, 0xa7, 0xe2, 0xb0, 0x35, 0x70, 0xb6, 0xe3, 0x38, 0x72, 0x17, 0x3b, 0x56, 0xb7, 0x89, 0xa7, - 0x21, 0xb5, 0xdd, 0x80, 0x3a, 0x19, 0xf6, 0x7e, 0x84, 0xb5, 0x62, 0x70, 0x59, 0xba, 0x18, 0xd8, - 0x68, 0xcf, 0xca, 0xec, 0x21, 0xc1, 0x56, 0x29, 0x85, 0xb5, 0xec, 0x7c, 0x4c, 0xe2, 0x03, 0x58, - 0x24, 0x33, 0xb2, 0xc8, 0xdb, 0xf7, 0x6e, 0x57, 0x00, 0xae, 0x21, 0xe3, 0x99, 0xf2, 0x76, 0x8b, - 0x10, 0xff, 0x32, 0xf1, 0x7b, 0xde, 0x27, 0x65, 0x70, 0x29, 0x97, 0x98, 0x88, 0xfd, 0x60, 0x24, - 0xe4, 0xf9, 0x9c, 0xd6, 0xc8, 0x3b, 0x38, 0x9d, 0x08, 0x72, 0xa0, 0xc5, 0x69, 0xed, 0xfd, 0x64, - 0xc1, 0x4a, 0x71, 0x3f, 0xfa, 0x64, 0x54, 0xc7, 0x05, 0x3e, 0x91, 0x56, 0x5e, 0x3c, 0x8f, 0xca, - 0xc5, 0xb3, 0x7e, 0xde, 0xbe, 0x72, 0xfd, 0x7c, 0x0a, 0xce, 0x8b, 0x20, 0x4c, 0xce, 0x54, 0xf8, - 0xaa, 0x84, 0xd0, 0x26, 0x77, 0x6d, 0x99, 0x8b, 0xfa, 0xd3, 0x78, 0x36, 0x4e, 0x25, 0x86, 0x5c, - 0x12, 0xde, 0x0e, 0xb4, 0x70, 0xbf, 0x0c, 0xdc, 0x93, 0xc6, 0xb2, 0xb2, 0x32, 0xfa, 0x03, 0x72, - 0xb9, 0x3c, 0x68, 0x0d, 0xea, 0xa4, 0x9c, 0x21, 0x21, 0x09, 0x6f, 0x17, 0x00, 0xa5, 0x53, 0x69, - 0x67, 0x03, 0xea, 0x44, 0x65, 0x20, 0x94, 0x0d, 0x49, 0xe1, 0x39, 0x96, 0x6e, 0x43, 0xdd, 0x1f, - 0xa7, 0x0f, 0xef, 0xa3, 0x58, 0x16, 0x24, 0x7a, 0x63, 0xf3, 0xac, 0x64, 0x66, 0xd0, 0x94, 0xd0, - 0xc5, 0x73, 0x6d, 0xc0, 0x32, 0x0c, 0x20, 0x17, 0xdb, 0x4a, 0x4f, 0xc5, 0x49, 0x04, 0x5e, 0x5b, - 0x1e, 0xcf, 0x35, 0x24, 0x19, 0xc5, 0xde, 0x53, 0xa7, 0x38, 0x14, 0xf3, 0x15, 0xe3, 0x2a, 0xa1, - 0x17, 0xea, 0xd8, 0x6f, 0x00, 0x3e, 0x4f, 0xe2, 0xd9, 0x84, 0x40, 0x63, 0x5d, 0xa8, 0x13, 0x95, - 0xc5, 0xc7, 0xf4, 0x26, 0xe5, 0x1b, 0x97, 0x0a, 0xd5, 0xa0, 0x63, 0x72, 0xfa, 0xb3, 0x91, 0xbc, - 0x69, 0x1c, 0x97, 0x58, 0x4a, 0xcd, 0xc3, 0x20, 0xca, 0xc5, 0x87, 0x41, 0x94, 0xc5, 0x8d, 0xcb, - 0xa2, 0x19, 0x5b, 0x99, 0xb9, 0x09, 0xcd, 0x67, 0x51, 0x1c, 0xa4, 0xa8, 0x8c, 0xb6, 0x2c, 0x9e, - 0xd3, 0x6c, 0x0b, 0xa0, 0x27, 0x06, 0xe1, 0x28, 0x88, 0x50, 0xea, 0x94, 0x1b, 0x40, 0x26, 0xe3, - 0x86, 0x92, 0xf7, 0x00, 0x1a, 0x19, 0x55, 0x8d, 0x3d, 0x72, 0xfb, 0x83, 0x20, 0x12, 0xca, 0x0b, - 0x22, 0xbc, 0xaf, 0x61, 0x59, 0x16, 0x23, 0x3e, 0x1f, 0x7d, 0x91, 0x5e, 0xa2, 0x14, 0x2f, 0xf5, - 0x10, 0x79, 0xbf, 0x5a, 0xe0, 0xe0, 0x4a, 0x19, 0xb0, 0xb4, 0x01, 0xf3, 0x36, 0x3a, 0xf2, 0x36, - 0xb2, 0x0e, 0xb4, 0xfb, 0x29, 0xbe, 0x53, 0xba, 0x8d, 0xb5, 0xb8, 0xc9, 0x42, 0xbc, 0xfc, 0x71, - 0xaa, 0xd3, 0x6d, 0xf3, 0x9c, 0x66, 0xb7, 0xa0, 0x85, 0xbd, 0x49, 0x0a, 0xb1, 0x91, 0x35, 0xb9, - 0x66, 0xb0, 0x75, 0x00, 0x85, 0xec, 0x4c, 0x50, 0x37, 0xb3, 0xb8, 0xc1, 0xf1, 0xee, 0x42, 0x03, - 0x3d, 0xdd, 0x0b, 0x26, 0x3a, 0x36, 0xeb, 0xa2, 0xd8, 0x5e, 0x5b, 0xb0, 0xf4, 0xd5, 0x4c, 0x24, - 0xa7, 0x5c, 0x7c, 0x3f, 0x13, 0xd3, 0x14, 0xb1, 0x25, 0x5a, 0xd5, 0x32, 0x11, 0x58, 0xb5, 0xfd, - 0x97, 0x41, 0x32, 0x94, 0x48, 0x39, 0x3c, 0xa3, 0x30, 0x56, 0x8d, 0xf9, 0x94, 0x62, 0x6d, 0x72, - 0x93, 0x45, 0xf5, 0x2e, 0x46, 0x71, 0xaa, 0x82, 0xc9, 0x28, 0xd6, 0x85, 0x2b, 0x3b, 0x27, 0x83, - 0x68, 0x36, 0x14, 0x3c, 0x9e, 0xcb, 0xdd, 0xd4, 0x9c, 0x79, 0x99, 0xcd, 0xde, 0xc7, 0xe6, 0x46, - 0x2c, 0xd5, 0x9a, 0x1a, 0xa4, 0x58, 0xe2, 0xb2, 0x2d, 0x58, 0xda, 0x19, 0x1d, 0x89, 0xe1, 0x50, - 0x0c, 0x7b, 0x41, 0x1a, 0xb8, 0x4d, 0x8a, 0xbb, 0xf4, 0xe0, 0x17, 0x54, 0xbc, 0x9f, 0x2d, 0x58, - 0xce, 0xa2, 0x9f, 0x4e, 0xe2, 0xf1, 0x54, 0x60, 0x8a, 0x77, 0x92, 0x44, 0xa5, 0x78, 0x27, 0x49, - 0xd8, 0x5d, 0x68, 0x70, 0x31, 0x9d, 0x45, 0xa9, 0xaa, 0x92, 0xeb, 0xda, 0xa2, 0xda, 0x3b, 0x8b, - 0x52, 0xae, 0xb4, 0xd8, 0x67, 0xb0, 0x52, 0xa8, 0x43, 0xf5, 0x2c, 0xbc, 0xa3, 0xf7, 0x15, 0xe4, - 0xbc, 0xa4, 0xee, 0xbd, 0x76, 0xa0, 0x6d, 0x58, 0xce, 0x8b, 0x0c, 0xf1, 0x59, 0xce, 0x8a, 0xec, - 0x0e, 0xcd, 0x5d, 0xe7, 0x4c, 0x3d, 0xd8, 0x93, 0x96, 0xc0, 0xda, 0xcf, 0xca, 0xd2, 0xda, 0xd7, - 0x8d, 0xd0, 0xbe, 0xa8, 0x11, 0xe2, 0x14, 0xf7, 0x32, 0x18, 0x1f, 0x8b, 0x21, 0x95, 0x65, 0x93, - 0x2b, 0x92, 0x6d, 0xea, 0xae, 0x40, 0x79, 0x2c, 0xf4, 0x1a, 0x25, 0xe1, 0xba, 0x73, 0xc8, 0x2e, - 0x87, 0x93, 0x41, 0x43, 0xd6, 0x8b, 0xa4, 0xd8, 0x43, 0x68, 0xeb, 0xf6, 0x35, 0xcd, 0x52, 0xb4, - 0xa6, 0x4d, 0x69, 0x21, 0x37, 0x15, 0xd9, 0xe3, 0xf2, 0x88, 0xe6, 0xb6, 0xc8, 0x0b, 0xb7, 0x10, - 0xb9, 0x21, 0xe7, 0xe5, 0x91, 0x6e, 0xcb, 0x98, 0x19, 0x5d, 0xa0, 0xcd, 0xd7, 0xf4, 0xe6, 0x5c, - 0xc4, 0x8d, 0xc9, 0xf2, 0xbe, 0xf9, 0x96, 0xb8, 0x6d, 0xda, 0xb3, 0x56, 0x44, 0x4e, 0xca, 0xb8, - 0xf9, 0xe6, 0x6c, 0x19, 0x0f, 0x99, 0xbb, 0x54, 0x3e, 0x28, 0x17, 0x71, 0xe3, 0xb9, 0xf3, 0x2b, - 0xe6, 0x3b, 0x77, 0x99, 0xb6, 0x56, 0x0f, 0x6f, 0x52, 0x85, 0x57, 0x4c, 0x85, 0x8f, 0xcb, 0x93, - 0x80, 0xbb, 0x52, 0x06, 0xaa, 0x28, 0xe7, 0x25, 0x7d, 0xef, 0xb7, 0x1a, 0x2c, 0xfb, 0xa3, 0x49, - 0x9c, 0xa4, 0x46, 0x4b, 0xf0, 0xc7, 0x43, 0x71, 0xa2, 0x5a, 0x02, 0x11, 0xd5, 0xaf, 0x26, 0xb5, - 0x66, 0x6c, 0x0d, 0xd4, 0x0a, 0x1c, 0x2e, 0x09, 0xa3, 0x1c, 0x9c, 0x42, 0x39, 0xdc, 0x82, 0x96, - 0xac, 0x7d, 0x14, 0xd5, 0x49, 0xa4, 0x19, 0xf2, 0x03, 0x60, 0x4e, 0x83, 0x63, 0x83, 0x46, 0x51, - 0x45, 0x62, 0x1b, 0x94, 0x6a, 0x24, 0x6c, 0x92, 0xd0, 0xe0, 0xa0, 0xfc, 0x20, 0x1c, 0x89, 0x69, - 0x1a, 0x8c, 0x26, 0xd8, 0x57, 0xec, 0xae, 0xcd, 0x0d, 0x0e, 0xb6, 0x14, 0x0a, 0xe2, 0x69, 0x22, - 0x82, 0x54, 0x0c, 0x9f, 0xa4, 0x54, 0x4e, 0x36, 0x2f, 0x71, 0x51, 0x8f, 0xc2, 0xd2, 0x7a, 0x20, - 0xf5, 0x8a, 0x5c, 0xef, 0xf7, 0x1a, 0x30, 0x89, 0x99, 0x1c, 0xf1, 0xfe, 0x37, 0xe0, 0x2e, 0x06, - 0xa8, 0x08, 0x43, 0xe3, 0x0c, 0x0c, 0x37, 0xf2, 0xc1, 0x54, 0x42, 0x90, 0x51, 0xd8, 0xb5, 0xf5, - 0x9b, 0x21, 0xf1, 0xb3, 0xb8, 0xc9, 0x62, 0x1e, 0x2c, 0x19, 0x0f, 0x16, 0xde, 0x36, 0xb4, 0x5d, - 0xe0, 0x55, 0x80, 0x08, 0x97, 0x04, 0xb1, 0x5d, 0x09, 0xe2, 0x21, 0xac, 0x1d, 0x24, 0xc1, 0x78, - 0x1a, 0x05, 0xa9, 0x40, 0xf7, 0xdf, 0x06, 0xc5, 0x8a, 0xaf, 0x4d, 0xef, 0x43, 0xb8, 0x5e, 0xb2, - 0xab, 0x7b, 0x3d, 0xc2, 0x6a, 0xeb, 0x6f, 0xb6, 0x3e, 0x5c, 0xcb, 0x55, 0xfd, 0xde, 0x5b, 0x79, - 0x70, 0xd6, 0xe8, 0x47, 0x46, 0x5c, 0x64, 0x34, 0x3b, 0xbe, 0xca, 0xd7, 0x6d, 0x70, 0xb3, 0xbb, - 0x27, 0x3f, 0x75, 0x33, 0x0f, 0x0e, 0x43, 0x31, 0x3f, 0xef, 0x6b, 0x80, 0xde, 0xba, 0x1a, 0x7d, - 0x20, 0xd3, 0xda, 0xfb, 0xd7, 0x82, 0xb5, 0x2a, 0x23, 0x34, 0xbc, 0x45, 0x22, 0x90, 0xaf, 0x5b, - 0x93, 0x4b, 0x82, 0x3d, 0x82, 0xfa, 0x0f, 0xa1, 0x98, 0xab, 0xd7, 0xcd, 0x33, 0x06, 0xcf, 0x73, - 0x3c, 0xe1, 0x72, 0x03, 0x96, 0xd7, 0x93, 0x41, 0x1a, 0xc6, 0x63, 0x35, 0xca, 0x4a, 0x0a, 0xcf, - 0xd9, 0x8e, 0xe2, 0xc1, 0x77, 0xf2, 0x23, 0x8d, 0x4b, 0xa2, 0xa2, 0x5c, 0xea, 0x97, 0x2c, 0x97, - 0xc5, 0xea, 0x3b, 0x67, 0x29, 0xac, 0x8c, 0x71, 0xe3, 0x8d, 0x19, 0x93, 0x77, 0x4c, 0xcd, 0x8d, - 0x74, 0xc7, 0x5c, 0x39, 0x33, 0xe9, 0xd1, 0x50, 0x91, 0x38, 0xa7, 0xe1, 0x92, 0xbe, 0xd0, 0x1d, - 0xca, 0x52, 0x4e, 0xbf, 0xe1, 0x66, 0x9e, 0x0d, 0x76, 0xb1, 0x2a, 0xd8, 0xed, 0xd5, 0x3f, 0x5e, - 0xad, 0x5b, 0x7f, 0xbe, 0x5a, 0xb7, 0xfe, 0x7e, 0xb5, 0x6e, 0xfd, 0xf2, 0xcf, 0xfa, 0xc2, 0xd1, - 0x22, 0xfd, 0x61, 0xf9, 0xf8, 0xbf, 0x00, 0x00, 0x00, 0xff, 0xff, 0xac, 0xd8, 0x28, 0xa3, 0x71, - 0x11, 0x00, 0x00, + // 1620 bytes of a gzipped FileDescriptorProto + 0x1f, 0x8b, 0x08, 0x00, 0x00, 0x00, 0x00, 0x00, 0x02, 0xff, 0xac, 0x58, 0x4f, 0x6f, 0xdb, 0x46, + 0x16, 0x37, 0x45, 0xca, 0x92, 0x9e, 0x64, 0xc7, 0x99, 0x38, 0x59, 0x22, 0xeb, 0x38, 0x02, 0xe1, + 0xdd, 0x68, 0xf7, 0xe0, 0xc0, 0xd9, 0x24, 0xc8, 0x65, 0x77, 0x63, 0x47, 0xce, 0x9a, 0xc8, 0xda, + 0x9b, 0x1d, 0x19, 0xde, 0xdb, 0x02, 0xb4, 0x34, 0x75, 0x88, 0x52, 0xa2, 0x4a, 0x51, 0x91, 0x7d, + 0x29, 0xd0, 0xcf, 0x90, 0x4b, 0x3f, 0x42, 0x3f, 0x47, 0x2f, 0xed, 0xb1, 0xc7, 0x02, 0xbd, 0x14, + 0x69, 0xbf, 0x45, 0x2e, 0xc5, 0x7b, 0xc3, 0xd1, 0x0c, 0x29, 0xda, 0x31, 0x82, 0xde, 0xe6, 0xfd, + 0x99, 0x37, 0xf3, 0x7e, 0xef, 0xc7, 0x37, 0x4f, 0x82, 0xd6, 0x78, 0x7a, 0x1a, 0x85, 0xfd, 0xed, + 0x71, 0x12, 0xa7, 0x31, 0xab, 0x87, 0xa3, 0x54, 0x24, 0xa3, 0x20, 0xf2, 0x26, 0x60, 0xf3, 0x78, + 0xc6, 0x5c, 0xa8, 0xbd, 0x88, 0xa3, 0xe9, 0x70, 0x34, 0x71, 0xad, 0xb6, 0xdd, 0x71, 0xb8, 0x12, + 0x19, 0x03, 0xe7, 0x95, 0xb8, 0x98, 0xb8, 0x76, 0xdb, 0xee, 0x34, 0x38, 0xad, 0xd9, 0x16, 0x54, + 0x77, 0xd3, 0x34, 0x99, 0xb8, 0x95, 0xb6, 0xdd, 0x69, 0x3e, 0x5a, 0xdd, 0x56, 0xe1, 0xb6, 0x51, + 0xcd, 0xa5, 0x11, 0x63, 0xf2, 0x38, 0x48, 0xc2, 0xd1, 0x99, 0xeb, 0xb4, 0xad, 0x4e, 0x8b, 0x2b, + 0xd1, 0x3b, 0x84, 0x46, 0x2f, 0x3c, 0x1b, 0x89, 0x01, 0x1e, 0x7d, 0x1f, 0xec, 0xd7, 0x31, 0x1e, + 0x6b, 0x75, 0x9a, 0x8f, 0x56, 0x74, 0x28, 0x1e, 0xcf, 0x38, 0x5a, 0xd0, 0xe1, 0x48, 0x9c, 0xb9, + 0x95, 0x52, 0x87, 0x23, 0x71, 0xe6, 0x3d, 0x83, 0x55, 0x1e, 0xcf, 0xfc, 0x81, 0x18, 0xa5, 0xe1, + 0x67, 0xa1, 0x48, 0xe8, 0xd2, 0x3c, 0x9e, 0xa9, 0x5c, 0x68, 0x3d, 0x4f, 0xa4, 0xa2, 0x13, 0xf1, + 0xee, 0xc2, 0xb2, 0xdf, 0xfd, 0x77, 0x38, 0x49, 0xd9, 0x1a, 0xd8, 0x7e, 0x57, 0x6d, 0xc0, 0xa5, + 0xe7, 0xc3, 0xcd, 0xfd, 0xf3, 0x34, 0x09, 0xfa, 0xa9, 0x18, 0xf8, 0x5d, 0x09, 0x07, 0x5b, 0x85, + 0x8a, 0xdf, 0xa5, 0xbb, 0x3a, 0xbc, 0xe2, 0x77, 0xd9, 0x16, 0x38, 0x27, 0x41, 0xa4, 0x80, 0x58, + 0xd3, 0x97, 0x93, 0x61, 0x39, 0x59, 0xbd, 0xd3, 0x5c, 0xa8, 0xc3, 0x20, 0x4d, 0xc2, 0x73, 0x76, + 0x07, 0x96, 0x5f, 0x86, 0x22, 0x1a, 0xc8, 0x43, 0x1b, 0x3c, 0x93, 0xd8, 0x13, 0x5d, 0x0a, 0x19, + 0xf5, 0x8f, 0x3a, 0xea, 0xc2, 0x85, 0xe6, 0x75, 0xf2, 0xee, 0x41, 0xed, 0x95, 0xb8, 0xa0, 0x5c, + 0x54, 0xa6, 0x96, 0x91, 0xe9, 0x4f, 0x16, 0xdc, 0x9a, 0xef, 0x3e, 0x0e, 0x4e, 0x23, 0x71, 0x12, + 0x44, 0x53, 0xc1, 0xb6, 0x54, 0xde, 0x56, 0xd9, 0xfd, 0x0f, 0x96, 0x08, 0x0b, 0xf6, 0x60, 0x8e, + 0x1d, 0xba, 0xdd, 0xd4, 0x6e, 0xd9, 0x91, 0x07, 0x4b, 0x19, 0x33, 0x36, 0xa0, 0xbe, 0xd7, 0xf3, + 0x29, 0xb4, 0x6b, 0xb7, 0xad, 0x8e, 0x7d, 0xb0, 0xc4, 0xe7, 0x1a, 0x76, 0x17, 0x6a, 0x87, 0xd3, + 0x54, 0x9c, 0xfb, 0x5d, 0x62, 0x84, 0x73, 0xb0, 0xc4, 0x95, 0x02, 0x77, 0xd2, 0xf2, 0x95, 0xb8, + 0x70, 0xab, 0x6d, 0xab, 0xd3, 0xc0, 0x9d, 0x4a, 0xc3, 0xd6, 0xc1, 0xd9, 0x8b, 0xe3, 0xc8, 0x5d, + 0x6e, 0x5b, 0x9d, 0x3a, 0x9e, 0x86, 0xd2, 0x5e, 0x0d, 0xaa, 0x14, 0xd8, 0xfb, 0x12, 0xd6, 0xf3, + 0xc9, 0x65, 0xe5, 0x62, 0x60, 0x63, 0x3c, 0x2b, 0x8b, 0x87, 0x02, 0x5b, 0xa3, 0x12, 0x56, 0xb2, + 0xf3, 0xb1, 0x88, 0x4f, 0x60, 0x99, 0xc2, 0x48, 0x92, 0x37, 0x1f, 0xdd, 0x2b, 0x01, 0x5c, 0x43, + 0xc6, 0x33, 0xe7, 0xbd, 0x06, 0x21, 0xfe, 0x9f, 0xc4, 0xef, 0x7a, 0x7f, 0x2f, 0x82, 0x4b, 0xb5, + 0xc4, 0x42, 0x1c, 0x05, 0x43, 0x21, 0xcf, 0xe7, 0xb4, 0x46, 0xdd, 0xf1, 0xc5, 0x58, 0xd0, 0x05, + 0x1a, 0x9c, 0xd6, 0xde, 0x57, 0x16, 0xac, 0xe6, 0xf7, 0xe3, 0x9d, 0x0c, 0x76, 0x5c, 0x71, 0x27, + 0xf2, 0x9a, 0x93, 0xe7, 0x59, 0x91, 0x3c, 0x9b, 0x97, 0xed, 0x2b, 0xf2, 0xe7, 0x1f, 0xe0, 0xbc, + 0x0e, 0xc2, 0x64, 0x81, 0xe1, 0x6b, 0x12, 0x42, 0x9b, 0xae, 0x6b, 0xcb, 0x5a, 0x54, 0x5f, 0xc4, + 0xd3, 0x51, 0x2a, 0x31, 0xe4, 0x52, 0xf0, 0xf6, 0xa1, 0x81, 0xfb, 0x65, 0xe2, 0x9e, 0x0c, 0x96, + 0xd1, 0xca, 0xe8, 0x0f, 0xa8, 0xe5, 0xf2, 0xa0, 0x75, 0xa8, 0x92, 0x73, 0x86, 0x84, 0x14, 0xbc, + 0x03, 0x00, 0xb4, 0x4e, 0x64, 0x9c, 0x2d, 0xa8, 0x92, 0x94, 0x81, 0x50, 0x0c, 0x24, 0x8d, 0x97, + 0x44, 0xba, 0x07, 0x55, 0x7f, 0x94, 0x3e, 0x7d, 0x8c, 0x66, 0x49, 0x48, 0xbc, 0x8d, 0xcd, 0x33, + 0xca, 0x4c, 0xa1, 0x2e, 0xa1, 0x8b, 0x67, 0x3a, 0x80, 0x65, 0x04, 0x40, 0x2d, 0xb6, 0x95, 0xae, + 0xca, 0x93, 0x04, 0xfc, 0x6c, 0x79, 0x3c, 0xd3, 0x90, 0x64, 0x12, 0xfb, 0x93, 0x3a, 0xc5, 0xa1, + 0x9c, 0x6f, 0x18, 0x9f, 0x12, 0xde, 0x42, 0x1d, 0xfb, 0x7f, 0x80, 0x7f, 0x25, 0xf1, 0x74, 0x4c, + 0xa0, 0xb1, 0x0e, 0x54, 0x49, 0xca, 0xf2, 0x63, 0x7a, 0x93, 0xba, 0x1b, 0x97, 0x0e, 0xe5, 0xa0, + 0x63, 0x71, 0x7a, 0xd3, 0xa1, 0xfc, 0xd2, 0x38, 0x2e, 0x91, 0x4a, 0xf5, 0x93, 0x20, 0x9a, 0x9b, + 0x4f, 0x82, 0x28, 0xcb, 0x1b, 0x97, 0xf9, 0x30, 0xb6, 0x0a, 0x73, 0x17, 0xea, 0x2f, 0xa3, 0x38, + 0x48, 0xd1, 0x19, 0x63, 0x59, 0x7c, 0x2e, 0xb3, 0x1d, 0x80, 0xae, 0xe8, 0x87, 0xc3, 0x20, 0x42, + 0xab, 0x53, 0x6c, 0x00, 0x99, 0x8d, 0x1b, 0x4e, 0xde, 0x13, 0xa8, 0x65, 0x52, 0x39, 0xf6, 0xa8, + 0xed, 0xf5, 0x83, 0x48, 0xa8, 0x5b, 0x90, 0xe0, 0xfd, 0x0f, 0x56, 0x24, 0x19, 0xf1, 0xf9, 0xe8, + 0x89, 0xf4, 0x1a, 0x54, 0xbc, 0xd6, 0x43, 0xe4, 0x7d, 0x63, 0x81, 0x83, 0x2b, 0x15, 0xc0, 0xd2, + 0x01, 0xcc, 0xaf, 0xd1, 0x91, 0x5f, 0x23, 0x6b, 0x43, 0xb3, 0x97, 0xe2, 0x3b, 0xa5, 0xdb, 0x58, + 0x83, 0x9b, 0x2a, 0xc4, 0xcb, 0x1f, 0xa5, 0xba, 0xdc, 0x36, 0x9f, 0xcb, 0x6c, 0x03, 0x1a, 0xd8, + 0x9b, 0xa4, 0x11, 0x1b, 0x59, 0x9d, 0x6b, 0x05, 0xdb, 0x04, 0x50, 0xc8, 0x4e, 0x05, 0x75, 0x33, + 0x8b, 0x1b, 0x1a, 0xef, 0x21, 0xd4, 0xf0, 0xa6, 0x87, 0xc1, 0x58, 0xe7, 0x66, 0x5d, 0x95, 0xdb, + 0x07, 0x0b, 0x5a, 0xff, 0x9d, 0x8a, 0xe4, 0x82, 0x8b, 0x2f, 0xa6, 0x62, 0x92, 0x22, 0xb6, 0x24, + 0x2b, 0x2e, 0x93, 0x80, 0xac, 0xed, 0xbd, 0x09, 0x92, 0x81, 0x44, 0xca, 0xe1, 0x99, 0x84, 0xb9, + 0x6a, 0xcc, 0x27, 0x94, 0x6b, 0x9d, 0x9b, 0x2a, 0xe2, 0xbb, 0x18, 0xc6, 0xa9, 0x4a, 0x26, 0x93, + 0x58, 0x07, 0x6e, 0xec, 0x9f, 0xf7, 0xa3, 0xe9, 0x40, 0xf0, 0x78, 0x26, 0x77, 0x53, 0x73, 0xe6, + 0x45, 0x35, 0xfb, 0x33, 0x36, 0x37, 0x52, 0xa9, 0xd6, 0x54, 0x23, 0xc7, 0x82, 0x96, 0xed, 0x40, + 0x6b, 0x7f, 0x78, 0x2a, 0x06, 0x03, 0x31, 0xe8, 0x06, 0x69, 0xe0, 0xd6, 0x29, 0xef, 0xc2, 0x83, + 0x9f, 0x73, 0xf1, 0xde, 0x59, 0xb0, 0x92, 0x65, 0x3f, 0x19, 0xc7, 0xa3, 0x89, 0xc0, 0x12, 0xef, + 0x27, 0x89, 0x2a, 0xf1, 0x7e, 0x92, 0xb0, 0x87, 0x50, 0xe3, 0x62, 0x32, 0x8d, 0x52, 0xc5, 0x92, + 0xdb, 0x3a, 0xa2, 0xda, 0x3b, 0x8d, 0x52, 0xae, 0xbc, 0xd8, 0x3f, 0x61, 0x35, 0xc7, 0x43, 0xf5, + 0x2c, 0xfc, 0x41, 0xef, 0xcb, 0xd9, 0x79, 0xc1, 0xdd, 0xfb, 0xe0, 0x40, 0xd3, 0x88, 0x3c, 0x27, + 0x19, 0xe2, 0xb3, 0x92, 0x91, 0xec, 0x3e, 0xcd, 0x5d, 0x97, 0x4c, 0x3d, 0xd8, 0x93, 0x5a, 0x60, + 0x1d, 0x65, 0xb4, 0xb4, 0x8e, 0x74, 0x23, 0xb4, 0xaf, 0x6a, 0x84, 0x38, 0xc5, 0xbd, 0x09, 0x46, + 0x67, 0x62, 0x40, 0xb4, 0xac, 0x73, 0x25, 0xb2, 0x6d, 0xdd, 0x15, 0xa8, 0x8e, 0xb9, 0x5e, 0xa3, + 0x2c, 0x5c, 0x77, 0x0e, 0xd9, 0xe5, 0x70, 0x32, 0xa8, 0x49, 0xbe, 0x48, 0x89, 0x3d, 0x85, 0xa6, + 0x6e, 0x5f, 0x93, 0xac, 0x44, 0xeb, 0x3a, 0x94, 0x36, 0x72, 0xd3, 0x91, 0x3d, 0x2f, 0x8e, 0x68, + 0x6e, 0x83, 0x6e, 0xe1, 0xe6, 0x32, 0x37, 0xec, 0xbc, 0x38, 0xd2, 0xed, 0x18, 0x33, 0xa3, 0x0b, + 0xb4, 0xf9, 0x96, 0xde, 0x3c, 0x37, 0x71, 0x63, 0xb2, 0x7c, 0x6c, 0xbe, 0x25, 0x6e, 0x93, 0xf6, + 0xac, 0xe7, 0x91, 0x93, 0x36, 0x6e, 0xbe, 0x39, 0x3b, 0xc6, 0x43, 0xe6, 0xb6, 0x8a, 0x07, 0xcd, + 0x4d, 0xdc, 0x78, 0xee, 0xfc, 0x92, 0xf9, 0xce, 0x5d, 0xa1, 0xad, 0xe5, 0xc3, 0x9b, 0x74, 0xe1, + 0x25, 0x53, 0xe1, 0xf3, 0xe2, 0x24, 0xe0, 0xae, 0x16, 0x81, 0xca, 0xdb, 0x79, 0xc1, 0xdf, 0xfb, + 0xae, 0x02, 0x2b, 0xfe, 0x70, 0x1c, 0x27, 0xa9, 0xd1, 0x12, 0xfc, 0xd1, 0x40, 0x9c, 0xab, 0x96, + 0x40, 0x42, 0xf9, 0xab, 0x49, 0xad, 0x19, 0x5b, 0x03, 0xb5, 0x02, 0x87, 0x4b, 0xc1, 0xa0, 0x83, + 0x93, 0xa3, 0xc3, 0x06, 0x34, 0x24, 0xf7, 0xd1, 0x54, 0x25, 0x93, 0x56, 0xc8, 0x1f, 0x00, 0x33, + 0x1a, 0x1c, 0x6b, 0x34, 0x8a, 0x2a, 0x11, 0xdb, 0xa0, 0x74, 0x23, 0x63, 0x9d, 0x8c, 0x86, 0x06, + 0xed, 0xc7, 0xe1, 0x50, 0x4c, 0xd2, 0x60, 0x38, 0xc6, 0xbe, 0x62, 0x77, 0x6c, 0x6e, 0x68, 0xb0, + 0xa5, 0x50, 0x12, 0x2f, 0x12, 0x11, 0xa4, 0x62, 0xb0, 0x9b, 0x12, 0x9d, 0x6c, 0x5e, 0xd0, 0xa2, + 0x1f, 0xa5, 0xa5, 0xfd, 0x40, 0xfa, 0xe5, 0xb5, 0xf4, 0x2c, 0x46, 0x22, 0x48, 0x88, 0x24, 0x75, + 0x2e, 0x05, 0xef, 0xc7, 0x0a, 0x30, 0x89, 0xa4, 0x1c, 0xfc, 0x7e, 0x37, 0x38, 0xaf, 0x86, 0x2d, + 0x0f, 0x4e, 0x6d, 0x01, 0x9c, 0x3b, 0xf3, 0x71, 0x55, 0x02, 0x93, 0x49, 0xd8, 0xcb, 0xf5, 0x4b, + 0x22, 0x51, 0xb5, 0xb8, 0xa9, 0x62, 0x1e, 0xb4, 0x8c, 0x67, 0x0c, 0xbf, 0x41, 0x8c, 0x9d, 0xd3, + 0x95, 0x40, 0x0b, 0xd7, 0x84, 0xb6, 0x79, 0x35, 0xb4, 0x2d, 0x13, 0xda, 0x77, 0x16, 0xb4, 0x76, + 0xd3, 0x78, 0x18, 0xf6, 0xb9, 0xe8, 0xc7, 0xc9, 0xe0, 0x72, 0x50, 0x25, 0x7c, 0x15, 0x13, 0xbe, + 0x6d, 0xb0, 0xfd, 0xb7, 0x49, 0xd6, 0x0a, 0x37, 0x8c, 0x41, 0x6b, 0xa1, 0x56, 0x1c, 0x1d, 0xd9, + 0x03, 0xa8, 0xf8, 0x09, 0x31, 0x37, 0xd7, 0xc4, 0x73, 0x1f, 0x09, 0xaf, 0xf8, 0x89, 0x77, 0x02, + 0xeb, 0xc7, 0x49, 0x30, 0x9a, 0x44, 0x41, 0x2a, 0x10, 0xea, 0x4f, 0xa9, 0x78, 0xc9, 0xef, 0x65, + 0xef, 0x2f, 0x70, 0xbb, 0x10, 0x57, 0xbf, 0x56, 0x48, 0x01, 0x5b, 0xff, 0xea, 0xec, 0xc1, 0xad, + 0xb9, 0xab, 0xdf, 0xfd, 0xa4, 0x1b, 0x2c, 0x06, 0xfd, 0xab, 0x91, 0x17, 0x05, 0xcd, 0x8e, 0x2f, + 0xbb, 0xeb, 0x1e, 0xb8, 0x19, 0x30, 0xf2, 0xc7, 0x7a, 0x76, 0x83, 0x93, 0x50, 0xcc, 0x2e, 0xfb, + 0x3d, 0x43, 0xaf, 0x75, 0x85, 0x7e, 0xe2, 0xd3, 0xda, 0xfb, 0xd5, 0x82, 0xf5, 0xb2, 0x20, 0x9a, + 0x0c, 0x96, 0x41, 0x06, 0xf6, 0x0c, 0xaa, 0x6f, 0x43, 0x31, 0x53, 0xef, 0xb3, 0xb7, 0x50, 0xa2, + 0x85, 0x9b, 0x70, 0xb9, 0x01, 0x3f, 0x85, 0xdd, 0x7e, 0x1a, 0xc6, 0x23, 0x35, 0x8c, 0x4b, 0x09, + 0xcf, 0xd9, 0x8b, 0xe2, 0xfe, 0xe7, 0xf2, 0x67, 0x26, 0x97, 0x42, 0x09, 0xb5, 0xab, 0xd7, 0xa4, + 0xf6, 0x72, 0x19, 0xb5, 0xbd, 0x6f, 0x2d, 0x85, 0x95, 0x31, 0x30, 0x7d, 0xb4, 0x62, 0x9a, 0xd0, + 0xb6, 0x22, 0xb4, 0x2b, 0xa7, 0x3e, 0x3d, 0xdc, 0x2a, 0x11, 0x27, 0x4d, 0x5c, 0xd2, 0x7f, 0x0c, + 0x0e, 0x55, 0x69, 0x2e, 0x7f, 0xa4, 0x8b, 0x2c, 0x26, 0xbb, 0x5c, 0x96, 0xec, 0xde, 0xda, 0xf7, + 0xef, 0x37, 0xad, 0x1f, 0xde, 0x6f, 0x5a, 0x3f, 0xbf, 0xdf, 0xb4, 0xbe, 0xfe, 0x65, 0x73, 0xe9, + 0x74, 0x99, 0xfe, 0x23, 0xfa, 0xdb, 0x6f, 0x01, 0x00, 0x00, 0xff, 0xff, 0x58, 0x5b, 0x7b, 0xe3, + 0x33, 0x12, 0x00, 0x00, } func (m *Row) Marshal() (dAtA []byte, err error) { @@ -4143,6 +4235,16 @@ func (m *ImportRequest) MarshalToSizedBuffer(dAtA []byte) (int, error) { i -= len(m.XXX_unrecognized) copy(dAtA[i:], m.XXX_unrecognized) } + if m.Clear { + i-- + if m.Clear { + dAtA[i] = 1 + } else { + dAtA[i] = 0 + } + i-- + dAtA[i] = 0x58 + } if m.FieldCreatedAt != 0 { i = encodeVarintPublic(dAtA, i, uint64(m.FieldCreatedAt)) i-- @@ -4272,6 +4374,16 @@ func (m *ImportValueRequest) MarshalToSizedBuffer(dAtA []byte) (int, error) { i -= len(m.XXX_unrecognized) copy(dAtA[i:], m.XXX_unrecognized) } + if m.Clear { + i-- + if m.Clear { + dAtA[i] = 1 + } else { + dAtA[i] = 0 + } + i-- + dAtA[i] = 0x60 + } if m.FieldCreatedAt != 0 { i = encodeVarintPublic(dAtA, i, uint64(m.FieldCreatedAt)) i-- @@ -4369,6 +4481,73 @@ func (m *ImportValueRequest) MarshalToSizedBuffer(dAtA []byte) (int, error) { return len(dAtA) - i, nil } +func (m *AtomicRecord) Marshal() (dAtA []byte, err error) { + size := m.Size() + dAtA = make([]byte, size) + n, err := m.MarshalToSizedBuffer(dAtA[:size]) + if err != nil { + return nil, err + } + return dAtA[:n], nil +} + +func (m *AtomicRecord) MarshalTo(dAtA []byte) (int, error) { + size := m.Size() + return m.MarshalToSizedBuffer(dAtA[:size]) +} + +func (m *AtomicRecord) MarshalToSizedBuffer(dAtA []byte) (int, error) { + i := len(dAtA) + _ = i + var l int + _ = l + if m.XXX_unrecognized != nil { + i -= len(m.XXX_unrecognized) + copy(dAtA[i:], m.XXX_unrecognized) + } + if len(m.Ir) > 0 { + for iNdEx := len(m.Ir) - 1; iNdEx >= 0; iNdEx-- { + { + size, err := m.Ir[iNdEx].MarshalToSizedBuffer(dAtA[:i]) + if err != nil { + return 0, err + } + i -= size + i = encodeVarintPublic(dAtA, i, uint64(size)) + } + i-- + dAtA[i] = 0x22 + } + } + if len(m.Ivr) > 0 { + for iNdEx := len(m.Ivr) - 1; iNdEx >= 0; iNdEx-- { + { + size, err := m.Ivr[iNdEx].MarshalToSizedBuffer(dAtA[:i]) + if err != nil { + return 0, err + } + i -= size + i = encodeVarintPublic(dAtA, i, uint64(size)) + } + i-- + dAtA[i] = 0x1a + } + } + if m.Shard != 0 { + i = encodeVarintPublic(dAtA, i, uint64(m.Shard)) + i-- + dAtA[i] = 0x10 + } + if len(m.Index) > 0 { + i -= len(m.Index) + copy(dAtA[i:], m.Index) + i = encodeVarintPublic(dAtA, i, uint64(len(m.Index))) + i-- + dAtA[i] = 0xa + } + return len(dAtA) - i, nil +} + func (m *TranslateKeysRequest) Marshal() (dAtA []byte, err error) { size := m.Size() dAtA = make([]byte, size) @@ -5529,6 +5708,9 @@ func (m *ImportRequest) Size() (n int) { if m.FieldCreatedAt != 0 { n += 1 + sovPublic(uint64(m.FieldCreatedAt)) } + if m.Clear { + n += 2 + } if m.XXX_unrecognized != nil { n += len(m.XXX_unrecognized) } @@ -5587,6 +5769,40 @@ func (m *ImportValueRequest) Size() (n int) { if m.FieldCreatedAt != 0 { n += 1 + sovPublic(uint64(m.FieldCreatedAt)) } + if m.Clear { + n += 2 + } + if m.XXX_unrecognized != nil { + n += len(m.XXX_unrecognized) + } + return n +} + +func (m *AtomicRecord) Size() (n int) { + if m == nil { + return 0 + } + var l int + _ = l + l = len(m.Index) + if l > 0 { + n += 1 + l + sovPublic(uint64(l)) + } + if m.Shard != 0 { + n += 1 + sovPublic(uint64(m.Shard)) + } + if len(m.Ivr) > 0 { + for _, e := range m.Ivr { + l = e.Size() + n += 1 + l + sovPublic(uint64(l)) + } + } + if len(m.Ir) > 0 { + for _, e := range m.Ir { + l = e.Size() + n += 1 + l + sovPublic(uint64(l)) + } + } if m.XXX_unrecognized != nil { n += len(m.XXX_unrecognized) } @@ -10139,6 +10355,26 @@ func (m *ImportRequest) Unmarshal(dAtA []byte) error { break } } + case 11: + if wireType != 0 { + return fmt.Errorf("proto: wrong wireType = %d for field Clear", wireType) + } + var v int + for shift := uint(0); ; shift += 7 { + if shift >= 64 { + return ErrIntOverflowPublic + } + if iNdEx >= l { + return io.ErrUnexpectedEOF + } + b := dAtA[iNdEx] + iNdEx++ + v |= int(b&0x7F) << shift + if b < 0x80 { + break + } + } + m.Clear = bool(v != 0) default: iNdEx = preIndex skippy, err := skipPublic(dAtA[iNdEx:]) @@ -10584,6 +10820,199 @@ func (m *ImportValueRequest) Unmarshal(dAtA []byte) error { break } } + case 12: + if wireType != 0 { + return fmt.Errorf("proto: wrong wireType = %d for field Clear", wireType) + } + var v int + for shift := uint(0); ; shift += 7 { + if shift >= 64 { + return ErrIntOverflowPublic + } + if iNdEx >= l { + return io.ErrUnexpectedEOF + } + b := dAtA[iNdEx] + iNdEx++ + v |= int(b&0x7F) << shift + if b < 0x80 { + break + } + } + m.Clear = bool(v != 0) + default: + iNdEx = preIndex + skippy, err := skipPublic(dAtA[iNdEx:]) + if err != nil { + return err + } + if skippy < 0 { + return ErrInvalidLengthPublic + } + if (iNdEx + skippy) < 0 { + return ErrInvalidLengthPublic + } + if (iNdEx + skippy) > l { + return io.ErrUnexpectedEOF + } + m.XXX_unrecognized = append(m.XXX_unrecognized, dAtA[iNdEx:iNdEx+skippy]...) + iNdEx += skippy + } + } + + if iNdEx > l { + return io.ErrUnexpectedEOF + } + return nil +} +func (m *AtomicRecord) Unmarshal(dAtA []byte) error { + l := len(dAtA) + iNdEx := 0 + for iNdEx < l { + preIndex := iNdEx + var wire uint64 + for shift := uint(0); ; shift += 7 { + if shift >= 64 { + return ErrIntOverflowPublic + } + if iNdEx >= l { + return io.ErrUnexpectedEOF + } + b := dAtA[iNdEx] + iNdEx++ + wire |= uint64(b&0x7F) << shift + if b < 0x80 { + break + } + } + fieldNum := int32(wire >> 3) + wireType := int(wire & 0x7) + if wireType == 4 { + return fmt.Errorf("proto: AtomicRecord: wiretype end group for non-group") + } + if fieldNum <= 0 { + return fmt.Errorf("proto: AtomicRecord: illegal tag %d (wire type %d)", fieldNum, wire) + } + switch fieldNum { + case 1: + if wireType != 2 { + return fmt.Errorf("proto: wrong wireType = %d for field Index", wireType) + } + var stringLen uint64 + for shift := uint(0); ; shift += 7 { + if shift >= 64 { + return ErrIntOverflowPublic + } + if iNdEx >= l { + return io.ErrUnexpectedEOF + } + b := dAtA[iNdEx] + iNdEx++ + stringLen |= uint64(b&0x7F) << shift + if b < 0x80 { + break + } + } + intStringLen := int(stringLen) + if intStringLen < 0 { + return ErrInvalidLengthPublic + } + postIndex := iNdEx + intStringLen + if postIndex < 0 { + return ErrInvalidLengthPublic + } + if postIndex > l { + return io.ErrUnexpectedEOF + } + m.Index = string(dAtA[iNdEx:postIndex]) + iNdEx = postIndex + case 2: + if wireType != 0 { + return fmt.Errorf("proto: wrong wireType = %d for field Shard", wireType) + } + m.Shard = 0 + for shift := uint(0); ; shift += 7 { + if shift >= 64 { + return ErrIntOverflowPublic + } + if iNdEx >= l { + return io.ErrUnexpectedEOF + } + b := dAtA[iNdEx] + iNdEx++ + m.Shard |= uint64(b&0x7F) << shift + if b < 0x80 { + break + } + } + case 3: + if wireType != 2 { + return fmt.Errorf("proto: wrong wireType = %d for field Ivr", wireType) + } + var msglen int + for shift := uint(0); ; shift += 7 { + if shift >= 64 { + return ErrIntOverflowPublic + } + if iNdEx >= l { + return io.ErrUnexpectedEOF + } + b := dAtA[iNdEx] + iNdEx++ + msglen |= int(b&0x7F) << shift + if b < 0x80 { + break + } + } + if msglen < 0 { + return ErrInvalidLengthPublic + } + postIndex := iNdEx + msglen + if postIndex < 0 { + return ErrInvalidLengthPublic + } + if postIndex > l { + return io.ErrUnexpectedEOF + } + m.Ivr = append(m.Ivr, &ImportValueRequest{}) + if err := m.Ivr[len(m.Ivr)-1].Unmarshal(dAtA[iNdEx:postIndex]); err != nil { + return err + } + iNdEx = postIndex + case 4: + if wireType != 2 { + return fmt.Errorf("proto: wrong wireType = %d for field Ir", wireType) + } + var msglen int + for shift := uint(0); ; shift += 7 { + if shift >= 64 { + return ErrIntOverflowPublic + } + if iNdEx >= l { + return io.ErrUnexpectedEOF + } + b := dAtA[iNdEx] + iNdEx++ + msglen |= int(b&0x7F) << shift + if b < 0x80 { + break + } + } + if msglen < 0 { + return ErrInvalidLengthPublic + } + postIndex := iNdEx + msglen + if postIndex < 0 { + return ErrInvalidLengthPublic + } + if postIndex > l { + return io.ErrUnexpectedEOF + } + m.Ir = append(m.Ir, &ImportRequest{}) + if err := m.Ir[len(m.Ir)-1].Unmarshal(dAtA[iNdEx:postIndex]); err != nil { + return err + } + iNdEx = postIndex default: iNdEx = preIndex skippy, err := skipPublic(dAtA[iNdEx:]) diff --git a/internal/public.proto b/internal/public.proto index 2bace4f67..630392e24 100644 --- a/internal/public.proto +++ b/internal/public.proto @@ -175,6 +175,7 @@ message ImportRequest { repeated int64 Timestamps = 6; int64 IndexCreatedAt = 9; int64 FieldCreatedAt = 10; + bool Clear = 11; } message ImportValueRequest { @@ -188,6 +189,14 @@ message ImportValueRequest { repeated string StringValues = 9; int64 IndexCreatedAt = 10; int64 FieldCreatedAt = 11; + bool Clear = 12; +} + +message AtomicRecord { + string Index = 1; + uint64 Shard = 2; + repeated ImportValueRequest Ivr = 3; + repeated ImportRequest Ir = 4; } message TranslateKeysRequest { diff --git a/lmdb_test.go b/lmdb_test.go index 89ba52f11..97b050a38 100644 --- a/lmdb_test.go +++ b/lmdb_test.go @@ -102,11 +102,8 @@ func mustOpenEmptyLMDBWrapper(path string) (w *LMDBWrapper, cleaner func()) { } return w, func() { - w.Close() // stop any started background GC goroutine. - os.RemoveAll(fn) - if FileExists(fn + "-lock") { - os.RemoveAll(fn) - } + w.Close() + panicOn(w.DeleteDBPath(fn)) } } diff --git a/pg/protocol.go b/pg/protocol.go index 57dc2168f..c55decbd7 100644 --- a/pg/protocol.go +++ b/pg/protocol.go @@ -38,7 +38,7 @@ type Protocol uint32 const ( // ProtocolPostgres30 is version 3.0 of the Postgres wire protocol. - ProtocolPostgres30 Protocol = (3 << 16) | 0 + ProtocolPostgres30 Protocol = (3 << 16) // ProtocolCancel is the protocol used for query cancellation. ProtocolCancel Protocol = (1234 << 16) | 5678 diff --git a/row.go b/row.go index ac0ef56e3..9170926a7 100644 --- a/row.go +++ b/row.go @@ -474,7 +474,7 @@ func (r *Row) MarshalJSON() ([]byte, error) { func (r *Row) Columns() []uint64 { a := make([]uint64, 0, r.Count()) for i := range r.segments { - a = append(a, r.segments[i].Columns()...) // Accessing Tx memory that is now invalid. + a = append(a, r.segments[i].Columns()...) } return a } diff --git a/tx_test.go b/tx_test.go new file mode 100644 index 000000000..300e3444b --- /dev/null +++ b/tx_test.go @@ -0,0 +1,245 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package pilosa_test + +import ( + "context" + "fmt" + "os" + "strings" + "testing" + + "github.com/pilosa/pilosa/v2" + "github.com/pilosa/pilosa/v2/http" + "github.com/pilosa/pilosa/v2/server" + "github.com/pilosa/pilosa/v2/test" +) + +func queryIRABit(m0api *pilosa.API, acctOwnerID uint64, iraField string, iraRowID uint64, index string) (bit bool) { + query := fmt.Sprintf("Row(%v=%v)", iraField, iraRowID) // acctOwnerID) + res, err := m0api.Query(context.Background(), &pilosa.QueryRequest{Index: index, Query: query}) + panicOn(err) + cols := res.Results[0].(*pilosa.Row).Columns() + for i := range cols { + if cols[i] == acctOwnerID { + return true + } + } + return false +} + +func mustQueryAcct(m0api *pilosa.API, acctOwnerID uint64, fieldAcct0, index string) (acctBal int64) { + query := fmt.Sprintf("FieldValue(field=%v, column=%v)", fieldAcct0, acctOwnerID) + res, err := m0api.Query(context.Background(), &pilosa.QueryRequest{Index: index, Query: query}) + panicOn(err) + + if len(res.Results) == 0 { + return 0 + } + valCount := res.Results[0].(pilosa.ValCount) + return valCount.Val +} + +func queryBalances(m0api *pilosa.API, acctOwnerID uint64, fldAcct0, fldAcct1, index string) (acct0bal, acct1bal int64) { + + acct0bal = mustQueryAcct(m0api, acctOwnerID, fldAcct0, index) + acct1bal = mustQueryAcct(m0api, acctOwnerID, fldAcct1, index) + return +} +func skipForRoaring(t *testing.T) { + if strings.Contains(os.Getenv("PILOSA_TXSRC"), "roaring") { + t.Skip("skip if roaring pseudo-txn involved -- won't show transactional rollback") + } +} + +func TestAPI_ImportAIR(t *testing.T) { + skipForRoaring(t) + c := test.MustRunCluster(t, 1, + []server.CommandOption{ + server.OptCommandServerOptions( + pilosa.OptServerNodeID("node0"), + pilosa.OptServerClusterHasher(&offsetModHasher{}), + pilosa.OptServerOpenTranslateReader(http.GetOpenTranslateReaderFunc(nil)), + )}, + ) + defer c.Close() + + m0 := c[0] + m0api := m0.API + + ctx := context.Background() + index := "i" + + fieldAcct0 := "acct0" + fieldAcct1 := "acct1" + + transferUSD := int64(100) + _ = transferUSD + opts := pilosa.OptFieldTypeInt(-1000, 1000) + + _, err := m0api.CreateIndex(ctx, index, pilosa.IndexOptions{}) + if err != nil { + t.Fatalf("creating index: %v", err) + } + _, err = m0api.CreateField(ctx, index, fieldAcct0, opts) + if err != nil { + t.Fatalf("creating fieldAcct0: %v", err) + } + _, err = m0api.CreateField(ctx, index, fieldAcct1, opts) + if err != nil { + t.Fatalf("creating fieldAcct1: %v", err) + } + + iraField := "ira" // set field. + iraRowID := uint64(3) + _, err = m0api.CreateField(ctx, index, iraField) + if err != nil { + t.Fatalf("creating fieldIRA: %v", err) + } + + acctOwnerID := uint64(78) // ColumnID + shard := acctOwnerID / ShardWidth + + // setup 500 USD in acct1 and 700 USD in acct2. + // transfer 100 USD. + // should see 400 USD in acct, and 800 USD in acct2. + // + + // setup initial balances + + createAIRUpdate := func(acct0bal, acct1bal int64) (air *pilosa.AtomicRecord) { + ivr0 := &pilosa.ImportValueRequest{ + Index: index, + Field: fieldAcct0, + Shard: shard, + ColumnIDs: []uint64{acctOwnerID}, + Values: []int64{acct0bal}, + } + ivr1 := &pilosa.ImportValueRequest{ + Index: index, + Field: fieldAcct1, + Shard: shard, + ColumnIDs: []uint64{acctOwnerID}, + Values: []int64{acct1bal}, + } + + ir0 := &pilosa.ImportRequest{ + Index: index, + Field: iraField, + Shard: shard, + ColumnIDs: []uint64{acctOwnerID}, + RowIDs: []uint64{iraRowID}, + } + + air = &pilosa.AtomicRecord{ + Index: index, + Shard: shard, + Ivr: []*pilosa.ImportValueRequest{ + ivr0, ivr1, + }, + Ir: []*pilosa.ImportRequest{ir0}, + } + return + } + + expectedBalStartingAcct0 := int64(500) + expectedBalStartingAcct1 := int64(700) + + air := createAIRUpdate(expectedBalStartingAcct0, expectedBalStartingAcct1) + + if err := m0api.ImportAtomicRecord(ctx, air); err != nil { + t.Fatal(err) + } + + iraBit := queryIRABit(m0api, acctOwnerID, iraField, iraRowID, index) + if !iraBit { + panic("IRA bit should have been set") + } + + startingBalanceAcct0, startingBalanceAcct1 := queryBalances(m0api, acctOwnerID, fieldAcct0, fieldAcct1, index) + //vv("starting balance: acct0=%v, acct1=%v", startingBalanceAcct0, startingBalanceAcct1) + + if startingBalanceAcct0 != expectedBalStartingAcct0 { + panic(fmt.Sprintf("expected %v, observed %v starting acct0 balance", expectedBalStartingAcct0, startingBalanceAcct0)) + } + if startingBalanceAcct1 != expectedBalStartingAcct1 { + panic(fmt.Sprintf("expected %v, observed %v starting acct1 balance", expectedBalStartingAcct1, startingBalanceAcct1)) + } + + //vv("sad path: transferUSD %v from %v -> %v, with power loss half-way through", transferUSD, fieldAcct0, fieldAcct1) + + opt := func(o *pilosa.ImportOptions) error { + o.SimPowerLossAfter = 1 + return nil + } + expectedBalEndingAcct0 := expectedBalStartingAcct0 - 100 + expectedBalEndingAcct1 := expectedBalStartingAcct1 + 100 + + air = createAIRUpdate(expectedBalEndingAcct0, expectedBalEndingAcct1) + + err = m0api.ImportAtomicRecord(ctx, air, opt) + if err != pilosa.ErrAborted { + panic(fmt.Sprintf("expected ErrTxnAborted but got err='%#v'", err)) + } + + b0, b1 := queryBalances(m0api, acctOwnerID, fieldAcct0, fieldAcct1, index) + //vv("after power failure tx, balance: acct0=%v, acct1=%v", b0, b1) + + if b0 != expectedBalStartingAcct0 { + panic(fmt.Sprintf("expected %v, observed %v starting acct0 balance", expectedBalStartingAcct0, b0)) + } + if b1 != expectedBalStartingAcct1 { + panic(fmt.Sprintf("expected %v, observed %v starting acct1 balance", expectedBalStartingAcct1, b1)) + } + //vv("good: with power loss half-way, no change in account balances; acct0=%v; acct1=%v", b0, b1) + + // next part of the test, just make sure we do the update. + //vv("happy path: transferUSD %v from %v -> %v, with no interruption.", transferUSD, fieldAcct0, fieldAcct1) + + // happy path with no power failure half-way through. + + err = m0api.ImportAtomicRecord(ctx, air) + panicOn(err) + + eb0, eb1 := queryBalances(m0api, acctOwnerID, fieldAcct0, fieldAcct1, index) + + // should have been applied this time. + if eb0 != expectedBalEndingAcct0 || + eb1 != expectedBalEndingAcct1 { + panic(fmt.Sprintf("problem: transaction did not get committed/applied. transferUSD=%v, but we see: startingBalanceAcct0=%v -> endingBalanceAcct0=%v; startingBalanceAcct1=%v -> endingBalanceAcct1=%v", transferUSD, startingBalanceAcct0, eb0, startingBalanceAcct1, eb1)) + } + //vv("ending balance: acct0=%v, acct1=%v", eb0, eb1) + + // clear all the bits + air.Ivr[0].Clear = true + air.Ivr[1].Clear = true + air.Ir[0].Clear = true + + err = m0api.ImportAtomicRecord(ctx, air) + panicOn(err) + + eb0, eb1 = queryBalances(m0api, acctOwnerID, fieldAcct0, fieldAcct1, index) + if eb0 != 0 || + eb1 != 0 { + panic("problem: bits did not clear") + } + //vv("cleared balances: acct0=%v, acct1=%v", eb0, eb1) + + iraBit = queryIRABit(m0api, acctOwnerID, iraField, iraRowID, index) + if iraBit { + panic("IRA bit should have been cleared") + } + +} From 17fa1e578be7447d12a6d5a3e2ff646eacbd0faa Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Kuba=20Podg=C3=B3rski?= Date: Fri, 21 Aug 2020 14:22:14 +0200 Subject: [PATCH 11/17] Add benchmark for translation reader --- http/translator.go | 28 ++++++++++++-- http/translator_test.go | 82 +++++++++++++++++++++++++++++++++++++++++ server/handler_test.go | 5 ++- server/server.go | 3 +- 4 files changed, 111 insertions(+), 7 deletions(-) diff --git a/http/translator.go b/http/translator.go index f95067f70..0722dcbf9 100644 --- a/http/translator.go +++ b/http/translator.go @@ -22,19 +22,26 @@ import ( "io" "io/ioutil" "net/http" + "sync" "github.com/pilosa/pilosa/v2" "github.com/pilosa/pilosa/v2/logger" ) func GetOpenTranslateReaderFunc(client *http.Client) pilosa.OpenTranslateReaderFunc { + return GetOpenTranslateReaderWithLockerFunc(client, nopLocker{}) +} + +func GetOpenTranslateReaderWithLockerFunc(client *http.Client, locker sync.Locker) pilosa.OpenTranslateReaderFunc { return func(ctx context.Context, nodeURL string, offsets pilosa.TranslateOffsetMap) (pilosa.TranslateEntryReader, error) { - return openTranslateReader(ctx, nodeURL, offsets, client) + return openTranslateReader(ctx, nodeURL, offsets, client, locker) } } -func openTranslateReader(ctx context.Context, nodeURL string, offsets pilosa.TranslateOffsetMap, client *http.Client) (pilosa.TranslateEntryReader, error) { +func openTranslateReader(ctx context.Context, nodeURL string, offsets pilosa.TranslateOffsetMap, client *http.Client, locker sync.Locker) (pilosa.TranslateEntryReader, error) { r := NewTranslateEntryReader(ctx, client) + r.locker = locker + r.URL = nodeURL + "/internal/translate/data" r.Offsets = offsets if err := r.Open(); err != nil { @@ -43,9 +50,16 @@ func openTranslateReader(ctx context.Context, nodeURL string, offsets pilosa.Tra return r, nil } +type nopLocker struct{} + +func (nopLocker) Lock() {} +func (nopLocker) Unlock() {} + // TranslateEntryReader represents an implementation of pilosa.TranslateEntryReader. // It consolidates all index & field translate entries into a single reader. type TranslateEntryReader struct { + locker sync.Locker + ctx context.Context cancel func() @@ -70,7 +84,7 @@ func NewTranslateEntryReader(ctx context.Context, client *http.Client) *Translat if client == nil { client = http.DefaultClient } - r := &TranslateEntryReader{HTTPClient: client, Logger: logger.NopLogger} + r := &TranslateEntryReader{locker: nopLocker{}, HTTPClient: client, Logger: logger.NopLogger} r.ctx, r.cancel = context.WithCancel(ctx) return r } @@ -116,7 +130,10 @@ func (r *TranslateEntryReader) Close() error { r.cancel() } if r.body != nil { - return r.body.Close() + r.locker.Lock() + err := r.body.Close() + r.locker.Unlock() + return err } return nil } @@ -124,5 +141,8 @@ func (r *TranslateEntryReader) Close() error { // ReadEntry reads the next entry from the stream into entry. // Returns io.EOF at the end of the stream. func (r *TranslateEntryReader) ReadEntry(entry *pilosa.TranslateEntry) error { + r.locker.Lock() + defer r.locker.Unlock() + return r.dec.Decode(&entry) } diff --git a/http/translator_test.go b/http/translator_test.go index be6db423a..7038c1773 100644 --- a/http/translator_test.go +++ b/http/translator_test.go @@ -17,6 +17,7 @@ package http_test import ( "context" "fmt" + "sync" "testing" "time" @@ -151,3 +152,84 @@ func TestTranslateStore_EntryReader(t *testing.T) { }) */ } + +func benchmarkSetup(b *testing.B, ctx context.Context, key string, nkeys int) (string, pilosa.TranslateOffsetMap, func()) { + b.Helper() + + cluster := test.MustRunCluster(b, 1) + primary := cluster[0] + + idx := primary.MustCreateIndex(b, "i", pilosa.IndexOptions{}) + fld := primary.MustCreateField(b, idx.Name(), "f", pilosa.OptFieldKeys()) + offset := make(pilosa.TranslateOffsetMap) + offset.SetIndexPartitionOffset(idx.Name(), 0, 1) + offset.SetFieldOffset(idx.Name(), fld.Name(), 1) + + // Set data on the primary node. + for k := 0; k < nkeys; k++ { + if _, err := primary.API.Query(ctx, &pilosa.QueryRequest{ + Index: idx.Name(), + Query: fmt.Sprintf(`Set(%d, %s="%s%[1]d")`, k, fld.Name(), key), + }); err != nil { + b.Fatalf("quering api: %+v", err) + } + } + + return primary.URL(), offset, func() { + b.Helper() + + if err := primary.API.DeleteIndex(ctx, idx.Name()); err != nil { + panic(err) + } + if err := cluster.Close(); err != nil { + panic(err) + } + } +} + +func benchmarkReadEntry(b *testing.B, r pilosa.TranslateEntryReader, key string, nkeys int) { + var entry pilosa.TranslateEntry + for k := 0; k < nkeys; k++ { + if err := r.ReadEntry(&entry); err != nil { + b.Fatalf("reading entry: %+v", err) + } + if entry.Key != fmt.Sprintf("%s%d", key, k) { + b.Fatalf("got: %s, expected: %s%d", entry.Key, key, k) + } + } +} + +const ( + key = "foo" + nkeys = 1000 +) + +func BenchmarkReadEntryNoMutex(b *testing.B) { + ctx := context.Background() + url, offset, teardown := benchmarkSetup(b, ctx, key, nkeys) + defer teardown() + + for n := 0; n < b.N; n++ { + r, err := http.GetOpenTranslateReaderFunc(nil)(ctx, url, offset) + if err != nil { + b.Fatalf("openining translate reader: %+v", err) + } + benchmarkReadEntry(b, r, key, nkeys) + r.Close() + } +} + +func BenchmarkReadEntryWithMutex(b *testing.B) { + ctx := context.Background() + url, offset, teardown := benchmarkSetup(b, ctx, key, nkeys) + defer teardown() + + for n := 0; n < b.N; n++ { + r, err := http.GetOpenTranslateReaderWithLockerFunc(nil, &sync.Mutex{})(ctx, url, offset) + if err != nil { + b.Fatalf("openining translate reader: %+v", err) + } + benchmarkReadEntry(b, r, key, nkeys) + r.Close() + } +} diff --git a/server/handler_test.go b/server/handler_test.go index 31954eff8..ccd91bdcc 100644 --- a/server/handler_test.go +++ b/server/handler_test.go @@ -27,6 +27,7 @@ import ( "net/http/httptest" "reflect" "strings" + "sync" "testing" "time" @@ -1300,7 +1301,7 @@ func TestCluster_TranslateStore(t *testing.T) { cluster[0] = test.NewCommandNode(true, server.OptCommandServerOptions( pilosa.OptServerOpenTranslateStore(boltdb.OpenTranslateStore), - pilosa.OptServerOpenTranslateReader(http.GetOpenTranslateReaderFunc(nil)), + pilosa.OptServerOpenTranslateReader(http.GetOpenTranslateReaderWithLockerFunc(nil, &sync.Mutex{})), ), ) cluster[0].Config.Gossip.Port = "0" @@ -1329,7 +1330,7 @@ func TestClusterTranslator(t *testing.T) { cluster[1] = test.NewCommandNode(false, server.OptCommandServerOptions( pilosa.OptServerOpenTranslateStore(boltdb.OpenTranslateStore), - pilosa.OptServerOpenTranslateReader(http.GetOpenTranslateReaderFunc(nil)), + pilosa.OptServerOpenTranslateReader(http.GetOpenTranslateReaderWithLockerFunc(nil, &sync.Mutex{})), ), ) cluster[1].Config.Gossip.Port = "0" diff --git a/server/server.go b/server/server.go index 5578cf35f..e74d460e2 100644 --- a/server/server.go +++ b/server/server.go @@ -31,6 +31,7 @@ import ( "os/signal" "runtime" "strconv" + "sync" "syscall" "time" @@ -389,7 +390,7 @@ func (m *Command) SetupServer() error { pilosa.OptServerDiagnosticsInterval(diagnosticsInterval), pilosa.OptServerExecutorPoolSize(m.Config.WorkerPoolSize), pilosa.OptServerOpenTranslateStore(boltdb.OpenTranslateStore), - pilosa.OptServerOpenTranslateReader(http.GetOpenTranslateReaderFunc(c)), + pilosa.OptServerOpenTranslateReader(http.GetOpenTranslateReaderWithLockerFunc(c, &sync.Mutex{})), pilosa.OptServerLogger(m.logger), pilosa.OptServerAttrStoreFunc(boltdb.NewAttrStore), pilosa.OptServerSystemInfo(gopsutil.NewSystemInfo()), From 49c3bd97015ae77ff3fd6ca53c4e1ebed0c469c6 Mon Sep 17 00:00:00 2001 From: Nia Weiss Date: Fri, 21 Aug 2020 14:47:59 -0400 Subject: [PATCH 12/17] add SQL to postgres endpoint --- server/grpc.go | 20 +------------------- server/pg.go | 38 ++++++++++++++++++++++++++++++++++++-- server/sql.go | 48 ++++++++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 85 insertions(+), 21 deletions(-) create mode 100644 server/sql.go diff --git a/server/grpc.go b/server/grpc.go index 3dcc420a1..af11b6f3c 100644 --- a/server/grpc.go +++ b/server/grpc.go @@ -26,7 +26,6 @@ import ( "github.com/pilosa/pilosa/v2" "github.com/pilosa/pilosa/v2/logger" pb "github.com/pilosa/pilosa/v2/proto" - "github.com/pilosa/pilosa/v2/sql" "github.com/pilosa/pilosa/v2/stats" "github.com/pkg/errors" "google.golang.org/grpc" @@ -120,24 +119,7 @@ func (h *GRPCHandler) DeleteVDS(ctx context.Context, req *pb.DeleteVDSRequest) ( } func (h *GRPCHandler) execSQL(ctx context.Context, queryStr string) (pb.StreamClient, error) { - mapper := sql.NewMapper() - mapper.Logger = h.logger - query, err := mapper.MapSQL(queryStr) - if err != nil { - return nil, errors.Wrap(err, "failed to map SQL") - } - var results pb.StreamClient - switch query.SQLType { - case sql.SQLTypeSelect: - handler := sql.NewSelectHandler(h.api) - results, err = handler.Handle(ctx, query) - if err != nil { - return nil, errors.Wrap(err, "failed to start SQL query") - } - default: - return nil, status.Errorf(codes.Unimplemented, "query type not supported") - } - return results, nil + return execSQL(ctx, h.api, h.logger, queryStr) } // QuerySQL handles the SQL request and sends RowResponses to the stream. diff --git a/server/pg.go b/server/pg.go index 98117c374..28bbef7af 100644 --- a/server/pg.go +++ b/server/pg.go @@ -19,6 +19,7 @@ import ( "crypto/tls" "encoding/json" "fmt" + "io" "net" "strconv" "strings" @@ -50,7 +51,8 @@ func NewPostgresServer(api *pilosa.API, logger logger.Logger, tls *tls.Config) * s: pg.Server{ QueryHandler: &queryDecodeHandler{ child: &pilosaQueryHandler{ - api: api, + api: api, + logger: logger, }, }, TypeEngine: pg.PrimitiveTypeEngine{}, @@ -125,7 +127,8 @@ func pgDecodePQL(str string) (q pg.Query, err error) { } type pilosaQueryHandler struct { - api *pilosa.API + api *pilosa.API + logger logger.Logger } func pgWriteRow(w pg.QueryResultWriter, row *pilosa.Row) error { @@ -342,6 +345,28 @@ func pgWriteRowser(w pg.QueryResultWriter, result pb.ToRowser) error { }) } +type clientRowser struct { + pb.StreamClient +} + +func (cr *clientRowser) ToRows(f func(*pb.RowResponse) error) error { + for { + resp, err := cr.StreamClient.Recv() + if err != nil { + if err == io.EOF { + return nil + } + + return err + } + + err = f(resp) + if err != nil { + return err + } + } +} + func pgWriteResult(w pg.QueryResultWriter, result interface{}) error { switch result := result.(type) { case *pilosa.Row: @@ -354,6 +379,8 @@ func pgWriteResult(w pg.QueryResultWriter, result interface{}) error { return pgWriteGroupCount(w, result) case pb.ToRowser: // we should avoid protobuf where we can... return pgWriteRowser(w, result) + case pb.StreamClient: + return pgWriteRowser(w, &clientRowser{result}) default: return errors.Errorf("result type %T not yet supported", result) } @@ -374,6 +401,13 @@ func (pqh *pilosaQueryHandler) HandleQuery(ctx context.Context, w pg.QueryResult } return errors.Wrap(pgWriteResult(w, resp.Results[0]), "writing query result") + case pg.SimpleQuery: + resp, err := execSQL(ctx, pqh.api, pqh.logger, string(q)) + if err != nil { + return errors.Wrap(err, "executing query") + } + return errors.Wrap(pgWriteResult(w, resp), "writing query result") + default: return errors.Errorf("query type %T not yet supported (query: %s)", q, q) } diff --git a/server/sql.go b/server/sql.go new file mode 100644 index 000000000..1e648b9e3 --- /dev/null +++ b/server/sql.go @@ -0,0 +1,48 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package server + +import ( + "context" + + "github.com/pilosa/pilosa/v2" + "github.com/pilosa/pilosa/v2/logger" + pb "github.com/pilosa/pilosa/v2/proto" + "github.com/pilosa/pilosa/v2/sql" + "github.com/pkg/errors" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" +) + +func execSQL(ctx context.Context, api *pilosa.API, logger logger.Logger, queryStr string) (pb.StreamClient, error) { + mapper := sql.NewMapper() + mapper.Logger = logger + query, err := mapper.MapSQL(queryStr) + if err != nil { + return nil, errors.Wrap(err, "failed to map SQL") + } + var results pb.StreamClient + switch query.SQLType { + case sql.SQLTypeSelect: + handler := sql.NewSelectHandler(api) + results, err = handler.Handle(ctx, query) + if err != nil { + return nil, errors.Wrap(err, "failed to start SQL query") + } + default: + return nil, status.Errorf(codes.Unimplemented, "query type not supported") + } + return results, nil +} From 40a9d01f4627d5e2de0631d552231fb034d15bf6 Mon Sep 17 00:00:00 2001 From: Jason Aten Date: Thu, 20 Aug 2020 19:11:07 -0500 Subject: [PATCH 13/17] ImportRequest.Clear and ImportValuesRequest.Clear respected by api.Import() and api.ImportValues() - tested in TestAPI_ClearFlagForImportAndImportValues api_test.go --- api.go | 24 ++++++++ api_test.go | 124 ++++++++++++++++++++++++++++++++++++++ fragment_internal_test.go | 13 ++++ lmdb.go | 8 +-- tx_test.go | 5 +- 5 files changed, 169 insertions(+), 5 deletions(-) diff --git a/api.go b/api.go index d8cdef042..8d3bfcea5 100644 --- a/api.go +++ b/api.go @@ -1114,7 +1114,28 @@ func (api *API) ImportAtomicRecord(ctx context.Context, req *AtomicRecord, opts return tx.Commit() } +// This is a hide your face ugly hack, forced upon +// us by the horrible invention of function based options +// by the usually brilliant Rob Pike. - JEA +func addClearToImportOptions(opts []ImportOption) []ImportOption { + var opt ImportOptions + for _, o := range opts { + // check for side-effect of setting io.Clear; that is + // how we know it is present. + _ = o(&opt) + if opt.Clear { + // we already have the clear flag set, so nothing more to do. + return opts + } + } + // no clear flag being set, add that option now. + return append(opts, OptImportOptionsClear(true)) +} + func (api *API) Import(ctx context.Context, req *ImportRequest, opts ...ImportOption) error { + if req.Clear { + opts = addClearToImportOptions(opts) + } return api.ImportWithTx(ctx, nil, req, opts...) } @@ -1254,6 +1275,9 @@ func (api *API) ImportWithTx(ctx context.Context, tx Tx, req *ImportRequest, opt } func (api *API) ImportValue(ctx context.Context, req *ImportValueRequest, opts ...ImportOption) error { + if req.Clear { + opts = addClearToImportOptions(opts) + } return api.ImportValueWithTx(ctx, nil, req, opts...) } diff --git a/api_test.go b/api_test.go index f592afb6e..88bb65510 100644 --- a/api_test.go +++ b/api_test.go @@ -474,3 +474,127 @@ type offsetModHasher struct{} func (*offsetModHasher) Hash(key uint64, n int) int { return int(key+1) % n } + +func TestAPI_ClearFlagForImportAndImportValues(t *testing.T) { + c := test.MustRunCluster(t, 1, + []server.CommandOption{ + server.OptCommandServerOptions( + pilosa.OptServerNodeID("node0"), + pilosa.OptServerClusterHasher(&offsetModHasher{}), + pilosa.OptServerOpenTranslateReader(http.GetOpenTranslateReaderFunc(nil)), + )}, + ) + defer c.Close() + + // plan: + // 1. set a bit + // 2. clear with Import() using the ImportRequest.Clear flag + // 3. verifiy the clear is done. + // repeat for ImportValueRequest and ImportValues() + + m0 := c[0] + m0api := m0.API + + ctx := context.Background() + index := "i" + fieldAcct0 := "acct0" + + opts := pilosa.OptFieldTypeInt(-1000, 1000) + + _, err := m0api.CreateIndex(ctx, index, pilosa.IndexOptions{}) + if err != nil { + t.Fatalf("creating index: %v", err) + } + _, err = m0api.CreateField(ctx, index, fieldAcct0, opts) + if err != nil { + t.Fatalf("creating fieldAcct0: %v", err) + } + + iraField := "ira" // set field. + iraRowID := uint64(3) + _, err = m0api.CreateField(ctx, index, iraField) + if err != nil { + t.Fatalf("creating fieldIRA: %v", err) + } + + acctOwnerID := uint64(78) // ColumnID + shard := acctOwnerID / ShardWidth + acct0bal := int64(500) + + ivr0 := &pilosa.ImportValueRequest{ + Index: index, + Field: fieldAcct0, + Shard: shard, + ColumnIDs: []uint64{acctOwnerID}, + Values: []int64{acct0bal}, + } + ir0 := &pilosa.ImportRequest{ + Index: index, + Field: iraField, + Shard: shard, + ColumnIDs: []uint64{acctOwnerID}, + RowIDs: []uint64{iraRowID}, + } + + if err := m0api.Import(ctx, ir0); err != nil { + t.Fatal(err) + } + if err := m0api.ImportValue(ctx, ivr0); err != nil { + t.Fatal(err) + } + + bitIsSet := func() bool { + query := fmt.Sprintf("Row(%v=%v)", iraField, iraRowID) + res, err := m0api.Query(context.Background(), &pilosa.QueryRequest{Index: index, Query: query}) + panicOn(err) + cols := res.Results[0].(*pilosa.Row).Columns() + for i := range cols { + if cols[i] == acctOwnerID { + return true + } + } + return false + } + + if !bitIsSet() { + panic("IRA bit should have been set") + } + + queryAcct := func(m0api *pilosa.API, acctOwnerID uint64, fieldAcct0, index string) (acctBal int64) { + query := fmt.Sprintf("FieldValue(field=%v, column=%v)", fieldAcct0, acctOwnerID) + res, err := m0api.Query(context.Background(), &pilosa.QueryRequest{Index: index, Query: query}) + panicOn(err) + + if len(res.Results) == 0 { + return 0 + } + valCount := res.Results[0].(pilosa.ValCount) + return valCount.Val + } + + bal := queryAcct(m0api, acctOwnerID, fieldAcct0, index) + + if bal != acct0bal { + panic(fmt.Sprintf("expected %v, observed %v starting acct0 balance", acct0bal, bal)) + } + + // clear the bit + ir0.Clear = true + if err := m0api.Import(ctx, ir0); err != nil { + t.Fatal(err) + } + + if bitIsSet() { + panic("IRA bit should have been cleared") + } + + // clear the BSI + ivr0.Clear = true + if err := m0api.ImportValue(ctx, ivr0); err != nil { + t.Fatal(err) + } + bal = queryAcct(m0api, acctOwnerID, fieldAcct0, index) + if bal != 0 { + panic(fmt.Sprintf("expected %v, observed %v starting acct0 balance", acct0bal, 0)) + } +} diff --git a/fragment_internal_test.go b/fragment_internal_test.go index 8ef6d5909..d033a24e6 100644 --- a/fragment_internal_test.go +++ b/fragment_internal_test.go @@ -29,6 +29,7 @@ import ( "runtime" "runtime/debug" "sort" + "strings" "testing" "testing/quick" @@ -190,6 +191,8 @@ func TestFragment_RowcacheMap(t *testing.T) { // Ensure a fragment can clear a row. func TestFragment_ClearRow(t *testing.T) { + notBlueGreenTest(t) + f, idx := mustOpenFragment("i", "f", viewStandard, 0, "") _ = idx defer f.Clean(t) @@ -225,6 +228,7 @@ func TestFragment_ClearRow(t *testing.T) { // Ensure a fragment can set a row. func TestFragment_SetRow(t *testing.T) { + notBlueGreenTest(t) f, idx := mustOpenFragment("i", "f", viewStandard, 7, "") _ = idx defer f.Clean(t) @@ -5644,3 +5648,12 @@ func TestFragment_Bug_Q2DoubleDelete(t *testing.T) { t.Fatalf("expected nothing got %v", res) } } + +func notBlueGreenTest(t *testing.T) { + src := os.Getenv("PILOSA_TXSRC") + if strings.Contains(src, "_") { + if strings.Contains(src, "roaring") { + t.Skip("skip under blue green with roaring") + } + } +} diff --git a/lmdb.go b/lmdb.go index 9e0d30ed9..b67436abf 100644 --- a/lmdb.go +++ b/lmdb.go @@ -147,9 +147,9 @@ func (r *lmdbRegistrar) openLMDBWrapper(path0 string) (*LMDBWrapper, error) { flags = flags | lmdb.WriteMap | // Use a writable memory map. - lmdb.NoMetaSync | // Don't fsync metapage after commit. - lmdb.NoSync | // Don't fsync after commit. - lmdb.MapAsync | // Flush asynchronously when using the WriteMap flag. + //lmdb.NoMetaSync | // Don't fsync metapage after commit. + //lmdb.NoSync | // Don't fsync after commit. + //lmdb.MapAsync | // Flush asynchronously when using the WriteMap flag. lmdb.NoMemInit // Disable LMDB memory initialization err = env.Open(path, flags, 0644) @@ -344,7 +344,7 @@ func (tx *LMDBTx) Type() string { } func (tx *LMDBTx) UseRowCache() bool { - return false + return true } // Pointer gives us a memory address for the underlying transaction for debugging. diff --git a/tx_test.go b/tx_test.go index 300e3444b..ae5e0b984 100644 --- a/tx_test.go +++ b/tx_test.go @@ -59,7 +59,10 @@ func queryBalances(m0api *pilosa.API, acctOwnerID uint64, fldAcct0, fldAcct1, in return } func skipForRoaring(t *testing.T) { - if strings.Contains(os.Getenv("PILOSA_TXSRC"), "roaring") { + src := os.Getenv("PILOSA_TXSRC") + // once txfactory.go DefaultTxsrc != RoaringTxn, this + // will break, of course. Take out the src == "" below. + if src == "" || strings.Contains(src, "roaring") { t.Skip("skip if roaring pseudo-txn involved -- won't show transactional rollback") } } From dd73bf396007035e8689571857ad7560d6a69521 Mon Sep 17 00:00:00 2001 From: Nia Weiss Date: Fri, 21 Aug 2020 17:20:10 -0400 Subject: [PATCH 14/17] stop logging in TestStartupInvalidLength --- pg/server_test.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pg/server_test.go b/pg/server_test.go index 4903fdb60..9eda44bcf 100644 --- a/pg/server_test.go +++ b/pg/server_test.go @@ -59,7 +59,7 @@ func TestStartupInvalidLength(t *testing.T) { res := testing.Benchmark(func(b *testing.B) { connect, shutdown, err := pgtest.ServeMem(&pg.Server{ MaxStartupSize: 1024, - Logger: logger.NewLogfLogger(t), + Logger: logger.NopLogger, }) if err != nil { t.Fatalf("starting in-memory postgres server: %v", err) From cc66775ea49f20445b108cee2ab3504499dd9260 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Kuba=20Podg=C3=B3rski?= Date: Sat, 22 Aug 2020 00:18:19 +0200 Subject: [PATCH 15/17] Update http/translator_test.go Co-authored-by: Travis Turner --- http/translator_test.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/http/translator_test.go b/http/translator_test.go index 7038c1773..1fcfb42d1 100644 --- a/http/translator_test.go +++ b/http/translator_test.go @@ -212,7 +212,7 @@ func BenchmarkReadEntryNoMutex(b *testing.B) { for n := 0; n < b.N; n++ { r, err := http.GetOpenTranslateReaderFunc(nil)(ctx, url, offset) if err != nil { - b.Fatalf("openining translate reader: %+v", err) + b.Fatalf("opening translate reader: %+v", err) } benchmarkReadEntry(b, r, key, nkeys) r.Close() From 6d0baa82b3a71e476a2788a133263afdd3b2b69d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Kuba=20Podg=C3=B3rski?= Date: Sat, 22 Aug 2020 00:18:25 +0200 Subject: [PATCH 16/17] Update http/translator_test.go Co-authored-by: Travis Turner --- http/translator_test.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/http/translator_test.go b/http/translator_test.go index 1fcfb42d1..8d4a9362d 100644 --- a/http/translator_test.go +++ b/http/translator_test.go @@ -227,7 +227,7 @@ func BenchmarkReadEntryWithMutex(b *testing.B) { for n := 0; n < b.N; n++ { r, err := http.GetOpenTranslateReaderWithLockerFunc(nil, &sync.Mutex{})(ctx, url, offset) if err != nil { - b.Fatalf("openining translate reader: %+v", err) + b.Fatalf("opening translate reader: %+v", err) } benchmarkReadEntry(b, r, key, nkeys) r.Close() From ad9e3338b5ac416e702fa1f3f2cdc80ddea3843a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Kuba=20Podg=C3=B3rski?= Date: Sat, 22 Aug 2020 01:08:33 +0200 Subject: [PATCH 17/17] Add support for SHOW queries --- proto/interface.go | 9 +++ server/grpc_test.go | 30 ++++++++- server/sql.go | 6 ++ sql/select.go | 9 ++- sql/show.go | 148 ++++++++++++++++++++++++++++++++++++++++++++ 5 files changed, 198 insertions(+), 4 deletions(-) create mode 100644 sql/show.go diff --git a/proto/interface.go b/proto/interface.go index 660aa6273..6a6df7213 100644 --- a/proto/interface.go +++ b/proto/interface.go @@ -31,6 +31,15 @@ type StreamClient interface { Recv() (*RowResponse, error) } +// EmptyStream implements StreamClient interface. +// It always returns empty RowResponse +type EmptyStream struct{} + +// Recv returns io.EOF +func (EmptyStream) Recv() (*RowResponse, error) { + return nil, io.EOF +} + // ReadIntoTable reads from a StreamClient and stores the result into a table response. func ReadIntoTable(cli StreamClient) (*TableResponse, error) { var headers []*ColumnInfo diff --git a/server/grpc_test.go b/server/grpc_test.go index e143c9168..4529c0df8 100644 --- a/server/grpc_test.go +++ b/server/grpc_test.go @@ -686,7 +686,6 @@ func TestQuerySQLUnary(t *testing.T) { }, eq: equalUnordered, }, - { // GroupBy(Rows(field='age'),limit=3) sql: "select age, count(*) as cnt from grouper group by age order by cnt desc, age desc limit 3", @@ -703,6 +702,35 @@ func TestQuerySQLUnary(t *testing.T) { }, eq: equal, }, + { + sql: "show tables", + exp: tableResponse{ + headers: []columnInfo{ + {"Table", "string"}, + }, + rows: []row{ + {[]columnResponse{"grouper"}}, + {[]columnResponse{"joiner"}}, + }, + }, + eq: equal, + }, + { + sql: "show fields from grouper", + exp: tableResponse{ + headers: []columnInfo{ + {"Field", "string"}, + {"Type", "string"}, + }, + rows: []row{ + {[]columnResponse{"age", "int64"}}, + {[]columnResponse{"color", "[]string"}}, + {[]columnResponse{"height", "int64"}}, + {[]columnResponse{"score", "int64"}}, + }, + }, + eq: equal, + }, } for i, test := range tests { diff --git a/server/sql.go b/server/sql.go index 1e648b9e3..0f1de916c 100644 --- a/server/sql.go +++ b/server/sql.go @@ -41,6 +41,12 @@ func execSQL(ctx context.Context, api *pilosa.API, logger logger.Logger, querySt if err != nil { return nil, errors.Wrap(err, "failed to start SQL query") } + case sql.SQLTypeShow: + handler := sql.NewShowHandler(api) + results, err = handler.Handle(ctx, query) + if err != nil { + return nil, errors.Wrap(err, "failed to start SQL query") + } default: return nil, status.Errorf(codes.Unimplemented, "query type not supported") } diff --git a/sql/select.go b/sql/select.go index 1c11cdfca..a195a7c85 100644 --- a/sql/select.go +++ b/sql/select.go @@ -42,7 +42,11 @@ func NewSelectHandler(api *pilosa.API) *SelectHandler { // Handle executes mapped SQL func (s *SelectHandler) Handle(ctx context.Context, mapped *MappedSQL) (pproto.StreamClient, error) { - mr, err := s.mapSelect(ctx, mapped.Statement.(*sqlparser.Select), mapped.Mask) + stmt, ok := mapped.Statement.(*sqlparser.Select) + if !ok { + return nil, fmt.Errorf("statement is not type select: %T", mapped.Statement) + } + mr, err := s.mapSelect(ctx, stmt, mapped.Mask) if err != nil { return nil, errors.Wrap(err, "mapping select") } @@ -70,12 +74,11 @@ func (s *SelectHandler) mapSelect(ctx context.Context, selectStmt *sqlparser.Sel return mr, nil } -func (s *SelectHandler) execMappingResult(ctx context.Context, mr *MappingResult) (stream pproto.StreamClient, err error) { +func (s *SelectHandler) execMappingResult(ctx context.Context, mr *MappingResult) (pproto.StreamClient, error) { if mr.Query == "" { return nil, errors.New("no pql query created") } - fmt.Println("PQL:", mr.Query) resp, err := s.api.Query(ctx, &pilosa.QueryRequest{Index: mr.IndexName, Query: mr.Query}) if err != nil { return nil, errors.Wrap(err, "doing pql query") diff --git a/sql/show.go b/sql/show.go new file mode 100644 index 000000000..4ac3c390d --- /dev/null +++ b/sql/show.go @@ -0,0 +1,148 @@ +// Copyright 2020 Pilosa Corp. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package sql + +import ( + "context" + "fmt" + + "github.com/pilosa/pilosa/v2" + pproto "github.com/pilosa/pilosa/v2/proto" + "github.com/pkg/errors" + "vitess.io/vitess/go/vt/sqlparser" +) + +// ShowHandler executes SQL show table/field statements +type ShowHandler struct { + api *pilosa.API +} + +// NewShowHandler constructor +func NewShowHandler(api *pilosa.API) *ShowHandler { + return &ShowHandler{ + api: api, + } +} + +// Handle executes mapped SQL +func (s *ShowHandler) Handle(ctx context.Context, mapped *MappedSQL) (pproto.StreamClient, error) { + stmt, ok := mapped.Statement.(*sqlparser.Show) + if !ok { + return nil, fmt.Errorf("statement is not type show: %T", mapped.Statement) + } + + switch stmt.Type { + case "tables": + return s.execShowTables(ctx, stmt) + case "fields": + return s.execShowFields(ctx, stmt) + default: + return nil, fmt.Errorf("cannot show: %s", stmt.Type) + } +} + +func (s *ShowHandler) execShowTables(ctx context.Context, showStmt *sqlparser.Show) (pproto.StreamClient, error) { + indexInfo := s.api.Schema(ctx) + sz := len(indexInfo) + // If there aren't any indexes, don't bother creating + // a result row buffer. + if sz == 0 { + return pproto.EmptyStream{}, nil + } + + // Create a buffer large enough to hold the entire result + // set. This way we don't have to use a goroutine. + result := pproto.NewRowBuffer(sz) + for _, ii := range indexInfo { + rr := &pproto.RowResponse{ + Headers: []*pproto.ColumnInfo{ + {Name: "Table", Datatype: "string"}, + }, + Columns: []*pproto.ColumnResponse{ + {ColumnVal: &pproto.ColumnResponse_StringVal{StringVal: ii.Name}}, + }, + } + if err := result.Send(rr); err != nil { + return nil, errors.Wrap(err, "sending row response") + } + } + if err := result.Send(pproto.EOF); err != nil { + return nil, errors.Wrap(err, "sending EOF") + } + + // Apply Sort Reducer + out := pproto.NewRowBuffer(0) + red := NewOrderByReducer([]string{"Table"}, []string{"asc"}, 0, 0) + go red.Reduce(result, out) //nolint:errcheck + + result = out + return result, nil +} + +func (s *ShowHandler) execShowFields(ctx context.Context, showStmt *sqlparser.Show) (pproto.StreamClient, error) { + indexName := showStmt.OnTable.ToViewName().Name.String() + index, err := s.api.Index(ctx, indexName) + if err != nil { + return nil, errors.Wrap(err, "getting schema") + } + if index == nil { + return nil, pilosa.ErrIndexNotFound + } + fields := index.Fields() + sz := len(fields) + // If there aren't any fields, don't bother creating + // a result row buffer. + if sz == 0 { + return pproto.EmptyStream{}, nil + } + + // Create a buffer large enough to hold the entire result + // set. This way we don't have to use a goroutine. + result := pproto.NewRowBuffer(sz) + for _, f := range fields { + if f.Name() == "_exists" { + continue + } + + dt, err := f.Datatype() + if err != nil { + return nil, errors.Wrapf(err, "field %s", f.Name()) + } + rr := &pproto.RowResponse{ + Headers: []*pproto.ColumnInfo{ + {Name: "Field", Datatype: "string"}, + {Name: "Type", Datatype: "string"}, + }, + Columns: []*pproto.ColumnResponse{ + {ColumnVal: &pproto.ColumnResponse_StringVal{StringVal: f.Name()}}, + {ColumnVal: &pproto.ColumnResponse_StringVal{StringVal: dt}}, + }, + } + if err := result.Send(rr); err != nil { + return nil, errors.Wrap(err, "sending row response") + } + } + if err := result.Send(pproto.EOF); err != nil { + return nil, errors.Wrap(err, "sending EOF") + } + + // Apply Sort Reducer + out := pproto.NewRowBuffer(0) + red := NewOrderByReducer([]string{"Field"}, []string{"asc"}, 0, 0) + go red.Reduce(result, out) //nolint:errcheck + + result = out + return result, nil +}