From 0918e5503e3ef7dc660aca5d05ff234ba6e508c9 Mon Sep 17 00:00:00 2001 From: go-jet Date: Thu, 7 Mar 2024 18:01:31 +0100 Subject: [PATCH 1/2] Add support for prepared statement caching. --- internal/jet/db/db.go | 182 ++++++++++++++++++++++++++ internal/jet/db/tx.go | 97 ++++++++++++++ internal/testutils/test_utils.go | 10 +- mysql/statement.go | 14 +- postgres/statement.go | 14 +- sqlite/statement.go | 14 +- tests/mysql/delete_test.go | 4 +- tests/mysql/insert_test.go | 36 ++--- tests/mysql/lock_test.go | 28 +++- tests/mysql/main_test.go | 15 ++- tests/mysql/select_test.go | 39 ++---- tests/mysql/update_test.go | 8 +- tests/postgres/alltypes_test.go | 8 +- tests/postgres/chinook_db_test.go | 6 +- tests/postgres/delete_test.go | 10 +- tests/postgres/insert_test.go | 10 +- tests/postgres/lock_test.go | 24 +--- tests/postgres/main_test.go | 19 ++- tests/postgres/northwind_test.go | 53 ++++---- tests/postgres/range_test.go | 8 +- tests/postgres/raw_statements_test.go | 4 +- tests/postgres/sample_test.go | 4 +- tests/postgres/select_test.go | 153 +++++++++------------- tests/postgres/update_test.go | 17 ++- tests/sqlite/delete_test.go | 17 ++- tests/sqlite/insert_test.go | 28 ++-- tests/sqlite/main_test.go | 31 +++-- tests/sqlite/sample_test.go | 4 +- tests/sqlite/select_test.go | 16 +-- tests/sqlite/update_test.go | 19 +-- 30 files changed, 603 insertions(+), 289 deletions(-) create mode 100644 internal/jet/db/db.go create mode 100644 internal/jet/db/tx.go diff --git a/internal/jet/db/db.go b/internal/jet/db/db.go new file mode 100644 index 0000000..94dfeca --- /dev/null +++ b/internal/jet/db/db.go @@ -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 +} diff --git a/internal/jet/db/tx.go b/internal/jet/db/tx.go new file mode 100644 index 0000000..b64b231 --- /dev/null +++ b/internal/jet/db/tx.go @@ -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 +} diff --git a/internal/testutils/test_utils.go b/internal/testutils/test_utils.go index 22c0d9c..3e03506 100644 --- a/internal/testutils/test_utils.go +++ b/internal/testutils/test_utils.go @@ -3,10 +3,10 @@ package testutils import ( "bytes" "context" - "database/sql" "encoding/json" "fmt" "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/qrm" "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 -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() require.NoError(t, err) 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 -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() require.NoError(t, err) 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 -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() 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 -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() if len(expectedArgs) > 0 { diff --git a/mysql/statement.go b/mysql/statement.go index 073adce..a0bb0a7 100644 --- a/mysql/statement.go +++ b/mysql/statement.go @@ -1,8 +1,20 @@ 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 func RawStatement(rawQuery string, namedArguments ...RawArgs) Statement { 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 diff --git a/postgres/statement.go b/postgres/statement.go index d10bd65..4199fa9 100644 --- a/postgres/statement.go +++ b/postgres/statement.go @@ -1,8 +1,20 @@ 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 func RawStatement(rawQuery string, namedArguments ...RawArgs) Statement { 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 diff --git a/sqlite/statement.go b/sqlite/statement.go index 754ae41..eb701e7 100644 --- a/sqlite/statement.go +++ b/sqlite/statement.go @@ -1,8 +1,20 @@ 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 func RawStatement(rawQuery string, namedArguments ...RawArgs) Statement { 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 diff --git a/tests/mysql/delete_test.go b/tests/mysql/delete_test.go index 7bb2422..6ba2fc2 100644 --- a/tests/mysql/delete_test.go +++ b/tests/mysql/delete_test.go @@ -2,9 +2,9 @@ package mysql import ( "context" - "database/sql" "github.com/go-jet/jet/v2/internal/testutils" . "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/test_sample/model" . "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'); `) - testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { + testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) { _, err := stmt.Exec(tx) require.NoError(t, err) }) diff --git a/tests/mysql/insert_test.go b/tests/mysql/insert_test.go index b05c91d..7eb6ec8 100644 --- a/tests/mysql/insert_test.go +++ b/tests/mysql/insert_test.go @@ -2,9 +2,9 @@ package mysql import ( "context" - "database/sql" "github.com/go-jet/jet/v2/internal/testutils" . "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/table" "github.com/stretchr/testify/require" @@ -29,7 +29,7 @@ VALUES (100, 'http://www.postgresqltutorial.com', 'PostgreSQL Tutorial', DEFAULT 101, "http://www.google.com", "Google", 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) require.NoError(t, err) requireLogged(t, insertQuery) @@ -74,7 +74,7 @@ VALUES (100, 'http://www.postgresqltutorial.com', 'PostgreSQL Tutorial', DEFAULT `, 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) require.NoError(t, err) requireLogged(t, stmt) @@ -108,7 +108,7 @@ VALUES ('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) require.NoError(t, err) }) @@ -130,7 +130,7 @@ INSERT INTO test_sample.link VALUES (1000, 'http://www.duckduckgo.com', 'Duck Duck go', NULL); `, 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) require.NoError(t, err) }) @@ -166,7 +166,7 @@ VALUES ('http://www.postgresqltutorial.com', 'PostgreSQL Tutorial'), "http://www.google.com", "Google", "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) require.NoError(t, err) }) @@ -201,7 +201,7 @@ VALUES ('http://www.postgresqltutorial.com', 'PostgreSQL Tutorial', DEFAULT), "http://www.google.com", "Google", 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) require.NoError(t, err) }) @@ -225,7 +225,7 @@ INSERT INTO test_sample.link (url, name) ( ); `, int64(1)) - testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { + testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) { _, err := query.Exec(tx) require.NoError(t, err) @@ -261,7 +261,7 @@ ON DUPLICATE KEY UPDATE id = (link.id + ?), randId, "http://www.postgresqltutorial.com", "PostgreSQL Tutorial", 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) require.NoError(t, err) @@ -320,7 +320,7 @@ ON DUPLICATE KEY UPDATE id = (link.id + ?), description = new.description; `) - testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { + testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) { _, err := stmt.Exec(tx) require.NoError(t, err) @@ -352,9 +352,12 @@ func TestInsertWithQueryContext(t *testing.T) { time.Sleep(10 * time.Millisecond) var dest []model.Link - err := stmt.QueryContext(ctx, db, &dest) - require.Error(t, err, "context deadline exceeded") + testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) { + err := stmt.QueryContext(ctx, tx, &dest) + require.Error(t, err, "context deadline exceeded") + }) + } func TestInsertWithExecContext(t *testing.T) { @@ -366,9 +369,10 @@ func TestInsertWithExecContext(t *testing.T) { time.Sleep(10 * time.Millisecond) - _, err := stmt.ExecContext(ctx, db) - - require.Error(t, err, "context deadline exceeded") + testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) { + _, err := stmt.ExecContext(ctx, tx) + require.Error(t, err, "context deadline exceeded") + }) } 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); `) - testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { + testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) { _, err := stmt.Exec(tx) require.NoError(t, err) }) diff --git a/tests/mysql/lock_test.go b/tests/mysql/lock_test.go index 939fb5a..f4f78d5 100644 --- a/tests/mysql/lock_test.go +++ b/tests/mysql/lock_test.go @@ -14,10 +14,14 @@ func TestLockRead(t *testing.T) { testutils.AssertStatementSql(t, query, ` LOCK TABLES dvds.customer READ; `) - - _, err := query.Exec(db) + tx, err := db.DB.Begin() // can't prepare LOCK statement 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) { @@ -27,9 +31,14 @@ func TestLockWrite(t *testing.T) { LOCK TABLES dvds.customer WRITE; `) - _, err := query.Exec(db) + tx, err := db.DB.Begin() // can't prepare LOCK statement 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) { @@ -39,7 +48,12 @@ func TestUnlockTables(t *testing.T) { UNLOCK TABLES; `) - _, err := query.Exec(db) + tx, err := db.DB.Begin() // can't prepare LOCK statement require.NoError(t, err) - requireLogged(t, query) + defer func() { + err := tx.Rollback() + require.NoError(t, err) + }() + + testutils.AssertExec(t, query, tx) } diff --git a/tests/mysql/main_test.go b/tests/mysql/main_test.go index f6ce57d..fd711b4 100644 --- a/tests/mysql/main_test.go +++ b/tests/mysql/main_test.go @@ -18,7 +18,7 @@ import ( "testing" ) -var db *sql.DB +var db *jetmysql.DB var source string @@ -37,15 +37,20 @@ func TestMain(m *testing.M) { defer profile.Start().Stop() var err error - db, err = sql.Open("mysql", dbconfig.MySQLConnectionString(sourceIsMariaDB(), "")) + sqlDB, err := sql.Open("mysql", dbconfig.MySQLConnectionString(sourceIsMariaDB(), "")) if err != nil { panic("Failed to connect to test db" + err.Error()) } + + db = jetmysql.NewDB(sqlDB).WithStatementsCaching(true) defer db.Close() - ret := m.Run() - - os.Exit(ret) + for i := 0; i < 2; i++ { + ret := m.Run() + if ret != 0 { + os.Exit(ret) + } + } } var loggedSQL string diff --git a/tests/mysql/select_test.go b/tests/mysql/select_test.go index 94967c8..6e4507d 100644 --- a/tests/mysql/select_test.go +++ b/tests/mysql/select_test.go @@ -2,8 +2,8 @@ package mysql import ( "context" - "database/sql" "github.com/go-jet/jet/v2/postgres" + "github.com/go-jet/jet/v2/qrm" "strings" "testing" "time" @@ -868,11 +868,13 @@ func TestRowLock(t *testing.T) { expectedSQL := ` SELECT * FROM dvds.address +ORDER BY address.address_id LIMIT 3 OFFSET 1 FOR` - query := Address. - SELECT(STAR). + query := SELECT(STAR). + FROM(Address). + ORDER_BY(Address.AddressID). LIMIT(3). OFFSET(1) @@ -881,28 +883,14 @@ FOR` expectedQuery := expectedSQL + " " + lockTypeStr + ";\n" testutils.AssertDebugStatementSql(t, query, expectedQuery, int64(3), int64(1)) - - tx, _ := db.Begin() - - _, err := query.Exec(tx) - require.NoError(t, err) - - err = tx.Rollback() - require.NoError(t, err) + testutils.AssertExecAndRollback(t, query, db) } for lockType, lockTypeStr := range getRowLockTestData() { query.FOR(lockType.NOWAIT()) testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+" NOWAIT;\n", int64(3), int64(1)) - - tx, _ := db.Begin() - - _, err := query.Exec(tx) - require.NoError(t, err) - - err = tx.Rollback() - require.NoError(t, err) + testutils.AssertExecAndRollback(t, query, db) } if sourceIsMariaDB() { @@ -913,14 +901,7 @@ FOR` query.FOR(lockType.SKIP_LOCKED()) testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+" SKIP LOCKED;\n", int64(3), int64(1)) - - tx, _ := db.Begin() - - _, err := query.Exec(tx) - require.NoError(t, err) - - err = tx.Rollback() - require.NoError(t, err) + testutils.AssertExecAndRollback(t, query, db) } } @@ -956,7 +937,7 @@ LIMIT 1 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 { model.Film CategoryID int @@ -993,7 +974,7 @@ LIMIT 1 FOR UPDATE OF ''myFilm''; `, "''", "`")) - testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { + testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) { var dest []struct { model.Film `alias:"myFilm.*"` } diff --git a/tests/mysql/update_test.go b/tests/mysql/update_test.go index 909d884..cf9a7ea 100644 --- a/tests/mysql/update_test.go +++ b/tests/mysql/update_test.go @@ -2,7 +2,7 @@ package mysql import ( "context" - "database/sql" + "github.com/go-jet/jet/v2/qrm" "testing" "time" @@ -28,7 +28,7 @@ WHERE link.name = 'Bing'; WHERE(Link.Name.EQ(String("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) requireLogged(t, query) @@ -59,7 +59,7 @@ WHERE link.name = 'Bing'; WHERE(Link.Name.EQ(String("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) requireLogged(t, stmt) @@ -280,7 +280,7 @@ SET id = 501, WHERE link.name = 'Bing'; `) - testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { + testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) { _, err := stmt.Exec(tx) require.NoError(t, err) }) diff --git a/tests/postgres/alltypes_test.go b/tests/postgres/alltypes_test.go index d41feee..25a29d2 100644 --- a/tests/postgres/alltypes_test.go +++ b/tests/postgres/alltypes_test.go @@ -1,7 +1,7 @@ package postgres import ( - "database/sql" + "github.com/go-jet/jet/v2/qrm" "testing" "time" @@ -71,7 +71,7 @@ func TestAllTypesInsertModel(t *testing.T) { MODEL(&allTypesRow1). RETURNING(AllTypes.AllColumns) - testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { + testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) { var dest []model.AllTypes err := query.Query(tx, &dest) require.NoError(t, err) @@ -94,7 +94,7 @@ func TestAllTypesInsertQuery(t *testing.T) { ). RETURNING(AllTypesAllColumns) - testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { + testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) { var dest []model.AllTypes err := query.Query(tx, &dest) @@ -1293,7 +1293,7 @@ WHERE all_types.small_int = 14 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 err := stmt.Query(tx, &res) diff --git a/tests/postgres/chinook_db_test.go b/tests/postgres/chinook_db_test.go index afc4842..462e205 100644 --- a/tests/postgres/chinook_db_test.go +++ b/tests/postgres/chinook_db_test.go @@ -12,7 +12,7 @@ import ( "time" ) -func TestSelect(t *testing.T) { +func TestSelectAlbum(t *testing.T) { stmt := SELECT(Album.AllColumns). FROM(Album). ORDER_BY(Album.AlbumId.ASC()) @@ -782,9 +782,11 @@ func TestQueryWithContext(t *testing.T) { 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() + time.Sleep(1 * time.Millisecond) + var dest []model.Album err := Album. diff --git a/tests/postgres/delete_test.go b/tests/postgres/delete_test.go index ee8d320..d055101 100644 --- a/tests/postgres/delete_test.go +++ b/tests/postgres/delete_test.go @@ -2,9 +2,9 @@ package postgres import ( "context" - "database/sql" "github.com/go-jet/jet/v2/internal/testutils" . "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" "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/table" "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"; `, "Gmail", "Outlook") - testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { + testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) { var dest []model.Link err := deleteStmt.Query(tx, &dest) @@ -67,7 +67,7 @@ func TestDeleteQueryContext(t *testing.T) { time.Sleep(10 * time.Millisecond) - testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { + testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) { dest := []model.Link{} err := deleteStmt.QueryContext(ctx, tx, &dest) @@ -88,7 +88,7 @@ func TestDeleteExecContext(t *testing.T) { 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) 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"; `) - testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { + testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) { var dest []struct { Rental model2.Rental Store model2.Store diff --git a/tests/postgres/insert_test.go b/tests/postgres/insert_test.go index e34405e..bec1fe2 100644 --- a/tests/postgres/insert_test.go +++ b/tests/postgres/insert_test.go @@ -2,9 +2,9 @@ package postgres import ( "context" - "database/sql" "github.com/go-jet/jet/v2/internal/testutils" . "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/table" "github.com/stretchr/testify/require" @@ -34,7 +34,7 @@ RETURNING link.id AS "link.id", 101, "http://www.google.com", "Google", 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 err := insertQuery.Query(tx, &insertedLinks) @@ -335,7 +335,7 @@ RETURNING link.id AS "link.id", link.description AS "link.description"; `, int64(0)) - testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { + testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) { var dest []model.Link err := query.Query(tx, &dest) @@ -362,7 +362,7 @@ func TestInsertWithQueryContext(t *testing.T) { time.Sleep(10 * time.Millisecond) - testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { + testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) { dest := []model.Link{} err := stmt.QueryContext(ctx, tx, &dest) @@ -379,7 +379,7 @@ func TestInsertWithExecContext(t *testing.T) { 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") }) } diff --git a/tests/postgres/lock_test.go b/tests/postgres/lock_test.go index 4caed10..bf162d4 100644 --- a/tests/postgres/lock_test.go +++ b/tests/postgres/lock_test.go @@ -32,34 +32,14 @@ LOCK TABLE dvds.address IN` query := Address.LOCK().IN(lockMode) testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+string(lockMode)+" MODE;\n") - - tx, _ := db.Begin() - - _, err := query.Exec(tx) - - require.NoError(t, err) - - err = tx.Rollback() - - require.NoError(t, err) - requireLogged(t, query) + testutils.AssertExecAndRollback(t, query, db) } for _, lockMode := range testData { query := Address.LOCK().IN(lockMode).NOWAIT() testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+string(lockMode)+" MODE NOWAIT;\n") - - tx, _ := db.Begin() - - _, err := query.Exec(tx) - - require.NoError(t, err) - - err = tx.Rollback() - - require.NoError(t, err) - requireLogged(t, query) + testutils.AssertExecAndRollback(t, query, db) } } diff --git a/tests/postgres/main_test.go b/tests/postgres/main_test.go index 08af67a..1251c93 100644 --- a/tests/postgres/main_test.go +++ b/tests/postgres/main_test.go @@ -22,7 +22,7 @@ import ( _ "github.com/jackc/pgx/v4/stdlib" ) -var db *sql.DB +var db *postgres.DB var testRoot string var source string @@ -60,18 +60,25 @@ func TestMain(m *testing.M) { connectionString = dbconfig.CockroachConnectString } - var err error - db, err = sql.Open(driverName, connectionString) + sqlDB, err := sql.Open(driverName, connectionString) if err != nil { fmt.Println(err.Error()) panic("Failed to connect to test db") } + db = postgres.NewDB(sqlDB).WithStatementsCaching(true) defer db.Close() - ret := m.Run() + for i := 0; i < 2; i++ { + ret := m.Run() + if ret != 0 { + os.Exit(ret) + } + } - if ret != 0 { - os.Exit(ret) + err = db.Clear() + + if err != nil { + os.Exit(-2) } }() } diff --git a/tests/postgres/northwind_test.go b/tests/postgres/northwind_test.go index 2d6784d..3084b0d 100644 --- a/tests/postgres/northwind_test.go +++ b/tests/postgres/northwind_test.go @@ -2,6 +2,7 @@ package postgres import ( "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/table" "github.com/stretchr/testify/require" @@ -10,30 +11,34 @@ import ( func TestNorthwindJoinEverything(t *testing.T) { - stmt := Customers. - LEFT_JOIN(CustomerCustomerDemo, Customers.CustomerID.EQ(CustomerCustomerDemo.CustomerID)). - LEFT_JOIN(CustomerDemographics, CustomerCustomerDemo.CustomerTypeID.EQ(CustomerDemographics.CustomerTypeID)). - LEFT_JOIN(Orders, Orders.CustomerID.EQ(Customers.CustomerID)). - LEFT_JOIN(Shippers, Orders.ShipVia.EQ(Shippers.ShipperID)). - LEFT_JOIN(OrderDetails, Orders.OrderID.EQ(OrderDetails.OrderID)). - LEFT_JOIN(Products, OrderDetails.ProductID.EQ(Products.ProductID)). - LEFT_JOIN(Categories, Products.CategoryID.EQ(Categories.CategoryID)). - LEFT_JOIN(Suppliers, Products.SupplierID.EQ(Suppliers.SupplierID)). - LEFT_JOIN(Employees, Orders.EmployeeID.EQ(Employees.EmployeeID)). - LEFT_JOIN(EmployeeTerritories, EmployeeTerritories.EmployeeID.EQ(Employees.EmployeeID)). - LEFT_JOIN(Territories, EmployeeTerritories.TerritoryID.EQ(Territories.TerritoryID)). - LEFT_JOIN(Region, Territories.RegionID.EQ(Region.RegionID)). - SELECT( - Customers.AllColumns, - CustomerDemographics.AllColumns, - Orders.AllColumns, - Shippers.AllColumns, - OrderDetails.AllColumns, - Products.AllColumns, - Categories.AllColumns, - Suppliers.AllColumns, - ). - ORDER_BY(Customers.CustomerID, Orders.OrderID, Products.ProductID) + 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(CustomerDemographics, CustomerCustomerDemo.CustomerTypeID.EQ(CustomerDemographics.CustomerTypeID)). + LEFT_JOIN(Orders, Orders.CustomerID.EQ(Customers.CustomerID)). + LEFT_JOIN(Shippers, Orders.ShipVia.EQ(Shippers.ShipperID)). + LEFT_JOIN(OrderDetails, Orders.OrderID.EQ(OrderDetails.OrderID)). + LEFT_JOIN(Products, OrderDetails.ProductID.EQ(Products.ProductID)). + LEFT_JOIN(Categories, Products.CategoryID.EQ(Categories.CategoryID)). + LEFT_JOIN(Suppliers, Products.SupplierID.EQ(Suppliers.SupplierID)). + LEFT_JOIN(Employees, Orders.EmployeeID.EQ(Employees.EmployeeID)). + LEFT_JOIN(EmployeeTerritories, EmployeeTerritories.EmployeeID.EQ(Employees.EmployeeID)). + LEFT_JOIN(Territories, EmployeeTerritories.TerritoryID.EQ(Territories.TerritoryID)). + LEFT_JOIN(Region, Territories.RegionID.EQ(Region.RegionID)), + ).ORDER_BY( + Customers.CustomerID, + Orders.OrderID, + Products.ProductID, + ) var dest []struct { model.Customers diff --git a/tests/postgres/range_test.go b/tests/postgres/range_test.go index b4bb712..196c4d6 100644 --- a/tests/postgres/range_test.go +++ b/tests/postgres/range_test.go @@ -1,8 +1,8 @@ package postgres import ( - "database/sql" "github.com/go-jet/jet/v2/internal/testutils" + "github.com/go-jet/jet/v2/qrm" "github.com/google/go-cmp/cmp" "github.com/jackc/pgtype" "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.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { + testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) { var dest []model.SampleRanges err := insertQuery.Query(tx, &dest) 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"; `) - testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { + testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) { var dest []model.SampleRanges err := stmt.Query(tx, &dest) @@ -351,7 +351,7 @@ SET int4_range = int4range(-12::integer, 78::integer), 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) }) }) diff --git a/tests/postgres/raw_statements_test.go b/tests/postgres/raw_statements_test.go index b60a011..9d7dd35 100644 --- a/tests/postgres/raw_statements_test.go +++ b/tests/postgres/raw_statements_test.go @@ -2,7 +2,7 @@ package postgres import ( "context" - "database/sql" + "github.com/go-jet/jet/v2/qrm" "testing" "time" @@ -128,7 +128,7 @@ RETURNING link.id AS "link.id", 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 err := stmt.Query(tx, &links) require.NoError(t, err) diff --git a/tests/postgres/sample_test.go b/tests/postgres/sample_test.go index 1df72dd..3d4d011 100644 --- a/tests/postgres/sample_test.go +++ b/tests/postgres/sample_test.go @@ -1,7 +1,7 @@ package postgres import ( - "database/sql" + "github.com/go-jet/jet/v2/qrm" "github.com/google/uuid" "testing" @@ -512,7 +512,7 @@ func TestMutableColumnsExcludeGeneratedColumn(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( People.MutableColumns, ).MODEL( diff --git a/tests/postgres/select_test.go b/tests/postgres/select_test.go index 048db14..0bcbffd 100644 --- a/tests/postgres/select_test.go +++ b/tests/postgres/select_test.go @@ -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 := ` SELECT payment.payment_id AS "payment.payment_id", payment.customer_id AS "payment.customer_id", @@ -145,7 +145,7 @@ LIMIT 30; testutils.AssertDebugStatementSql(t, query, expectedSQL, int64(30)) - dest := []model.Payment{} + var dest []model.Payment 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)) - dest := []struct{}{} + var dest []struct{} err := query.Query(db, &dest) require.NoError(t, err) } -func TestFetchFirst(t *testing.T) { +func TestSelectFetchFirst(t *testing.T) { t.Run("rows only", func(t *testing.T) { stmt := SELECT(Actor.AllColumns). @@ -320,7 +320,7 @@ FETCH FIRST ( }) } -func TestOffsetExpression(t *testing.T) { +func TestSelectOffsetExpression(t *testing.T) { stmt := SELECT(Actor.AllColumns). FROM(Actor). @@ -352,7 +352,7 @@ OFFSET ( require.Equal(t, dest[0].ActorID, int32(3)) } -func TestJoinQueryStruct(t *testing.T) { +func TestSelectJoinQueryStruct(t *testing.T) { expectedSQL := ` 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 := ` SELECT language.language_id AS "language.language_id", language.name AS "language.name", @@ -474,7 +474,7 @@ LIMIT 15; Film []model.Film } - filmsPerLanguage := []FilmsPerLanguage{} + var filmsPerLanguage []FilmsPerLanguage limit := 15 query := Film. @@ -504,7 +504,7 @@ LIMIT 15; } // https://github.com/go-jet/jet/issues/226 -func TestDuplicateSlicesInDestination(t *testing.T) { +func TestSelectDuplicateSlicesInDestination(t *testing.T) { type Staffs struct { StaffList []model.Staff @@ -645,7 +645,7 @@ func TestDuplicateSlicesInDestination(t *testing.T) { `) } -func TestExecution1(t *testing.T) { +func TestSelectExecution1(t *testing.T) { stmt := City. INNER_JOIN(Address, Address.CityID.EQ(City.CityID)). 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 { 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") } -func TestExecution3(t *testing.T) { +func TestSelectExecution3(t *testing.T) { var dest []struct { 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") } -func TestExecution4(t *testing.T) { +func TestSelectExecution4(t *testing.T) { var dest []struct { 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) -func TestExecutionCustomPKTypes1(t *testing.T) { +func TestSelectExecutionCustomPKTypes1(t *testing.T) { var dest []struct { 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) -func TestExecutionCustomPKTypes2(t *testing.T) { +func TestSelectExecutionCustomPKTypes2(t *testing.T) { var dest []struct { 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 { Language model.Language Film *[]*model.Film @@ -1160,7 +1160,7 @@ func TestJoinQuerySliceWithPtrs(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 @@ -1219,7 +1219,7 @@ func TestSelectOrderByAscDesc(t *testing.T) { testutils.AssertDeepEqual(t, customerAscDesc327, customersAscDesc[327]) } -func TestOrderBy(t *testing.T) { +func TestSelectOrderBy(t *testing.T) { t.Run("default", func(t *testing.T) { stmt := SELECT( @@ -1650,7 +1650,7 @@ LIMIT 1000; testutils.AssertDeepEqual(t, films[0], thesameLengthFilms{"Alien Center", "Iron Moon", 46}) } -func TestSubQuery(t *testing.T) { +func TestSelectSubQuery(t *testing.T) { rRatingFilms := SELECT( Film.FilmID, @@ -1816,7 +1816,7 @@ ORDER BY film.film_id ASC; testutils.AssertDebugStatementSql(t, query, expectedSQL) - maxRentalRateFilms := []model.Film{} + var maxRentalRateFilms []model.Film err := query.Query(db, &maxRentalRateFilms) 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") } -func TestGroupByGroupingSets(t *testing.T) { +func TestSelectGroupByGroupingSets(t *testing.T) { skipForCockroachDB(t) 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) stmt := SELECT( @@ -2079,7 +2079,7 @@ ORDER BY country.country, city.city; `) } -func TestGroupByRollup(t *testing.T) { +func TestSelectGroupByRollup(t *testing.T) { skipForCockroachDB(t) 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( Payment.CustomerID, @@ -2297,7 +2297,7 @@ ORDER BY customer_payment_sum.amount_sum ASC; } func TestSelectStaff(t *testing.T) { - staffs := []model.Staff{} + var staffs []model.Staff 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 := ` ( 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( SELECT(Rental.AllColumns). FROM(Rental). @@ -2474,7 +2474,7 @@ OFFSET ( 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 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 := ` SELECT * FROM dvds.address -LIMIT 3 +WHERE address.address_id < 10 FOR` - query := Address. - SELECT(STAR). - LIMIT(3) for lockType, lockTypeStr := range getRowLockTestData() { query.FOR(lockType) - testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+";\n", int64(3)) - - 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) + testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+";\n", int64(10)) + testutils.AssertExecAndRollback(t, query, db, 9) } for lockType, lockTypeStr := range getRowLockTestData() { query.FOR(lockType.NOWAIT()) - testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+" NOWAIT;\n", int64(3)) - - 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) + testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+" NOWAIT;\n", int64(10)) + testutils.AssertExecAndRollback(t, query, db, 9) } if sourceIsCockroachDB() { @@ -2628,21 +2612,13 @@ FOR` for lockType, lockTypeStr := range getRowLockTestData() { query.FOR(lockType.SKIP_LOCKED()) - testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+" SKIP LOCKED;\n", int64(3)) - - 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) + testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+" SKIP LOCKED;\n", int64(10)) + testutils.AssertExecAndRollback(t, query, db, 9) } } -func TestRowLockWithUpdateOf(t *testing.T) { +func TestSelectRowLockWithUpdateOf(t *testing.T) { + stmt := SELECT( Film.FilmID, Film.Title, @@ -2672,9 +2648,10 @@ LIMIT 1 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 { model.Film + CategoryID int 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") @@ -2708,18 +2685,18 @@ LIMIT 1 FOR UPDATE OF "myFilm"; `) - testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { + testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) { var dest []struct { model.Film `alias:"myFilm.*"` } - err := stmt.Query(db, &dest) + err := stmt.Query(tx, &dest) require.NoError(t, err) require.Len(t, dest, 1) }) } -func TestQuickStart(t *testing.T) { +func TestSelectQuickStart(t *testing.T) { var expectedSQL = ` 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") } -func TestQuickStartWithSubQueries(t *testing.T) { +func TestSelectQuickStartWithSubQueries(t *testing.T) { filmLogerThan180 := Film. SELECT(Film.AllColumns). @@ -2875,7 +2852,7 @@ func TestQuickStartWithSubQueries(t *testing.T) { testutils.AssertJSONFile(t, dest2, "./testdata/results/postgres/quick-start-dest2.json") } -func TestExpressionWrappers(t *testing.T) { +func TestSelectExpressionWrappers(t *testing.T) { query := SELECT( BoolExp(Raw("true")), IntExp(Raw("11")), @@ -2905,7 +2882,7 @@ SELECT true, require.NoError(t, err) } -func TestWindowFunction(t *testing.T) { +func TestSelectWindowFunction(t *testing.T) { var expectedSQL = ` SELECT AVG(payment.amount) OVER (), 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) } -func TestWindowClause(t *testing.T) { +func TestSelectWindowClause(t *testing.T) { var expectedSQL = ` SELECT AVG(payment.amount) OVER (), AVG(payment.amount) OVER (w1), @@ -3014,7 +2991,7 @@ ORDER BY payment.customer_id; require.NoError(t, err) } -func TestSimpleView(t *testing.T) { +func TestSelectSimpleView(t *testing.T) { query := SELECT( view.ActorInfo.AllColumns, @@ -3053,7 +3030,7 @@ func TestSimpleView(t *testing.T) { } -func TestJoinViewWithTable(t *testing.T) { +func TestSelectJoinViewWithTable(t *testing.T) { query := SELECT( view.CustomerList.AllColumns, Rental.AllColumns, @@ -3079,7 +3056,7 @@ func TestJoinViewWithTable(t *testing.T) { require.Equal(t, len(dest[1].Rentals), 27) } -func TestDynamicProjectionList(t *testing.T) { +func TestSelectDynamicProjectionList(t *testing.T) { var request struct { ColumnsToSelect []string @@ -3126,7 +3103,7 @@ LIMIT 3; require.Equal(t, len(dest), 3) } -func TestDynamicCondition(t *testing.T) { +func TestSelectDynamicCondition(t *testing.T) { var request struct { CustomerID *int64 Email *string @@ -3176,7 +3153,7 @@ WHERE ($1::boolean AND (customer.customer_id = $2)) AND (customer.activebool = $ testutils.AssertDeepEqual(t, dest[0], customer0) } -func TestLateral(t *testing.T) { +func TestSelectLateral(t *testing.T) { languages := LATERAL( SELECT( @@ -3348,7 +3325,7 @@ type ActorWrap struct { Films []FilmWrap } -func TestRecursionScanNxM(t *testing.T) { +func TestSelectRecursionScanNxM(t *testing.T) { stmt := SELECT( Actor.AllColumns, @@ -3491,7 +3468,7 @@ type StaffWrap struct { Store StoreWrap } -func TestRecursionScanNx1(t *testing.T) { +func TestSelectRecursionScanNx1(t *testing.T) { stmt := SELECT( Store.AllColumns, Staff.AllColumns, @@ -3638,7 +3615,7 @@ type ManagerInfo struct { Store *StoreInfo } -func TestRecursionScan1x1(t *testing.T) { +func TestSelectRecursionScan1x1(t *testing.T) { stmt := SELECT( 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, // 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. -func TestLiteralTypeDeduction(t *testing.T) { +func TestSelectLiteralTypeDeduction(t *testing.T) { stmt := SELECT( SUM( 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)) } -func TestCustomFunctionCall(t *testing.T) { +func TestSelectCustomFunctionCall(t *testing.T) { skipForCockroachDB(t) stmt := SELECT( @@ -3761,7 +3738,7 @@ SELECT dvds.get_film_count(100, 120) AS "film_count"; require.Equal(t, dest.FilmCount, 165) } -func TestScanUsingConn(t *testing.T) { +func TestSelectScanUsingConn(t *testing.T) { conn, err := db.Conn(context.Background()) require.NoError(t, err) defer conn.Close() @@ -3803,7 +3780,7 @@ func TestScanUsingConn(t *testing.T) { }) } -func TestConditionalFunctions(t *testing.T) { +func TestSelectConditionalFunctions(t *testing.T) { stmt := SELECT( EXISTS( Film.SELECT(Film.FilmID).WHERE(Film.RentalDuration.GT(Int(100))), diff --git a/tests/postgres/update_test.go b/tests/postgres/update_test.go index e0c7da2..91a212e 100644 --- a/tests/postgres/update_test.go +++ b/tests/postgres/update_test.go @@ -2,9 +2,9 @@ package postgres import ( "context" - "database/sql" "github.com/go-jet/jet/v2/internal/testutils" . "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" "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/table" "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; `, "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) requireLogged(t, query) @@ -142,7 +142,7 @@ RETURNING link.id AS "link.id", link.description AS "link.description"; `, "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{} err := stmt.Query(tx, &links) @@ -325,8 +325,8 @@ func TestUpdateQueryContext(t *testing.T) { time.Sleep(10 * time.Millisecond) - testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { - dest := []model.Link{} + testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) { + var dest []model.Link err := updateStmt.QueryContext(ctx, tx, &dest) require.Error(t, err, "context deadline exceeded") @@ -344,7 +344,10 @@ func TestUpdateExecContext(t *testing.T) { 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) { @@ -385,7 +388,7 @@ RETURNING rental.rental_id AS "rental.rental_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 { Rental model2.Rental Store model2.Store diff --git a/tests/sqlite/delete_test.go b/tests/sqlite/delete_test.go index 7045772..3c8c0c4 100644 --- a/tests/sqlite/delete_test.go +++ b/tests/sqlite/delete_test.go @@ -2,6 +2,7 @@ package sqlite import ( "context" + "github.com/go-jet/jet/v2/qrm" "testing" "time" @@ -60,8 +61,6 @@ LIMIT 1; } func TestDeleteContextDeadlineExceeded(t *testing.T) { - tx := beginSampleDBTx(t) - defer tx.Rollback() deleteStmt := Link. DELETE(). @@ -72,12 +71,16 @@ func TestDeleteContextDeadlineExceeded(t *testing.T) { time.Sleep(10 * time.Millisecond) - dest := []model.Link{} - err := deleteStmt.QueryContext(ctx, tx, &dest) - require.Error(t, err, "context deadline exceeded") + testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx qrm.DB) { + var dest []model.Link + err := deleteStmt.QueryContext(ctx, tx, &dest) + require.Error(t, err, "context deadline exceeded") + }) - _, err = deleteStmt.ExecContext(ctx, tx) - require.Error(t, err, "context deadline exceeded") + testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx qrm.DB) { + _, err := deleteStmt.ExecContext(ctx, tx) + require.Error(t, err, "context deadline exceeded") + }) requireLogged(t, deleteStmt) } diff --git a/tests/sqlite/insert_test.go b/tests/sqlite/insert_test.go index e1b3a54..4352ce2 100644 --- a/tests/sqlite/insert_test.go +++ b/tests/sqlite/insert_test.go @@ -2,7 +2,7 @@ package sqlite import ( "context" - "database/sql" + "github.com/go-jet/jet/v2/qrm" "math/rand" "testing" @@ -30,7 +30,7 @@ VALUES (?, ?, ?, ?), 101, "http://www.google.com", "Google", "Search engine", 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) requireLogged(t, insertQuery) @@ -75,7 +75,7 @@ VALUES (100, 'http://www.postgresqltutorial.com', 'PostgreSQL Tutorial', NULL); `, 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) require.NoError(t, err) requireLogged(t, stmt) @@ -108,7 +108,7 @@ INSERT INTO link (url, name) VALUES ('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) require.NoError(t, err) }) @@ -130,7 +130,7 @@ INSERT INTO link VALUES (1000, 'http://www.duckduckgo.com', 'Duck Duck go', NULL); `, 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) require.NoError(t, err) }) @@ -219,7 +219,7 @@ WHERE link.id = 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) require.NoError(t, err) @@ -248,7 +248,7 @@ RETURNING link.id AS "link.id", 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 err := stmt.Query(tx, &link) require.NoError(t, err) @@ -353,9 +353,15 @@ func TestInsertContextDeadlineExceeded(t *testing.T) { time.Sleep(10 * time.Millisecond) var dest []model.Link - err := stmt.QueryContext(ctx, sampleDB, &dest) - require.Error(t, err, "context deadline exceeded") - _, err = stmt.ExecContext(ctx, db) - require.Error(t, err, "context deadline exceeded") + testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx qrm.DB) { + err := stmt.QueryContext(ctx, tx, &dest) + 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") + }) + } diff --git a/tests/sqlite/main_test.go b/tests/sqlite/main_test.go index 4975845..3e06124 100644 --- a/tests/sqlite/main_test.go +++ b/tests/sqlite/main_test.go @@ -22,8 +22,8 @@ import ( _ "github.com/mattn/go-sqlite3" ) -var db *sql.DB -var sampleDB *sql.DB +var db *sqlite.DB +var sampleDB *sqlite.DB var testRoot string func TestMain(m *testing.M) { @@ -32,21 +32,30 @@ func TestMain(m *testing.M) { setTestRoot() - var err error - db, err = sql.Open("sqlite3", "file:"+dbconfig.SakilaDBPath) + sqlDB, err := sql.Open("sqlite3", "file:"+dbconfig.SakilaDBPath) throw.OnError(err) + db = sqlite.NewDB(sqlDB).WithStatementsCaching(true) defer db.Close() _, err = db.Exec(fmt.Sprintf("ATTACH DATABASE '%s' as 'chinook';", dbconfig.ChinookDBPath)) throw.OnError(err) - sampleDB, err = sql.Open("sqlite3", dbconfig.TestSampleDBPath) + sqlSampleDB, err := sql.Open("sqlite3", dbconfig.TestSampleDBPath) throw.OnError(err) + sampleDB = sqlite.NewDB(sqlSampleDB).WithStatementsCaching(true) + defer sampleDB.Close() - ret := m.Run() + for i := 0; i < 2; i++ { + ret := m.Run() + if ret != 0 { + os.Exit(ret) + } + } - if ret != 0 { - os.Exit(ret) + err = sampleDB.Clear() + + if err != nil { + panic(err) } } @@ -103,13 +112,13 @@ func requireLogged(t *testing.T, statement sqlite.Statement) { require.Equal(t, loggedDebugSQL, statement.DebugSql()) } -func beginSampleDBTx(t *testing.T) *sql.Tx { - tx, err := sampleDB.Begin() +func beginSampleDBTx(t *testing.T) *sqlite.Tx { + tx, err := sampleDB.BeginTx(context.Background(), nil) require.NoError(t, err) return tx } -func beginDBTx(t *testing.T) *sql.Tx { +func beginDBTx(t *testing.T) *sqlite.Tx { tx, err := db.Begin() require.NoError(t, err) return tx diff --git a/tests/sqlite/sample_test.go b/tests/sqlite/sample_test.go index 4775eeb..1a68b4c 100644 --- a/tests/sqlite/sample_test.go +++ b/tests/sqlite/sample_test.go @@ -1,8 +1,8 @@ package sqlite import ( - "database/sql" "github.com/go-jet/jet/v2/internal/testutils" + "github.com/go-jet/jet/v2/qrm" "github.com/stretchr/testify/require" "testing" @@ -48,7 +48,7 @@ WHERE people.people_id = ?; }) 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( People.MutableColumns, ).MODEL( diff --git a/tests/sqlite/select_test.go b/tests/sqlite/select_test.go index 63d43a9..57ee587 100644 --- a/tests/sqlite/select_test.go +++ b/tests/sqlite/select_test.go @@ -986,14 +986,14 @@ WHERE artists.''ArtistId'' = 11; } func TestRowsScan(t *testing.T) { - stmt := - SELECT( - Inventory.AllColumns, - ).FROM( - Inventory, - ).ORDER_BY( - Inventory.InventoryID.ASC(), - ) + + stmt := SELECT( + Inventory.AllColumns, + ).FROM( + Inventory, + ).ORDER_BY( + Inventory.InventoryID.ASC(), + ) rows, err := stmt.Rows(context.Background(), db) require.NoError(t, err) diff --git a/tests/sqlite/update_test.go b/tests/sqlite/update_test.go index 110c659..ad491d2 100644 --- a/tests/sqlite/update_test.go +++ b/tests/sqlite/update_test.go @@ -2,6 +2,7 @@ package sqlite import ( "context" + "github.com/go-jet/jet/v2/qrm" model2 "github.com/go-jet/jet/v2/tests/.gentestdata/sqlite/sakila/model" "github.com/go-jet/jet/v2/tests/.gentestdata/sqlite/sakila/table" "testing" @@ -215,7 +216,7 @@ WHERE link.id = 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) } @@ -271,8 +272,6 @@ func TestUpdateWithInvalidModelData(t *testing.T) { } func TestUpdateContextDeadlineExceeded(t *testing.T) { - tx := beginSampleDBTx(t) - defer tx.Rollback() updateStmt := Link.UPDATE(Link.Name, Link.URL). SET("Bong", "http://bong.com"). @@ -283,12 +282,16 @@ func TestUpdateContextDeadlineExceeded(t *testing.T) { time.Sleep(10 * time.Millisecond) - dest := []model.Link{} - err := updateStmt.QueryContext(ctx, tx, &dest) - require.Error(t, err, "context deadline exceeded") + testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) { + var dest []model.Link + err := updateStmt.QueryContext(ctx, tx, &dest) + require.Error(t, err, "context deadline exceeded") + }) - _, err = updateStmt.ExecContext(ctx, tx) - require.Error(t, err, "context deadline exceeded") + testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) { + _, err := updateStmt.ExecContext(ctx, tx) + require.Error(t, err, "context deadline exceeded") + }) } func TestUpdateFrom(t *testing.T) { From 5f220569dda43c5f75b8cf54e1d049abd8505b8b Mon Sep 17 00:00:00 2001 From: go-jet Date: Sat, 19 Oct 2024 14:06:12 +0200 Subject: [PATCH 2/2] Add support for prepared statements caching. --- .circleci/config.yml | 9 +- go.mod | 2 +- internal/testutils/test_utils.go | 19 +++- mysql/statement.go | 10 -- postgres/statement.go | 10 -- sqlite/statement.go | 10 -- {internal/jet/db => stmtcache}/db.go | 63 ++++++++---- {internal/jet/db => stmtcache}/tx.go | 11 ++- tests/mysql/main_test.go | 53 +++++++--- tests/mysql/stmtcache_test.go | 128 ++++++++++++++++++++++++ tests/postgres/alltypes_test.go | 5 +- tests/postgres/main_test.go | 54 +++++++---- tests/postgres/sample_test.go | 2 +- tests/postgres/stmtcache_test.go | 139 +++++++++++++++++++++++++++ tests/postgres/values_test.go | 4 +- tests/sqlite/delete_test.go | 2 +- tests/sqlite/insert_test.go | 3 +- tests/sqlite/main_test.go | 66 ++++++++----- tests/sqlite/stmtcache_test.go | 131 +++++++++++++++++++++++++ tests/sqlite/values_test.go | 4 +- 20 files changed, 591 insertions(+), 134 deletions(-) rename {internal/jet/db => stmtcache}/db.go (72%) rename {internal/jet/db => stmtcache}/tx.go (93%) create mode 100644 tests/mysql/stmtcache_test.go create mode 100644 tests/postgres/stmtcache_test.go create mode 100644 tests/sqlite/stmtcache_test.go diff --git a/.circleci/config.yml b/.circleci/config.yml index 74f2bd8..5db44b9 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -126,15 +126,20 @@ jobs: cd tests go run ./init/init.go -testsuite all - # to create test results report - run: name: Install gotestsum command: go install gotest.tools/gotestsum@latest + + # to create test results report - run: mkdir -p $TEST_RESULTS - run: name: Running tests - command: gotestsum --junitfile $TEST_RESULTS/report.xml --format testname -- -coverprofile=cover.out -covermode=atomic -coverpkg=github.com/go-jet/jet/v2/postgres/...,github.com/go-jet/jet/v2/mysql/...,github.com/go-jet/jet/v2/sqlite/...,github.com/go-jet/jet/v2/qrm/...,github.com/go-jet/jet/v2/generator/...,github.com/go-jet/jet/v2/internal/... ./... + command: gotestsum --junitfile $TEST_RESULTS/report.xml --format testname -- -coverprofile=cover.out -covermode=atomic -coverpkg=github.com/go-jet/jet/v2/postgres/...,github.com/go-jet/jet/v2/mysql/...,github.com/go-jet/jet/v2/sqlite/...,github.com/go-jet/jet/v2/qrm/...,github.com/go-jet/jet/v2/generator/...,github.com/go-jet/jet/v2/internal/...,github.com/go-jet/jet/v2/stmtcache/... ./... + + - run: + name: Running tests with statement caching enabled + command: JET_TESTS_WITH_STMT_CACHE=true go test -v ./tests/... # run mariaDB and cockroachdb tests. No need to collect coverage, because coverage is already included with mysql and postgres tests - run: MY_SQL_SOURCE=MariaDB go test -v ./tests/mysql/ diff --git a/go.mod b/go.mod index 890dc0d..4b48c66 100644 --- a/go.mod +++ b/go.mod @@ -1,6 +1,6 @@ module github.com/go-jet/jet/v2 -go 1.18 +go 1.20 require ( github.com/go-sql-driver/mysql v1.8.1 diff --git a/internal/testutils/test_utils.go b/internal/testutils/test_utils.go index a4a06ee..817de32 100644 --- a/internal/testutils/test_utils.go +++ b/internal/testutils/test_utils.go @@ -6,9 +6,9 @@ import ( "encoding/json" "fmt" "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/qrm" + "github.com/go-jet/jet/v2/stmtcache" "github.com/google/go-cmp/cmp" "github.com/google/uuid" "github.com/stretchr/testify/assert" @@ -26,7 +26,7 @@ var UnixTimeComparer = cmp.Comparer(func(t1, t2 time.Time) bool { }) // AssertExecAndRollback will execute and rollback statement in sql transaction -func AssertExecAndRollback(t *testing.T, stmt jet.Statement, db *jet2.DB, rowsAffected ...int64) { +func AssertExecAndRollback(t *testing.T, stmt jet.Statement, db *stmtcache.DB, rowsAffected ...int64) { tx, err := db.Begin() require.NoError(t, err) defer func() { @@ -50,8 +50,21 @@ func AssertExec(t *testing.T, stmt jet.Statement, db qrm.DB, rowsAffected ...int } } +// AssertExecContext assert statement execution for successful execution and number of rows affected +func AssertExecContext(t *testing.T, stmt jet.Statement, ctx context.Context, db qrm.DB, rowsAffected ...int64) { + res, err := stmt.ExecContext(ctx, db) + + require.NoError(t, err) + rows, err := res.RowsAffected() + require.NoError(t, err) + + if len(rowsAffected) > 0 { + require.Equal(t, rowsAffected[0], rows) + } +} + // ExecuteInTxAndRollback will execute function in sql transaction and then rollback transaction -func ExecuteInTxAndRollback(t *testing.T, db *jet2.DB, f func(tx qrm.DB)) { +func ExecuteInTxAndRollback(t *testing.T, db *stmtcache.DB, f func(tx qrm.DB)) { tx, err := db.Begin() require.NoError(t, err) defer func() { diff --git a/mysql/statement.go b/mysql/statement.go index 008ace5..5219ffb 100644 --- a/mysql/statement.go +++ b/mysql/statement.go @@ -2,19 +2,9 @@ package mysql 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 func RawStatement(rawQuery string, namedArguments ...RawArgs) jet.SerializerStatement { 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 diff --git a/postgres/statement.go b/postgres/statement.go index 4199fa9..be645c9 100644 --- a/postgres/statement.go +++ b/postgres/statement.go @@ -2,19 +2,9 @@ package postgres 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 func RawStatement(rawQuery string, namedArguments ...RawArgs) Statement { 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 diff --git a/sqlite/statement.go b/sqlite/statement.go index eb701e7..3e837cf 100644 --- a/sqlite/statement.go +++ b/sqlite/statement.go @@ -2,19 +2,9 @@ package sqlite 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 func RawStatement(rawQuery string, namedArguments ...RawArgs) Statement { 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 diff --git a/internal/jet/db/db.go b/stmtcache/db.go similarity index 72% rename from internal/jet/db/db.go rename to stmtcache/db.go index 94dfeca..a558993 100644 --- a/internal/jet/db/db.go +++ b/stmtcache/db.go @@ -1,39 +1,54 @@ -package db +package stmtcache import ( "context" "database/sql" + "errors" "fmt" "sync" ) -// DB is a wrapper around sql.DB, adding prepared statement caching capability. +// DB is a wrapper for sql.DB, providing an additional layer for caching prepared statements +// to optimize database interactions and improve performance. type DB struct { *sql.DB - statementsCaching bool + cachingEnabled bool lock sync.RWMutex statements map[string]*sql.Stmt } -// NewDB creates new DB wrapper with statements caching disabled -func NewDB(db *sql.DB) *DB { +// New creates new DB wrapper with statements caching enabled +func New(db *sql.DB) *DB { return &DB{ - DB: db, - statementsCaching: false, - statements: make(map[string]*sql.Stmt), + DB: db, + cachingEnabled: true, + statements: make(map[string]*sql.Stmt), } } -// WithStatementsCaching returns *DB wrapper with prepared statements caching enabled or disabled. This method should be +// SetCaching 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 +func (d *DB) SetCaching(enabled bool) *DB { + d.cachingEnabled = enabled return d } -// Begin starts sql transaction and returns wrapped Tx object. +// CachingEnabled returns true if statements caching is enabled +func (d *DB) CachingEnabled() bool { + return d.cachingEnabled +} + +// CacheSize returns the current number of prepared statements stored in the cache. +func (d *DB) CacheSize() int { + d.lock.RLock() + ret := len(d.statements) + d.lock.RUnlock() + return ret +} + +// Begin starts a new SQL transaction and returns a Tx object with statement caching capabilities. func (d *DB) Begin() (*Tx, error) { tx, err := d.DB.Begin() @@ -48,7 +63,7 @@ func (d *DB) Begin() (*Tx, error) { }, nil } -// BeginTx starts sql transaction and returns wrapped Tx object. +// BeginTx starts a new SQL transaction and returns a Tx object with statement caching capabilities. func (d *DB) BeginTx(ctx context.Context, opts *sql.TxOptions) (*Tx, error) { tx, err := d.DB.BeginTx(ctx, opts) @@ -73,7 +88,7 @@ func (d *DB) Exec(query string, args ...interface{}) (sql.Result, error) { // 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 { + if !d.cachingEnabled { return d.DB.ExecContext(ctx, query, args...) } @@ -95,7 +110,7 @@ func (d *DB) Query(query string, args ...interface{}) (*sql.Rows, error) { // 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 { + if !d.cachingEnabled { return d.DB.QueryContext(ctx, query, args...) } @@ -122,7 +137,7 @@ func (d *DB) Prepare(query string) (*sql.Stmt, error) { // 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 { + if !d.cachingEnabled { return d.DB.PrepareContext(ctx, query) } @@ -157,8 +172,8 @@ func (d *DB) PrepareContext(ctx context.Context, query string) (*sql.Stmt, error return prepStmt, nil } -// Clear will close all cached prepared statements -func (d *DB) Clear() error { +// ClearCache will close all cached prepared statements and clear statements cache map +func (d *DB) ClearCache() error { d.lock.Lock() defer d.lock.Unlock() @@ -168,15 +183,23 @@ func (d *DB) Clear() error { closeErr := statement.Close() if closeErr != nil { - err = closeErr + err = errors.Join(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 errors.Join(errors.New("jet: some of the prepared statements failed to close"), err) } return nil } + +// Close will clear the statements cache and close the underlying db connection +func (d *DB) Close() error { + clearErr := d.ClearCache() + closeErr := d.DB.Close() + + return errors.Join(clearErr, closeErr) +} diff --git a/internal/jet/db/tx.go b/stmtcache/tx.go similarity index 93% rename from internal/jet/db/tx.go rename to stmtcache/tx.go index b64b231..c02fb6b 100644 --- a/internal/jet/db/tx.go +++ b/stmtcache/tx.go @@ -1,4 +1,4 @@ -package db +package stmtcache import ( "context" @@ -7,6 +7,7 @@ import ( ) // Tx is a wrapper around *sql.Tx, adding prepared statement caching capability. +// Tx is not thread safe and should not be shared between goroutines. type Tx struct { *sql.Tx @@ -24,7 +25,7 @@ func (t *Tx) Exec(query string, args ...interface{}) (sql.Result, error) { // 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 { + if !t.db.cachingEnabled { return t.Tx.ExecContext(ctx, query, args...) } @@ -46,7 +47,7 @@ func (t *Tx) Query(query string, args ...interface{}) (*sql.Rows, error) { // 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 { + if !t.db.cachingEnabled { return t.Tx.QueryContext(ctx, query, args...) } @@ -73,8 +74,8 @@ func (t *Tx) Prepare(query string) (*sql.Stmt, error) { // 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) + if !t.db.cachingEnabled { + return t.Tx.PrepareContext(ctx, query) } prepStmt, ok := t.statements[query] diff --git a/tests/mysql/main_test.go b/tests/mysql/main_test.go index ee02a77..e4f4322 100644 --- a/tests/mysql/main_test.go +++ b/tests/mysql/main_test.go @@ -3,8 +3,10 @@ package mysql import ( "context" "database/sql" + "fmt" + "github.com/go-jet/jet/v2/mysql" jetmysql "github.com/go-jet/jet/v2/mysql" - "github.com/go-jet/jet/v2/postgres" + "github.com/go-jet/jet/v2/stmtcache" "github.com/go-jet/jet/v2/tests/dbconfig" _ "github.com/go-sql-driver/mysql" "github.com/stretchr/testify/require" @@ -15,14 +17,16 @@ import ( "testing" ) -var db *jetmysql.DB +var db *stmtcache.DB var source string +var withStatementCaching bool const MariaDB = "MariaDB" func init() { source = os.Getenv("MY_SQL_SOURCE") + withStatementCaching = os.Getenv("JET_TESTS_WITH_STMT_CACHE") == "true" } func sourceIsMariaDB() bool { @@ -32,21 +36,38 @@ func sourceIsMariaDB() bool { func TestMain(m *testing.M) { defer profile.Start().Stop() - var err error - sqlDB, err := sql.Open("mysql", dbconfig.MySQLConnectionString(sourceIsMariaDB(), "")) - if err != nil { - panic("Failed to connect to test db" + err.Error()) - } + func() { + fmt.Printf("\nRunning mysql tests caching enabled: %t \n", withStatementCaching) - db = jetmysql.NewDB(sqlDB).WithStatementsCaching(true) - defer db.Close() - - for i := 0; i < 2; i++ { - ret := m.Run() - if ret != 0 { - os.Exit(ret) + sqlDB, err := sql.Open("mysql", dbconfig.MySQLConnectionString(sourceIsMariaDB(), "")) + if err != nil { + panic("Failed to connect to test db" + err.Error()) } + + db = stmtcache.New(sqlDB).SetCaching(withStatementCaching) + defer db.Close() + + for i := 0; i < runCount(withStatementCaching); i++ { + ret := m.Run() + if ret != 0 { + fmt.Printf("\nFAIL: Running mysql tests failed, caching enabled: %t \n", withStatementCaching) + os.Exit(ret) + } + } + }() + +} + +func getConnectionString() string { + return dbconfig.MySQLConnectionString(sourceIsMariaDB(), "") +} + +func runCount(stmtCaching bool) int { + if stmtCaching { + return 3 } + + return 1 } var loggedSQL string @@ -70,14 +91,14 @@ func init() { }) } -func requireLogged(t *testing.T, statement postgres.Statement) { +func requireLogged(t *testing.T, statement mysql.Statement) { query, args := statement.Sql() require.Equal(t, loggedSQL, query) require.Equal(t, loggedSQLArgs, args) require.Equal(t, loggedDebugSQL, statement.DebugSql()) } -func requireQueryLogged(t *testing.T, statement postgres.Statement, rowsProcessed int64) { +func requireQueryLogged(t *testing.T, statement mysql.Statement, rowsProcessed int64) { query, args := statement.Sql() queryLogged, argsLogged := queryInfo.Statement.Sql() diff --git a/tests/mysql/stmtcache_test.go b/tests/mysql/stmtcache_test.go new file mode 100644 index 0000000..87774c1 --- /dev/null +++ b/tests/mysql/stmtcache_test.go @@ -0,0 +1,128 @@ +package mysql + +import ( + "context" + "database/sql" + "github.com/go-jet/jet/v2/internal/testutils" + . "github.com/go-jet/jet/v2/mysql" + "github.com/go-jet/jet/v2/stmtcache" + "github.com/go-jet/jet/v2/tests/.gentestdata/mysql/dvds/model" + . "github.com/go-jet/jet/v2/tests/.gentestdata/mysql/dvds/table" + "github.com/stretchr/testify/require" + "testing" +) + +func TestPreparedStatementCache(t *testing.T) { + sqlDB, err := sql.Open("mysql", getConnectionString()) + require.NoError(t, err) + stmtCachedDB := stmtcache.New(sqlDB) + defer func(db *stmtcache.DB) { + err := db.Close() + require.NoError(t, err) + require.Equal(t, db.CacheSize(), 0) + }(stmtCachedDB) + + require.True(t, stmtCachedDB.CachingEnabled()) + require.Equal(t, stmtCachedDB.CacheSize(), 0) + + testStatementCaching := func(cachingEnabled bool) { + + stmtCachedDB.SetCaching(cachingEnabled) + require.Equal(t, stmtCachedDB.CachingEnabled(), cachingEnabled) + + ctx := context.TODO() + + stmt := SELECT(Actor.AllColumns). + FROM(Actor). + WHERE(Actor.ActorID.BETWEEN(Int(1), Int(10))) + + query, args := stmt.Sql() + + preStmt, err := stmtCachedDB.Prepare(query) + require.NoError(t, err) + + preStmt2, err := stmtCachedDB.PrepareContext(ctx, query) + require.NoError(t, err) + require.Equal(t, preStmt == preStmt2, cachingEnabled) + + t.Run("Exec", func(t *testing.T) { + testutils.AssertExec(t, stmt, stmtCachedDB) + testutils.AssertExecContext(t, stmt, ctx, stmtCachedDB) + _, err := stmtCachedDB.Exec(query, args...) + require.NoError(t, err) + }) + + t.Run("Query", func(t *testing.T) { + var dest []model.Actor + + err := stmt.Query(stmtCachedDB, &dest) + require.NoError(t, err) + require.Len(t, dest, 10) + rows, err := stmtCachedDB.Query(query, args...) + rows.Close() + require.NoError(t, err) + + t.Run("ctx", func(t *testing.T) { + var dest []model.Actor + err := stmt.QueryContext(ctx, stmtCachedDB, &dest) + require.NoError(t, err) + require.Len(t, dest, 10) + }) + + }) + + t.Run("tx", func(t *testing.T) { + tx, err := stmtCachedDB.Begin() + require.NoError(t, err) + preStmtTx, err := tx.Prepare(query) + require.NoError(t, err) + _, err = preStmtTx.Exec(args...) + require.NoError(t, err) + preStmtTx2, err := tx.PrepareContext(ctx, query) + require.NoError(t, err) + require.Equal(t, preStmtTx == preStmtTx2, cachingEnabled) + _, err = preStmtTx2.ExecContext(ctx, args...) + require.NoError(t, err) + + t.Run("Exec", func(t *testing.T) { + testutils.AssertExec(t, stmt, tx) + testutils.AssertExecContext(t, stmt, ctx, tx) + + _, err := tx.Exec(query, args...) + require.NoError(t, err) + }) + + t.Run("Query", func(t *testing.T) { + var dest []model.Actor + err = stmt.QueryContext(ctx, tx, &dest) + require.NoError(t, err) + require.Len(t, dest, 10) + + rows, err := tx.Query(query, args...) + require.NoError(t, err) + require.NoError(t, rows.Close()) + }) + + t.Run("new tx", func(t *testing.T) { + txCtx, err := stmtCachedDB.BeginTx(ctx, nil) + require.NoError(t, err) + + preStmtTxCtx, err := txCtx.PrepareContext(ctx, query) + require.NoError(t, err) + require.NotEqual(t, preStmtTx, preStmtTxCtx) + + require.NoError(t, txCtx.Rollback()) + }) + + require.NoError(t, tx.Commit()) + }) + } + + testStatementCaching(true) + require.Equal(t, stmtCachedDB.CacheSize(), 1) + testStatementCaching(false) + require.Equal(t, stmtCachedDB.CacheSize(), 1) + + require.NoError(t, stmtCachedDB.ClearCache()) + require.Equal(t, stmtCachedDB.CacheSize(), 0) +} diff --git a/tests/postgres/alltypes_test.go b/tests/postgres/alltypes_test.go index 6d9c0b7..d6cdfce 100644 --- a/tests/postgres/alltypes_test.go +++ b/tests/postgres/alltypes_test.go @@ -1,7 +1,6 @@ package postgres import ( - "database/sql" "github.com/go-jet/jet/v2/internal/utils/ptr" "github.com/stretchr/testify/assert" @@ -944,7 +943,7 @@ RETURNING employee.employee_id AS "employee.employee_id", employee.manager_id AS "employee.manager_id", employee.pto_accrual AS "employee.pto_accrual"; ` - testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { + testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) { var windy model.Employee windy.PtoAccrual = ptr.Of("3h") stmt := Employee.UPDATE(Employee.PtoAccrual).SET( @@ -972,7 +971,7 @@ RETURNING employee.employee_id AS "employee.employee_id", employee.manager_id AS "employee.manager_id", employee.pto_accrual AS "employee.pto_accrual"; ` - testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { + testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) { var employee model.Employee employee.PtoAccrual = ptr.Of("5h") stmt := Employee.INSERT(Employee.AllColumns). diff --git a/tests/postgres/main_test.go b/tests/postgres/main_test.go index 1f02e69..46dfd94 100644 --- a/tests/postgres/main_test.go +++ b/tests/postgres/main_test.go @@ -4,6 +4,7 @@ import ( "context" "database/sql" "fmt" + "github.com/go-jet/jet/v2/stmtcache" "github.com/go-jet/jet/v2/tests/internal/utils/repo" "github.com/jackc/pgx/v4/stdlib" "os" @@ -19,15 +20,17 @@ import ( _ "github.com/jackc/pgx/v4/stdlib" ) -var db *postgres.DB +var db *stmtcache.DB var testRoot string var source string +var withStatementCaching bool const CockroachDB = "COCKROACH_DB" func init() { source = os.Getenv("PG_SOURCE") + withStatementCaching = os.Getenv("JET_TESTS_WITH_STMT_CACHE") == "true" } func sourceIsCockroachDB() bool { @@ -45,39 +48,50 @@ func TestMain(m *testing.M) { setTestRoot() - for _, driverName := range []string{"pgx", "postgres"} { - fmt.Printf("\nRunning postgres tests for '%s' driver\n", driverName) + for _, driverName := range []string{"postgres", "pgx"} { + + fmt.Printf("\nRunning postgres tests for driver: %s, caching enabled: %t \n", driverName, withStatementCaching) func() { - - connectionString := dbconfig.PostgresConnectString - - if sourceIsCockroachDB() { - connectionString = dbconfig.CockroachConnectString - } - - sqlDB, err := sql.Open(driverName, connectionString) + sqlDB, err := sql.Open(driverName, getConnectionString()) if err != nil { fmt.Println(err.Error()) panic("Failed to connect to test db") } - db = postgres.NewDB(sqlDB).WithStatementsCaching(true) - defer db.Close() + db = stmtcache.New(sqlDB).SetCaching(withStatementCaching) + defer func(db *stmtcache.DB) { + err := db.Close() + if err != nil { + fmt.Printf("ERROR: Failed to close db connection, %v", err) + } + }(db) - for i := 0; i < 2; i++ { + for i := 0; i < runCount(withStatementCaching); i++ { ret := m.Run() if ret != 0 { + fmt.Printf("\nFAIL: Running postgres tests failed for driver: %s, caching enabled: %t \n", driverName, withStatementCaching) os.Exit(ret) } } - - err = db.Clear() - - if err != nil { - os.Exit(-2) - } }() } + +} + +func runCount(stmtCaching bool) int { + if stmtCaching { + return 2 + } + + return 1 +} + +func getConnectionString() string { + if sourceIsCockroachDB() { + return dbconfig.CockroachConnectString + } + + return dbconfig.PostgresConnectString } func setTestRoot() { diff --git a/tests/postgres/sample_test.go b/tests/postgres/sample_test.go index 73b652c..7ac0506 100644 --- a/tests/postgres/sample_test.go +++ b/tests/postgres/sample_test.go @@ -1,8 +1,8 @@ package postgres import ( - "github.com/go-jet/jet/v2/qrm" "github.com/go-jet/jet/v2/internal/utils/ptr" + "github.com/go-jet/jet/v2/qrm" "github.com/google/uuid" "testing" diff --git a/tests/postgres/stmtcache_test.go b/tests/postgres/stmtcache_test.go new file mode 100644 index 0000000..99f12a2 --- /dev/null +++ b/tests/postgres/stmtcache_test.go @@ -0,0 +1,139 @@ +package postgres + +import ( + "context" + "database/sql" + "github.com/go-jet/jet/v2/internal/testutils" + . "github.com/go-jet/jet/v2/postgres" + "github.com/go-jet/jet/v2/stmtcache" + "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/model" + . "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/table" + "github.com/stretchr/testify/require" + "testing" +) + +func TestPreparedStatementCache(t *testing.T) { + sqlDB, err := sql.Open("postgres", getConnectionString()) + require.NoError(t, err) + stmtCachedDB := stmtcache.New(sqlDB) + defer func(db *stmtcache.DB) { + err := db.Close() + require.NoError(t, err) + require.Equal(t, db.CacheSize(), 0) + }(stmtCachedDB) + ctx := context.TODO() + + require.True(t, stmtCachedDB.CachingEnabled()) + require.Equal(t, stmtCachedDB.CacheSize(), 0) + + testStatementCaching := func(cachingEnabled bool) { + + stmtCachedDB.SetCaching(cachingEnabled) + require.Equal(t, stmtCachedDB.CachingEnabled(), cachingEnabled) + + stmt := Actor.UPDATE(). + SET(Actor.LastName.SET(Actor.LastName)). + WHERE(Actor.ActorID.BETWEEN(Int(1), Int(10))). + RETURNING(Actor.AllColumns) + + query, args := stmt.Sql() + + preStmt, err := stmtCachedDB.Prepare(query) + require.NoError(t, err) + + preStmt2, err := stmtCachedDB.PrepareContext(ctx, query) + require.NoError(t, err) + require.Equal(t, preStmt == preStmt2, cachingEnabled) + + t.Run("Exec", func(t *testing.T) { + testutils.AssertExec(t, stmt, stmtCachedDB, 10) + testutils.AssertExecContext(t, stmt, ctx, stmtCachedDB, 10) + _, err := stmtCachedDB.Exec(query, args...) + require.NoError(t, err) + }) + + t.Run("Query", func(t *testing.T) { + var dest []model.Actor + + err := stmt.Query(stmtCachedDB, &dest) + require.NoError(t, err) + require.Len(t, dest, 10) + rows, err := stmtCachedDB.Query(query, args...) + rows.Close() + require.NoError(t, err) + + t.Run("ctx", func(t *testing.T) { + var dest []model.Actor + err := stmt.QueryContext(ctx, stmtCachedDB, &dest) + require.NoError(t, err) + require.Len(t, dest, 10) + }) + + }) + + t.Run("tx", func(t *testing.T) { + tx, err := stmtCachedDB.Begin() + require.NoError(t, err) + preStmtTx, err := tx.Prepare(query) + require.NoError(t, err) + _, err = preStmtTx.Exec(args...) + require.NoError(t, err) + preStmtTx2, err := tx.PrepareContext(ctx, query) + require.NoError(t, err) + require.Equal(t, preStmtTx == preStmtTx2, cachingEnabled) + _, err = preStmtTx2.ExecContext(ctx, args...) + require.NoError(t, err) + + t.Run("Exec", func(t *testing.T) { + testutils.AssertExec(t, stmt, tx, 10) + testutils.AssertExecContext(t, stmt, ctx, tx, 10) + + _, err := tx.Exec(query, args...) + require.NoError(t, err) + }) + + t.Run("Query", func(t *testing.T) { + var dest []model.Actor + err = stmt.QueryContext(ctx, tx, &dest) + require.NoError(t, err) + require.Len(t, dest, 10) + + rows, err := tx.Query(query, args...) + require.NoError(t, err) + require.NoError(t, rows.Close()) + }) + + t.Run("new tx", func(t *testing.T) { + txCtx, err := stmtCachedDB.BeginTx(ctx, nil) + require.NoError(t, err) + + preStmtTxCtx, err := txCtx.PrepareContext(ctx, query) + require.NoError(t, err) + require.NotEqual(t, preStmtTx, preStmtTxCtx) + + require.NoError(t, txCtx.Rollback()) + }) + + require.NoError(t, tx.Commit()) + }) + + // second prepared statement + stmt2 := SELECT(Actor.AllColumns). + FROM(Actor). + WHERE(Actor.ActorID.EQ(Int(11))) + + var actor model.Actor + + err = stmt2.Query(stmtCachedDB, &actor) + require.NoError(t, err) + } + + testStatementCaching(true) + require.Equal(t, stmtCachedDB.CacheSize(), 2) + testStatementCaching(false) + require.Equal(t, stmtCachedDB.CacheSize(), 2) + + // clear all + require.NoError(t, stmtCachedDB.ClearCache()) + require.Equal(t, stmtCachedDB.CacheSize(), 0) +} diff --git a/tests/postgres/values_test.go b/tests/postgres/values_test.go index 9e89f7e..ce26276 100644 --- a/tests/postgres/values_test.go +++ b/tests/postgres/values_test.go @@ -1,9 +1,9 @@ package postgres import ( - "database/sql" "github.com/go-jet/jet/v2/internal/testutils" . "github.com/go-jet/jet/v2/postgres" + "github.com/go-jet/jet/v2/qrm" "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/model" . "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/table" "github.com/stretchr/testify/assert" @@ -251,7 +251,7 @@ RETURNING payment.payment_id AS "payment.payment_id", payment.payment_date AS "payment.payment_date"; `) - testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { + testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) { var payments []model.Payment diff --git a/tests/sqlite/delete_test.go b/tests/sqlite/delete_test.go index 3c8c0c4..f700902 100644 --- a/tests/sqlite/delete_test.go +++ b/tests/sqlite/delete_test.go @@ -69,7 +69,7 @@ func TestDeleteContextDeadlineExceeded(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), 1*time.Microsecond) defer cancel() - time.Sleep(10 * time.Millisecond) + time.Sleep(20 * time.Millisecond) testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx qrm.DB) { var dest []model.Link diff --git a/tests/sqlite/insert_test.go b/tests/sqlite/insert_test.go index c504d90..0cd7804 100644 --- a/tests/sqlite/insert_test.go +++ b/tests/sqlite/insert_test.go @@ -2,9 +2,8 @@ package sqlite import ( "context" - "github.com/go-jet/jet/v2/qrm" - "database/sql" "github.com/go-jet/jet/v2/internal/utils/ptr" + "github.com/go-jet/jet/v2/qrm" "math/rand" "testing" diff --git a/tests/sqlite/main_test.go b/tests/sqlite/main_test.go index 593d630..a113c14 100644 --- a/tests/sqlite/main_test.go +++ b/tests/sqlite/main_test.go @@ -5,8 +5,8 @@ import ( "database/sql" "fmt" "github.com/go-jet/jet/v2/internal/utils/throw" - "github.com/go-jet/jet/v2/postgres" "github.com/go-jet/jet/v2/sqlite" + "github.com/go-jet/jet/v2/stmtcache" "github.com/go-jet/jet/v2/tests/dbconfig" "github.com/pkg/profile" "github.com/stretchr/testify/require" @@ -17,38 +17,52 @@ import ( _ "github.com/mattn/go-sqlite3" ) -var db *sqlite.DB -var sampleDB *sqlite.DB -var testRoot string +var db *stmtcache.DB +var sampleDB *stmtcache.DB + +var withStatementCaching bool + +func init() { + withStatementCaching = os.Getenv("JET_TESTS_WITH_STMT_CACHE") == "true" +} func TestMain(m *testing.M) { defer profile.Start().Stop() - sqlDB, err := sql.Open("sqlite3", "file:"+dbconfig.SakilaDBPath) - throw.OnError(err) - db = sqlite.NewDB(sqlDB).WithStatementsCaching(true) - defer db.Close() + func() { + fmt.Printf("\nRunning sqlite tests caching enabled: %t \n", withStatementCaching) - _, err = db.Exec(fmt.Sprintf("ATTACH DATABASE '%s' as 'chinook';", dbconfig.ChinookDBPath)) - throw.OnError(err) + sqlDB, err := sql.Open("sqlite3", "file:"+dbconfig.SakilaDBPath) + throw.OnError(err) + db = stmtcache.New(sqlDB).SetCaching(withStatementCaching) + defer db.Close() - sqlSampleDB, err := sql.Open("sqlite3", dbconfig.TestSampleDBPath) - throw.OnError(err) - sampleDB = sqlite.NewDB(sqlSampleDB).WithStatementsCaching(true) - defer sampleDB.Close() + _, err = db.Exec(fmt.Sprintf("ATTACH DATABASE '%s' as 'chinook';", dbconfig.ChinookDBPath)) + throw.OnError(err) - for i := 0; i < 2; i++ { - ret := m.Run() - if ret != 0 { - os.Exit(ret) + sqlSampleDB, err := sql.Open("sqlite3", dbconfig.TestSampleDBPath) + throw.OnError(err) + sampleDB = stmtcache.New(sqlSampleDB).SetCaching(withStatementCaching) + defer sampleDB.Close() + + for i := 0; i < runCount(withStatementCaching); i++ { + ret := m.Run() + if ret != 0 { + fmt.Printf("\nFAIL: Running sqlite tests failed, caching enabled: %t \n", withStatementCaching) + os.Exit(ret) + } } + + }() + +} + +func runCount(stmtCaching bool) int { + if stmtCaching { + return 4 } - err = sampleDB.Clear() - - if err != nil { - panic(err) - } + return 1 } var loggedSQL string @@ -72,7 +86,7 @@ func init() { }) } -func requireQueryLogged(t *testing.T, statement postgres.Statement, rowsProcessed int64) { +func requireQueryLogged(t *testing.T, statement sqlite.Statement, rowsProcessed int64) { query, args := statement.Sql() queryLogged, argsLogged := queryInfo.Statement.Sql() @@ -94,13 +108,13 @@ func requireLogged(t *testing.T, statement sqlite.Statement) { require.Equal(t, loggedDebugSQL, statement.DebugSql()) } -func beginSampleDBTx(t *testing.T) *sqlite.Tx { +func beginSampleDBTx(t *testing.T) *stmtcache.Tx { tx, err := sampleDB.BeginTx(context.Background(), nil) require.NoError(t, err) return tx } -func beginDBTx(t *testing.T) *sqlite.Tx { +func beginDBTx(t *testing.T) *stmtcache.Tx { tx, err := db.Begin() require.NoError(t, err) return tx diff --git a/tests/sqlite/stmtcache_test.go b/tests/sqlite/stmtcache_test.go new file mode 100644 index 0000000..be8cb91 --- /dev/null +++ b/tests/sqlite/stmtcache_test.go @@ -0,0 +1,131 @@ +package sqlite + +import ( + "context" + "database/sql" + "github.com/go-jet/jet/v2/internal/testutils" + . "github.com/go-jet/jet/v2/sqlite" + "github.com/go-jet/jet/v2/stmtcache" + "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/dbconfig" + "github.com/stretchr/testify/require" + "testing" +) + +func TestPreparedStatementCache(t *testing.T) { + sqlDB, err := sql.Open("sqlite3", "file:"+dbconfig.SakilaDBPath) + require.NoError(t, err) + stmtCachedDB := stmtcache.New(sqlDB) + defer func(db *stmtcache.DB) { + err := db.Close() + require.NoError(t, err) + require.Equal(t, db.CacheSize(), 0) + }(stmtCachedDB) + + require.True(t, stmtCachedDB.CachingEnabled()) + require.Equal(t, stmtCachedDB.CacheSize(), 0) + + testStatementCaching := func(cachingEnabled bool) { + + stmtCachedDB.SetCaching(cachingEnabled) + require.Equal(t, stmtCachedDB.CachingEnabled(), cachingEnabled) + + ctx := context.TODO() + + stmt := SELECT(Actor.AllColumns). + FROM(Actor). + WHERE(Actor.ActorID.BETWEEN(Int(1), Int(10))) + + query, args := stmt.Sql() + + preStmt, err := stmtCachedDB.Prepare(query) + require.NoError(t, err) + + preStmt2, err := stmtCachedDB.PrepareContext(ctx, query) + require.NoError(t, err) + require.Equal(t, preStmt == preStmt2, cachingEnabled) + + t.Run("Exec", func(t *testing.T) { + testutils.AssertExec(t, stmt, stmtCachedDB) + testutils.AssertExecContext(t, stmt, ctx, stmtCachedDB) + _, err := stmtCachedDB.Exec(query, args...) + require.NoError(t, err) + }) + + t.Run("Query", func(t *testing.T) { + var dest []model.Actor + + err := stmt.Query(stmtCachedDB, &dest) + require.NoError(t, err) + require.Len(t, dest, 10) + rows, err := stmtCachedDB.Query(query, args...) + rows.Close() + require.NoError(t, err) + + t.Run("ctx", func(t *testing.T) { + var dest []model.Actor + err := stmt.QueryContext(ctx, stmtCachedDB, &dest) + require.NoError(t, err) + require.Len(t, dest, 10) + }) + + }) + + t.Run("tx", func(t *testing.T) { + tx, err := stmtCachedDB.Begin() + require.NoError(t, err) + preStmtTx, err := tx.Prepare(query) + require.NoError(t, err) + _, err = preStmtTx.Exec(args...) + require.NoError(t, err) + preStmtTx2, err := tx.PrepareContext(ctx, query) + require.NoError(t, err) + require.Equal(t, preStmtTx == preStmtTx2, cachingEnabled) + _, err = preStmtTx2.ExecContext(ctx, args...) + require.NoError(t, err) + + t.Run("Exec", func(t *testing.T) { + testutils.AssertExec(t, stmt, tx) + testutils.AssertExecContext(t, stmt, ctx, tx) + + _, err := tx.Exec(query, args...) + require.NoError(t, err) + }) + + t.Run("Query", func(t *testing.T) { + var dest []model.Actor + err = stmt.QueryContext(ctx, tx, &dest) + require.NoError(t, err) + require.Len(t, dest, 10) + + rows, err := tx.Query(query, args...) + require.NoError(t, err) + require.NoError(t, rows.Close()) + }) + + t.Run("new tx", func(t *testing.T) { + txCtx, err := stmtCachedDB.BeginTx(ctx, nil) + require.NoError(t, err) + + preStmtTxCtx, err := txCtx.PrepareContext(ctx, query) + require.NoError(t, err) + require.NotEqual(t, preStmtTx, preStmtTxCtx) + + require.NoError(t, txCtx.Rollback()) + }) + + require.NoError(t, preStmtTx.Close()) + require.NoError(t, preStmtTx2.Close()) + require.NoError(t, tx.Commit()) + }) + } + + testStatementCaching(true) + require.Equal(t, stmtCachedDB.CacheSize(), 1) + testStatementCaching(false) + require.Equal(t, stmtCachedDB.CacheSize(), 1) + + require.NoError(t, stmtCachedDB.ClearCache()) + require.Equal(t, stmtCachedDB.CacheSize(), 0) +} diff --git a/tests/sqlite/values_test.go b/tests/sqlite/values_test.go index 0793397..948d211 100644 --- a/tests/sqlite/values_test.go +++ b/tests/sqlite/values_test.go @@ -1,8 +1,8 @@ package sqlite import ( - "database/sql" "github.com/go-jet/jet/v2/internal/testutils" + "github.com/go-jet/jet/v2/qrm" "github.com/stretchr/testify/require" "strings" "testing" @@ -293,7 +293,7 @@ RETURNING payment.payment_id AS "payment.payment_id", payment.last_update AS "payment.last_update"; `, "''", "`")) - testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) { + testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) { var payments []model.Payment err := stmt.Query(tx, &payments)