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
51 changes: 48 additions & 3 deletions pkg/frontend/mysql_cmd_executor.go
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,7 @@ import (
"strings"
"sync"
"time"
"unicode"

"github.com/confluentinc/confluent-kafka-go/v2/kafka"
"github.com/google/uuid"
Expand Down Expand Up @@ -1551,10 +1552,54 @@ func handleExplainStmt(ses FeSession, execCtx *ExecCtx, stmt *tree.ExplainStmt)
return doExplainStmt(execCtx.reqCtx, ses.(*Session), stmt)
}

func extractPrepareStmtSQL(ctx context.Context, sql, sqlMode string) (string, error) {
scanner := mysql.NewScannerWithSQLMode(dialect.MYSQL, sql, mysql.ParseSQLModeFlags(sqlMode))
defer mysql.PutScanner(scanner)

if token, _ := scanner.Scan(); token != mysql.PREPARE {
return "", moerr.NewInvalidInput(ctx, "invalid PREPARE statement")
}
if token, _ := scanner.Scan(); token == mysql.EofChar() || token == mysql.LEX_ERROR {
return "", moerr.NewInvalidInput(ctx, "invalid PREPARE statement name")
}
if token, _ := scanner.Scan(); token != mysql.FROM {
return "", moerr.NewInvalidInput(ctx, "invalid PREPARE statement delimiter")
}

preparedStart := scanner.Pos
preparedSQL := sql[preparedStart:]
if scanner.CommentFlag {
scanner.TakeExecutableCommentEnd()
var commentEnd int
for commentEnd == 0 {
previousPos := scanner.Pos
token, _ := scanner.Scan()
commentEnd = scanner.TakeExecutableCommentEnd()
if commentEnd == 0 &&
(token == mysql.EofChar() || token == mysql.LEX_ERROR || scanner.Pos == previousPos) {
return "", moerr.NewInvalidInput(ctx, "invalid PREPARE executable comment")
}
}
insideComment := strings.TrimSpace(sql[preparedStart : commentEnd-2])
afterComment := strings.TrimLeftFunc(sql[commentEnd:], unicode.IsSpace)
switch {
case insideComment == "":
preparedSQL = afterComment
case afterComment == "":
preparedSQL = insideComment
default:
preparedSQL = insideComment + " " + afterComment
}
}

return strings.TrimLeftFunc(preparedSQL, unicode.IsSpace), nil
}

func doPrepareStmt(execCtx *ExecCtx, ses *Session, st *tree.PrepareStmt, sql string, paramTypes []byte) (*PrepareStmt, error) {
idx := strings.Index(strings.ToLower(sql[:(len(st.Name)+20)]), "from") + 5
originSql := strings.TrimLeft(sql[idx:], " ")
// fmt.Print(originSql)
originSql, err := extractPrepareStmtSQL(execCtx.reqCtx, sql, sessionSQLModeForParser(ses))
if err != nil {
return nil, err
}
prepareStmt, err := createPrepareStmt(execCtx, ses, originSql, st, st.Stmt)
if err != nil {
return nil, err
Expand Down
144 changes: 144 additions & 0 deletions pkg/frontend/mysql_cmd_executor_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -1568,6 +1568,150 @@ func Test_HandlePrepareStmt(t *testing.T) {
})
}

func TestHandlePrepareStmtNameContainingFrom(t *testing.T) {
setSessionAlloc("", NewLeakCheckAllocator())
ctx := defines.AttachAccountId(context.TODO(), catalog.System_Account)
const sql = "prepare fromx from select 1"
stmt, err := parsers.ParseOne(ctx, dialect.MYSQL, sql, 1)
require.NoError(t, err)

ctrl := gomock.NewController(t)
defer ctrl.Finish()
execCtx := newTestExecCtx(ctx, ctrl)

runTestHandle("handlePrepareStmt name containing from", t, func(ses *Session) error {
execCtx.resper = ses.respr
prepared, err := handlePrepareStmt(ses, execCtx, stmt.(*tree.PrepareStmt), sql)
if err != nil {
return err
}
require.Equal(t, "select 1", prepared.Sql)
return nil
})
}

func TestHandlePrepareStmtExecutableCommentDelimiter(t *testing.T) {
setSessionAlloc("", NewLeakCheckAllocator())
ctx := defines.AttachAccountId(context.TODO(), catalog.System_Account)
const sql = "prepare fromx /*! from */ select 1"
stmt, err := parsers.ParseOne(ctx, dialect.MYSQL, sql, 1)
require.NoError(t, err)

ctrl := gomock.NewController(t)
defer ctrl.Finish()
execCtx := newTestExecCtx(ctx, ctrl)

runTestHandle("handlePrepareStmt executable comment delimiter", t, func(ses *Session) error {
execCtx.resper = ses.respr
prepared, err := handlePrepareStmt(ses, execCtx, stmt.(*tree.PrepareStmt), sql)
if err != nil {
return err
}
require.Equal(t, "select 1", prepared.Sql)
return nil
})
}

func TestHandlePrepareStmtQuotedCommentTerminator(t *testing.T) {
setSessionAlloc("", NewLeakCheckAllocator())
ctx := defines.AttachAccountId(context.TODO(), catalog.System_Account)
const sql = "prepare fromx /*! from select 'x*/y' */"
stmt, err := parsers.ParseOne(ctx, dialect.MYSQL, sql, 1)
require.NoError(t, err)

ctrl := gomock.NewController(t)
defer ctrl.Finish()
execCtx := newTestExecCtx(ctx, ctrl)

runTestHandle("handlePrepareStmt quoted comment terminator", t, func(ses *Session) error {
execCtx.resper = ses.respr
prepared, err := handlePrepareStmt(ses, execCtx, stmt.(*tree.PrepareStmt), sql)
if err != nil {
return err
}
require.Equal(t, "select 'x*/y'", prepared.Sql)
return nil
})
}

func TestExtractPrepareStmtSQL(t *testing.T) {
testCases := []struct {
name string
sql string
sqlMode string
want string
}{
{
name: "name contains delimiter text",
sql: "prepare fromx from select 1",
want: "select 1",
},
{
name: "quoted delimiter name",
sql: "prepare `from` from\n\tselect 2",
want: "select 2",
},
{
name: "ansi quoted delimiter name",
sql: `prepare "from" from select 3`,
sqlMode: "ANSI_QUOTES",
want: "select 3",
},
{
name: "preserve inner comment",
sql: "prepare from_name /* before delimiter */ from /* inner */ select 4",
want: "/* inner */ select 4",
},
{
name: "leading comment",
sql: "/* rewrite hint */ prepare fromx from select 5",
want: "select 5",
},
{
name: "executable comment delimiter",
sql: "prepare fromx /*! from */ select 1",
want: "select 1",
},
{
name: "statement inside executable comment",
sql: "prepare fromx /*! from select 2 */",
want: "select 2",
},
{
name: "quoted comment terminator",
sql: "prepare fromx /*! from select 'x*/y' */",
want: "select 'x*/y'",
},
{
name: "preserve comment after executable delimiter",
sql: "prepare fromx /*! from */ /* inner */ select 3",
want: "/* inner */ select 3",
},
}

for _, testCase := range testCases {
t.Run(testCase.name, func(t *testing.T) {
got, err := extractPrepareStmtSQL(context.Background(), testCase.sql, testCase.sqlMode)
require.NoError(t, err)
require.Equal(t, testCase.want, got)
})
}
}

func TestExtractPrepareStmtSQLRejectsInvalidInput(t *testing.T) {
for _, sql := range []string{
"select 1",
"prepare",
"prepare stmt select 1",
"prepare stmt /*! from select 1",
} {
t.Run(sql, func(t *testing.T) {
_, err := extractPrepareStmtSQL(context.Background(), sql, "")
require.Error(t, err)
})
}
}

func TestHandlePrepareStmtStoresRemapPolicy(t *testing.T) {
setSessionAlloc("", NewLeakCheckAllocator())
ctx := defines.AttachAccountId(context.Background(), catalog.System_Account)
Expand Down
25 changes: 19 additions & 6 deletions pkg/sql/parsers/dialect/mysql/scanner.go
Original file line number Diff line number Diff line change
Expand Up @@ -45,12 +45,13 @@ type Scanner struct {
sqlMode SQLModeFlags
MysqlSpecialComment *Scanner

CommentFlag bool
Pos int
Line int
Col int
PrePos int
buf string
CommentFlag bool
Pos int
Line int
Col int
PrePos int
buf string
executableCommentEnd int

strBuilder bytes.Buffer
}
Expand All @@ -66,6 +67,7 @@ func (s *Scanner) reset(clearLargeOnly bool, oversized bool) {
s.Line = 0
s.Col = 0
s.PrePos = 0
s.executableCommentEnd = 0
s.sqlMode = 0

if clearLargeOnly {
Expand Down Expand Up @@ -266,6 +268,9 @@ func (s *Scanner) Scan() (int, string) {
case '/':
s.CommentFlag = false
s.inc()
if s.executableCommentEnd == 0 {
s.executableCommentEnd = s.Pos
}
return s.Scan()
default:
return s.stepBackOneChar(ch)
Expand Down Expand Up @@ -309,6 +314,14 @@ func (s *Scanner) Scan() (int, string) {
}
}

// TakeExecutableCommentEnd returns the byte offset immediately after the first
// executable-comment terminator scanned since the previous call.
func (s *Scanner) TakeExecutableCommentEnd() int {
end := s.executableCommentEnd
s.executableCommentEnd = 0
return end
}

func (s *Scanner) isCollate() bool {
if s.peek(1) == 'u' && s.peek(2) == 't' && s.peek(3) == 'f' && s.peek(4) == '8' && s.peek(5) == 'm' && s.peek(6) == 'b' && s.peek(7) == '4' {
return true
Expand Down
28 changes: 27 additions & 1 deletion pkg/sql/parsers/dialect/mysql/scanner_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -448,11 +448,12 @@ func TestScannerPoolCleanupAndThreshold(t *testing.T) {
s := NewScanner(dialect.MYSQL, "select 1")
// grow strBuilder a little to ensure it is cleared on Put
s.strBuilder.WriteString("abc")
s.executableCommentEnd = 42
PutScanner(s)

// Fetch again to see if we receive a cleared scanner from pool
s2 := NewScanner(dialect.MYSQL, "select 2")
if s2.LastToken != "" || s2.LastError != nil || s2.MysqlSpecialComment != nil || s2.Pos != 0 || s2.Line != 0 || s2.Col != 0 || s2.PrePos != 0 {
if s2.LastToken != "" || s2.LastError != nil || s2.MysqlSpecialComment != nil || s2.Pos != 0 || s2.Line != 0 || s2.Col != 0 || s2.PrePos != 0 || s2.executableCommentEnd != 0 {
t.Fatalf("pooled scanner should be reset: %+v", s2)
}
if s2.strBuilder.Len() != 0 {
Expand All @@ -478,6 +479,31 @@ func TestScannerPoolCleanupAndThreshold(t *testing.T) {
PutScanner(s3)
}

func TestExecutableCommentEndSkipsQuotedTerminator(t *testing.T) {
const sql = "prepare fromx /*! from select 'x*/y' */"
s := NewScanner(dialect.MYSQL, sql)
defer PutScanner(s)

for i := 0; i < 3; i++ {
if token, _ := s.Scan(); token == LEX_ERROR || token == EofChar() {
t.Fatalf("unexpected token %d", token)
}
}
s.TakeExecutableCommentEnd()
for {
token, _ := s.Scan()
if end := s.TakeExecutableCommentEnd(); end != 0 {
if end != len(sql) {
t.Fatalf("comment end = %d, want %d", end, len(sql))
}
return
}
if token == LEX_ERROR || token == EofChar() {
t.Fatal("executable comment terminator not found")
}
}
}

func TestPutScannerSmallKeepsBuffers(t *testing.T) {
// Small SQL should keep buf and builder content when returned to pool
sql := "select 1"
Expand Down
5 changes: 5 additions & 0 deletions test/distributed/cases/prepare/prepare.result
Original file line number Diff line number Diff line change
@@ -1,3 +1,8 @@
prepare fromx from select 1;
execute fromx;
1
1
deallocate prepare fromx;
drop table if exists t1;
create table t1 (a int, b int);
prepare stmt1 from 'select * from t1 where a > ?';
Expand Down
4 changes: 4 additions & 0 deletions test/distributed/cases/prepare/prepare.test
Original file line number Diff line number Diff line change
@@ -1,5 +1,9 @@

-- @label:bvt
prepare fromx from select 1;
execute fromx;
deallocate prepare fromx;

drop table if exists t1;
create table t1 (a int, b int);
prepare stmt1 from 'select * from t1 where a > ?';
Expand Down
Loading