diff --git a/data/test/vtgate/filter_cases.txt b/data/test/vtgate/filter_cases.txt index 8e021fdb1fb..ac56a1dd440 100644 --- a/data/test/vtgate/filter_cases.txt +++ b/data/test/vtgate/filter_cases.txt @@ -747,6 +747,21 @@ } } +# database() call in where clause. +"select id from user where database()" +{ + "Original": "select id from user where database()", + "Instructions": { + "Opcode": "SelectScatter", + "Keyspace": { + "Name": "user", + "Sharded": true + }, + "Query": "select id from user where 'targetString'", + "FieldQuery": "select id from user where 1 != 1" + } +} + # outer and inner subquery route reference the same "uu.id" name # but they refer to different things. The first reference is to the outermost query, # and the second reference is to the the innermost 'from' subquery. diff --git a/data/test/vtgate/from_cases.txt b/data/test/vtgate/from_cases.txt index 46eb385ef7d..5d61fc3a729 100644 --- a/data/test/vtgate/from_cases.txt +++ b/data/test/vtgate/from_cases.txt @@ -1100,6 +1100,22 @@ } } +# database call in ON clause. +# The on clause is weird because the substitution must even for root expressions. +"select u1.a from unsharded u1 join unsharded u2 on database()" +{ + "Original": "select u1.a from unsharded u1 join unsharded u2 on database()", + "Instructions": { + "Opcode": "SelectUnsharded", + "Keyspace": { + "Name": "main", + "Sharded": false + }, + "Query": "select u1.a from unsharded as u1 join unsharded as u2 on 'targetString'", + "FieldQuery": "select u1.a from unsharded as u1 join unsharded as u2 on 'targetString' where 1 != 1" + } +} + # verify ',' vs JOIN precedence "select u1.a from unsharded u1, unsharded u2 join unsharded u3 on u1.a = u2.a" "symbol u1.a not found" diff --git a/data/test/vtgate/select_cases.txt b/data/test/vtgate/select_cases.txt index fb0f1209f09..d19f8034268 100644 --- a/data/test/vtgate/select_cases.txt +++ b/data/test/vtgate/select_cases.txt @@ -143,6 +143,22 @@ } } +# database calls should be substituted +"select database() from dual" +{ + "Original": "select database() from dual", + "Instructions": { + "Opcode": "SelectUnsharded", + "Keyspace": { + "Name": "main", + "Sharded": false + }, + "Query": "select 'targetString' from dual", + "FieldQuery": "select 'targetString' from dual where 1 != 1" + } +} + + # nextval for simple route "select next value from user" "NEXT used on a sharded table" diff --git a/go/vt/vtgate/planbuilder/builder.go b/go/vt/vtgate/planbuilder/builder.go index 22ee2185397..221151d892f 100644 --- a/go/vt/vtgate/planbuilder/builder.go +++ b/go/vt/vtgate/planbuilder/builder.go @@ -111,6 +111,7 @@ type ContextVSchema interface { FindTable(tablename sqlparser.TableName) (*vindexes.Table, string, topodatapb.TabletType, key.Destination, error) FindTableOrVindex(tablename sqlparser.TableName) (*vindexes.Table, vindexes.Vindex, string, topodatapb.TabletType, key.Destination, error) DefaultKeyspace() (*vindexes.Keyspace, error) + TargetString() string } // Build builds a plan for a query based on the specified vschema. diff --git a/go/vt/vtgate/planbuilder/expr.go b/go/vt/vtgate/planbuilder/expr.go index cb044532cbc..672a9e2675b 100644 --- a/go/vt/vtgate/planbuilder/expr.go +++ b/go/vt/vtgate/planbuilder/expr.go @@ -79,7 +79,7 @@ func skipParenthesis(node sqlparser.Expr) sqlparser.Expr { // // If an expression has no references to the current query, then the left-most // origin is chosen as the default. -func (pb *primitiveBuilder) findOrigin(expr sqlparser.Expr) (origin builder, err error) { +func (pb *primitiveBuilder) findOrigin(expr sqlparser.Expr) (origin builder, pushExpr sqlparser.Expr, err error) { highestOrigin := pb.bldr.First() var subroutes []*route err = sqlparser.Walk(func(node sqlparser.SQLNode) (kontinue bool, err error) { @@ -120,30 +120,33 @@ func (pb *primitiveBuilder) findOrigin(expr sqlparser.Expr) (origin builder, err subroutes = append(subroutes, subroute) return false, nil case *sqlparser.FuncExpr: + switch { // If it's last_insert_id, ensure it's a single unsharded route. - if !node.Name.EqualString("last_insert_id") { - return true, nil - } - if rb, isRoute := pb.bldr.(*route); !isRoute || rb.ERoute.Keyspace.Sharded { - return false, errors.New("unsupported: LAST_INSERT_ID is only allowed for unsharded keyspaces") + case node.Name.EqualString("last_insert_id"): + if rb, isRoute := pb.bldr.(*route); !isRoute || rb.ERoute.Keyspace.Sharded { + return false, errors.New("unsupported: LAST_INSERT_ID is only allowed for unsharded keyspaces") + } + case node.Name.EqualString("database"): + expr = sqlparser.ReplaceExpr(expr, node, sqlparser.NewStrVal([]byte(pb.vschema.TargetString()))) } + return true, nil } return true, nil }, expr) if err != nil { - return nil, err + return nil, nil, err } highestRoute, isRoute := highestOrigin.(*route) if !isRoute && len(subroutes) > 0 { - return nil, errors.New("unsupported: subquery cannot be merged with cross-shard subquery") + return nil, nil, errors.New("unsupported: subquery cannot be merged with cross-shard subquery") } for _, subroute := range subroutes { if err := highestRoute.SubqueryCanMerge(pb, subroute); err != nil { - return nil, err + return nil, nil, err } subroute.Redirect = highestRoute } - return highestOrigin, nil + return highestOrigin, expr, nil } func hasSubquery(node sqlparser.SQLNode) bool { diff --git a/go/vt/vtgate/planbuilder/from.go b/go/vt/vtgate/planbuilder/from.go index e6541b6fe2b..cc0571821fe 100644 --- a/go/vt/vtgate/planbuilder/from.go +++ b/go/vt/vtgate/planbuilder/from.go @@ -315,13 +315,12 @@ func (pb *primitiveBuilder) mergeRoutes(rpb *primitiveBuilder, ajoin *sqlparser. if ajoin == nil { return nil } + _, expr, err := pb.findOrigin(ajoin.Condition.On) + if err != nil { + return err + } + ajoin.Condition.On = expr for _, filter := range splitAndExpression(nil, ajoin.Condition.On) { - // If VTGate evolves, this section should be rewritten - // to use processExpr. - _, err = pb.findOrigin(filter) - if err != nil { - return err - } lRoute.UpdatePlan(pb, filter) } return nil diff --git a/go/vt/vtgate/planbuilder/plan_test.go b/go/vt/vtgate/planbuilder/plan_test.go index 409a40f869b..d6bcd90d667 100644 --- a/go/vt/vtgate/planbuilder/plan_test.go +++ b/go/vt/vtgate/planbuilder/plan_test.go @@ -210,6 +210,10 @@ func (vw *vschemaWrapper) DefaultKeyspace() (*vindexes.Keyspace, error) { return vw.v.Keyspaces["main"].Keyspace, nil } +func (vw *vschemaWrapper) TargetString() string { + return "targetString" +} + // For the purposes of this set of tests, just compare the actual plan // and ignore all the metrics. type testPlan struct { diff --git a/go/vt/vtgate/planbuilder/select.go b/go/vt/vtgate/planbuilder/select.go index cb1e854ba4e..76d5f73ec26 100644 --- a/go/vt/vtgate/planbuilder/select.go +++ b/go/vt/vtgate/planbuilder/select.go @@ -119,11 +119,11 @@ func (pb *primitiveBuilder) pushFilter(boolExpr sqlparser.Expr, whereType string filters := splitAndExpression(nil, boolExpr) reorderBySubquery(filters) for _, filter := range filters { - origin, err := pb.findOrigin(filter) + origin, expr, err := pb.findOrigin(filter) if err != nil { return err } - if err := pb.bldr.PushFilter(pb, filter, whereType, origin); err != nil { + if err := pb.bldr.PushFilter(pb, expr, whereType, origin); err != nil { return err } } @@ -168,10 +168,11 @@ func (pb *primitiveBuilder) pushSelectRoutes(selectExprs sqlparser.SelectExprs) for i, node := range selectExprs { switch node := node.(type) { case *sqlparser.AliasedExpr: - origin, err := pb.findOrigin(node.Expr) + origin, expr, err := pb.findOrigin(node.Expr) if err != nil { return nil, err } + node.Expr = expr resultColumns[i], _, err = pb.bldr.PushSelect(node, origin) if err != nil { return nil, err diff --git a/go/vt/vtgate/vcursor_impl.go b/go/vt/vtgate/vcursor_impl.go index 8c0f02c8c02..60b7328231d 100644 --- a/go/vt/vtgate/vcursor_impl.go +++ b/go/vt/vtgate/vcursor_impl.go @@ -128,6 +128,11 @@ func (vc *vcursorImpl) DefaultKeyspace() (*vindexes.Keyspace, error) { return ks.Keyspace, nil } +// TargetString returns the current TargetString of the session. +func (vc *vcursorImpl) TargetString() string { + return vc.safeSession.TargetString +} + // Execute performs a V3 level execution of the query. func (vc *vcursorImpl) Execute(method string, query string, BindVars map[string]*querypb.BindVariable, isDML bool) (*sqltypes.Result, error) { qr, err := vc.executor.Execute(vc.ctx, method, vc.safeSession, vc.marginComments.Leading+query+vc.marginComments.Trailing, BindVars)