diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..1fba180 --- /dev/null +++ b/.gitignore @@ -0,0 +1,2 @@ +*.out +*.html diff --git a/delete.go b/delete.go index 7cd4fe9..9b8ea9f 100644 --- a/delete.go +++ b/delete.go @@ -32,7 +32,7 @@ func (s delete_) Compile() (string, []any, error) { var sql strings.Builder var args []any - sql.WriteString("DELETE ") + sql.WriteString("DELETE") if s.from == "" { return "", nil, fmt.Errorf("FROM clause is required") @@ -50,34 +50,23 @@ func (s delete_) Compile() (string, []any, error) { return sql.String(), args, nil } -func (s delete_) Build(p Database) (*sql.Stmt, []any, error) { - return Build(s, p) -} +func (s delete_) MustCompile() (string, []any) { + sql, args, err := s.Compile() + if err != nil { + panic(err) + } -func (s delete_) MustBuild(p Database) (*sql.Stmt, []any) { - return MustBuild(s, p) + return sql, args } -func (s delete_) Exec(p Database) (sql.Result, error) { - return Exec(s, p) -} +func (s delete_) Build(p Database) (*sql.Stmt, []any, error) { return Build(s, p) } +func (s delete_) MustBuild(p Database) (*sql.Stmt, []any) { return MustBuild(s, p) } +func (s delete_) Exec(p Database) (sql.Result, error) { return Exec(s, p) } func (s delete_) ExecContext(ctx context.Context, p Database) (sql.Result, error) { return ExecContext(ctx, s, p) } - -func (s delete_) Query(p Database) (*sql.Rows, error) { - return Query(s, p) -} - -func (s delete_) QueryContext(ctx context.Context, p Database) (*sql.Rows, error) { - return QueryContext(ctx, s, p) -} - -func (s delete_) QueryRow(p Database) (*sql.Row, error) { - return QueryRow(s, p) -} - -func (s delete_) QueryRowContext(ctx context.Context, p Database) (*sql.Row, error) { - return QueryRowContext(ctx, s, p) +func (s delete_) MustExec(p Database) sql.Result { return MustExec(s, p) } +func (s delete_) MustExecContext(ctx context.Context, p Database) sql.Result { + return MustExecContext(ctx, s, p) } diff --git a/delete_test.go b/delete_test.go index ec7b7dd..87c9833 100644 --- a/delete_test.go +++ b/delete_test.go @@ -5,7 +5,101 @@ import ( "testing" ) -func TestDeleteIntegration_BasicQueries(t *testing.T) { +func TestDeleteCompileSuccess(t *testing.T) { + tests := []struct { + name string + stmt Compiler + expectedSql string + expectedArgs []any + }{ + { + name: "Delete all users", + stmt: Delete().From("users"), + expectedSql: "DELETE FROM users", + }, + { + name: "Delete active users only", + stmt: Delete(). + From("users"). + Where(Eq("active", true)), + expectedSql: "DELETE FROM users WHERE (active) = (?)", + expectedArgs: []any{true}, + }, + { + name: "Delete users in Engineering", + stmt: Delete(). + From("users"). + Where(Eq("department", "Engineering")), + expectedSql: "DELETE FROM users WHERE (department) = (?)", + expectedArgs: []any{"Engineering"}, + }, + { + name: "Delete users with age > 30", + stmt: Delete(). + From("users"). + Where(Gt("age", 30)), + expectedSql: "DELETE FROM users WHERE (age) > (?)", + expectedArgs: []any{30}, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + sql, args := test.stmt.MustCompile() + + if sql != test.expectedSql { + t.Errorf("Expected '%s', got '%s'", test.expectedSql, sql) + } + + if len(args) != len(test.expectedArgs) { + t.Errorf("Expected '%d' args, got '%d' args", len(test.expectedArgs), len(args)) + } + + for i := range len(args) { + if args[i] != test.expectedArgs[i] { + t.Errorf("Expected '%s', got '%s' at index %d", test.expectedArgs[i], args[i], i) + } + } + }) + } +} + +func TestDeleteCompileFail(t *testing.T) { + tests := []struct { + name string + stmt Compiler + expectedError string + }{ + { + name: "No from clause", + stmt: Delete(), + expectedError: "FROM clause is required", + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + sql, args, err := test.stmt.Compile() + if err == nil { + t.Error("Expected error, got nil") + } + + if err.Error() != test.expectedError { + t.Errorf("Expected error '%s', got '%s'", test.expectedError, err.Error()) + } + + if sql != "" { + t.Errorf("Expected empty SQL on error, got '%s'", sql) + } + + if args != nil { + t.Errorf("Expected empty args on error, got '%q'", args) + } + }) + } +} + +func TestDeleteIntegration(t *testing.T) { tests := []struct { name string stmt Execer diff --git a/expr.go b/expr.go index a0779c6..03d7d0a 100644 --- a/expr.go +++ b/expr.go @@ -36,6 +36,7 @@ func (e Expr) Binds() []any { return e.value.Binds() } + // unreachable return nil } diff --git a/expr_test.go b/expr_test.go index b30961f..8038cf5 100644 --- a/expr_test.go +++ b/expr_test.go @@ -100,3 +100,20 @@ func TestComplexExpressions(t *testing.T) { t.Errorf("Expected '%s', got '%s'", expected, complex.String()) } } + +// this just needs to compile +func TestImplsInterfaces(t *testing.T) { + sel := Select() + _ = isCompiler(sel) & isBuilder(sel) & isExecer(sel) & isQuerier(sel) + + del := Delete() + _ = isCompiler(del) & isBuilder(del) & isExecer(del) + + ins := Insert() + _ = isCompiler(ins) & isBuilder(ins) & isExecer(ins) +} + +func isCompiler[S Compiler](S) int { return 0 } +func isBuilder[S Builder](S) int { return 0 } +func isExecer[S Execer](S) int { return 0 } +func isQuerier[S Querier](S) int { return 0 } diff --git a/insert.go b/insert.go index 098c0b2..7e9e514 100644 --- a/insert.go +++ b/insert.go @@ -76,7 +76,8 @@ func (s insert) Compile() (string, []any, error) { sql.WriteString("INSERT ") - if s.or != None { + orKw := s.or.String() + if orKw != "" { sql.WriteString("OR ") sql.WriteString(s.or.String()) sql.WriteString(" ") @@ -110,34 +111,23 @@ func (s insert) Compile() (string, []any, error) { return sql.String(), args, nil } -func (s insert) Build(p Database) (*sql.Stmt, []any, error) { - return Build(s, p) -} +func (s insert) MustCompile() (string, []any) { + sql, args, err := s.Compile() + if err != nil { + panic(err) + } -func (s insert) MustBuild(p Database) (*sql.Stmt, []any) { - return MustBuild(s, p) + return sql, args } -func (s insert) Exec(p Database) (sql.Result, error) { - return Exec(s, p) -} +func (s insert) Build(p Database) (*sql.Stmt, []any, error) { return Build(s, p) } +func (s insert) MustBuild(p Database) (*sql.Stmt, []any) { return MustBuild(s, p) } +func (s insert) Exec(p Database) (sql.Result, error) { return Exec(s, p) } func (s insert) ExecContext(ctx context.Context, p Database) (sql.Result, error) { return ExecContext(ctx, s, p) } - -func (s insert) Query(p Database) (*sql.Rows, error) { - return Query(s, p) -} - -func (s insert) QueryContext(ctx context.Context, p Database) (*sql.Rows, error) { - return QueryContext(ctx, s, p) -} - -func (s insert) QueryRow(p Database) (*sql.Row, error) { - return QueryRow(s, p) -} - -func (s insert) QueryRowContext(ctx context.Context, p Database) (*sql.Row, error) { - return QueryRowContext(ctx, s, p) +func (s insert) MustExec(p Database) sql.Result { return MustExec(s, p) } +func (s insert) MustExecContext(ctx context.Context, p Database) sql.Result { + return MustExecContext(ctx, s, p) } diff --git a/insert_test.go b/insert_test.go index b2f1157..5db36f1 100644 --- a/insert_test.go +++ b/insert_test.go @@ -18,12 +18,42 @@ func TestInsertBuild_Success(t *testing.T) { expectedSql: "INSERT INTO users (name) VALUES (?)", expectedArgs: []any{"John"}, }, + { + name: "Abort clause", + stmt: Insert().Or(Abort).Into("users").Value("name", "John"), + expectedSql: "INSERT OR ABORT INTO users (name) VALUES (?)", + expectedArgs: []any{"John"}, + }, + { + name: "Ignore clause", + stmt: Insert().Or(Ignore).Into("users").Value("name", "John"), + expectedSql: "INSERT OR IGNORE INTO users (name) VALUES (?)", + expectedArgs: []any{"John"}, + }, + { + name: "Fail clause", + stmt: Insert().Or(Fail).Into("users").Value("name", "John"), + expectedSql: "INSERT OR FAIL INTO users (name) VALUES (?)", + expectedArgs: []any{"John"}, + }, { name: "Replace clause", stmt: Insert().Or(Replace).Into("users").Value("name", "John"), expectedSql: "INSERT OR REPLACE INTO users (name) VALUES (?)", expectedArgs: []any{"John"}, }, + { + name: "Rollback clause", + stmt: Insert().Or(Rollback).Into("users").Value("name", "John"), + expectedSql: "INSERT OR ROLLBACK INTO users (name) VALUES (?)", + expectedArgs: []any{"John"}, + }, + { + name: "Default clause", + stmt: Insert().Or(InsertOr(10)).Into("users").Value("name", "John"), + expectedSql: "INSERT INTO users (name) VALUES (?)", + expectedArgs: []any{"John"}, + }, { name: "More values", stmt: Insert().Into("users").Value("name", "John").Value("age", 35), @@ -34,11 +64,7 @@ func TestInsertBuild_Success(t *testing.T) { for _, test := range tests { t.Run(test.name, func(t *testing.T) { - sql, args, err := test.stmt.Compile() - - if err != nil { - t.Errorf("Expected no error, got %v", err) - } + sql, args := test.stmt.MustCompile() if sql != test.expectedSql { t.Errorf("Expected '%s', got '%s'", test.expectedSql, sql) @@ -56,3 +82,106 @@ func TestInsertBuild_Success(t *testing.T) { }) } } + +func TestInsertCompileFail(t *testing.T) { + tests := []struct { + name string + stmt Compiler + expectedError string + }{ + { + name: "No into clause", + stmt: Insert(), + expectedError: "INTO clause is required", + }, + { + name: "No into clause", + stmt: Insert().Into("users"), + expectedError: "no values supplied", + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + sql, args, err := test.stmt.Compile() + if err == nil { + t.Error("Expected error, got nil") + } + + if err.Error() != test.expectedError { + t.Errorf("Expected error '%s', got '%s'", test.expectedError, err.Error()) + } + + if sql != "" { + t.Errorf("Expected empty SQL on error, got '%s'", sql) + } + + if args != nil { + t.Errorf("Expected empty args on error, got '%q'", args) + } + }) + } +} + +func TestInsertIntegration(t *testing.T) { + tests := []struct { + name string + stmt Execer + expectedRows int64 + }{ + { + name: "Delete all users", + stmt: Delete().From("users"), + expectedRows: 6, + }, + { + name: "Delete active users only", + stmt: Delete(). + From("users"). + Where(Eq("active", true)), + expectedRows: 4, + }, + { + name: "Select users in Engineering", + stmt: Delete(). + From("users"). + Where(Eq("department", "Engineering")), + expectedRows: 3, + }, + { + name: "Delete users with age > 30", + stmt: Delete(). + From("users"). + Where(Gt("age", 30)), + expectedRows: 2, + }, + { + name: "Delete users with salary between 70000 and 80000", + stmt: Delete(). + From("users"). + Where(Gte("salary", 70000.0).And(Lte("salary", 80000.0))), + expectedRows: 3, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + db := setupTestDB(t) + defer db.Close() + + res, err := test.stmt.Exec(db) + if err != nil { + t.Fatalf("Failed to execute query: %v", err) + } + + count, err := res.RowsAffected() + if err != nil { + t.Fatalf("Failed to execute query: %v", err) + } + + if count != test.expectedRows { + t.Errorf("Expected %d rows, got %d", test.expectedRows, count) + } + }) + } +} diff --git a/select.go b/select.go index d6986a5..181c673 100644 --- a/select.go +++ b/select.go @@ -135,34 +135,41 @@ func (s select_) Compile() (string, []any, error) { return sql.String(), args, nil } -func (s select_) Build(p Database) (*sql.Stmt, []any, error) { - return Build(s, p) -} +func (s select_) MustCompile() (string, []any) { + sql, args, err := s.Compile() + if err != nil { + panic(err) + } -func (s select_) MustBuild(p Database) (*sql.Stmt, []any) { - return MustBuild(s, p) + return sql, args } -func (s select_) Exec(p Database) (sql.Result, error) { - return Exec(s, p) -} +func (s select_) Build(p Database) (*sql.Stmt, []any, error) { return Build(s, p) } +func (s select_) MustBuild(p Database) (*sql.Stmt, []any) { return MustBuild(s, p) } +func (s select_) Exec(p Database) (sql.Result, error) { return Exec(s, p) } func (s select_) ExecContext(ctx context.Context, p Database) (sql.Result, error) { return ExecContext(ctx, s, p) } - -func (s select_) Query(p Database) (*sql.Rows, error) { - return Query(s, p) +func (s select_) MustExec(p Database) sql.Result { return MustExec(s, p) } +func (s select_) MustExecContext(ctx context.Context, p Database) sql.Result { + return MustExecContext(ctx, s, p) } +func (s select_) Query(p Database) (*sql.Rows, error) { return Query(s, p) } func (s select_) QueryContext(ctx context.Context, p Database) (*sql.Rows, error) { return QueryContext(ctx, s, p) } - -func (s select_) QueryRow(p Database) (*sql.Row, error) { - return QueryRow(s, p) -} - +func (s select_) QueryRow(p Database) (*sql.Row, error) { return QueryRow(s, p) } func (s select_) QueryRowContext(ctx context.Context, p Database) (*sql.Row, error) { return QueryRowContext(ctx, s, p) } + +func (s select_) MustQuery(p Database) *sql.Rows { return MustQuery(s, p) } +func (s select_) MustQueryContext(ctx context.Context, p Database) *sql.Rows { + return MustQueryContext(ctx, s, p) +} +func (s select_) MustQueryRow(p Database) *sql.Row { return MustQueryRow(s, p) } +func (s select_) MustQueryRowContext(ctx context.Context, p Database) *sql.Row { + return MustQueryRowContext(ctx, s, p) +} diff --git a/select_test.go b/select_test.go index 07c22ee..7772987 100644 --- a/select_test.go +++ b/select_test.go @@ -5,60 +5,7 @@ import ( "testing" ) -func TestSelectBasic(t *testing.T) { - s := Select("name", "age") - - if len(s.resultColumns) != 2 { - t.Errorf("Expected 2 columns, got %d", len(s.resultColumns)) - } - - if s.resultColumns[0] != "name" || s.resultColumns[1] != "age" { - t.Errorf("Expected columns [name, age], got %v", s.resultColumns) - } -} - -func TestSelectAPI(t *testing.T) { - s := Select("*"). - From("users"). - Where(Eq("id", 1)). - OrderBy("name", Ascending). - GroupBy("department"). - Limit(10) - - if s.from != "users" { - t.Errorf("Expected from to be 'users', got '%s'", s.from) - } - - if s.where == nil { - t.Error("Expected where clause to be set") - } - - if len(s.orderBy) != 1 { - t.Errorf("Expected 1 order by clause, got %d", len(s.orderBy)) - } - - if s.orderBy[0].field != "name" || s.orderBy[0].direction != Ascending { - t.Errorf("Expected order by name ASC, got %v", s.orderBy[0]) - } - - if len(s.groupBy) != 1 { - t.Errorf("Expected 1 group by clause, got %d", len(s.groupBy)) - } - - if s.groupBy[0].field != "department" { - t.Errorf("Expected group by department, got %s", s.groupBy[0].field) - } - - if s.limit == nil { - t.Error("Expected limit to be set") - } - - if s.limit.limit != 10 { - t.Errorf("Expected limit 10, got %d", s.limit.limit) - } -} - -func TestSelectBuild_Success(t *testing.T) { +func TestSelectCompileSuccess(t *testing.T) { tests := []struct { name string stmt Compiler @@ -127,11 +74,7 @@ func TestSelectBuild_Success(t *testing.T) { for _, test := range tests { t.Run(test.name, func(t *testing.T) { - sql, args, err := test.stmt.Compile() - - if err != nil { - t.Errorf("Expected no error, got %v", err) - } + sql, args := test.stmt.MustCompile() if sql != test.expectedSql { t.Errorf("Expected '%s', got '%s'", test.expectedSql, sql) @@ -150,7 +93,7 @@ func TestSelectBuild_Success(t *testing.T) { } } -func TestSelectBuild_Errors(t *testing.T) { +func TestSelectCompileFail(t *testing.T) { tests := []struct { name string stmt Compiler @@ -205,50 +148,7 @@ func TestSelectBuild_Errors(t *testing.T) { } } -func TestMultipleOrderBy(t *testing.T) { - s := Select("name", "age"). - From("users"). - OrderBy("name", Ascending). - OrderBy("age", Descending). - OrderBy("created_at", Descending) - - if len(s.orderBy) != 3 { - t.Errorf("Expected 3 order by clauses, got %d", len(s.orderBy)) - } - - expected := []orderBy{ - {"name", Ascending}, - {"age", Descending}, - {"created_at", Descending}, - } - - for i, ob := range s.orderBy { - if ob.field != expected[i].field || ob.direction != expected[i].direction { - t.Errorf("Expected order by %v, got %v", expected[i], ob) - } - } -} - -func TestMultipleGroupBy(t *testing.T) { - s := Select("department", "status", "COUNT(*)"). - From("users"). - GroupBy("department"). - GroupBy("status") - - if len(s.groupBy) != 2 { - t.Errorf("Expected 2 group by clauses, got %d", len(s.groupBy)) - } - - expected := []string{"department", "status"} - - for i, gb := range s.groupBy { - if gb.field != expected[i] { - t.Errorf("Expected group by %s, got %s", expected[i], gb.field) - } - } -} - -func TestSelectIntegration_BasicQueries(t *testing.T) { +func TestSelectIntegration(t *testing.T) { db := setupTestDB(t) defer db.Close() @@ -310,239 +210,3 @@ func TestSelectIntegration_BasicQueries(t *testing.T) { }) } } - -func TestSelectIntegration_ComplexQueries(t *testing.T) { - db := setupTestDB(t) - defer db.Close() - - t.Run("Complex WHERE with AND/OR", func(t *testing.T) { - s := Select("name", "department", "age"). - From("users"). - Where( - Eq("department", "Engineering"). - And(Gt("age", 30)). - Or(Eq("department", "Marketing").And(Lt("age", 30))), - ) - - rows, err := s.Query(db) - if err != nil { - t.Fatalf("Failed to execute query: %v", err) - } - defer rows.Close() - - var results []struct { - name, department string - age int - } - - for rows.Next() { - var name, department string - var age int - if err := rows.Scan(&name, &department, &age); err != nil { - t.Fatalf("Failed to scan row: %v", err) - } - results = append(results, struct { - name, department string - age int - }{name, department, age}) - } - - // Should get: - // - Bob Johnson (Engineering, 35) - // - Charlie Wilson (Engineering, 32) - // - Jane Smith (Marketing, 25) - // - Diana Prince (Marketing, 29) - if len(results) != 4 { - t.Errorf("Expected 4 results, got %d", len(results)) - } - }) - - t.Run("ORDER BY multiple columns", func(t *testing.T) { - s := Select("name", "department", "age"). - From("users"). - OrderBy("department", Ascending). - OrderBy("age", Descending) - - rows, err := s.Query(db) - if err != nil { - t.Fatalf("Failed to execute query: %v", err) - } - defer rows.Close() - - var results []struct { - name, department string - age int - } - - for rows.Next() { - var name, department string - var age int - if err := rows.Scan(&name, &department, &age); err != nil { - t.Fatalf("Failed to scan row: %v", err) - } - results = append(results, struct { - name, department string - age int - }{name, department, age}) - } - - if len(results) != 6 { - t.Fatalf("Expected 6 results, got %d", len(results)) - } - - if results[0].department != "Engineering" { - t.Errorf("Expected first result to be from Engineering, got %s", results[0].department) - } - if results[0].age != 35 { - t.Errorf("Expected first result age to be 35, got %d", results[0].age) - } - }) - - t.Run("GROUP BY with COUNT", func(t *testing.T) { - s := Select("department", "COUNT(*) as user_count"). - From("users"). - GroupBy("department"). - OrderBy("user_count", Descending) - - rows, err := s.Query(db) - if err != nil { - t.Fatalf("Failed to execute query: %v", err) - } - defer rows.Close() - - var results []struct { - department string - count int - } - - for rows.Next() { - var department string - var count int - if err := rows.Scan(&department, &count); err != nil { - t.Fatalf("Failed to scan row: %v", err) - } - results = append(results, struct { - department string - count int - }{department, count}) - } - - if len(results) != 3 { - t.Errorf("Expected 3 departments, got %d", len(results)) - } - - if results[0].department != "Engineering" || results[0].count != 3 { - t.Errorf("Expected Engineering with 3 users first, got %s with %d users", results[0].department, results[0].count) - } - }) - - t.Run("LIMIT with ORDER BY", func(t *testing.T) { - s := Select("name", "salary"). - From("users"). - Where(Eq("active", true)). - OrderBy("salary", Descending). - Limit(2) - - rows, err := s.Query(db) - if err != nil { - t.Fatalf("Failed to execute query: %v", err) - } - defer rows.Close() - - var results []struct { - name string - salary float64 - } - - for rows.Next() { - var name string - var salary float64 - if err := rows.Scan(&name, &salary); err != nil { - t.Fatalf("Failed to scan row: %v", err) - } - results = append(results, struct { - name string - salary float64 - }{name, salary}) - } - - if len(results) != 2 { - t.Errorf("Expected 2 results, got %d", len(results)) - } - - if results[0].salary != 85000.0 { - t.Errorf("Expected highest salary to be 85000, got %f", results[0].salary) - } - if results[1].salary != 75000.0 { - t.Errorf("Expected second highest salary to be 75000, got %f", results[1].salary) - } - }) -} - -func TestSelectIntegration_DataTypes(t *testing.T) { - db := setupTestDB(t) - defer db.Close() - - t.Run("String comparisons", func(t *testing.T) { - s := Select("name", "email"). - From("users"). - Where(Eq("name", "John Doe")) - - rows, err := s.Query(db) - if err != nil { - t.Fatalf("Failed to execute query: %v", err) - } - defer rows.Close() - - count := 0 - for rows.Next() { - count++ - } - - if count != 1 { - t.Errorf("Expected 1 result, got %d", count) - } - }) - - t.Run("Boolean comparisons", func(t *testing.T) { - s := Select("name"). - From("users"). - Where(Eq("active", false)) - - rows, err := s.Query(db) - if err != nil { - t.Fatalf("Failed to execute query: %v", err) - } - defer rows.Close() - - count := 0 - for rows.Next() { - count++ - } - - if count != 2 { - t.Errorf("Expected 2 inactive users, got %d", count) - } - }) - - t.Run("Float comparisons", func(t *testing.T) { - s := Select("name", "salary"). - From("users"). - Where(Gte("salary", 80000.0)) - - rows, err := s.Query(db) - if err != nil { - t.Fatalf("Failed to execute query: %v", err) - } - defer rows.Close() - - count := 0 - for rows.Next() { - count++ - } - - if count != 2 { - t.Errorf("Expected 2 users with salary >= 80000, got %d", count) - } - }) -} diff --git a/types.go b/types.go index ae4dbce..7c76d95 100644 --- a/types.go +++ b/types.go @@ -16,6 +16,7 @@ type Database interface { type Compiler interface { Compile() (string, []any, error) + MustCompile() (string, []any) } type Builder interface { @@ -49,6 +50,9 @@ func MustBuild(c Compiler, p Database) (*sql.Stmt, []any) { type Execer interface { Exec(db Database) (sql.Result, error) ExecContext(ctx context.Context, db Database) (sql.Result, error) + + MustExec(db Database) sql.Result + MustExecContext(ctx context.Context, db Database) sql.Result } // anything that is a Builder can also be an Execer @@ -65,20 +69,43 @@ func ExecContext(ctx context.Context, b Builder, p Database) (sql.Result, error) return stmt.ExecContext(ctx, args...) } +func MustExec(b Builder, p Database) sql.Result { + res, err := Exec(b, p) + if err != nil { + panic(err) + } + + return res +} + +func MustExecContext(ctx context.Context, b Builder, p Database) sql.Result { + res, err := ExecContext(ctx, b, p) + if err != nil { + panic(err) + } + + return res +} + type Querier interface { Query(db Database) (*sql.Rows, error) QueryContext(ctx context.Context, db Database) (*sql.Rows, error) QueryRow(db Database) (*sql.Row, error) QueryRowContext(ctx context.Context, db Database) (*sql.Row, error) + + MustQuery(db Database) *sql.Rows + MustQueryContext(ctx context.Context, db Database) *sql.Rows + MustQueryRow(db Database) *sql.Row + MustQueryRowContext(ctx context.Context, db Database) *sql.Row } // anything that is a Builder can also be an Querier -func Query(b Builder, p Database) (*sql.Rows, error) { - return QueryContext(context.Background(), b, p) +func Query(b Builder, db Database) (*sql.Rows, error) { + return QueryContext(context.Background(), b, db) } -func QueryContext(ctx context.Context, b Builder, p Database) (*sql.Rows, error) { - stmt, args, err := b.Build(p) +func QueryContext(ctx context.Context, b Builder, db Database) (*sql.Rows, error) { + stmt, args, err := b.Build(db) if err != nil { return nil, err } @@ -86,12 +113,12 @@ func QueryContext(ctx context.Context, b Builder, p Database) (*sql.Rows, error) return stmt.QueryContext(ctx, args...) } -func QueryRow(b Builder, p Database) (*sql.Row, error) { - return QueryRowContext(context.Background(), b, p) +func QueryRow(b Builder, db Database) (*sql.Row, error) { + return QueryRowContext(context.Background(), b, db) } -func QueryRowContext(ctx context.Context, b Builder, p Database) (*sql.Row, error) { - stmt, args, err := b.Build(p) +func QueryRowContext(ctx context.Context, b Builder, db Database) (*sql.Row, error) { + stmt, args, err := b.Build(db) if err != nil { return nil, err } @@ -99,6 +126,42 @@ func QueryRowContext(ctx context.Context, b Builder, p Database) (*sql.Row, erro return stmt.QueryRowContext(ctx, args...), nil } +func MustQuery(b Builder, db Database) *sql.Rows { + rows, err := Query(b, db) + if err != nil { + panic(err) + } + + return rows +} + +func MustQueryContext(ctx context.Context, b Builder, db Database) *sql.Rows { + rows, err := QueryContext(ctx, b, db) + if err != nil { + panic(err) + } + + return rows +} + +func MustQueryRow(b Builder, db Database) *sql.Row { + row, err := QueryRow(b, db) + if err != nil { + panic(err) + } + + return row +} + +func MustQueryRowContext(ctx context.Context, b Builder, db Database) *sql.Row { + row, err := QueryRowContext(ctx, b, db) + if err != nil { + panic(err) + } + + return row +} + type Direction string const (