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
147 changes: 147 additions & 0 deletions pkg/embed/prepared_specialized_domains_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -89,6 +89,153 @@ func TestPreparedSpecializedDomains(t *testing.T) {
require.Equal(t, [][]string{{tc.want}}, query(t, "execute numeric_sum using @numeric_source"), tc.assignment)
}
})
t.Run("prepared_round_truncate_value_domains", func(t *testing.T) {
for _, tc := range []struct {
name string
fn string
}{
{name: "round", fn: "round"},
{name: "truncate", fn: "truncate"},
} {
fractionalResult := "1.4"
if tc.fn == "round" {
fractionalResult = "1.5"
}
t.Run(tc.name+"/sql_execute", func(t *testing.T) {
stmtName := "numeric_" + tc.name
exec(t, "prepare "+stmtName+" from 'select "+tc.fn+"(?,?)'")
defer conn.ExecContext(ctx, "deallocate prepare "+stmtName)
for _, value := range []struct {
assignment string
precision string
want string
}{
{"'1.46'", "1", fractionalResult + "0"},
{"cast(1.46 as decimal(10,2))", "1", fractionalResult + "0"},
{"2", "0", "2"},
{"null", "1", "NULL"},
} {
exec(t, "set @numeric_value="+value.assignment+", @numeric_precision="+value.precision)
require.Equal(t, [][]string{{value.want}}, query(t,
"execute "+stmtName+" using @numeric_value,@numeric_precision"), value)
}
exec(t, "set @numeric_value='not-a-number', @numeric_precision=1")
err := func() error {
rows, err := conn.QueryContext(ctx,
"execute "+stmtName+" using @numeric_value,@numeric_precision")
if rows == nil {
return err
}
defer rows.Close()
for rows.Next() {
}
return rows.Err()
}()
require.Error(t, err, "invalid text should fail as on the GOOD baseline")
exec(t, "set @numeric_value='1.46'")
require.Equal(t, [][]string{{fractionalResult + "0"}},
query(t, "execute "+stmtName+" using @numeric_value,@numeric_precision"),
"a failed execution must not poison the next binding")
})

t.Run(tc.name+"/binary_protocol", func(t *testing.T) {
stmt, err := conn.PrepareContext(ctx, "select "+tc.fn+"(?,?)")
require.NoError(t, err)
defer stmt.Close()
for _, value := range []struct {
input any
want string
}{
{"1.46", fractionalResult + "0"},
{int64(2), "2"},
{float64(1.46), fractionalResult},
{[]byte("1.46"), fractionalResult + "0"},
{nil, "NULL"},
} {
var got sql.NullString
require.NoError(t, stmt.QueryRowContext(ctx, value.input, 1).Scan(&got), value)
gotString := "NULL"
if got.Valid {
gotString = got.String
}
require.Equal(t, value.want, gotString, value)
}
})
for _, shape := range []struct {
name string
value string
}{
{"scalar", "(select ?)"},
{"derived", "x"},
} {
t.Run(tc.name+"/"+shape.name, func(t *testing.T) {
statement := "select cast(" + tc.fn + "(" + shape.value + ",1) as double)"
if shape.name == "derived" {
statement += " from (select ? x limit 1) d"
}
stmt, err := conn.PrepareContext(ctx, statement)
require.NoError(t, err)
defer stmt.Close()
var got sql.NullString
require.NoError(t, stmt.QueryRowContext(ctx, "1.46").Scan(&got))
require.Equal(t, fractionalResult, got.String)
tieStatement := "select cast(" + tc.fn + "(" + shape.value + ",0) as double)"
if shape.name == "derived" {
tieStatement += " from (select ? x limit 1) d"
}
tie, tieErr := conn.PrepareContext(ctx, tieStatement)
require.NoError(t, tieErr)
defer tie.Close()
require.NoError(t, tie.QueryRowContext(ctx, "2.5").Scan(&got))
if tc.name == "round" {
require.Equal(t, "3", got.String)
} else {
require.Equal(t, "2", got.String)
}
if shape.name == "derived" {
require.Error(t, stmt.QueryRowContext(ctx, "not-a-number").Scan(&got))
}
require.NoError(t, stmt.QueryRowContext(ctx, "1.46").Scan(&got))
require.Equal(t, fractionalResult, got.String)
})
}
t.Run(tc.name+"/set_operation_domain", func(t *testing.T) {
stmt, err := conn.PrepareContext(ctx,
"select cast("+tc.fn+"(x,1) as double) from "+
"(select ? x union all select cast(1.46 as decimal(10,2))) d")
require.NoError(t, err)
defer stmt.Close()
rows, err := stmt.QueryContext(ctx, int64(1))
require.NoError(t, err)
defer rows.Close()
var got []string
for rows.Next() {
var value string
require.NoError(t, rows.Scan(&value))
got = append(got, value)
}
require.NoError(t, rows.Err())
fraction := "1.5"
if tc.name == "truncate" {
fraction = "1.4"
}
require.ElementsMatch(t, []string{"1", fraction}, got)
})
}

exec(t, "prepare explicit_decimal_round from 'select round(cast(? as decimal(10,2)),1)'")
defer conn.ExecContext(ctx, "deallocate prepare explicit_decimal_round")
exec(t, "set @explicit_decimal='1.46'")
require.Equal(t, [][]string{{"1.5"}}, query(t,
"execute explicit_decimal_round using @explicit_decimal"))
explicitDerived, err := conn.PrepareContext(ctx,
"select cast(round(x,0) as double) from (select cast(? as decimal(10,2)) x) d")
require.NoError(t, err)
defer explicitDerived.Close()
var explicitGot sql.NullString
require.NoError(t, explicitDerived.QueryRowContext(ctx, "2.5").Scan(&explicitGot))
require.Equal(t, "3", explicitGot.String)
})
t.Run("ntile_null_runtime_error", func(t *testing.T) {
exec(t, "create table ntile_source(id int)")
exec(t, "insert into ntile_source values (1),(2)")
Expand Down
7 changes: 6 additions & 1 deletion pkg/frontend/computation_wrapper.go
Original file line number Diff line number Diff line change
Expand Up @@ -643,7 +643,12 @@ func (cwft *TxnComputationWrapper) Compile(any any, fill func(*batch.Batch, *per
*/
} else {
var planSnapshotTS *timestamp.Timestamp
if cwft.hasPlanSnapshotTS {
if cwft.preparedStmt != nil {
// Executable EXPLAIN uses the already bound prepared plan. Carry the
// binding's diagnostic proof and snapshot into Compile as EXECUTE does.
preparedExprRetry = cwft.preparedExecutionRetry()
planSnapshotTS = &cwft.preparedStmt.Ts
} else if cwft.hasPlanSnapshotTS {
planSnapshotTS = &cwft.planSnapshotTS
}
cwft.compile, err = createCompile(
Expand Down
Loading
Loading