(sql string, limitCount int)
| 122 | } |
| 123 | |
| 124 | func rewriteMySQLSelectProcedureLimit(sql string, limitCount int) (string, error) { |
| 125 | statements, err := mysqlparser.SplitSQL(sql) |
| 126 | if err != nil { |
| 127 | return "", err |
| 128 | } |
| 129 | statementCount := 0 |
| 130 | for _, statement := range statements { |
| 131 | if !statement.Empty { |
| 132 | statementCount++ |
| 133 | } |
| 134 | } |
| 135 | if statementCount != 1 { |
| 136 | return "", errors.Errorf("expected exactly one statement, got %d", statementCount) |
| 137 | } |
| 138 | |
| 139 | tokens := omnimysqlparser.Tokenize(sql) |
| 140 | if len(tokens) == 0 || omnimysqlparser.TokenName(tokens[0].Type) != "SELECT" { |
| 141 | return "", errors.New("not a SELECT PROCEDURE statement") |
| 142 | } |
| 143 | for i, token := range tokens[1:] { |
| 144 | if omnimysqlparser.TokenName(token.Type) == "PROCEDURE" { |
| 145 | procedureTokenIndex := i + 1 |
| 146 | if stmt, ok := rewriteMySQLLimitBeforeProcedure(sql, tokens[:procedureTokenIndex], limitCount); ok { |
| 147 | return stmt, nil |
| 148 | } |
| 149 | return sql[:token.Loc] + fmt.Sprintf("LIMIT %d ", limitCount) + sql[token.Loc:], nil |
| 150 | } |
| 151 | } |
| 152 | return "", errors.New("SELECT PROCEDURE clause not found") |
| 153 | } |
| 154 | |
| 155 | func rewriteMySQLLimitBeforeProcedure(sql string, tokens []omnimysqlparser.Token, limitCount int) (string, bool) { |
| 156 | depth := 0 |
no test coverage detected