(statement string, limit int)
| 96 | } |
| 97 | |
| 98 | func getStatementWithResultLimitInline(statement string, limit int) (string, error) { |
| 99 | stmtList, err := tidbparser.ParseTiDB(statement, "", "") |
| 100 | if err != nil { |
| 101 | return "", errors.Wrapf(err, "failed to parse tidb statement: %s", statement) |
| 102 | } |
| 103 | if len(stmtList) != 1 { |
| 104 | return "", errors.Errorf("expect one single statement in the query, %s", statement) |
| 105 | } |
| 106 | restoreFlags := format.DefaultRestoreFlags | format.RestoreStringWithoutDefaultCharset |
| 107 | stmt := stmtList[0] |
| 108 | switch stmt := stmt.(type) { |
| 109 | case *tidbast.SelectStmt: |
| 110 | if stmt.Limit != nil && stmt.Limit.Count != nil { |
| 111 | if v, ok := stmt.Limit.Count.(*tidbdriver.ValueExpr); ok { |
| 112 | userLimit := int(v.GetInt64()) |
| 113 | if limit < userLimit { |
| 114 | userLimit = limit |
| 115 | } |
| 116 | stmt.Limit.Count = tidbast.NewValueExpr(int64(userLimit), "", "") |
| 117 | } |
| 118 | } else { |
| 119 | stmt.Limit = &tidbast.Limit{ |
| 120 | Count: tidbast.NewValueExpr(int64(limit), "", ""), |
| 121 | } |
| 122 | } |
| 123 | var buffer strings.Builder |
| 124 | ctx := format.NewRestoreCtx(restoreFlags, &buffer) |
| 125 | if err := stmt.Restore(ctx); err != nil { |
| 126 | return "", err |
| 127 | } |
| 128 | return buffer.String(), nil |
| 129 | case *tidbast.SetOprStmt: |
| 130 | if stmt.Limit != nil && stmt.Limit.Count != nil { |
| 131 | if v, ok := stmt.Limit.Count.(*tidbdriver.ValueExpr); ok { |
| 132 | userLimit := int(v.GetInt64()) |
| 133 | if limit < userLimit { |
| 134 | userLimit = limit |
| 135 | } |
| 136 | stmt.Limit.Count = tidbast.NewValueExpr(int64(userLimit), "", "") |
| 137 | } |
| 138 | } else { |
| 139 | stmt.Limit = &tidbast.Limit{ |
| 140 | Count: tidbast.NewValueExpr(int64(limit), "", ""), |
| 141 | } |
| 142 | } |
| 143 | var buffer strings.Builder |
| 144 | ctx := format.NewRestoreCtx(restoreFlags, &buffer) |
| 145 | if err := stmt.Restore(ctx); err != nil { |
| 146 | return "", err |
| 147 | } |
| 148 | return buffer.String(), nil |
| 149 | default: |
| 150 | } |
| 151 | return statement, nil |
| 152 | } |
no test coverage detected