Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
47 changes: 47 additions & 0 deletions pkg/embed/cluster_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -228,6 +228,53 @@ func TestRowCountOverMySQLProtocol(t *testing.T) {
require.NoError(t, err)
require.Equal(t, int64(1), affectedRows)

_, err = conn.ExecContext(ctx, "create procedure insert_rows() 'begin insert into t values (4), (5); end'")
require.NoError(t, err)
result, err = conn.ExecContext(ctx, "call insert_rows()")
require.NoError(t, err)
affectedRows, err = result.RowsAffected()
require.NoError(t, err)
require.Equal(t, int64(2), affectedRows)
require.NoError(t, stmt.QueryRowContext(ctx).Scan(&rowCount))
require.Equal(t, int64(2), rowCount)

_, err = conn.ExecContext(ctx, "create procedure caller_count() 'begin select row_count(); end'")
require.NoError(t, err)
_, err = conn.ExecContext(ctx, "insert into t values (6), (7), (8), (9), (10), (11)")
require.NoError(t, err)
func() {
rows, err := conn.QueryContext(ctx, "call caller_count()")
require.NoError(t, err)
defer rows.Close()
require.True(t, rows.Next())
require.NoError(t, rows.Scan(&rowCount))
require.NoError(t, rows.Err())
require.Equal(t, int64(6), rowCount)
}()

_, err = conn.ExecContext(ctx, "create procedure inner_results() 'begin select 20; select 21; end'")
require.NoError(t, err)
_, err = conn.ExecContext(ctx, "create procedure outer_results() 'begin select 10; call inner_results(); select 30; end'")
require.NoError(t, err)
func() {
rows, err := conn.QueryContext(ctx, "call outer_results()")
require.NoError(t, err)
defer rows.Close()
var got []int64
for {
for rows.Next() {
var value int64
require.NoError(t, rows.Scan(&value))
got = append(got, value)
}
require.NoError(t, rows.Err())
if !rows.NextResultSet() {
break
}
}
require.Equal(t, []int64{10, 20, 21, 30}, got)
}()

_, err = conn.ExecContext(ctx, "insert into t values (1)")
require.Error(t, err)
require.NoError(t, stmt.QueryRowContext(ctx).Scan(&rowCount))
Expand Down
17 changes: 14 additions & 3 deletions pkg/frontend/authenticate.go
Original file line number Diff line number Diff line change
Expand Up @@ -11513,12 +11513,20 @@ func GetVersionCompatibility(ctx context.Context, ses *Session, dbName string) (
return resultConfig, err
}

func doInterpretCall(ctx context.Context, ses FeSession, call *tree.CallStmt, bg bool) ([]ExecResult, error) {
func doInterpretCall(
ctx context.Context,
ses FeSession,
call *tree.CallStmt,
bg bool,
callerAffectedRows int64,
affectedRows *int64,
) ([]ExecResult, error) {
if parsed, ok, err := parseIcebergBuiltinCall(ctx, call); ok || err != nil {
if err != nil {
return nil, err
}
return executeIcebergBuiltinCall(ctx, ses, parsed)
erArray, err := executeIcebergBuiltinCall(ctx, ses, parsed)
return erArray, err
}
// fetch related
var spLang string
Expand Down Expand Up @@ -11625,6 +11633,7 @@ func doInterpretCall(ctx context.Context, ses FeSession, call *tree.CallStmt, bg
interpreter.argsMap = argsMap
interpreter.argsAttr = argsAttr
interpreter.outParamMap = make(map[string]interface{})
interpreter.initialAffectedRows = callerAffectedRows

switch spLang {
case "sql":
Expand All @@ -11634,17 +11643,19 @@ func doInterpretCall(ctx context.Context, ses FeSession, call *tree.CallStmt, bg
}
defer freeStatements(stmt)

err = interpreter.ExecuteSp(stmt[0], dbName)
err = interpreter.ExecuteSp(stmt[0], dbName, bg)
Comment thread
ck89119 marked this conversation as resolved.
if err != nil {
return nil, err
}
*affectedRows = interpreter.lastAffectedRows
return interpreter.GetResult(), nil

case "starlark":
err = interpreter.ExecuteStarlark(spBody, dbName, bg)
if err != nil {
return nil, err
}
*affectedRows = interpreter.lastAffectedRows
return interpreter.GetResult(), nil

default:
Expand Down
89 changes: 71 additions & 18 deletions pkg/frontend/authenticate_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,7 @@ import (
"github.com/matrixorigin/matrixone/pkg/pb/timestamp"
"github.com/matrixorigin/matrixone/pkg/pb/txn"
"github.com/matrixorigin/matrixone/pkg/queryservice"
"github.com/matrixorigin/matrixone/pkg/sql/parsers"
"github.com/matrixorigin/matrixone/pkg/sql/parsers/dialect"
mysqlparser "github.com/matrixorigin/matrixone/pkg/sql/parsers/dialect/mysql"
"github.com/matrixorigin/matrixone/pkg/sql/parsers/tree"
Expand Down Expand Up @@ -8982,7 +8983,19 @@ func Test_doDropUser(t *testing.T) {
}

func Test_doInterpretCall(t *testing.T) {
t.Skip("skip doInterpretCall")
convey.Convey("call procedure without database fails", t, func() {
ctrl := gomock.NewController(t)
defer ctrl.Finish()

call := &tree.CallStmt{
Name: tree.NewProcedureName("test_without_database", tree.ObjectNamePrefix{}),
}
ses := newSes(determinePrivilegeSetOfStatement(call), ctrl)

_, err := doInterpretCall(context.Background(), ses, call, false, 0, new(int64))
convey.So(err, convey.ShouldNotBeNil)
})

convey.Convey("call precedure (not exist)fail", t, func() {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
Expand Down Expand Up @@ -9022,7 +9035,7 @@ func Test_doInterpretCall(t *testing.T) {
mrs := newMrsForPasswordOfUser([][]interface{}{})
bh.sql2result[sql] = mrs

_, err = doInterpretCall(ctx, ses, call, false)
_, err = doInterpretCall(ctx, ses, call, false, 0, new(int64))
convey.So(err, convey.ShouldNotBeNil)
})

Expand Down Expand Up @@ -9060,8 +9073,8 @@ func Test_doInterpretCall(t *testing.T) {

sql, err := getSqlForSpBody(ses.GetTxnHandler().GetConnCtx(), string(call.Name.Name.ObjectName), ses.GetDatabaseName())
convey.So(err, convey.ShouldBeNil)
mrs := newMrsForPasswordOfUser([][]interface{}{
{"begin set sid = 1000; end", "{}"},
mrs := newMrsForStoredProcedure([][]interface{}{
{"unsupported", "begin set sid = 1000; end", "[]", ""},
})
bh.sql2result[sql] = mrs

Expand All @@ -9075,7 +9088,7 @@ func Test_doInterpretCall(t *testing.T) {
})
bh.sql2result[sql] = mrs

_, err = doInterpretCall(ctx, ses, call, false)
_, err = doInterpretCall(ctx, ses, call, false, 0, new(int64))
convey.So(err, convey.ShouldNotBeNil)
})

Expand Down Expand Up @@ -9114,8 +9127,8 @@ func Test_doInterpretCall(t *testing.T) {

sql, err := getSqlForSpBody(ses.GetTxnHandler().GetConnCtx(), string(call.Name.Name.ObjectName), ses.GetDatabaseName())
convey.So(err, convey.ShouldBeNil)
mrs := newMrsForPasswordOfUser([][]interface{}{
{"begin DECLARE v1 INT; SET v1 = 10; IF v1 > 5 THEN select * from tbh1; ELSEIF v1 = 5 THEN select * from tbh2; ELSEIF v1 = 4 THEN select * from tbh2 limit 1; ELSE select * from tbh3; END IF; end", "{}"},
mrs := newMrsForStoredProcedure([][]interface{}{
{"sql", "begin DECLARE v1 INT; SET v1 = 10; end", "[]", ""},
})
bh.sql2result[sql] = mrs

Expand All @@ -9129,21 +9142,47 @@ func Test_doInterpretCall(t *testing.T) {
})
bh.sql2result[sql] = mrs

sql = "select v1 > 5"
mrs = newMrsForPasswordOfUser([][]interface{}{
{"1"},
})
bh.sql2result[sql] = mrs

sql = "select * from tbh1"
mrs = newMrsForPasswordOfUser([][]interface{}{})
bh.sql2result[sql] = mrs

_, err = doInterpretCall(ctx, ses, call, false)
var affectedRows int64
_, err = doInterpretCall(ctx, ses, call, false, 7, &affectedRows)
convey.So(err, convey.ShouldBeNil)
convey.So(affectedRows, convey.ShouldEqual, int64(7))
})
}

func TestProceduralOnlyStatementsPreserveAffectedRows(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
ses := newTestSession(t, ctrl)
ses.GetTxnCompileCtx().execCtx = &ExecCtx{
reqCtx: context.Background(),
proc: testutil.NewProcess(t),
ses: ses,
}
stmt, err := parsers.ParseOne(
context.Background(),
dialect.MYSQL,
"begin declare x int default 1; set x = 2; end",
1,
)
require.NoError(t, err)
varScope := []map[string]interface{}{}
back := &evalCondBackgroundExec{}
interpreter := &Interpreter{
ctx: context.Background(),
ses: ses,
bh: back,
varScope: &varScope,
fmtctx: tree.NewFmtCtx(dialect.MYSQL),
argsMap: map[string]tree.Expr{},
argsAttr: map[string]tree.InOutArgType{},
outParamMap: map[string]interface{}{},
initialAffectedRows: 7,
}

require.NoError(t, interpreter.ExecuteSp(stmt, "db", false))
require.Equal(t, int64(7), interpreter.lastAffectedRows)
}

func TestParseStoredProcedureBodyUsesCreationSQLMode(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()
Expand Down Expand Up @@ -11522,6 +11561,20 @@ func newMrsForPasswordOfUser(rows [][]interface{}) *MysqlResultSet {
return mrs
}

func newMrsForStoredProcedure(rows [][]interface{}) *MysqlResultSet {
mrs := &MysqlResultSet{}
for _, name := range []string{"language", "body", "args", "sql_mode"} {
col := &MysqlColumn{}
col.SetName(name)
col.SetColumnType(defines.MYSQL_TYPE_VARCHAR)
mrs.AddColumn(col)
}
for _, row := range rows {
mrs.AddRow(row)
}
return mrs
}

func newMrsForFeatureRegistry(rows [][]interface{}) *MysqlResultSet {
mrs := &MysqlResultSet{}

Expand Down
56 changes: 55 additions & 1 deletion pkg/frontend/back_exec.go
Original file line number Diff line number Diff line change
Expand Up @@ -54,6 +54,11 @@ type backExec struct {
statsArray *statistic.StatsArray
}

type backgroundExecRowCount interface {
GetLastAffectedRows() int64
SetLastAffectedRows(int64)
}

func (back *backExec) init(ses FeSession, txnOp TxnOperator, db string, callBack outputCallBackFunc) {
back.backSes = newBackSession(ses, txnOp, db, callBack)
if back.statsArray != nil {
Expand All @@ -67,6 +72,14 @@ func (back *backExec) Service() string {
return back.backSes.GetService()
}

func (back *backExec) GetLastAffectedRows() int64 {
return back.backSes.lastAffectedRows
}

func (back *backExec) SetLastAffectedRows(rows int64) {
back.backSes.lastAffectedRows = rows
}

func (back *backExec) Close() {
if back == nil {
return
Expand Down Expand Up @@ -108,6 +121,11 @@ func (back *backExec) ExecWithSQLMode(ctx context.Context, sql string, sqlMode s
func (back *backExec) exec(ctx context.Context, sql string, sqlMode string, useSQLMode bool) (retErr error) {
back.backSes.EnterFPrint(FPBackExecExec)
defer back.backSes.ExitFPrint(FPBackExecExec)
defer func() {
if retErr != nil {
back.SetLastAffectedRows(-1)
}
}()
if ctx == nil {
return moerr.NewInternalError(context.Background(), "context is nil")
}
Expand Down Expand Up @@ -378,6 +396,7 @@ func doComQueryInBack(
StorageEngine: pu.StorageEngine,
Buf: backSes.buf,
}
proc.SetAffectedRows(backSes.lastAffectedRows)
proc.SetStmtProfile(&backSes.stmtProfile)
proc.SetResolveVariableFunc(backSes.txnCompileCtx.ResolveVariable)
// Frontend back-exec — session-bound resolver. backSession is a
Expand Down Expand Up @@ -509,11 +528,44 @@ func doComQueryInBack(
if err != nil {
return err
}
if _, ok := stmt.(*tree.CallStmt); ok {
if err = appendNestedCallResults(execCtx.reqCtx, backSes, execCtx.results); err != nil {
return err
}
}
backSes.lastAffectedRows = affectedRowsForStatement(execCtx)
Comment thread
ck89119 marked this conversation as resolved.
} // end of for

return nil
}

func appendNestedCallResults(ctx context.Context, backSes *backSession, results []ExecResult) error {
for _, result := range results {
mrs, ok := result.(*MysqlResultSet)
if !ok {
return moerr.NewInternalError(ctx, "nested CALL returned an unsupported result type")
}
backSes.allResultSet = append(backSes.allResultSet, mrs)
}
return nil
}

func affectedRowsForStatement(execCtx *ExecCtx) int64 {
switch execCtx.stmt.StmtKind().OutputType() {
case tree.OUTPUT_RESULT_ROW:
return -1
case tree.OUTPUT_STATUS:
if execCtx.runResult != nil {
return int64(execCtx.runResult.AffectRows)
}
case tree.OUTPUT_UNDEFINED:
if _, ok := execCtx.stmt.(*tree.CallStmt); ok && execCtx.runResult != nil {
return int64(execCtx.runResult.AffectRows)
}
}
return 0
}

func executeStmtInBack(backSes *backSession,
statsArr *statistic.StatsArray,
execCtx *ExecCtx,
Expand Down Expand Up @@ -893,7 +945,9 @@ func getResultSet(ctx context.Context, bh BackgroundExec) ([]ExecResult, error)

type backSession struct {
feSessionImpl
//ep *ExportConfig
// lastAffectedRows carries the previous statement's ROW_COUNT() value into
// the next process created by this background executor.
lastAffectedRows int64
}

func newBackSession(ses FeSession, txnOp TxnOperator, db string, callBack outputCallBackFunc) *backSession {
Expand Down
2 changes: 1 addition & 1 deletion pkg/frontend/back_status_stmt.go
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,7 @@ func executeStatusStmtInBack(backSes *backSession,
}

runBegin := time.Now()
if _, err = execCtx.runner.Run(0); err != nil {
if execCtx.runResult, err = execCtx.runner.Run(0); err != nil {
return
}

Expand Down
Loading
Loading