Add support for prepared statement caching.

This commit is contained in:
go-jet 2024-03-07 18:01:31 +01:00
parent 1b63280b74
commit 0918e5503e
30 changed files with 603 additions and 289 deletions

182
internal/jet/db/db.go Normal file
View file

@ -0,0 +1,182 @@
package db
import (
"context"
"database/sql"
"fmt"
"sync"
)
// DB is a wrapper around sql.DB, adding prepared statement caching capability.
type DB struct {
*sql.DB
statementsCaching bool
lock sync.RWMutex
statements map[string]*sql.Stmt
}
// NewDB creates new DB wrapper with statements caching disabled
func NewDB(db *sql.DB) *DB {
return &DB{
DB: db,
statementsCaching: false,
statements: make(map[string]*sql.Stmt),
}
}
// WithStatementsCaching returns *DB wrapper with prepared statements caching enabled or disabled. This method should be
// called only once. It is not concurrency-safe.
func (d *DB) WithStatementsCaching(enabled bool) *DB {
d.statementsCaching = enabled
return d
}
// Begin starts sql transaction and returns wrapped Tx object.
func (d *DB) Begin() (*Tx, error) {
tx, err := d.DB.Begin()
if err != nil {
return nil, err
}
return &Tx{
Tx: tx,
db: d,
statements: make(map[string]*sql.Stmt),
}, nil
}
// BeginTx starts sql transaction and returns wrapped Tx object.
func (d *DB) BeginTx(ctx context.Context, opts *sql.TxOptions) (*Tx, error) {
tx, err := d.DB.BeginTx(ctx, opts)
if err != nil {
return nil, err
}
return &Tx{
Tx: tx,
db: d,
statements: make(map[string]*sql.Stmt),
}, nil
}
// Exec executes a query that doesn't return rows. Exec delegates call to ExecContext with contex.Background()
// as parameter.
func (d *DB) Exec(query string, args ...interface{}) (sql.Result, error) {
return d.ExecContext(context.Background(), query, args...)
}
// ExecContext executes a query that doesn't return rows. If statement caching is enabled, ExecContext will
// first call PrepareContext to retrieve a prepared statement, and then execute a query using a prepared statement.
// If statement caching is disabled, this method delegates the call to the *sql.DB ExecContext method.
func (d *DB) ExecContext(ctx context.Context, query string, args ...interface{}) (sql.Result, error) {
if !d.statementsCaching {
return d.DB.ExecContext(ctx, query, args...)
}
prepStmt, err := d.PrepareContext(ctx, query)
if err != nil {
return nil, err
}
return prepStmt.ExecContext(ctx, args...)
}
// Query delegates call to QueryContext using context.Background() as parameter.
func (d *DB) Query(query string, args ...interface{}) (*sql.Rows, error) {
return d.QueryContext(context.Background(), query, args...)
}
// QueryContext executes a query that returns rows. If statement caching is enabled, QueryContext will
// first call PrepareContext to retrieve a prepared statement, and then execute a query using a prepared statement.
// If statement caching is disabled, this method delegates the call to the *sql.DB QueryContext method.
func (d *DB) QueryContext(ctx context.Context, query string, args ...interface{}) (*sql.Rows, error) {
if !d.statementsCaching {
return d.DB.QueryContext(ctx, query, args...)
}
prepStmt, err := d.PrepareContext(ctx, query)
if err != nil {
return nil, err
}
return prepStmt.QueryContext(ctx, args...)
}
// Prepare delegates call to PrepareContext using context.Background as a parameter.
func (d *DB) Prepare(query string) (*sql.Stmt, error) {
return d.PrepareContext(context.Background(), query)
}
// PrepareContext returns database prepared statement for a query. When statement caching is enabled, it returns a cached
// prepared statement if available; otherwise, it creates a new prepared statement and adds it to the cache.
// Invoking this method directly is unnecessary, as wrapper methods like Exec/ExecContext and Query/QueryContext
// will call PrepareContext before executing a query on it.
// If statement caching is disabled, this method delegates the call to the *sql.DB PrepareContext method.
//
// There's no need to manually close the returned statement; it operates within the transaction scope and will be closed
// automatically upon the completion of the transaction, whether it's committed or rolled back.
func (d *DB) PrepareContext(ctx context.Context, query string) (*sql.Stmt, error) {
if !d.statementsCaching {
return d.DB.PrepareContext(ctx, query)
}
d.lock.RLock()
prepStmt, ok := d.statements[query]
d.lock.RUnlock()
if ok {
return prepStmt, nil
}
prepStmt, err := d.DB.PrepareContext(ctx, query)
if err != nil {
return nil, fmt.Errorf("failed to prepare statement %s: %w", query, err)
}
d.lock.Lock()
existingPrepStmt, exist := d.statements[query]
// if in the meantime, another goroutine created prepared statements for this query, we will close this
// prepared statement and return the existing one.
if exist {
_ = prepStmt.Close()
d.lock.Unlock()
return existingPrepStmt, nil
}
d.statements[query] = prepStmt
d.lock.Unlock()
return prepStmt, nil
}
// Clear will close all cached prepared statements
func (d *DB) Clear() error {
d.lock.Lock()
defer d.lock.Unlock()
var err error
for _, statement := range d.statements {
closeErr := statement.Close()
if closeErr != nil {
err = closeErr
}
}
d.statements = make(map[string]*sql.Stmt)
if err != nil {
return fmt.Errorf("some of the prepared statements failed to close, last err: %w", err)
}
return nil
}

97
internal/jet/db/tx.go Normal file
View file

@ -0,0 +1,97 @@
package db
import (
"context"
"database/sql"
"fmt"
)
// Tx is a wrapper around *sql.Tx, adding prepared statement caching capability.
type Tx struct {
*sql.Tx
db *DB
statements map[string]*sql.Stmt
}
// Exec executes a query that doesn't return rows. Exec delegates call to ExecContext with contex.Background()
// as parameter.
func (t *Tx) Exec(query string, args ...interface{}) (sql.Result, error) {
return t.ExecContext(context.Background(), query, args...)
}
// ExecContext executes a query that doesn't return rows. If statement caching is enabled, ExecContext will
// first call PrepareContext to retrieve a prepared statement, and then execute a query using a prepared statement.
// If statement caching is disabled, this method delegates the call to the *sql.Tx ExecContext method.
func (t *Tx) ExecContext(ctx context.Context, query string, args ...interface{}) (sql.Result, error) {
if !t.db.statementsCaching {
return t.Tx.ExecContext(ctx, query, args...)
}
prepStmt, err := t.PrepareContext(ctx, query)
if err != nil {
return nil, err
}
return prepStmt.ExecContext(ctx, args...)
}
// Query delegates call to QueryContext using context.Background() as parameter.
func (t *Tx) Query(query string, args ...interface{}) (*sql.Rows, error) {
return t.QueryContext(context.Background(), query, args...)
}
// QueryContext executes a query that returns rows. If statement caching is enabled, QueryContext will
// first call PrepareContext to retrieve a prepared statement, and then execute a query using a prepared statement.
// If statement caching is disabled, this method delegates the call to the *sql.Tx QueryContext method.
func (t *Tx) QueryContext(ctx context.Context, query string, args ...interface{}) (*sql.Rows, error) {
if !t.db.statementsCaching {
return t.Tx.QueryContext(ctx, query, args...)
}
prepStmt, err := t.PrepareContext(ctx, query)
if err != nil {
return nil, err
}
return prepStmt.Query(args...)
}
// Prepare delegates call to PrepareContext using context.Background as a parameter.
func (t *Tx) Prepare(query string) (*sql.Stmt, error) {
return t.PrepareContext(context.Background(), query)
}
// PrepareContext returns database prepared statement for a query. When statement caching is enabled, it returns a cached
// prepared statement if available; otherwise, it creates a new prepared statement and adds it to the cache.
// Invoking this method directly is unnecessary, as wrapper methods like Exec/ExecContext and Query/QueryContext
// will call PrepareContext before executing a query on it.
// If statement caching is disabled, this method delegates the call to the *sql.Tx PrepareContext method.
//
// There's no need to manually close the returned statement; it operates within the transaction scope and will be closed
// automatically upon the completion of the transaction, whether it's committed or rolled back.
func (t *Tx) PrepareContext(ctx context.Context, query string) (*sql.Stmt, error) {
if !t.db.statementsCaching {
return t.PrepareContext(ctx, query)
}
prepStmt, ok := t.statements[query]
if ok {
return prepStmt, nil
}
dbPrepStmt, err := t.db.PrepareContext(ctx, query)
if err != nil {
return nil, fmt.Errorf("failed to prepare statement, %w", err)
}
prepStmt = t.Tx.StmtContext(ctx, dbPrepStmt)
t.statements[query] = prepStmt
return prepStmt, nil
}

View file

@ -3,10 +3,10 @@ package testutils
import ( import (
"bytes" "bytes"
"context" "context"
"database/sql"
"encoding/json" "encoding/json"
"fmt" "fmt"
"github.com/go-jet/jet/v2/internal/jet" "github.com/go-jet/jet/v2/internal/jet"
jet2 "github.com/go-jet/jet/v2/internal/jet/db"
"github.com/go-jet/jet/v2/internal/utils/throw" "github.com/go-jet/jet/v2/internal/utils/throw"
"github.com/go-jet/jet/v2/qrm" "github.com/go-jet/jet/v2/qrm"
"github.com/google/uuid" "github.com/google/uuid"
@ -28,7 +28,7 @@ var UnixTimeComparer = cmp.Comparer(func(t1, t2 time.Time) bool {
}) })
// AssertExecAndRollback will execute and rollback statement in sql transaction // AssertExecAndRollback will execute and rollback statement in sql transaction
func AssertExecAndRollback(t *testing.T, stmt jet.Statement, db *sql.DB, rowsAffected ...int64) { func AssertExecAndRollback(t *testing.T, stmt jet.Statement, db *jet2.DB, rowsAffected ...int64) {
tx, err := db.Begin() tx, err := db.Begin()
require.NoError(t, err) require.NoError(t, err)
defer func() { defer func() {
@ -53,7 +53,7 @@ func AssertExec(t *testing.T, stmt jet.Statement, db qrm.DB, rowsAffected ...int
} }
// ExecuteInTxAndRollback will execute function in sql transaction and then rollback transaction // ExecuteInTxAndRollback will execute function in sql transaction and then rollback transaction
func ExecuteInTxAndRollback(t *testing.T, db *sql.DB, f func(tx *sql.Tx)) { func ExecuteInTxAndRollback(t *testing.T, db *jet2.DB, f func(tx qrm.DB)) {
tx, err := db.Begin() tx, err := db.Begin()
require.NoError(t, err) require.NoError(t, err)
defer func() { defer func() {
@ -133,7 +133,7 @@ func AssertJSONFile(t *testing.T, data interface{}, testRelativePath string) {
} }
// AssertStatementSql check if statement Sql() is the same as expectedQuery and expectedArgs // AssertStatementSql check if statement Sql() is the same as expectedQuery and expectedArgs
func AssertStatementSql(t *testing.T, query jet.Statement, expectedQuery string, expectedArgs ...interface{}) { func AssertStatementSql(t *testing.T, query jet.PrintableStatement, expectedQuery string, expectedArgs ...interface{}) {
queryStr, args := query.Sql() queryStr, args := query.Sql()
assertQueryString(t, queryStr, expectedQuery) assertQueryString(t, queryStr, expectedQuery)
@ -154,7 +154,7 @@ func AssertStatementSqlErr(t *testing.T, stmt jet.Statement, errorStr string) {
} }
// AssertDebugStatementSql check if statement Sql() is the same as expectedQuery // AssertDebugStatementSql check if statement Sql() is the same as expectedQuery
func AssertDebugStatementSql(t *testing.T, query jet.Statement, expectedQuery string, expectedArgs ...interface{}) { func AssertDebugStatementSql(t *testing.T, query jet.PrintableStatement, expectedQuery string, expectedArgs ...interface{}) {
_, args := query.Sql() _, args := query.Sql()
if len(expectedArgs) > 0 { if len(expectedArgs) > 0 {

View file

@ -1,8 +1,20 @@
package mysql package mysql
import "github.com/go-jet/jet/v2/internal/jet" import (
"github.com/go-jet/jet/v2/internal/jet"
"github.com/go-jet/jet/v2/internal/jet/db"
)
// RawStatement creates new sql statements from raw query and optional map of named arguments // RawStatement creates new sql statements from raw query and optional map of named arguments
func RawStatement(rawQuery string, namedArguments ...RawArgs) Statement { func RawStatement(rawQuery string, namedArguments ...RawArgs) Statement {
return jet.RawStatement(Dialect, rawQuery, namedArguments...) return jet.RawStatement(Dialect, rawQuery, namedArguments...)
} }
// DB is a wrapper around sql.DB, adding prepared statement caching capability.
type DB = db.DB
// NewDB creates new DB wrapper with statements caching disabled
var NewDB = db.NewDB
// Tx is a wrapper around *sql.Tx, adding prepared statement caching capability.
type Tx = db.Tx

View file

@ -1,8 +1,20 @@
package postgres package postgres
import "github.com/go-jet/jet/v2/internal/jet" import (
"github.com/go-jet/jet/v2/internal/jet"
"github.com/go-jet/jet/v2/internal/jet/db"
)
// RawStatement creates new sql statements from raw query and optional map of named arguments // RawStatement creates new sql statements from raw query and optional map of named arguments
func RawStatement(rawQuery string, namedArguments ...RawArgs) Statement { func RawStatement(rawQuery string, namedArguments ...RawArgs) Statement {
return jet.RawStatement(Dialect, rawQuery, namedArguments...) return jet.RawStatement(Dialect, rawQuery, namedArguments...)
} }
// DB is a wrapper around sql.DB, adding prepared statement caching capability.
type DB = db.DB
// NewDB creates new DB wrapper with statements caching disabled
var NewDB = db.NewDB
// Tx is a wrapper around *sql.Tx, adding prepared statement caching capability.
type Tx = db.Tx

View file

@ -1,8 +1,20 @@
package sqlite package sqlite
import "github.com/go-jet/jet/v2/internal/jet" import (
"github.com/go-jet/jet/v2/internal/jet"
"github.com/go-jet/jet/v2/internal/jet/db"
)
// RawStatement creates new sql statements from raw query and optional map of named arguments // RawStatement creates new sql statements from raw query and optional map of named arguments
func RawStatement(rawQuery string, namedArguments ...RawArgs) Statement { func RawStatement(rawQuery string, namedArguments ...RawArgs) Statement {
return jet.RawStatement(Dialect, rawQuery, namedArguments...) return jet.RawStatement(Dialect, rawQuery, namedArguments...)
} }
// DB is a wrapper around sql.DB, adding prepared statement caching capability.
type DB = db.DB
// NewDB creates new DB wrapper with statements caching disabled
var NewDB = db.NewDB
// Tx is a wrapper around *sql.Tx, adding prepared statement caching capability.
type Tx = db.Tx

View file

@ -2,9 +2,9 @@ package mysql
import ( import (
"context" "context"
"database/sql"
"github.com/go-jet/jet/v2/internal/testutils" "github.com/go-jet/jet/v2/internal/testutils"
. "github.com/go-jet/jet/v2/mysql" . "github.com/go-jet/jet/v2/mysql"
"github.com/go-jet/jet/v2/qrm"
"github.com/go-jet/jet/v2/tests/.gentestdata/mysql/dvds/table" "github.com/go-jet/jet/v2/tests/.gentestdata/mysql/dvds/table"
"github.com/go-jet/jet/v2/tests/.gentestdata/mysql/test_sample/model" "github.com/go-jet/jet/v2/tests/.gentestdata/mysql/test_sample/model"
. "github.com/go-jet/jet/v2/tests/.gentestdata/mysql/test_sample/table" . "github.com/go-jet/jet/v2/tests/.gentestdata/mysql/test_sample/table"
@ -113,7 +113,7 @@ DELETE /*+ QB_NAME(deleteIns) MRR(link) */ FROM test_sample.link
WHERE link.name IN ('Gmail', 'Outlook'); WHERE link.name IN ('Gmail', 'Outlook');
`) `)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
_, err := stmt.Exec(tx) _, err := stmt.Exec(tx)
require.NoError(t, err) require.NoError(t, err)
}) })

View file

@ -2,9 +2,9 @@ package mysql
import ( import (
"context" "context"
"database/sql"
"github.com/go-jet/jet/v2/internal/testutils" "github.com/go-jet/jet/v2/internal/testutils"
. "github.com/go-jet/jet/v2/mysql" . "github.com/go-jet/jet/v2/mysql"
"github.com/go-jet/jet/v2/qrm"
"github.com/go-jet/jet/v2/tests/.gentestdata/mysql/test_sample/model" "github.com/go-jet/jet/v2/tests/.gentestdata/mysql/test_sample/model"
. "github.com/go-jet/jet/v2/tests/.gentestdata/mysql/test_sample/table" . "github.com/go-jet/jet/v2/tests/.gentestdata/mysql/test_sample/table"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
@ -29,7 +29,7 @@ VALUES (100, 'http://www.postgresqltutorial.com', 'PostgreSQL Tutorial', DEFAULT
101, "http://www.google.com", "Google", 101, "http://www.google.com", "Google",
102, "http://www.yahoo.com", "Yahoo", nil) 102, "http://www.yahoo.com", "Yahoo", nil)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
_, err := insertQuery.Exec(tx) _, err := insertQuery.Exec(tx)
require.NoError(t, err) require.NoError(t, err)
requireLogged(t, insertQuery) requireLogged(t, insertQuery)
@ -74,7 +74,7 @@ VALUES (100, 'http://www.postgresqltutorial.com', 'PostgreSQL Tutorial', DEFAULT
`, `,
100, "http://www.postgresqltutorial.com", "PostgreSQL Tutorial") 100, "http://www.postgresqltutorial.com", "PostgreSQL Tutorial")
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
_, err := stmt.Exec(tx) _, err := stmt.Exec(tx)
require.NoError(t, err) require.NoError(t, err)
requireLogged(t, stmt) requireLogged(t, stmt)
@ -108,7 +108,7 @@ VALUES ('http://www.duckduckgo.com', 'Duck Duck go');
`, `,
"http://www.duckduckgo.com", "Duck Duck go") "http://www.duckduckgo.com", "Duck Duck go")
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
_, err := query.Exec(tx) _, err := query.Exec(tx)
require.NoError(t, err) require.NoError(t, err)
}) })
@ -130,7 +130,7 @@ INSERT INTO test_sample.link
VALUES (1000, 'http://www.duckduckgo.com', 'Duck Duck go', NULL); VALUES (1000, 'http://www.duckduckgo.com', 'Duck Duck go', NULL);
`, int32(1000), "http://www.duckduckgo.com", "Duck Duck go", nil) `, int32(1000), "http://www.duckduckgo.com", "Duck Duck go", nil)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
_, err := query.Exec(tx) _, err := query.Exec(tx)
require.NoError(t, err) require.NoError(t, err)
}) })
@ -166,7 +166,7 @@ VALUES ('http://www.postgresqltutorial.com', 'PostgreSQL Tutorial'),
"http://www.google.com", "Google", "http://www.google.com", "Google",
"http://www.yahoo.com", "Yahoo") "http://www.yahoo.com", "Yahoo")
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
_, err := query.Exec(tx) _, err := query.Exec(tx)
require.NoError(t, err) require.NoError(t, err)
}) })
@ -201,7 +201,7 @@ VALUES ('http://www.postgresqltutorial.com', 'PostgreSQL Tutorial', DEFAULT),
"http://www.google.com", "Google", nil, "http://www.google.com", "Google", nil,
"http://www.yahoo.com", "Yahoo", nil) "http://www.yahoo.com", "Yahoo", nil)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
_, err := stmt.Exec(tx) _, err := stmt.Exec(tx)
require.NoError(t, err) require.NoError(t, err)
}) })
@ -225,7 +225,7 @@ INSERT INTO test_sample.link (url, name) (
); );
`, int64(1)) `, int64(1))
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
_, err := query.Exec(tx) _, err := query.Exec(tx)
require.NoError(t, err) require.NoError(t, err)
@ -261,7 +261,7 @@ ON DUPLICATE KEY UPDATE id = (link.id + ?),
randId, "http://www.postgresqltutorial.com", "PostgreSQL Tutorial", randId, "http://www.postgresqltutorial.com", "PostgreSQL Tutorial",
int64(11), "PostgreSQL Tutorial 2") int64(11), "PostgreSQL Tutorial 2")
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
_, err := stmt.Exec(tx) _, err := stmt.Exec(tx)
require.NoError(t, err) require.NoError(t, err)
@ -320,7 +320,7 @@ ON DUPLICATE KEY UPDATE id = (link.id + ?),
description = new.description; description = new.description;
`) `)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
_, err := stmt.Exec(tx) _, err := stmt.Exec(tx)
require.NoError(t, err) require.NoError(t, err)
@ -352,9 +352,12 @@ func TestInsertWithQueryContext(t *testing.T) {
time.Sleep(10 * time.Millisecond) time.Sleep(10 * time.Millisecond)
var dest []model.Link var dest []model.Link
err := stmt.QueryContext(ctx, db, &dest)
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
err := stmt.QueryContext(ctx, tx, &dest)
require.Error(t, err, "context deadline exceeded") require.Error(t, err, "context deadline exceeded")
})
} }
func TestInsertWithExecContext(t *testing.T) { func TestInsertWithExecContext(t *testing.T) {
@ -366,9 +369,10 @@ func TestInsertWithExecContext(t *testing.T) {
time.Sleep(10 * time.Millisecond) time.Sleep(10 * time.Millisecond)
_, err := stmt.ExecContext(ctx, db) testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
_, err := stmt.ExecContext(ctx, tx)
require.Error(t, err, "context deadline exceeded") require.Error(t, err, "context deadline exceeded")
})
} }
func TestInsertOptimizerHints(t *testing.T) { func TestInsertOptimizerHints(t *testing.T) {
@ -385,7 +389,7 @@ INSERT /*+ QB_NAME(qbIns) NO_ICP(link) */ INTO test_sample.link (url, name, desc
VALUES ('http://www.google.com', 'Google', NULL); VALUES ('http://www.google.com', 'Google', NULL);
`) `)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
_, err := stmt.Exec(tx) _, err := stmt.Exec(tx)
require.NoError(t, err) require.NoError(t, err)
}) })

View file

@ -14,10 +14,14 @@ func TestLockRead(t *testing.T) {
testutils.AssertStatementSql(t, query, ` testutils.AssertStatementSql(t, query, `
LOCK TABLES dvds.customer READ; LOCK TABLES dvds.customer READ;
`) `)
tx, err := db.DB.Begin() // can't prepare LOCK statement
_, err := query.Exec(db)
require.NoError(t, err) require.NoError(t, err)
requireLogged(t, query) defer func() {
err := tx.Rollback()
require.NoError(t, err)
}()
testutils.AssertExec(t, query, tx)
} }
func TestLockWrite(t *testing.T) { func TestLockWrite(t *testing.T) {
@ -27,9 +31,14 @@ func TestLockWrite(t *testing.T) {
LOCK TABLES dvds.customer WRITE; LOCK TABLES dvds.customer WRITE;
`) `)
_, err := query.Exec(db) tx, err := db.DB.Begin() // can't prepare LOCK statement
require.NoError(t, err) require.NoError(t, err)
requireLogged(t, query) defer func() {
err := tx.Rollback()
require.NoError(t, err)
}()
testutils.AssertExec(t, query, tx)
} }
func TestUnlockTables(t *testing.T) { func TestUnlockTables(t *testing.T) {
@ -39,7 +48,12 @@ func TestUnlockTables(t *testing.T) {
UNLOCK TABLES; UNLOCK TABLES;
`) `)
_, err := query.Exec(db) tx, err := db.DB.Begin() // can't prepare LOCK statement
require.NoError(t, err) require.NoError(t, err)
requireLogged(t, query) defer func() {
err := tx.Rollback()
require.NoError(t, err)
}()
testutils.AssertExec(t, query, tx)
} }

View file

@ -18,7 +18,7 @@ import (
"testing" "testing"
) )
var db *sql.DB var db *jetmysql.DB
var source string var source string
@ -37,16 +37,21 @@ func TestMain(m *testing.M) {
defer profile.Start().Stop() defer profile.Start().Stop()
var err error var err error
db, err = sql.Open("mysql", dbconfig.MySQLConnectionString(sourceIsMariaDB(), "")) sqlDB, err := sql.Open("mysql", dbconfig.MySQLConnectionString(sourceIsMariaDB(), ""))
if err != nil { if err != nil {
panic("Failed to connect to test db" + err.Error()) panic("Failed to connect to test db" + err.Error())
} }
db = jetmysql.NewDB(sqlDB).WithStatementsCaching(true)
defer db.Close() defer db.Close()
for i := 0; i < 2; i++ {
ret := m.Run() ret := m.Run()
if ret != 0 {
os.Exit(ret) os.Exit(ret)
} }
}
}
var loggedSQL string var loggedSQL string
var loggedSQLArgs []interface{} var loggedSQLArgs []interface{}

View file

@ -2,8 +2,8 @@ package mysql
import ( import (
"context" "context"
"database/sql"
"github.com/go-jet/jet/v2/postgres" "github.com/go-jet/jet/v2/postgres"
"github.com/go-jet/jet/v2/qrm"
"strings" "strings"
"testing" "testing"
"time" "time"
@ -868,11 +868,13 @@ func TestRowLock(t *testing.T) {
expectedSQL := ` expectedSQL := `
SELECT * SELECT *
FROM dvds.address FROM dvds.address
ORDER BY address.address_id
LIMIT 3 LIMIT 3
OFFSET 1 OFFSET 1
FOR` FOR`
query := Address. query := SELECT(STAR).
SELECT(STAR). FROM(Address).
ORDER_BY(Address.AddressID).
LIMIT(3). LIMIT(3).
OFFSET(1) OFFSET(1)
@ -881,28 +883,14 @@ FOR`
expectedQuery := expectedSQL + " " + lockTypeStr + ";\n" expectedQuery := expectedSQL + " " + lockTypeStr + ";\n"
testutils.AssertDebugStatementSql(t, query, expectedQuery, int64(3), int64(1)) testutils.AssertDebugStatementSql(t, query, expectedQuery, int64(3), int64(1))
testutils.AssertExecAndRollback(t, query, db)
tx, _ := db.Begin()
_, err := query.Exec(tx)
require.NoError(t, err)
err = tx.Rollback()
require.NoError(t, err)
} }
for lockType, lockTypeStr := range getRowLockTestData() { for lockType, lockTypeStr := range getRowLockTestData() {
query.FOR(lockType.NOWAIT()) query.FOR(lockType.NOWAIT())
testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+" NOWAIT;\n", int64(3), int64(1)) testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+" NOWAIT;\n", int64(3), int64(1))
testutils.AssertExecAndRollback(t, query, db)
tx, _ := db.Begin()
_, err := query.Exec(tx)
require.NoError(t, err)
err = tx.Rollback()
require.NoError(t, err)
} }
if sourceIsMariaDB() { if sourceIsMariaDB() {
@ -913,14 +901,7 @@ FOR`
query.FOR(lockType.SKIP_LOCKED()) query.FOR(lockType.SKIP_LOCKED())
testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+" SKIP LOCKED;\n", int64(3), int64(1)) testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+" SKIP LOCKED;\n", int64(3), int64(1))
testutils.AssertExecAndRollback(t, query, db)
tx, _ := db.Begin()
_, err := query.Exec(tx)
require.NoError(t, err)
err = tx.Rollback()
require.NoError(t, err)
} }
} }
@ -956,7 +937,7 @@ LIMIT 1
FOR UPDATE OF film, actor NOWAIT; FOR UPDATE OF film, actor NOWAIT;
`) `)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
var dest []struct { var dest []struct {
model.Film model.Film
CategoryID int CategoryID int
@ -993,7 +974,7 @@ LIMIT 1
FOR UPDATE OF ''myFilm''; FOR UPDATE OF ''myFilm'';
`, "''", "`")) `, "''", "`"))
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
var dest []struct { var dest []struct {
model.Film `alias:"myFilm.*"` model.Film `alias:"myFilm.*"`
} }

View file

@ -2,7 +2,7 @@ package mysql
import ( import (
"context" "context"
"database/sql" "github.com/go-jet/jet/v2/qrm"
"testing" "testing"
"time" "time"
@ -28,7 +28,7 @@ WHERE link.name = 'Bing';
WHERE(Link.Name.EQ(String("Bing"))) WHERE(Link.Name.EQ(String("Bing")))
testutils.AssertDebugStatementSql(t, query, expectedSQL, "Bong", "http://bong.com", "Bing") testutils.AssertDebugStatementSql(t, query, expectedSQL, "Bong", "http://bong.com", "Bing")
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
testutils.AssertExec(t, query, tx) testutils.AssertExec(t, query, tx)
requireLogged(t, query) requireLogged(t, query)
@ -59,7 +59,7 @@ WHERE link.name = 'Bing';
WHERE(Link.Name.EQ(String("Bing"))) WHERE(Link.Name.EQ(String("Bing")))
testutils.AssertDebugStatementSql(t, stmt, expectedSQL, "Bong", "http://bong.com", "Bing") testutils.AssertDebugStatementSql(t, stmt, expectedSQL, "Bong", "http://bong.com", "Bing")
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
testutils.AssertExec(t, stmt, tx) testutils.AssertExec(t, stmt, tx)
requireLogged(t, stmt) requireLogged(t, stmt)
@ -280,7 +280,7 @@ SET id = 501,
WHERE link.name = 'Bing'; WHERE link.name = 'Bing';
`) `)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
_, err := stmt.Exec(tx) _, err := stmt.Exec(tx)
require.NoError(t, err) require.NoError(t, err)
}) })

View file

@ -1,7 +1,7 @@
package postgres package postgres
import ( import (
"database/sql" "github.com/go-jet/jet/v2/qrm"
"testing" "testing"
"time" "time"
@ -71,7 +71,7 @@ func TestAllTypesInsertModel(t *testing.T) {
MODEL(&allTypesRow1). MODEL(&allTypesRow1).
RETURNING(AllTypes.AllColumns) RETURNING(AllTypes.AllColumns)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
var dest []model.AllTypes var dest []model.AllTypes
err := query.Query(tx, &dest) err := query.Query(tx, &dest)
require.NoError(t, err) require.NoError(t, err)
@ -94,7 +94,7 @@ func TestAllTypesInsertQuery(t *testing.T) {
). ).
RETURNING(AllTypesAllColumns) RETURNING(AllTypesAllColumns)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
var dest []model.AllTypes var dest []model.AllTypes
err := query.Query(tx, &dest) err := query.Query(tx, &dest)
@ -1293,7 +1293,7 @@ WHERE all_types.small_int = 14
RETURNING all_types.json AS "all_types.json"; RETURNING all_types.json AS "all_types.json";
`) `)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
var res model.AllTypes var res model.AllTypes
err := stmt.Query(tx, &res) err := stmt.Query(tx, &res)

View file

@ -12,7 +12,7 @@ import (
"time" "time"
) )
func TestSelect(t *testing.T) { func TestSelectAlbum(t *testing.T) {
stmt := SELECT(Album.AllColumns). stmt := SELECT(Album.AllColumns).
FROM(Album). FROM(Album).
ORDER_BY(Album.AlbumId.ASC()) ORDER_BY(Album.AlbumId.ASC())
@ -782,9 +782,11 @@ func TestQueryWithContext(t *testing.T) {
return // context cancellation doesn't work for pq driver return // context cancellation doesn't work for pq driver
} }
ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) ctx, cancel := context.WithTimeout(context.Background(), 1*time.Microsecond)
defer cancel() defer cancel()
time.Sleep(1 * time.Millisecond)
var dest []model.Album var dest []model.Album
err := Album. err := Album.

View file

@ -2,9 +2,9 @@ package postgres
import ( import (
"context" "context"
"database/sql"
"github.com/go-jet/jet/v2/internal/testutils" "github.com/go-jet/jet/v2/internal/testutils"
. "github.com/go-jet/jet/v2/postgres" . "github.com/go-jet/jet/v2/postgres"
"github.com/go-jet/jet/v2/qrm"
model2 "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/model" model2 "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/model"
"github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/table" "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/table"
"github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/test_sample/model" "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/test_sample/model"
@ -43,7 +43,7 @@ RETURNING link.id AS "link.id",
link.description AS "link.description"; link.description AS "link.description";
`, "Gmail", "Outlook") `, "Gmail", "Outlook")
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
var dest []model.Link var dest []model.Link
err := deleteStmt.Query(tx, &dest) err := deleteStmt.Query(tx, &dest)
@ -67,7 +67,7 @@ func TestDeleteQueryContext(t *testing.T) {
time.Sleep(10 * time.Millisecond) time.Sleep(10 * time.Millisecond)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
dest := []model.Link{} dest := []model.Link{}
err := deleteStmt.QueryContext(ctx, tx, &dest) err := deleteStmt.QueryContext(ctx, tx, &dest)
@ -88,7 +88,7 @@ func TestDeleteExecContext(t *testing.T) {
time.Sleep(10 * time.Millisecond) time.Sleep(10 * time.Millisecond)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
_, err := deleteStmt.ExecContext(ctx, tx) _, err := deleteStmt.ExecContext(ctx, tx)
require.Error(t, err, "context deadline exceeded") require.Error(t, err, "context deadline exceeded")
@ -140,7 +140,7 @@ RETURNING rental.rental_id AS "rental.rental_id",
store.last_update AS "store.last_update"; store.last_update AS "store.last_update";
`) `)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
var dest []struct { var dest []struct {
Rental model2.Rental Rental model2.Rental
Store model2.Store Store model2.Store

View file

@ -2,9 +2,9 @@ package postgres
import ( import (
"context" "context"
"database/sql"
"github.com/go-jet/jet/v2/internal/testutils" "github.com/go-jet/jet/v2/internal/testutils"
. "github.com/go-jet/jet/v2/postgres" . "github.com/go-jet/jet/v2/postgres"
"github.com/go-jet/jet/v2/qrm"
"github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/test_sample/model" "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/test_sample/model"
. "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/test_sample/table" . "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/test_sample/table"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
@ -34,7 +34,7 @@ RETURNING link.id AS "link.id",
101, "http://www.google.com", "Google", 101, "http://www.google.com", "Google",
102, "http://www.yahoo.com", "Yahoo", nil) 102, "http://www.yahoo.com", "Yahoo", nil)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
var insertedLinks []model.Link var insertedLinks []model.Link
err := insertQuery.Query(tx, &insertedLinks) err := insertQuery.Query(tx, &insertedLinks)
@ -335,7 +335,7 @@ RETURNING link.id AS "link.id",
link.description AS "link.description"; link.description AS "link.description";
`, int64(0)) `, int64(0))
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
var dest []model.Link var dest []model.Link
err := query.Query(tx, &dest) err := query.Query(tx, &dest)
@ -362,7 +362,7 @@ func TestInsertWithQueryContext(t *testing.T) {
time.Sleep(10 * time.Millisecond) time.Sleep(10 * time.Millisecond)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
dest := []model.Link{} dest := []model.Link{}
err := stmt.QueryContext(ctx, tx, &dest) err := stmt.QueryContext(ctx, tx, &dest)
@ -379,7 +379,7 @@ func TestInsertWithExecContext(t *testing.T) {
time.Sleep(10 * time.Millisecond) time.Sleep(10 * time.Millisecond)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
testutils.AssertExecContextErr(ctx, t, stmt, tx, "context deadline exceeded") testutils.AssertExecContextErr(ctx, t, stmt, tx, "context deadline exceeded")
}) })
} }

View file

@ -32,34 +32,14 @@ LOCK TABLE dvds.address IN`
query := Address.LOCK().IN(lockMode) query := Address.LOCK().IN(lockMode)
testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+string(lockMode)+" MODE;\n") testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+string(lockMode)+" MODE;\n")
testutils.AssertExecAndRollback(t, query, db)
tx, _ := db.Begin()
_, err := query.Exec(tx)
require.NoError(t, err)
err = tx.Rollback()
require.NoError(t, err)
requireLogged(t, query)
} }
for _, lockMode := range testData { for _, lockMode := range testData {
query := Address.LOCK().IN(lockMode).NOWAIT() query := Address.LOCK().IN(lockMode).NOWAIT()
testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+string(lockMode)+" MODE NOWAIT;\n") testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+string(lockMode)+" MODE NOWAIT;\n")
testutils.AssertExecAndRollback(t, query, db)
tx, _ := db.Begin()
_, err := query.Exec(tx)
require.NoError(t, err)
err = tx.Rollback()
require.NoError(t, err)
requireLogged(t, query)
} }
} }

View file

@ -22,7 +22,7 @@ import (
_ "github.com/jackc/pgx/v4/stdlib" _ "github.com/jackc/pgx/v4/stdlib"
) )
var db *sql.DB var db *postgres.DB
var testRoot string var testRoot string
var source string var source string
@ -60,19 +60,26 @@ func TestMain(m *testing.M) {
connectionString = dbconfig.CockroachConnectString connectionString = dbconfig.CockroachConnectString
} }
var err error sqlDB, err := sql.Open(driverName, connectionString)
db, err = sql.Open(driverName, connectionString)
if err != nil { if err != nil {
fmt.Println(err.Error()) fmt.Println(err.Error())
panic("Failed to connect to test db") panic("Failed to connect to test db")
} }
db = postgres.NewDB(sqlDB).WithStatementsCaching(true)
defer db.Close() defer db.Close()
for i := 0; i < 2; i++ {
ret := m.Run() ret := m.Run()
if ret != 0 { if ret != 0 {
os.Exit(ret) os.Exit(ret)
} }
}
err = db.Clear()
if err != nil {
os.Exit(-2)
}
}() }()
} }
} }

View file

@ -2,6 +2,7 @@ package postgres
import ( import (
"github.com/go-jet/jet/v2/internal/testutils" "github.com/go-jet/jet/v2/internal/testutils"
. "github.com/go-jet/jet/v2/postgres"
"github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/northwind/model" "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/northwind/model"
. "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/northwind/table" . "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/northwind/table"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
@ -10,7 +11,17 @@ import (
func TestNorthwindJoinEverything(t *testing.T) { func TestNorthwindJoinEverything(t *testing.T) {
stmt := Customers. stmt := SELECT(
Customers.AllColumns,
CustomerDemographics.AllColumns,
Orders.AllColumns,
Shippers.AllColumns,
OrderDetails.AllColumns,
Products.AllColumns,
Categories.AllColumns,
Suppliers.AllColumns,
).FROM(
Customers.
LEFT_JOIN(CustomerCustomerDemo, Customers.CustomerID.EQ(CustomerCustomerDemo.CustomerID)). LEFT_JOIN(CustomerCustomerDemo, Customers.CustomerID.EQ(CustomerCustomerDemo.CustomerID)).
LEFT_JOIN(CustomerDemographics, CustomerCustomerDemo.CustomerTypeID.EQ(CustomerDemographics.CustomerTypeID)). LEFT_JOIN(CustomerDemographics, CustomerCustomerDemo.CustomerTypeID.EQ(CustomerDemographics.CustomerTypeID)).
LEFT_JOIN(Orders, Orders.CustomerID.EQ(Customers.CustomerID)). LEFT_JOIN(Orders, Orders.CustomerID.EQ(Customers.CustomerID)).
@ -22,18 +33,12 @@ func TestNorthwindJoinEverything(t *testing.T) {
LEFT_JOIN(Employees, Orders.EmployeeID.EQ(Employees.EmployeeID)). LEFT_JOIN(Employees, Orders.EmployeeID.EQ(Employees.EmployeeID)).
LEFT_JOIN(EmployeeTerritories, EmployeeTerritories.EmployeeID.EQ(Employees.EmployeeID)). LEFT_JOIN(EmployeeTerritories, EmployeeTerritories.EmployeeID.EQ(Employees.EmployeeID)).
LEFT_JOIN(Territories, EmployeeTerritories.TerritoryID.EQ(Territories.TerritoryID)). LEFT_JOIN(Territories, EmployeeTerritories.TerritoryID.EQ(Territories.TerritoryID)).
LEFT_JOIN(Region, Territories.RegionID.EQ(Region.RegionID)). LEFT_JOIN(Region, Territories.RegionID.EQ(Region.RegionID)),
SELECT( ).ORDER_BY(
Customers.AllColumns, Customers.CustomerID,
CustomerDemographics.AllColumns, Orders.OrderID,
Orders.AllColumns, Products.ProductID,
Shippers.AllColumns, )
OrderDetails.AllColumns,
Products.AllColumns,
Categories.AllColumns,
Suppliers.AllColumns,
).
ORDER_BY(Customers.CustomerID, Orders.OrderID, Products.ProductID)
var dest []struct { var dest []struct {
model.Customers model.Customers

View file

@ -1,8 +1,8 @@
package postgres package postgres
import ( import (
"database/sql"
"github.com/go-jet/jet/v2/internal/testutils" "github.com/go-jet/jet/v2/internal/testutils"
"github.com/go-jet/jet/v2/qrm"
"github.com/google/go-cmp/cmp" "github.com/google/go-cmp/cmp"
"github.com/jackc/pgtype" "github.com/jackc/pgtype"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
@ -294,7 +294,7 @@ RETURNING sample_ranges.date_range AS "sample_ranges.date_range",
` `
testutils.AssertDebugStatementSql(t, insertQuery, expectedQuery) testutils.AssertDebugStatementSql(t, insertQuery, expectedQuery)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
var dest []model.SampleRanges var dest []model.SampleRanges
err := insertQuery.Query(tx, &dest) err := insertQuery.Query(tx, &dest)
require.NoError(t, err) require.NoError(t, err)
@ -324,7 +324,7 @@ RETURNING sample_ranges.date_range AS "sample_ranges.date_range",
sample_ranges.num_range AS "sample_ranges.num_range"; sample_ranges.num_range AS "sample_ranges.num_range";
`) `)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
var dest []model.SampleRanges var dest []model.SampleRanges
err := stmt.Query(tx, &dest) err := stmt.Query(tx, &dest)
@ -351,7 +351,7 @@ SET int4_range = int4range(-12::integer, 78::integer),
WHERE LOWER(sample_ranges.timestampz_range) > NOW(); WHERE LOWER(sample_ranges.timestampz_range) > NOW();
`) `)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
testutils.AssertExec(t, stmt, tx, 1) testutils.AssertExec(t, stmt, tx, 1)
}) })
}) })

View file

@ -2,7 +2,7 @@ package postgres
import ( import (
"context" "context"
"database/sql" "github.com/go-jet/jet/v2/qrm"
"testing" "testing"
"time" "time"
@ -128,7 +128,7 @@ RETURNING link.id AS "link.id",
link.description AS "link.description"; link.description AS "link.description";
`) `)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
var links []model2.Link var links []model2.Link
err := stmt.Query(tx, &links) err := stmt.Query(tx, &links)
require.NoError(t, err) require.NoError(t, err)

View file

@ -1,7 +1,7 @@
package postgres package postgres
import ( import (
"database/sql" "github.com/go-jet/jet/v2/qrm"
"github.com/google/uuid" "github.com/google/uuid"
"testing" "testing"
@ -512,7 +512,7 @@ func TestMutableColumnsExcludeGeneratedColumn(t *testing.T) {
}) })
t.Run("should insert without generated columns", func(t *testing.T) { t.Run("should insert without generated columns", func(t *testing.T) {
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
insertQuery := People.INSERT( insertQuery := People.INSERT(
People.MutableColumns, People.MutableColumns,
).MODEL( ).MODEL(

View file

@ -109,7 +109,7 @@ ORDER BY rental.staff_id ASC, rental.customer_id ASC, rental.rental_id ASC;
`) `)
} }
func TestClassicSelect(t *testing.T) { func TestSelectClassic(t *testing.T) {
expectedSQL := ` expectedSQL := `
SELECT payment.payment_id AS "payment.payment_id", SELECT payment.payment_id AS "payment.payment_id",
payment.customer_id AS "payment.customer_id", payment.customer_id AS "payment.customer_id",
@ -145,7 +145,7 @@ LIMIT 30;
testutils.AssertDebugStatementSql(t, query, expectedSQL, int64(30)) testutils.AssertDebugStatementSql(t, query, expectedSQL, int64(30))
dest := []model.Payment{} var dest []model.Payment
err := query.Query(db, &dest) err := query.Query(db, &dest)
@ -231,12 +231,12 @@ LIMIT 12;
testutils.AssertDebugStatementSql(t, query, expectedSQL, int64(1), int64(1), int64(10), int64(1), int64(2), int64(1), int64(12)) testutils.AssertDebugStatementSql(t, query, expectedSQL, int64(1), int64(1), int64(10), int64(1), int64(2), int64(1), int64(12))
dest := []struct{}{} var dest []struct{}
err := query.Query(db, &dest) err := query.Query(db, &dest)
require.NoError(t, err) require.NoError(t, err)
} }
func TestFetchFirst(t *testing.T) { func TestSelectFetchFirst(t *testing.T) {
t.Run("rows only", func(t *testing.T) { t.Run("rows only", func(t *testing.T) {
stmt := SELECT(Actor.AllColumns). stmt := SELECT(Actor.AllColumns).
@ -320,7 +320,7 @@ FETCH FIRST (
}) })
} }
func TestOffsetExpression(t *testing.T) { func TestSelectOffsetExpression(t *testing.T) {
stmt := SELECT(Actor.AllColumns). stmt := SELECT(Actor.AllColumns).
FROM(Actor). FROM(Actor).
@ -352,7 +352,7 @@ OFFSET (
require.Equal(t, dest[0].ActorID, int32(3)) require.Equal(t, dest[0].ActorID, int32(3))
} }
func TestJoinQueryStruct(t *testing.T) { func TestSelectJoinQueryStruct(t *testing.T) {
expectedSQL := ` expectedSQL := `
SELECT film_actor.actor_id AS "film_actor.actor_id", SELECT film_actor.actor_id AS "film_actor.actor_id",
@ -445,7 +445,7 @@ LIMIT 1000;
} }
} }
func TestJoinQuerySlice(t *testing.T) { func TestSelectJoinQuerySlice(t *testing.T) {
expectedSQL := ` expectedSQL := `
SELECT language.language_id AS "language.language_id", SELECT language.language_id AS "language.language_id",
language.name AS "language.name", language.name AS "language.name",
@ -474,7 +474,7 @@ LIMIT 15;
Film []model.Film Film []model.Film
} }
filmsPerLanguage := []FilmsPerLanguage{} var filmsPerLanguage []FilmsPerLanguage
limit := 15 limit := 15
query := Film. query := Film.
@ -504,7 +504,7 @@ LIMIT 15;
} }
// https://github.com/go-jet/jet/issues/226 // https://github.com/go-jet/jet/issues/226
func TestDuplicateSlicesInDestination(t *testing.T) { func TestSelectDuplicateSlicesInDestination(t *testing.T) {
type Staffs struct { type Staffs struct {
StaffList []model.Staff StaffList []model.Staff
@ -645,7 +645,7 @@ func TestDuplicateSlicesInDestination(t *testing.T) {
`) `)
} }
func TestExecution1(t *testing.T) { func TestSelectExecution1(t *testing.T) {
stmt := City. stmt := City.
INNER_JOIN(Address, Address.CityID.EQ(City.CityID)). INNER_JOIN(Address, Address.CityID.EQ(City.CityID)).
INNER_JOIN(Customer, Customer.AddressID.EQ(Address.AddressID)). INNER_JOIN(Customer, Customer.AddressID.EQ(Address.AddressID)).
@ -706,7 +706,7 @@ ORDER BY city.city_id, address.address_id, customer.customer_id;
} }
func TestExecution2(t *testing.T) { func TestSelectExecution2(t *testing.T) {
type MyAddress struct { type MyAddress struct {
ID int32 `sql:"primary_key"` ID int32 `sql:"primary_key"`
@ -769,7 +769,7 @@ ORDER BY city.city_id, address.address_id, customer.customer_id;
require.Equal(t, *dest[0].Customers[1].LastName, "Vines") require.Equal(t, *dest[0].Customers[1].LastName, "Vines")
} }
func TestExecution3(t *testing.T) { func TestSelectExecution3(t *testing.T) {
var dest []struct { var dest []struct {
CityID int32 `sql:"primary_key"` CityID int32 `sql:"primary_key"`
@ -826,7 +826,7 @@ ORDER BY city.city_id, address.address_id, customer.customer_id;
require.Equal(t, *dest[0].Customers[1].LastName, "Vines") require.Equal(t, *dest[0].Customers[1].LastName, "Vines")
} }
func TestExecution4(t *testing.T) { func TestSelectExecution4(t *testing.T) {
var dest []struct { var dest []struct {
CityID int32 `sql:"primary_key" alias:"city.city_id"` CityID int32 `sql:"primary_key" alias:"city.city_id"`
@ -918,7 +918,7 @@ ORDER BY city.city_id, address.address_id, customer.customer_id;
} }
// Test join with custom primary keys (sql.NullInt64) // Test join with custom primary keys (sql.NullInt64)
func TestExecutionCustomPKTypes1(t *testing.T) { func TestSelectExecutionCustomPKTypes1(t *testing.T) {
var dest []struct { var dest []struct {
CityID sql.NullInt64 `sql:"primary_key" alias:"city.city_id"` CityID sql.NullInt64 `sql:"primary_key" alias:"city.city_id"`
@ -1039,7 +1039,7 @@ ORDER BY city.city_id, address.address_id, customer.customer_id;
} }
// Test join with custom primary keys (null.Int) // Test join with custom primary keys (null.Int)
func TestExecutionCustomPKTypes2(t *testing.T) { func TestSelectExecutionCustomPKTypes2(t *testing.T) {
var dest []struct { var dest []struct {
CityID null.Int `sql:"primary_key" alias:"city.city_id"` CityID null.Int `sql:"primary_key" alias:"city.city_id"`
@ -1135,7 +1135,7 @@ ORDER BY city.city_id, address.address_id, customer.customer_id;
`) `)
} }
func TestJoinQuerySliceWithPtrs(t *testing.T) { func TestSelectJoinQuerySliceWithPtrs(t *testing.T) {
type FilmsPerLanguage struct { type FilmsPerLanguage struct {
Language model.Language Language model.Language
Film *[]*model.Film Film *[]*model.Film
@ -1160,7 +1160,7 @@ func TestJoinQuerySliceWithPtrs(t *testing.T) {
} }
func TestSelect_WithoutUniqueColumnSelected(t *testing.T) { func TestSelect_WithoutUniqueColumnSelected(t *testing.T) {
query := Customer.SELECT(Customer.FirstName, Customer.LastName, Customer.Email) query := Customer.SELECT(Customer.FirstName, Customer.LastName, Customer.Email).ORDER_BY(Customer.Email)
var customers []model.Customer var customers []model.Customer
@ -1219,7 +1219,7 @@ func TestSelectOrderByAscDesc(t *testing.T) {
testutils.AssertDeepEqual(t, customerAscDesc327, customersAscDesc[327]) testutils.AssertDeepEqual(t, customerAscDesc327, customersAscDesc[327])
} }
func TestOrderBy(t *testing.T) { func TestSelectOrderBy(t *testing.T) {
t.Run("default", func(t *testing.T) { t.Run("default", func(t *testing.T) {
stmt := SELECT( stmt := SELECT(
@ -1650,7 +1650,7 @@ LIMIT 1000;
testutils.AssertDeepEqual(t, films[0], thesameLengthFilms{"Alien Center", "Iron Moon", 46}) testutils.AssertDeepEqual(t, films[0], thesameLengthFilms{"Alien Center", "Iron Moon", 46})
} }
func TestSubQuery(t *testing.T) { func TestSelectSubQuery(t *testing.T) {
rRatingFilms := rRatingFilms :=
SELECT( SELECT(
Film.FilmID, Film.FilmID,
@ -1816,7 +1816,7 @@ ORDER BY film.film_id ASC;
testutils.AssertDebugStatementSql(t, query, expectedSQL) testutils.AssertDebugStatementSql(t, query, expectedSQL)
maxRentalRateFilms := []model.Film{} var maxRentalRateFilms []model.Film
err := query.Query(db, &maxRentalRateFilms) err := query.Query(db, &maxRentalRateFilms)
require.NoError(t, err) require.NoError(t, err)
@ -1920,7 +1920,7 @@ ORDER BY customer.customer_id, SUM(payment.amount) ASC;
testutils.AssertJSONFile(t, dest, "./testdata/results/postgres/customer_payment_sum.json") testutils.AssertJSONFile(t, dest, "./testdata/results/postgres/customer_payment_sum.json")
} }
func TestGroupByGroupingSets(t *testing.T) { func TestSelectGroupByGroupingSets(t *testing.T) {
skipForCockroachDB(t) skipForCockroachDB(t)
stmt := SELECT( stmt := SELECT(
@ -2002,7 +2002,7 @@ ORDER BY inventory.film_id, inventory.store_id;
`) `)
} }
func TestGroupByCube(t *testing.T) { func TestSelectGroupByCube(t *testing.T) {
skipForCockroachDB(t) skipForCockroachDB(t)
stmt := SELECT( stmt := SELECT(
@ -2079,7 +2079,7 @@ ORDER BY country.country, city.city;
`) `)
} }
func TestGroupByRollup(t *testing.T) { func TestSelectGroupByRollup(t *testing.T) {
skipForCockroachDB(t) skipForCockroachDB(t)
stmt := SELECT( stmt := SELECT(
@ -2166,7 +2166,7 @@ ORDER BY year ASC, EXTRACT(MONTH FROM rental.rental_date) ASC, day ASC;
`) `)
} }
func TestAggregateFunctionDistinct(t *testing.T) { func TestSelectAggregateFunctionDistinct(t *testing.T) {
stmt := SELECT( stmt := SELECT(
Payment.CustomerID, Payment.CustomerID,
@ -2297,7 +2297,7 @@ ORDER BY customer_payment_sum.amount_sum ASC;
} }
func TestSelectStaff(t *testing.T) { func TestSelectStaff(t *testing.T) {
staffs := []model.Staff{} var staffs []model.Staff
err := Staff.SELECT(Staff.AllColumns).Query(db, &staffs) err := Staff.SELECT(Staff.AllColumns).Query(db, &staffs)
@ -2371,7 +2371,7 @@ ORDER BY payment.payment_date ASC;
}) })
} }
func TestUnion(t *testing.T) { func TestSelectUnion(t *testing.T) {
expectedQuery := ` expectedQuery := `
( (
SELECT payment.payment_id AS "payment.payment_id", SELECT payment.payment_id AS "payment.payment_id",
@ -2424,7 +2424,7 @@ OFFSET 20;
}) })
} }
func TestUnionOffsetWithExpression(t *testing.T) { func TestSelectUnionOffsetWithExpression(t *testing.T) {
stmt := UNION( stmt := UNION(
SELECT(Rental.AllColumns). SELECT(Rental.AllColumns).
FROM(Rental). FROM(Rental).
@ -2474,7 +2474,7 @@ OFFSET (
require.Len(t, dest, 10) require.Len(t, dest, 10)
} }
func TestAllSetOperators(t *testing.T) { func TestSelectAllSetOperators(t *testing.T) {
var select1 = Payment.SELECT(Payment.AllColumns).WHERE(Payment.PaymentID.GT_EQ(Int(17600)).AND(Payment.PaymentID.LT(Int(17610)))) var select1 = Payment.SELECT(Payment.AllColumns).WHERE(Payment.PaymentID.GT_EQ(Int(17600)).AND(Payment.PaymentID.LT(Int(17610))))
var select2 = Payment.SELECT(Payment.AllColumns).WHERE(Payment.PaymentID.GT_EQ(Int(17620)).AND(Payment.PaymentID.LT(Int(17630)))) var select2 = Payment.SELECT(Payment.AllColumns).WHERE(Payment.PaymentID.GT_EQ(Int(17620)).AND(Payment.PaymentID.LT(Int(17630))))
@ -2579,46 +2579,30 @@ func getRowLockTestData() map[RowLock]string {
} }
} }
func TestRowLock(t *testing.T) { func TestSelectRowLock(t *testing.T) {
query := SELECT(STAR).
FROM(Address).
WHERE(Address.AddressID.LT(Int(10)))
expectedSQL := ` expectedSQL := `
SELECT * SELECT *
FROM dvds.address FROM dvds.address
LIMIT 3 WHERE address.address_id < 10
FOR` FOR`
query := Address.
SELECT(STAR).
LIMIT(3)
for lockType, lockTypeStr := range getRowLockTestData() { for lockType, lockTypeStr := range getRowLockTestData() {
query.FOR(lockType) query.FOR(lockType)
testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+";\n", int64(3)) testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+";\n", int64(10))
testutils.AssertExecAndRollback(t, query, db, 9)
tx, _ := db.Begin()
res, err := query.Exec(tx)
require.NoError(t, err)
rowsAffected, _ := res.RowsAffected()
require.Equal(t, rowsAffected, int64(3))
err = tx.Rollback()
require.NoError(t, err)
} }
for lockType, lockTypeStr := range getRowLockTestData() { for lockType, lockTypeStr := range getRowLockTestData() {
query.FOR(lockType.NOWAIT()) query.FOR(lockType.NOWAIT())
testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+" NOWAIT;\n", int64(3)) testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+" NOWAIT;\n", int64(10))
testutils.AssertExecAndRollback(t, query, db, 9)
tx, _ := db.Begin()
res, err := query.Exec(tx)
require.NoError(t, err)
rowsAffected, _ := res.RowsAffected()
require.Equal(t, rowsAffected, int64(3))
err = tx.Rollback()
require.NoError(t, err)
} }
if sourceIsCockroachDB() { if sourceIsCockroachDB() {
@ -2628,21 +2612,13 @@ FOR`
for lockType, lockTypeStr := range getRowLockTestData() { for lockType, lockTypeStr := range getRowLockTestData() {
query.FOR(lockType.SKIP_LOCKED()) query.FOR(lockType.SKIP_LOCKED())
testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+" SKIP LOCKED;\n", int64(3)) testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+" SKIP LOCKED;\n", int64(10))
testutils.AssertExecAndRollback(t, query, db, 9)
tx, _ := db.Begin()
res, err := query.Exec(tx)
require.NoError(t, err)
rowsAffected, _ := res.RowsAffected()
require.Equal(t, rowsAffected, int64(3))
err = tx.Rollback()
require.NoError(t, err)
} }
} }
func TestRowLockWithUpdateOf(t *testing.T) { func TestSelectRowLockWithUpdateOf(t *testing.T) {
stmt := SELECT( stmt := SELECT(
Film.FilmID, Film.FilmID,
Film.Title, Film.Title,
@ -2672,9 +2648,10 @@ LIMIT 1
FOR UPDATE OF film, actor NOWAIT; FOR UPDATE OF film, actor NOWAIT;
`) `)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
var dest []struct { var dest []struct {
model.Film model.Film
CategoryID int CategoryID int
Actor []model.Actor Actor []model.Actor
} }
@ -2685,7 +2662,7 @@ FOR UPDATE OF film, actor NOWAIT;
}) })
} }
func TestRowLockWithUpdateOfAliasedTable(t *testing.T) { func TestSelectRowLockWithUpdateOfAliasedTable(t *testing.T) {
myFilm := Film.AS("myFilm") myFilm := Film.AS("myFilm")
@ -2708,18 +2685,18 @@ LIMIT 1
FOR UPDATE OF "myFilm"; FOR UPDATE OF "myFilm";
`) `)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
var dest []struct { var dest []struct {
model.Film `alias:"myFilm.*"` model.Film `alias:"myFilm.*"`
} }
err := stmt.Query(db, &dest) err := stmt.Query(tx, &dest)
require.NoError(t, err) require.NoError(t, err)
require.Len(t, dest, 1) require.Len(t, dest, 1)
}) })
} }
func TestQuickStart(t *testing.T) { func TestSelectQuickStart(t *testing.T) {
var expectedSQL = ` var expectedSQL = `
SELECT actor.actor_id AS "actor.actor_id", SELECT actor.actor_id AS "actor.actor_id",
@ -2810,7 +2787,7 @@ ORDER BY actor.actor_id ASC, film.film_id ASC;
testutils.AssertJSONFile(t, dest2, "./testdata/results/postgres/quick-start-dest2.json") testutils.AssertJSONFile(t, dest2, "./testdata/results/postgres/quick-start-dest2.json")
} }
func TestQuickStartWithSubQueries(t *testing.T) { func TestSelectQuickStartWithSubQueries(t *testing.T) {
filmLogerThan180 := Film. filmLogerThan180 := Film.
SELECT(Film.AllColumns). SELECT(Film.AllColumns).
@ -2875,7 +2852,7 @@ func TestQuickStartWithSubQueries(t *testing.T) {
testutils.AssertJSONFile(t, dest2, "./testdata/results/postgres/quick-start-dest2.json") testutils.AssertJSONFile(t, dest2, "./testdata/results/postgres/quick-start-dest2.json")
} }
func TestExpressionWrappers(t *testing.T) { func TestSelectExpressionWrappers(t *testing.T) {
query := SELECT( query := SELECT(
BoolExp(Raw("true")), BoolExp(Raw("true")),
IntExp(Raw("11")), IntExp(Raw("11")),
@ -2905,7 +2882,7 @@ SELECT true,
require.NoError(t, err) require.NoError(t, err)
} }
func TestWindowFunction(t *testing.T) { func TestSelectWindowFunction(t *testing.T) {
var expectedSQL = ` var expectedSQL = `
SELECT AVG(payment.amount) OVER (), SELECT AVG(payment.amount) OVER (),
AVG(payment.amount) OVER (PARTITION BY payment.customer_id), AVG(payment.amount) OVER (PARTITION BY payment.customer_id),
@ -2977,7 +2954,7 @@ GROUP BY payment.amount, payment.customer_id, payment.payment_date;
require.NoError(t, err) require.NoError(t, err)
} }
func TestWindowClause(t *testing.T) { func TestSelectWindowClause(t *testing.T) {
var expectedSQL = ` var expectedSQL = `
SELECT AVG(payment.amount) OVER (), SELECT AVG(payment.amount) OVER (),
AVG(payment.amount) OVER (w1), AVG(payment.amount) OVER (w1),
@ -3014,7 +2991,7 @@ ORDER BY payment.customer_id;
require.NoError(t, err) require.NoError(t, err)
} }
func TestSimpleView(t *testing.T) { func TestSelectSimpleView(t *testing.T) {
query := SELECT( query := SELECT(
view.ActorInfo.AllColumns, view.ActorInfo.AllColumns,
@ -3053,7 +3030,7 @@ func TestSimpleView(t *testing.T) {
} }
func TestJoinViewWithTable(t *testing.T) { func TestSelectJoinViewWithTable(t *testing.T) {
query := SELECT( query := SELECT(
view.CustomerList.AllColumns, view.CustomerList.AllColumns,
Rental.AllColumns, Rental.AllColumns,
@ -3079,7 +3056,7 @@ func TestJoinViewWithTable(t *testing.T) {
require.Equal(t, len(dest[1].Rentals), 27) require.Equal(t, len(dest[1].Rentals), 27)
} }
func TestDynamicProjectionList(t *testing.T) { func TestSelectDynamicProjectionList(t *testing.T) {
var request struct { var request struct {
ColumnsToSelect []string ColumnsToSelect []string
@ -3126,7 +3103,7 @@ LIMIT 3;
require.Equal(t, len(dest), 3) require.Equal(t, len(dest), 3)
} }
func TestDynamicCondition(t *testing.T) { func TestSelectDynamicCondition(t *testing.T) {
var request struct { var request struct {
CustomerID *int64 CustomerID *int64
Email *string Email *string
@ -3176,7 +3153,7 @@ WHERE ($1::boolean AND (customer.customer_id = $2)) AND (customer.activebool = $
testutils.AssertDeepEqual(t, dest[0], customer0) testutils.AssertDeepEqual(t, dest[0], customer0)
} }
func TestLateral(t *testing.T) { func TestSelectLateral(t *testing.T) {
languages := LATERAL( languages := LATERAL(
SELECT( SELECT(
@ -3348,7 +3325,7 @@ type ActorWrap struct {
Films []FilmWrap Films []FilmWrap
} }
func TestRecursionScanNxM(t *testing.T) { func TestSelectRecursionScanNxM(t *testing.T) {
stmt := SELECT( stmt := SELECT(
Actor.AllColumns, Actor.AllColumns,
@ -3491,7 +3468,7 @@ type StaffWrap struct {
Store StoreWrap Store StoreWrap
} }
func TestRecursionScanNx1(t *testing.T) { func TestSelectRecursionScanNx1(t *testing.T) {
stmt := SELECT( stmt := SELECT(
Store.AllColumns, Store.AllColumns,
Staff.AllColumns, Staff.AllColumns,
@ -3638,7 +3615,7 @@ type ManagerInfo struct {
Store *StoreInfo Store *StoreInfo
} }
func TestRecursionScan1x1(t *testing.T) { func TestSelectRecursionScan1x1(t *testing.T) {
stmt := SELECT( stmt := SELECT(
Store.AllColumns, Store.AllColumns,
@ -3703,7 +3680,7 @@ func TestRecursionScan1x1(t *testing.T) {
// In parameterized statements integer literals, like Int(num), are replaced with a placeholders. For some expressions, // In parameterized statements integer literals, like Int(num), are replaced with a placeholders. For some expressions,
// postgres interpreter will not have enough information to deduce the type. If this is the case postgres returns an error. // postgres interpreter will not have enough information to deduce the type. If this is the case postgres returns an error.
// Int8, Int16, .... functions will add automatic type cast over placeholder, so type deduction is always possible. // Int8, Int16, .... functions will add automatic type cast over placeholder, so type deduction is always possible.
func TestLiteralTypeDeduction(t *testing.T) { func TestSelectLiteralTypeDeduction(t *testing.T) {
stmt := SELECT( stmt := SELECT(
SUM( SUM(
CASE().WHEN(Staff.Active.IS_TRUE()). CASE().WHEN(Staff.Active.IS_TRUE()).
@ -3725,7 +3702,7 @@ func GET_FILM_COUNT(lenFrom, lenTo IntegerExpression) IntegerExpression {
return IntExp(Func("dvds.get_film_count", lenFrom, lenTo)) return IntExp(Func("dvds.get_film_count", lenFrom, lenTo))
} }
func TestCustomFunctionCall(t *testing.T) { func TestSelectCustomFunctionCall(t *testing.T) {
skipForCockroachDB(t) skipForCockroachDB(t)
stmt := SELECT( stmt := SELECT(
@ -3761,7 +3738,7 @@ SELECT dvds.get_film_count(100, 120) AS "film_count";
require.Equal(t, dest.FilmCount, 165) require.Equal(t, dest.FilmCount, 165)
} }
func TestScanUsingConn(t *testing.T) { func TestSelectScanUsingConn(t *testing.T) {
conn, err := db.Conn(context.Background()) conn, err := db.Conn(context.Background())
require.NoError(t, err) require.NoError(t, err)
defer conn.Close() defer conn.Close()
@ -3803,7 +3780,7 @@ func TestScanUsingConn(t *testing.T) {
}) })
} }
func TestConditionalFunctions(t *testing.T) { func TestSelectConditionalFunctions(t *testing.T) {
stmt := SELECT( stmt := SELECT(
EXISTS( EXISTS(
Film.SELECT(Film.FilmID).WHERE(Film.RentalDuration.GT(Int(100))), Film.SELECT(Film.FilmID).WHERE(Film.RentalDuration.GT(Int(100))),

View file

@ -2,9 +2,9 @@ package postgres
import ( import (
"context" "context"
"database/sql"
"github.com/go-jet/jet/v2/internal/testutils" "github.com/go-jet/jet/v2/internal/testutils"
. "github.com/go-jet/jet/v2/postgres" . "github.com/go-jet/jet/v2/postgres"
"github.com/go-jet/jet/v2/qrm"
model2 "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/model" model2 "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/model"
"github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/table" "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/table"
"github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/test_sample/model" "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/test_sample/model"
@ -27,7 +27,7 @@ SET (name, url) = ('Bong', 'http://bong.com')
WHERE link.name = 'Bing'::text; WHERE link.name = 'Bing'::text;
`, "Bong", "http://bong.com", "Bing") `, "Bong", "http://bong.com", "Bing")
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
testutils.AssertExec(t, query, tx, 1) testutils.AssertExec(t, query, tx, 1)
requireLogged(t, query) requireLogged(t, query)
@ -142,7 +142,7 @@ RETURNING link.id AS "link.id",
link.description AS "link.description"; link.description AS "link.description";
`, "DuckDuckGo", "http://www.duckduckgo.com", "Ask") `, "DuckDuckGo", "http://www.duckduckgo.com", "Ask")
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
links := []model.Link{} links := []model.Link{}
err := stmt.Query(tx, &links) err := stmt.Query(tx, &links)
@ -325,8 +325,8 @@ func TestUpdateQueryContext(t *testing.T) {
time.Sleep(10 * time.Millisecond) time.Sleep(10 * time.Millisecond)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
dest := []model.Link{} var dest []model.Link
err := updateStmt.QueryContext(ctx, tx, &dest) err := updateStmt.QueryContext(ctx, tx, &dest)
require.Error(t, err, "context deadline exceeded") require.Error(t, err, "context deadline exceeded")
@ -344,7 +344,10 @@ func TestUpdateExecContext(t *testing.T) {
time.Sleep(10 * time.Millisecond) time.Sleep(10 * time.Millisecond)
testutils.AssertExecContextErr(ctx, t, updateStmt, db, "context deadline exceeded") testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
_, err := updateStmt.ExecContext(ctx, db)
require.Error(t, err, "context deadline exceeded")
})
} }
func TestUpdateFrom(t *testing.T) { func TestUpdateFrom(t *testing.T) {
@ -385,7 +388,7 @@ RETURNING rental.rental_id AS "rental.rental_id",
store.address_id AS "store.address_id"; store.address_id AS "store.address_id";
`) `)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
var dest []struct { var dest []struct {
Rental model2.Rental Rental model2.Rental
Store model2.Store Store model2.Store

View file

@ -2,6 +2,7 @@ package sqlite
import ( import (
"context" "context"
"github.com/go-jet/jet/v2/qrm"
"testing" "testing"
"time" "time"
@ -60,8 +61,6 @@ LIMIT 1;
} }
func TestDeleteContextDeadlineExceeded(t *testing.T) { func TestDeleteContextDeadlineExceeded(t *testing.T) {
tx := beginSampleDBTx(t)
defer tx.Rollback()
deleteStmt := Link. deleteStmt := Link.
DELETE(). DELETE().
@ -72,12 +71,16 @@ func TestDeleteContextDeadlineExceeded(t *testing.T) {
time.Sleep(10 * time.Millisecond) time.Sleep(10 * time.Millisecond)
dest := []model.Link{} testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx qrm.DB) {
var dest []model.Link
err := deleteStmt.QueryContext(ctx, tx, &dest) err := deleteStmt.QueryContext(ctx, tx, &dest)
require.Error(t, err, "context deadline exceeded") require.Error(t, err, "context deadline exceeded")
})
_, err = deleteStmt.ExecContext(ctx, tx) testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx qrm.DB) {
_, err := deleteStmt.ExecContext(ctx, tx)
require.Error(t, err, "context deadline exceeded") require.Error(t, err, "context deadline exceeded")
})
requireLogged(t, deleteStmt) requireLogged(t, deleteStmt)
} }

View file

@ -2,7 +2,7 @@ package sqlite
import ( import (
"context" "context"
"database/sql" "github.com/go-jet/jet/v2/qrm"
"math/rand" "math/rand"
"testing" "testing"
@ -30,7 +30,7 @@ VALUES (?, ?, ?, ?),
101, "http://www.google.com", "Google", "Search engine", 101, "http://www.google.com", "Google", "Search engine",
102, "http://www.yahoo.com", "Yahoo", nil) 102, "http://www.yahoo.com", "Yahoo", nil)
testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx qrm.DB) {
testutils.AssertExec(t, insertQuery, tx) testutils.AssertExec(t, insertQuery, tx)
requireLogged(t, insertQuery) requireLogged(t, insertQuery)
@ -75,7 +75,7 @@ VALUES (100, 'http://www.postgresqltutorial.com', 'PostgreSQL Tutorial', NULL);
`, `,
100, "http://www.postgresqltutorial.com", "PostgreSQL Tutorial", nil) 100, "http://www.postgresqltutorial.com", "PostgreSQL Tutorial", nil)
testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx qrm.DB) {
_, err := stmt.Exec(tx) _, err := stmt.Exec(tx)
require.NoError(t, err) require.NoError(t, err)
requireLogged(t, stmt) requireLogged(t, stmt)
@ -108,7 +108,7 @@ INSERT INTO link (url, name)
VALUES ('http://www.duckduckgo.com', 'Duck Duck go'); VALUES ('http://www.duckduckgo.com', 'Duck Duck go');
`, "http://www.duckduckgo.com", "Duck Duck go") `, "http://www.duckduckgo.com", "Duck Duck go")
testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx qrm.DB) {
_, err := query.Exec(tx) _, err := query.Exec(tx)
require.NoError(t, err) require.NoError(t, err)
}) })
@ -130,7 +130,7 @@ INSERT INTO link
VALUES (1000, 'http://www.duckduckgo.com', 'Duck Duck go', NULL); VALUES (1000, 'http://www.duckduckgo.com', 'Duck Duck go', NULL);
`, int32(1000), "http://www.duckduckgo.com", "Duck Duck go", nil) `, int32(1000), "http://www.duckduckgo.com", "Duck Duck go", nil)
testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx qrm.DB) {
_, err := query.Exec(tx) _, err := query.Exec(tx)
require.NoError(t, err) require.NoError(t, err)
}) })
@ -219,7 +219,7 @@ WHERE link.id = 24;
) )
testutils.AssertDebugStatementSql(t, query, expectedSQL, int64(24)) testutils.AssertDebugStatementSql(t, query, expectedSQL, int64(24))
testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx qrm.DB) {
_, err := query.Exec(tx) _, err := query.Exec(tx)
require.NoError(t, err) require.NoError(t, err)
@ -248,7 +248,7 @@ RETURNING link.id AS "link.id",
link.description AS "link.description"; link.description AS "link.description";
`) `)
testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx qrm.DB) {
var link model.Link var link model.Link
err := stmt.Query(tx, &link) err := stmt.Query(tx, &link)
require.NoError(t, err) require.NoError(t, err)
@ -353,9 +353,15 @@ func TestInsertContextDeadlineExceeded(t *testing.T) {
time.Sleep(10 * time.Millisecond) time.Sleep(10 * time.Millisecond)
var dest []model.Link var dest []model.Link
err := stmt.QueryContext(ctx, sampleDB, &dest)
require.Error(t, err, "context deadline exceeded")
_, err = stmt.ExecContext(ctx, db) testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx qrm.DB) {
err := stmt.QueryContext(ctx, tx, &dest)
require.Error(t, err, "context deadline exceeded") require.Error(t, err, "context deadline exceeded")
})
testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx qrm.DB) {
_, err := stmt.ExecContext(ctx, tx)
require.Error(t, err, "context deadline exceeded")
})
} }

View file

@ -22,8 +22,8 @@ import (
_ "github.com/mattn/go-sqlite3" _ "github.com/mattn/go-sqlite3"
) )
var db *sql.DB var db *sqlite.DB
var sampleDB *sql.DB var sampleDB *sqlite.DB
var testRoot string var testRoot string
func TestMain(m *testing.M) { func TestMain(m *testing.M) {
@ -32,24 +32,33 @@ func TestMain(m *testing.M) {
setTestRoot() setTestRoot()
var err error sqlDB, err := sql.Open("sqlite3", "file:"+dbconfig.SakilaDBPath)
db, err = sql.Open("sqlite3", "file:"+dbconfig.SakilaDBPath)
throw.OnError(err) throw.OnError(err)
db = sqlite.NewDB(sqlDB).WithStatementsCaching(true)
defer db.Close() defer db.Close()
_, err = db.Exec(fmt.Sprintf("ATTACH DATABASE '%s' as 'chinook';", dbconfig.ChinookDBPath)) _, err = db.Exec(fmt.Sprintf("ATTACH DATABASE '%s' as 'chinook';", dbconfig.ChinookDBPath))
throw.OnError(err) throw.OnError(err)
sampleDB, err = sql.Open("sqlite3", dbconfig.TestSampleDBPath) sqlSampleDB, err := sql.Open("sqlite3", dbconfig.TestSampleDBPath)
throw.OnError(err) throw.OnError(err)
sampleDB = sqlite.NewDB(sqlSampleDB).WithStatementsCaching(true)
defer sampleDB.Close()
for i := 0; i < 2; i++ {
ret := m.Run() ret := m.Run()
if ret != 0 { if ret != 0 {
os.Exit(ret) os.Exit(ret)
} }
} }
err = sampleDB.Clear()
if err != nil {
panic(err)
}
}
func setTestRoot() { func setTestRoot() {
cmd := exec.Command("git", "rev-parse", "--show-toplevel") cmd := exec.Command("git", "rev-parse", "--show-toplevel")
byteArr, err := cmd.Output() byteArr, err := cmd.Output()
@ -103,13 +112,13 @@ func requireLogged(t *testing.T, statement sqlite.Statement) {
require.Equal(t, loggedDebugSQL, statement.DebugSql()) require.Equal(t, loggedDebugSQL, statement.DebugSql())
} }
func beginSampleDBTx(t *testing.T) *sql.Tx { func beginSampleDBTx(t *testing.T) *sqlite.Tx {
tx, err := sampleDB.Begin() tx, err := sampleDB.BeginTx(context.Background(), nil)
require.NoError(t, err) require.NoError(t, err)
return tx return tx
} }
func beginDBTx(t *testing.T) *sql.Tx { func beginDBTx(t *testing.T) *sqlite.Tx {
tx, err := db.Begin() tx, err := db.Begin()
require.NoError(t, err) require.NoError(t, err)
return tx return tx

View file

@ -1,8 +1,8 @@
package sqlite package sqlite
import ( import (
"database/sql"
"github.com/go-jet/jet/v2/internal/testutils" "github.com/go-jet/jet/v2/internal/testutils"
"github.com/go-jet/jet/v2/qrm"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
"testing" "testing"
@ -48,7 +48,7 @@ WHERE people.people_id = ?;
}) })
t.Run("should insert without generated columns", func(t *testing.T) { t.Run("should insert without generated columns", func(t *testing.T) {
testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx *sql.Tx) { testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx qrm.DB) {
insertQuery := People.INSERT( insertQuery := People.INSERT(
People.MutableColumns, People.MutableColumns,
).MODEL( ).MODEL(

View file

@ -986,8 +986,8 @@ WHERE artists.''ArtistId'' = 11;
} }
func TestRowsScan(t *testing.T) { func TestRowsScan(t *testing.T) {
stmt :=
SELECT( stmt := SELECT(
Inventory.AllColumns, Inventory.AllColumns,
).FROM( ).FROM(
Inventory, Inventory,

View file

@ -2,6 +2,7 @@ package sqlite
import ( import (
"context" "context"
"github.com/go-jet/jet/v2/qrm"
model2 "github.com/go-jet/jet/v2/tests/.gentestdata/sqlite/sakila/model" model2 "github.com/go-jet/jet/v2/tests/.gentestdata/sqlite/sakila/model"
"github.com/go-jet/jet/v2/tests/.gentestdata/sqlite/sakila/table" "github.com/go-jet/jet/v2/tests/.gentestdata/sqlite/sakila/table"
"testing" "testing"
@ -215,7 +216,7 @@ WHERE link.id = 20;
testutils.AssertDebugStatementSql(t, stmt, expectedSQL, nil, "DuckDuckGo", "http://www.duckduckgo.com", int32(20)) testutils.AssertDebugStatementSql(t, stmt, expectedSQL, nil, "DuckDuckGo", "http://www.duckduckgo.com", int32(20))
testutils.AssertExec(t, stmt, tx) testutils.AssertExec(t, stmt, tx, 1)
requireLogged(t, stmt) requireLogged(t, stmt)
} }
@ -271,8 +272,6 @@ func TestUpdateWithInvalidModelData(t *testing.T) {
} }
func TestUpdateContextDeadlineExceeded(t *testing.T) { func TestUpdateContextDeadlineExceeded(t *testing.T) {
tx := beginSampleDBTx(t)
defer tx.Rollback()
updateStmt := Link.UPDATE(Link.Name, Link.URL). updateStmt := Link.UPDATE(Link.Name, Link.URL).
SET("Bong", "http://bong.com"). SET("Bong", "http://bong.com").
@ -283,12 +282,16 @@ func TestUpdateContextDeadlineExceeded(t *testing.T) {
time.Sleep(10 * time.Millisecond) time.Sleep(10 * time.Millisecond)
dest := []model.Link{} testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
var dest []model.Link
err := updateStmt.QueryContext(ctx, tx, &dest) err := updateStmt.QueryContext(ctx, tx, &dest)
require.Error(t, err, "context deadline exceeded") require.Error(t, err, "context deadline exceeded")
})
_, err = updateStmt.ExecContext(ctx, tx) testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
_, err := updateStmt.ExecContext(ctx, tx)
require.Error(t, err, "context deadline exceeded") require.Error(t, err, "context deadline exceeded")
})
} }
func TestUpdateFrom(t *testing.T) { func TestUpdateFrom(t *testing.T) {