Merge pull request #416 from go-jet/stmt-cache2
Add support for prepared statement caching
This commit is contained in:
commit
49104d1969
36 changed files with 1067 additions and 301 deletions
|
|
@ -126,15 +126,20 @@ jobs:
|
||||||
cd tests
|
cd tests
|
||||||
go run ./init/init.go -testsuite all
|
go run ./init/init.go -testsuite all
|
||||||
|
|
||||||
# to create test results report
|
|
||||||
- run:
|
- run:
|
||||||
name: Install gotestsum
|
name: Install gotestsum
|
||||||
command: go install gotest.tools/gotestsum@latest
|
command: go install gotest.tools/gotestsum@latest
|
||||||
|
|
||||||
|
# to create test results report
|
||||||
- run: mkdir -p $TEST_RESULTS
|
- run: mkdir -p $TEST_RESULTS
|
||||||
|
|
||||||
- run:
|
- run:
|
||||||
name: Running tests
|
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 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/
|
- run: MY_SQL_SOURCE=MariaDB go test -v ./tests/mysql/
|
||||||
|
|
|
||||||
2
go.mod
2
go.mod
|
|
@ -1,6 +1,6 @@
|
||||||
module github.com/go-jet/jet/v2
|
module github.com/go-jet/jet/v2
|
||||||
|
|
||||||
go 1.18
|
go 1.20
|
||||||
|
|
||||||
require (
|
require (
|
||||||
github.com/go-sql-driver/mysql v1.8.1
|
github.com/go-sql-driver/mysql v1.8.1
|
||||||
|
|
|
||||||
|
|
@ -3,12 +3,12 @@ package testutils
|
||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
"context"
|
"context"
|
||||||
"database/sql"
|
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
"github.com/go-jet/jet/v2/internal/jet"
|
"github.com/go-jet/jet/v2/internal/jet"
|
||||||
"github.com/go-jet/jet/v2/internal/utils/throw"
|
"github.com/go-jet/jet/v2/internal/utils/throw"
|
||||||
"github.com/go-jet/jet/v2/qrm"
|
"github.com/go-jet/jet/v2/qrm"
|
||||||
|
"github.com/go-jet/jet/v2/stmtcache"
|
||||||
"github.com/google/go-cmp/cmp"
|
"github.com/google/go-cmp/cmp"
|
||||||
"github.com/google/uuid"
|
"github.com/google/uuid"
|
||||||
"github.com/stretchr/testify/assert"
|
"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
|
// 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 *stmtcache.DB, rowsAffected ...int64) {
|
||||||
tx, err := db.Begin()
|
tx, err := db.Begin()
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
defer func() {
|
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
|
// 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 *stmtcache.DB, f func(tx qrm.DB)) {
|
||||||
tx, err := db.Begin()
|
tx, err := db.Begin()
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
defer func() {
|
defer func() {
|
||||||
|
|
@ -132,7 +145,7 @@ func AssertJSONFile(t *testing.T, data interface{}, testRelativePath string) {
|
||||||
}
|
}
|
||||||
|
|
||||||
// AssertStatementSql check if statement Sql() is the same as expectedQuery and expectedArgs
|
// AssertStatementSql check if statement Sql() is the same as expectedQuery and expectedArgs
|
||||||
func AssertStatementSql(t *testing.T, query jet.Statement, expectedQuery string, expectedArgs ...interface{}) {
|
func AssertStatementSql(t *testing.T, query jet.PrintableStatement, expectedQuery string, expectedArgs ...interface{}) {
|
||||||
queryStr, args := query.Sql()
|
queryStr, args := query.Sql()
|
||||||
assertQueryString(t, queryStr, expectedQuery)
|
assertQueryString(t, queryStr, expectedQuery)
|
||||||
|
|
||||||
|
|
@ -153,7 +166,7 @@ func AssertStatementSqlErr(t *testing.T, stmt jet.Statement, errorStr string) {
|
||||||
}
|
}
|
||||||
|
|
||||||
// AssertDebugStatementSql check if statement Sql() is the same as expectedQuery
|
// AssertDebugStatementSql check if statement Sql() is the same as expectedQuery
|
||||||
func AssertDebugStatementSql(t *testing.T, query jet.Statement, expectedQuery string, expectedArgs ...interface{}) {
|
func AssertDebugStatementSql(t *testing.T, query jet.PrintableStatement, expectedQuery string, expectedArgs ...interface{}) {
|
||||||
_, args := query.Sql()
|
_, args := query.Sql()
|
||||||
|
|
||||||
if len(expectedArgs) > 0 {
|
if len(expectedArgs) > 0 {
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,8 @@
|
||||||
package mysql
|
package mysql
|
||||||
|
|
||||||
import "github.com/go-jet/jet/v2/internal/jet"
|
import (
|
||||||
|
"github.com/go-jet/jet/v2/internal/jet"
|
||||||
|
)
|
||||||
|
|
||||||
// RawStatement creates new sql statements from raw query and optional map of named arguments
|
// RawStatement creates new sql statements from raw query and optional map of named arguments
|
||||||
func RawStatement(rawQuery string, namedArguments ...RawArgs) jet.SerializerStatement {
|
func RawStatement(rawQuery string, namedArguments ...RawArgs) jet.SerializerStatement {
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,8 @@
|
||||||
package postgres
|
package postgres
|
||||||
|
|
||||||
import "github.com/go-jet/jet/v2/internal/jet"
|
import (
|
||||||
|
"github.com/go-jet/jet/v2/internal/jet"
|
||||||
|
)
|
||||||
|
|
||||||
// RawStatement creates new sql statements from raw query and optional map of named arguments
|
// RawStatement creates new sql statements from raw query and optional map of named arguments
|
||||||
func RawStatement(rawQuery string, namedArguments ...RawArgs) Statement {
|
func RawStatement(rawQuery string, namedArguments ...RawArgs) Statement {
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,8 @@
|
||||||
package sqlite
|
package sqlite
|
||||||
|
|
||||||
import "github.com/go-jet/jet/v2/internal/jet"
|
import (
|
||||||
|
"github.com/go-jet/jet/v2/internal/jet"
|
||||||
|
)
|
||||||
|
|
||||||
// RawStatement creates new sql statements from raw query and optional map of named arguments
|
// RawStatement creates new sql statements from raw query and optional map of named arguments
|
||||||
func RawStatement(rawQuery string, namedArguments ...RawArgs) Statement {
|
func RawStatement(rawQuery string, namedArguments ...RawArgs) Statement {
|
||||||
|
|
|
||||||
205
stmtcache/db.go
Normal file
205
stmtcache/db.go
Normal file
|
|
@ -0,0 +1,205 @@
|
||||||
|
package stmtcache
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"sync"
|
||||||
|
)
|
||||||
|
|
||||||
|
// 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
|
||||||
|
|
||||||
|
cachingEnabled bool
|
||||||
|
|
||||||
|
lock sync.RWMutex
|
||||||
|
statements map[string]*sql.Stmt
|
||||||
|
}
|
||||||
|
|
||||||
|
// New creates new DB wrapper with statements caching enabled
|
||||||
|
func New(db *sql.DB) *DB {
|
||||||
|
return &DB{
|
||||||
|
DB: db,
|
||||||
|
cachingEnabled: true,
|
||||||
|
statements: make(map[string]*sql.Stmt),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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) SetCaching(enabled bool) *DB {
|
||||||
|
d.cachingEnabled = enabled
|
||||||
|
return d
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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()
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
return &Tx{
|
||||||
|
Tx: tx,
|
||||||
|
db: d,
|
||||||
|
statements: make(map[string]*sql.Stmt),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// 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)
|
||||||
|
|
||||||
|
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.cachingEnabled {
|
||||||
|
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.cachingEnabled {
|
||||||
|
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.cachingEnabled {
|
||||||
|
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
|
||||||
|
}
|
||||||
|
|
||||||
|
// ClearCache will close all cached prepared statements and clear statements cache map
|
||||||
|
func (d *DB) ClearCache() error {
|
||||||
|
d.lock.Lock()
|
||||||
|
defer d.lock.Unlock()
|
||||||
|
|
||||||
|
var err error
|
||||||
|
|
||||||
|
for _, statement := range d.statements {
|
||||||
|
closeErr := statement.Close()
|
||||||
|
|
||||||
|
if closeErr != nil {
|
||||||
|
err = errors.Join(err, closeErr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
d.statements = make(map[string]*sql.Stmt)
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
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)
|
||||||
|
}
|
||||||
98
stmtcache/tx.go
Normal file
98
stmtcache/tx.go
Normal file
|
|
@ -0,0 +1,98 @@
|
||||||
|
package stmtcache
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"fmt"
|
||||||
|
)
|
||||||
|
|
||||||
|
// 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
|
||||||
|
|
||||||
|
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.cachingEnabled {
|
||||||
|
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.cachingEnabled {
|
||||||
|
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.cachingEnabled {
|
||||||
|
return t.Tx.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
|
||||||
|
}
|
||||||
|
|
@ -2,9 +2,9 @@ package mysql
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"database/sql"
|
|
||||||
"github.com/go-jet/jet/v2/internal/testutils"
|
"github.com/go-jet/jet/v2/internal/testutils"
|
||||||
. "github.com/go-jet/jet/v2/mysql"
|
. "github.com/go-jet/jet/v2/mysql"
|
||||||
|
"github.com/go-jet/jet/v2/qrm"
|
||||||
"github.com/go-jet/jet/v2/tests/.gentestdata/mysql/dvds/table"
|
"github.com/go-jet/jet/v2/tests/.gentestdata/mysql/dvds/table"
|
||||||
"github.com/go-jet/jet/v2/tests/.gentestdata/mysql/test_sample/model"
|
"github.com/go-jet/jet/v2/tests/.gentestdata/mysql/test_sample/model"
|
||||||
. "github.com/go-jet/jet/v2/tests/.gentestdata/mysql/test_sample/table"
|
. "github.com/go-jet/jet/v2/tests/.gentestdata/mysql/test_sample/table"
|
||||||
|
|
@ -113,7 +113,7 @@ DELETE /*+ QB_NAME(deleteIns) MRR(link) */ FROM test_sample.link
|
||||||
WHERE link.name IN ('Gmail', 'Outlook');
|
WHERE link.name IN ('Gmail', 'Outlook');
|
||||||
`)
|
`)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
_, err := stmt.Exec(tx)
|
_, err := stmt.Exec(tx)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
})
|
})
|
||||||
|
|
|
||||||
|
|
@ -2,10 +2,10 @@ package mysql
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"database/sql"
|
|
||||||
"github.com/go-jet/jet/v2/internal/testutils"
|
"github.com/go-jet/jet/v2/internal/testutils"
|
||||||
"github.com/go-jet/jet/v2/internal/utils/ptr"
|
"github.com/go-jet/jet/v2/internal/utils/ptr"
|
||||||
. "github.com/go-jet/jet/v2/mysql"
|
. "github.com/go-jet/jet/v2/mysql"
|
||||||
|
"github.com/go-jet/jet/v2/qrm"
|
||||||
"github.com/go-jet/jet/v2/tests/.gentestdata/mysql/test_sample/model"
|
"github.com/go-jet/jet/v2/tests/.gentestdata/mysql/test_sample/model"
|
||||||
. "github.com/go-jet/jet/v2/tests/.gentestdata/mysql/test_sample/table"
|
. "github.com/go-jet/jet/v2/tests/.gentestdata/mysql/test_sample/table"
|
||||||
|
|
||||||
|
|
@ -31,7 +31,7 @@ VALUES (100, 'http://www.postgresqltutorial.com', 'PostgreSQL Tutorial', DEFAULT
|
||||||
101, "http://www.google.com", "Google",
|
101, "http://www.google.com", "Google",
|
||||||
102, "http://www.yahoo.com", "Yahoo", nil)
|
102, "http://www.yahoo.com", "Yahoo", nil)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
_, err := insertQuery.Exec(tx)
|
_, err := insertQuery.Exec(tx)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
requireLogged(t, insertQuery)
|
requireLogged(t, insertQuery)
|
||||||
|
|
@ -76,7 +76,7 @@ VALUES (100, 'http://www.postgresqltutorial.com', 'PostgreSQL Tutorial', DEFAULT
|
||||||
`,
|
`,
|
||||||
100, "http://www.postgresqltutorial.com", "PostgreSQL Tutorial")
|
100, "http://www.postgresqltutorial.com", "PostgreSQL Tutorial")
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
_, err := stmt.Exec(tx)
|
_, err := stmt.Exec(tx)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
requireLogged(t, stmt)
|
requireLogged(t, stmt)
|
||||||
|
|
@ -110,7 +110,7 @@ VALUES ('http://www.duckduckgo.com', 'Duck Duck go');
|
||||||
`,
|
`,
|
||||||
"http://www.duckduckgo.com", "Duck Duck go")
|
"http://www.duckduckgo.com", "Duck Duck go")
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
_, err := query.Exec(tx)
|
_, err := query.Exec(tx)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
})
|
})
|
||||||
|
|
@ -132,7 +132,7 @@ INSERT INTO test_sample.link
|
||||||
VALUES (1000, 'http://www.duckduckgo.com', 'Duck Duck go', NULL);
|
VALUES (1000, 'http://www.duckduckgo.com', 'Duck Duck go', NULL);
|
||||||
`, int32(1000), "http://www.duckduckgo.com", "Duck Duck go", nil)
|
`, int32(1000), "http://www.duckduckgo.com", "Duck Duck go", nil)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
_, err := query.Exec(tx)
|
_, err := query.Exec(tx)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
})
|
})
|
||||||
|
|
@ -168,7 +168,7 @@ VALUES ('http://www.postgresqltutorial.com', 'PostgreSQL Tutorial'),
|
||||||
"http://www.google.com", "Google",
|
"http://www.google.com", "Google",
|
||||||
"http://www.yahoo.com", "Yahoo")
|
"http://www.yahoo.com", "Yahoo")
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
_, err := query.Exec(tx)
|
_, err := query.Exec(tx)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
})
|
})
|
||||||
|
|
@ -203,7 +203,7 @@ VALUES ('http://www.postgresqltutorial.com', 'PostgreSQL Tutorial', DEFAULT),
|
||||||
"http://www.google.com", "Google", nil,
|
"http://www.google.com", "Google", nil,
|
||||||
"http://www.yahoo.com", "Yahoo", nil)
|
"http://www.yahoo.com", "Yahoo", nil)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
_, err := stmt.Exec(tx)
|
_, err := stmt.Exec(tx)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
})
|
})
|
||||||
|
|
@ -227,7 +227,7 @@ INSERT INTO test_sample.link (url, name) (
|
||||||
);
|
);
|
||||||
`, int64(1))
|
`, int64(1))
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
_, err := query.Exec(tx)
|
_, err := query.Exec(tx)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
|
@ -263,7 +263,7 @@ ON DUPLICATE KEY UPDATE id = (link.id + ?),
|
||||||
randId, "http://www.postgresqltutorial.com", "PostgreSQL Tutorial",
|
randId, "http://www.postgresqltutorial.com", "PostgreSQL Tutorial",
|
||||||
int64(11), "PostgreSQL Tutorial 2")
|
int64(11), "PostgreSQL Tutorial 2")
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
_, err := stmt.Exec(tx)
|
_, err := stmt.Exec(tx)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
|
@ -322,7 +322,7 @@ ON DUPLICATE KEY UPDATE id = (link.id + ?),
|
||||||
description = new.description;
|
description = new.description;
|
||||||
`)
|
`)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
_, err := stmt.Exec(tx)
|
_, err := stmt.Exec(tx)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
|
@ -354,9 +354,12 @@ func TestInsertWithQueryContext(t *testing.T) {
|
||||||
time.Sleep(10 * time.Millisecond)
|
time.Sleep(10 * time.Millisecond)
|
||||||
|
|
||||||
var dest []model.Link
|
var dest []model.Link
|
||||||
err := stmt.QueryContext(ctx, db, &dest)
|
|
||||||
|
|
||||||
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) {
|
func TestInsertWithExecContext(t *testing.T) {
|
||||||
|
|
@ -368,9 +371,10 @@ func TestInsertWithExecContext(t *testing.T) {
|
||||||
|
|
||||||
time.Sleep(10 * time.Millisecond)
|
time.Sleep(10 * time.Millisecond)
|
||||||
|
|
||||||
_, err := stmt.ExecContext(ctx, db)
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
|
_, err := stmt.ExecContext(ctx, tx)
|
||||||
require.Error(t, err, "context deadline exceeded")
|
require.Error(t, err, "context deadline exceeded")
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestInsertOptimizerHints(t *testing.T) {
|
func TestInsertOptimizerHints(t *testing.T) {
|
||||||
|
|
@ -387,7 +391,7 @@ INSERT /*+ QB_NAME(qbIns) NO_ICP(link) */ INTO test_sample.link (url, name, desc
|
||||||
VALUES ('http://www.google.com', 'Google', NULL);
|
VALUES ('http://www.google.com', 'Google', NULL);
|
||||||
`)
|
`)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
_, err := stmt.Exec(tx)
|
_, err := stmt.Exec(tx)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
})
|
})
|
||||||
|
|
|
||||||
|
|
@ -14,10 +14,14 @@ func TestLockRead(t *testing.T) {
|
||||||
testutils.AssertStatementSql(t, query, `
|
testutils.AssertStatementSql(t, query, `
|
||||||
LOCK TABLES dvds.customer READ;
|
LOCK TABLES dvds.customer READ;
|
||||||
`)
|
`)
|
||||||
|
tx, err := db.DB.Begin() // can't prepare LOCK statement
|
||||||
_, err := query.Exec(db)
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
requireLogged(t, query)
|
defer func() {
|
||||||
|
err := tx.Rollback()
|
||||||
|
require.NoError(t, err)
|
||||||
|
}()
|
||||||
|
|
||||||
|
testutils.AssertExec(t, query, tx)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestLockWrite(t *testing.T) {
|
func TestLockWrite(t *testing.T) {
|
||||||
|
|
@ -27,9 +31,14 @@ func TestLockWrite(t *testing.T) {
|
||||||
LOCK TABLES dvds.customer WRITE;
|
LOCK TABLES dvds.customer WRITE;
|
||||||
`)
|
`)
|
||||||
|
|
||||||
_, err := query.Exec(db)
|
tx, err := db.DB.Begin() // can't prepare LOCK statement
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
requireLogged(t, query)
|
defer func() {
|
||||||
|
err := tx.Rollback()
|
||||||
|
require.NoError(t, err)
|
||||||
|
}()
|
||||||
|
|
||||||
|
testutils.AssertExec(t, query, tx)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestUnlockTables(t *testing.T) {
|
func TestUnlockTables(t *testing.T) {
|
||||||
|
|
@ -39,7 +48,12 @@ func TestUnlockTables(t *testing.T) {
|
||||||
UNLOCK TABLES;
|
UNLOCK TABLES;
|
||||||
`)
|
`)
|
||||||
|
|
||||||
_, err := query.Exec(db)
|
tx, err := db.DB.Begin() // can't prepare LOCK statement
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
requireLogged(t, query)
|
defer func() {
|
||||||
|
err := tx.Rollback()
|
||||||
|
require.NoError(t, err)
|
||||||
|
}()
|
||||||
|
|
||||||
|
testutils.AssertExec(t, query, tx)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -3,8 +3,10 @@ package mysql
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"database/sql"
|
"database/sql"
|
||||||
|
"fmt"
|
||||||
|
"github.com/go-jet/jet/v2/mysql"
|
||||||
jetmysql "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-jet/jet/v2/tests/dbconfig"
|
||||||
_ "github.com/go-sql-driver/mysql"
|
_ "github.com/go-sql-driver/mysql"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
|
|
@ -15,14 +17,16 @@ import (
|
||||||
"testing"
|
"testing"
|
||||||
)
|
)
|
||||||
|
|
||||||
var db *sql.DB
|
var db *stmtcache.DB
|
||||||
|
|
||||||
var source string
|
var source string
|
||||||
|
var withStatementCaching bool
|
||||||
|
|
||||||
const MariaDB = "MariaDB"
|
const MariaDB = "MariaDB"
|
||||||
|
|
||||||
func init() {
|
func init() {
|
||||||
source = os.Getenv("MY_SQL_SOURCE")
|
source = os.Getenv("MY_SQL_SOURCE")
|
||||||
|
withStatementCaching = os.Getenv("JET_TESTS_WITH_STMT_CACHE") == "true"
|
||||||
}
|
}
|
||||||
|
|
||||||
func sourceIsMariaDB() bool {
|
func sourceIsMariaDB() bool {
|
||||||
|
|
@ -32,16 +36,38 @@ func sourceIsMariaDB() bool {
|
||||||
func TestMain(m *testing.M) {
|
func TestMain(m *testing.M) {
|
||||||
defer profile.Start().Stop()
|
defer profile.Start().Stop()
|
||||||
|
|
||||||
var err error
|
func() {
|
||||||
db, err = sql.Open("mysql", dbconfig.MySQLConnectionString(sourceIsMariaDB(), ""))
|
fmt.Printf("\nRunning mysql tests caching enabled: %t \n", withStatementCaching)
|
||||||
if err != nil {
|
|
||||||
panic("Failed to connect to test db" + err.Error())
|
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
|
||||||
}
|
}
|
||||||
defer db.Close()
|
|
||||||
|
|
||||||
ret := m.Run()
|
return 1
|
||||||
|
|
||||||
os.Exit(ret)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
var loggedSQL string
|
var loggedSQL string
|
||||||
|
|
@ -65,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()
|
query, args := statement.Sql()
|
||||||
require.Equal(t, loggedSQL, query)
|
require.Equal(t, loggedSQL, query)
|
||||||
require.Equal(t, loggedSQLArgs, args)
|
require.Equal(t, loggedSQLArgs, args)
|
||||||
require.Equal(t, loggedDebugSQL, statement.DebugSql())
|
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()
|
query, args := statement.Sql()
|
||||||
queryLogged, argsLogged := queryInfo.Statement.Sql()
|
queryLogged, argsLogged := queryInfo.Statement.Sql()
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -2,8 +2,8 @@ package mysql
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"database/sql"
|
|
||||||
"github.com/go-jet/jet/v2/postgres"
|
"github.com/go-jet/jet/v2/postgres"
|
||||||
|
"github.com/go-jet/jet/v2/qrm"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
@ -868,11 +868,13 @@ func TestRowLock(t *testing.T) {
|
||||||
expectedSQL := `
|
expectedSQL := `
|
||||||
SELECT *
|
SELECT *
|
||||||
FROM dvds.address
|
FROM dvds.address
|
||||||
|
ORDER BY address.address_id
|
||||||
LIMIT 3
|
LIMIT 3
|
||||||
OFFSET 1
|
OFFSET 1
|
||||||
FOR`
|
FOR`
|
||||||
query := Address.
|
query := SELECT(STAR).
|
||||||
SELECT(STAR).
|
FROM(Address).
|
||||||
|
ORDER_BY(Address.AddressID).
|
||||||
LIMIT(3).
|
LIMIT(3).
|
||||||
OFFSET(1)
|
OFFSET(1)
|
||||||
|
|
||||||
|
|
@ -881,28 +883,14 @@ FOR`
|
||||||
|
|
||||||
expectedQuery := expectedSQL + " " + lockTypeStr + ";\n"
|
expectedQuery := expectedSQL + " " + lockTypeStr + ";\n"
|
||||||
testutils.AssertDebugStatementSql(t, query, expectedQuery, int64(3), int64(1))
|
testutils.AssertDebugStatementSql(t, query, expectedQuery, int64(3), int64(1))
|
||||||
|
testutils.AssertExecAndRollback(t, query, db)
|
||||||
tx, _ := db.Begin()
|
|
||||||
|
|
||||||
_, err := query.Exec(tx)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
err = tx.Rollback()
|
|
||||||
require.NoError(t, err)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
for lockType, lockTypeStr := range getRowLockTestData() {
|
for lockType, lockTypeStr := range getRowLockTestData() {
|
||||||
query.FOR(lockType.NOWAIT())
|
query.FOR(lockType.NOWAIT())
|
||||||
|
|
||||||
testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+" NOWAIT;\n", int64(3), int64(1))
|
testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+" NOWAIT;\n", int64(3), int64(1))
|
||||||
|
testutils.AssertExecAndRollback(t, query, db)
|
||||||
tx, _ := db.Begin()
|
|
||||||
|
|
||||||
_, err := query.Exec(tx)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
err = tx.Rollback()
|
|
||||||
require.NoError(t, err)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if sourceIsMariaDB() {
|
if sourceIsMariaDB() {
|
||||||
|
|
@ -913,14 +901,7 @@ FOR`
|
||||||
query.FOR(lockType.SKIP_LOCKED())
|
query.FOR(lockType.SKIP_LOCKED())
|
||||||
|
|
||||||
testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+" SKIP LOCKED;\n", int64(3), int64(1))
|
testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+" SKIP LOCKED;\n", int64(3), int64(1))
|
||||||
|
testutils.AssertExecAndRollback(t, query, db)
|
||||||
tx, _ := db.Begin()
|
|
||||||
|
|
||||||
_, err := query.Exec(tx)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
err = tx.Rollback()
|
|
||||||
require.NoError(t, err)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -956,7 +937,7 @@ LIMIT 1
|
||||||
FOR UPDATE OF film, actor NOWAIT;
|
FOR UPDATE OF film, actor NOWAIT;
|
||||||
`)
|
`)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
var dest []struct {
|
var dest []struct {
|
||||||
model.Film
|
model.Film
|
||||||
CategoryID int
|
CategoryID int
|
||||||
|
|
@ -993,7 +974,7 @@ LIMIT 1
|
||||||
FOR UPDATE OF ''myFilm'';
|
FOR UPDATE OF ''myFilm'';
|
||||||
`, "''", "`"))
|
`, "''", "`"))
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
var dest []struct {
|
var dest []struct {
|
||||||
model.Film `alias:"myFilm.*"`
|
model.Film `alias:"myFilm.*"`
|
||||||
}
|
}
|
||||||
|
|
|
||||||
128
tests/mysql/stmtcache_test.go
Normal file
128
tests/mysql/stmtcache_test.go
Normal file
|
|
@ -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)
|
||||||
|
}
|
||||||
|
|
@ -2,7 +2,7 @@ package mysql
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"database/sql"
|
"github.com/go-jet/jet/v2/qrm"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
|
@ -28,7 +28,7 @@ WHERE link.name = 'Bing';
|
||||||
WHERE(Link.Name.EQ(String("Bing")))
|
WHERE(Link.Name.EQ(String("Bing")))
|
||||||
|
|
||||||
testutils.AssertDebugStatementSql(t, query, expectedSQL, "Bong", "http://bong.com", "Bing")
|
testutils.AssertDebugStatementSql(t, query, expectedSQL, "Bong", "http://bong.com", "Bing")
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
testutils.AssertExec(t, query, tx)
|
testutils.AssertExec(t, query, tx)
|
||||||
requireLogged(t, query)
|
requireLogged(t, query)
|
||||||
|
|
||||||
|
|
@ -59,7 +59,7 @@ WHERE link.name = 'Bing';
|
||||||
WHERE(Link.Name.EQ(String("Bing")))
|
WHERE(Link.Name.EQ(String("Bing")))
|
||||||
|
|
||||||
testutils.AssertDebugStatementSql(t, stmt, expectedSQL, "Bong", "http://bong.com", "Bing")
|
testutils.AssertDebugStatementSql(t, stmt, expectedSQL, "Bong", "http://bong.com", "Bing")
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
testutils.AssertExec(t, stmt, tx)
|
testutils.AssertExec(t, stmt, tx)
|
||||||
requireLogged(t, stmt)
|
requireLogged(t, stmt)
|
||||||
|
|
||||||
|
|
@ -280,7 +280,7 @@ SET id = 501,
|
||||||
WHERE link.name = 'Bing';
|
WHERE link.name = 'Bing';
|
||||||
`)
|
`)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
_, err := stmt.Exec(tx)
|
_, err := stmt.Exec(tx)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
})
|
})
|
||||||
|
|
|
||||||
|
|
@ -1,10 +1,10 @@
|
||||||
package postgres
|
package postgres
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"database/sql"
|
|
||||||
"github.com/go-jet/jet/v2/internal/utils/ptr"
|
"github.com/go-jet/jet/v2/internal/utils/ptr"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
|
|
||||||
|
"github.com/go-jet/jet/v2/qrm"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
|
@ -74,7 +74,7 @@ func TestAllTypesInsertModel(t *testing.T) {
|
||||||
MODEL(&allTypesRow1).
|
MODEL(&allTypesRow1).
|
||||||
RETURNING(AllTypes.AllColumns)
|
RETURNING(AllTypes.AllColumns)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
var dest []model.AllTypes
|
var dest []model.AllTypes
|
||||||
err := query.Query(tx, &dest)
|
err := query.Query(tx, &dest)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
@ -97,7 +97,7 @@ func TestAllTypesInsertQuery(t *testing.T) {
|
||||||
).
|
).
|
||||||
RETURNING(AllTypesAllColumns)
|
RETURNING(AllTypesAllColumns)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
var dest []model.AllTypes
|
var dest []model.AllTypes
|
||||||
err := query.Query(tx, &dest)
|
err := query.Query(tx, &dest)
|
||||||
|
|
||||||
|
|
@ -943,7 +943,7 @@ RETURNING employee.employee_id AS "employee.employee_id",
|
||||||
employee.manager_id AS "employee.manager_id",
|
employee.manager_id AS "employee.manager_id",
|
||||||
employee.pto_accrual AS "employee.pto_accrual";
|
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
|
var windy model.Employee
|
||||||
windy.PtoAccrual = ptr.Of("3h")
|
windy.PtoAccrual = ptr.Of("3h")
|
||||||
stmt := Employee.UPDATE(Employee.PtoAccrual).SET(
|
stmt := Employee.UPDATE(Employee.PtoAccrual).SET(
|
||||||
|
|
@ -971,7 +971,7 @@ RETURNING employee.employee_id AS "employee.employee_id",
|
||||||
employee.manager_id AS "employee.manager_id",
|
employee.manager_id AS "employee.manager_id",
|
||||||
employee.pto_accrual AS "employee.pto_accrual";
|
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
|
var employee model.Employee
|
||||||
employee.PtoAccrual = ptr.Of("5h")
|
employee.PtoAccrual = ptr.Of("5h")
|
||||||
stmt := Employee.INSERT(Employee.AllColumns).
|
stmt := Employee.INSERT(Employee.AllColumns).
|
||||||
|
|
@ -1393,7 +1393,7 @@ WHERE all_types.small_int = 14
|
||||||
RETURNING all_types.json AS "all_types.json";
|
RETURNING all_types.json AS "all_types.json";
|
||||||
`)
|
`)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
var res model.AllTypes
|
var res model.AllTypes
|
||||||
|
|
||||||
err := stmt.Query(tx, &res)
|
err := stmt.Query(tx, &res)
|
||||||
|
|
|
||||||
|
|
@ -13,7 +13,7 @@ import (
|
||||||
"time"
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestSelect(t *testing.T) {
|
func TestSelectAlbum(t *testing.T) {
|
||||||
stmt := SELECT(Album.AllColumns).
|
stmt := SELECT(Album.AllColumns).
|
||||||
FROM(Album).
|
FROM(Album).
|
||||||
ORDER_BY(Album.AlbumId.ASC())
|
ORDER_BY(Album.AlbumId.ASC())
|
||||||
|
|
@ -783,9 +783,11 @@ func TestQueryWithContext(t *testing.T) {
|
||||||
return // context cancellation doesn't work for pq driver
|
return // context cancellation doesn't work for pq driver
|
||||||
}
|
}
|
||||||
|
|
||||||
ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond)
|
ctx, cancel := context.WithTimeout(context.Background(), 1*time.Microsecond)
|
||||||
defer cancel()
|
defer cancel()
|
||||||
|
|
||||||
|
time.Sleep(1 * time.Millisecond)
|
||||||
|
|
||||||
var dest []model.Album
|
var dest []model.Album
|
||||||
|
|
||||||
err := Album.
|
err := Album.
|
||||||
|
|
|
||||||
|
|
@ -2,9 +2,9 @@ package postgres
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"database/sql"
|
|
||||||
"github.com/go-jet/jet/v2/internal/testutils"
|
"github.com/go-jet/jet/v2/internal/testutils"
|
||||||
. "github.com/go-jet/jet/v2/postgres"
|
. "github.com/go-jet/jet/v2/postgres"
|
||||||
|
"github.com/go-jet/jet/v2/qrm"
|
||||||
model2 "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/model"
|
model2 "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/model"
|
||||||
"github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/table"
|
"github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/table"
|
||||||
"github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/test_sample/model"
|
"github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/test_sample/model"
|
||||||
|
|
@ -43,7 +43,7 @@ RETURNING link.id AS "link.id",
|
||||||
link.description AS "link.description";
|
link.description AS "link.description";
|
||||||
`, "Gmail", "Outlook")
|
`, "Gmail", "Outlook")
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
var dest []model.Link
|
var dest []model.Link
|
||||||
|
|
||||||
err := deleteStmt.Query(tx, &dest)
|
err := deleteStmt.Query(tx, &dest)
|
||||||
|
|
@ -67,7 +67,7 @@ func TestDeleteQueryContext(t *testing.T) {
|
||||||
|
|
||||||
time.Sleep(10 * time.Millisecond)
|
time.Sleep(10 * time.Millisecond)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
dest := []model.Link{}
|
dest := []model.Link{}
|
||||||
err := deleteStmt.QueryContext(ctx, tx, &dest)
|
err := deleteStmt.QueryContext(ctx, tx, &dest)
|
||||||
|
|
||||||
|
|
@ -88,7 +88,7 @@ func TestDeleteExecContext(t *testing.T) {
|
||||||
|
|
||||||
time.Sleep(10 * time.Millisecond)
|
time.Sleep(10 * time.Millisecond)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
_, err := deleteStmt.ExecContext(ctx, tx)
|
_, err := deleteStmt.ExecContext(ctx, tx)
|
||||||
|
|
||||||
require.Error(t, err, "context deadline exceeded")
|
require.Error(t, err, "context deadline exceeded")
|
||||||
|
|
@ -140,7 +140,7 @@ RETURNING rental.rental_id AS "rental.rental_id",
|
||||||
store.last_update AS "store.last_update";
|
store.last_update AS "store.last_update";
|
||||||
`)
|
`)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
var dest []struct {
|
var dest []struct {
|
||||||
Rental model2.Rental
|
Rental model2.Rental
|
||||||
Store model2.Store
|
Store model2.Store
|
||||||
|
|
|
||||||
|
|
@ -2,9 +2,9 @@ package postgres
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"database/sql"
|
|
||||||
"github.com/go-jet/jet/v2/internal/testutils"
|
"github.com/go-jet/jet/v2/internal/testutils"
|
||||||
. "github.com/go-jet/jet/v2/postgres"
|
. "github.com/go-jet/jet/v2/postgres"
|
||||||
|
"github.com/go-jet/jet/v2/qrm"
|
||||||
"github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/test_sample/model"
|
"github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/test_sample/model"
|
||||||
. "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/test_sample/table"
|
. "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/test_sample/table"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
|
|
@ -34,7 +34,7 @@ RETURNING link.id AS "link.id",
|
||||||
101, "http://www.google.com", "Google",
|
101, "http://www.google.com", "Google",
|
||||||
102, "http://www.yahoo.com", "Yahoo", nil)
|
102, "http://www.yahoo.com", "Yahoo", nil)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
var insertedLinks []model.Link
|
var insertedLinks []model.Link
|
||||||
|
|
||||||
err := insertQuery.Query(tx, &insertedLinks)
|
err := insertQuery.Query(tx, &insertedLinks)
|
||||||
|
|
@ -368,7 +368,7 @@ RETURNING link.id AS "link.id",
|
||||||
link.description AS "link.description";
|
link.description AS "link.description";
|
||||||
`, int64(0))
|
`, int64(0))
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
var dest []model.Link
|
var dest []model.Link
|
||||||
|
|
||||||
err := query.Query(tx, &dest)
|
err := query.Query(tx, &dest)
|
||||||
|
|
@ -395,7 +395,7 @@ func TestInsertWithQueryContext(t *testing.T) {
|
||||||
|
|
||||||
time.Sleep(10 * time.Millisecond)
|
time.Sleep(10 * time.Millisecond)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
dest := []model.Link{}
|
dest := []model.Link{}
|
||||||
err := stmt.QueryContext(ctx, tx, &dest)
|
err := stmt.QueryContext(ctx, tx, &dest)
|
||||||
|
|
||||||
|
|
@ -412,7 +412,7 @@ func TestInsertWithExecContext(t *testing.T) {
|
||||||
|
|
||||||
time.Sleep(10 * time.Millisecond)
|
time.Sleep(10 * time.Millisecond)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
testutils.AssertExecContextErr(ctx, t, stmt, tx, "context deadline exceeded")
|
testutils.AssertExecContextErr(ctx, t, stmt, tx, "context deadline exceeded")
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -32,34 +32,14 @@ LOCK TABLE dvds.address IN`
|
||||||
query := Address.LOCK().IN(lockMode)
|
query := Address.LOCK().IN(lockMode)
|
||||||
|
|
||||||
testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+string(lockMode)+" MODE;\n")
|
testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+string(lockMode)+" MODE;\n")
|
||||||
|
testutils.AssertExecAndRollback(t, query, db)
|
||||||
tx, _ := db.Begin()
|
|
||||||
|
|
||||||
_, err := query.Exec(tx)
|
|
||||||
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
err = tx.Rollback()
|
|
||||||
|
|
||||||
require.NoError(t, err)
|
|
||||||
requireLogged(t, query)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
for _, lockMode := range testData {
|
for _, lockMode := range testData {
|
||||||
query := Address.LOCK().IN(lockMode).NOWAIT()
|
query := Address.LOCK().IN(lockMode).NOWAIT()
|
||||||
|
|
||||||
testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+string(lockMode)+" MODE NOWAIT;\n")
|
testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+string(lockMode)+" MODE NOWAIT;\n")
|
||||||
|
testutils.AssertExecAndRollback(t, query, db)
|
||||||
tx, _ := db.Begin()
|
|
||||||
|
|
||||||
_, err := query.Exec(tx)
|
|
||||||
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
err = tx.Rollback()
|
|
||||||
|
|
||||||
require.NoError(t, err)
|
|
||||||
requireLogged(t, query)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -4,6 +4,7 @@ import (
|
||||||
"context"
|
"context"
|
||||||
"database/sql"
|
"database/sql"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"github.com/go-jet/jet/v2/stmtcache"
|
||||||
"github.com/go-jet/jet/v2/tests/internal/utils/repo"
|
"github.com/go-jet/jet/v2/tests/internal/utils/repo"
|
||||||
"github.com/jackc/pgx/v4/stdlib"
|
"github.com/jackc/pgx/v4/stdlib"
|
||||||
"os"
|
"os"
|
||||||
|
|
@ -19,15 +20,17 @@ import (
|
||||||
_ "github.com/jackc/pgx/v4/stdlib"
|
_ "github.com/jackc/pgx/v4/stdlib"
|
||||||
)
|
)
|
||||||
|
|
||||||
var db *sql.DB
|
var db *stmtcache.DB
|
||||||
var testRoot string
|
var testRoot string
|
||||||
|
|
||||||
var source string
|
var source string
|
||||||
|
var withStatementCaching bool
|
||||||
|
|
||||||
const CockroachDB = "COCKROACH_DB"
|
const CockroachDB = "COCKROACH_DB"
|
||||||
|
|
||||||
func init() {
|
func init() {
|
||||||
source = os.Getenv("PG_SOURCE")
|
source = os.Getenv("PG_SOURCE")
|
||||||
|
withStatementCaching = os.Getenv("JET_TESTS_WITH_STMT_CACHE") == "true"
|
||||||
}
|
}
|
||||||
|
|
||||||
func sourceIsCockroachDB() bool {
|
func sourceIsCockroachDB() bool {
|
||||||
|
|
@ -45,32 +48,50 @@ func TestMain(m *testing.M) {
|
||||||
|
|
||||||
setTestRoot()
|
setTestRoot()
|
||||||
|
|
||||||
for _, driverName := range []string{"pgx", "postgres"} {
|
for _, driverName := range []string{"postgres", "pgx"} {
|
||||||
fmt.Printf("\nRunning postgres tests for '%s' driver\n", driverName)
|
|
||||||
|
fmt.Printf("\nRunning postgres tests for driver: %s, caching enabled: %t \n", driverName, withStatementCaching)
|
||||||
|
|
||||||
func() {
|
func() {
|
||||||
|
sqlDB, err := sql.Open(driverName, getConnectionString())
|
||||||
connectionString := dbconfig.PostgresConnectString
|
|
||||||
|
|
||||||
if sourceIsCockroachDB() {
|
|
||||||
connectionString = dbconfig.CockroachConnectString
|
|
||||||
}
|
|
||||||
|
|
||||||
var err error
|
|
||||||
db, err = sql.Open(driverName, connectionString)
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
fmt.Println(err.Error())
|
fmt.Println(err.Error())
|
||||||
panic("Failed to connect to test db")
|
panic("Failed to connect to test db")
|
||||||
}
|
}
|
||||||
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)
|
||||||
|
|
||||||
ret := m.Run()
|
for i := 0; i < runCount(withStatementCaching); i++ {
|
||||||
|
ret := m.Run()
|
||||||
if ret != 0 {
|
if ret != 0 {
|
||||||
os.Exit(ret)
|
fmt.Printf("\nFAIL: Running postgres tests failed for driver: %s, caching enabled: %t \n", driverName, withStatementCaching)
|
||||||
|
os.Exit(ret)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}()
|
}()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
}
|
||||||
|
|
||||||
|
func runCount(stmtCaching bool) int {
|
||||||
|
if stmtCaching {
|
||||||
|
return 2
|
||||||
|
}
|
||||||
|
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
|
||||||
|
func getConnectionString() string {
|
||||||
|
if sourceIsCockroachDB() {
|
||||||
|
return dbconfig.CockroachConnectString
|
||||||
|
}
|
||||||
|
|
||||||
|
return dbconfig.PostgresConnectString
|
||||||
}
|
}
|
||||||
|
|
||||||
func setTestRoot() {
|
func setTestRoot() {
|
||||||
|
|
|
||||||
|
|
@ -1,8 +1,8 @@
|
||||||
package postgres
|
package postgres
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"database/sql"
|
|
||||||
"github.com/go-jet/jet/v2/internal/testutils"
|
"github.com/go-jet/jet/v2/internal/testutils"
|
||||||
|
"github.com/go-jet/jet/v2/qrm"
|
||||||
"github.com/google/go-cmp/cmp"
|
"github.com/google/go-cmp/cmp"
|
||||||
"github.com/jackc/pgtype"
|
"github.com/jackc/pgtype"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
|
|
@ -294,7 +294,7 @@ RETURNING sample_ranges.date_range AS "sample_ranges.date_range",
|
||||||
`
|
`
|
||||||
testutils.AssertDebugStatementSql(t, insertQuery, expectedQuery)
|
testutils.AssertDebugStatementSql(t, insertQuery, expectedQuery)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
var dest []model.SampleRanges
|
var dest []model.SampleRanges
|
||||||
err := insertQuery.Query(tx, &dest)
|
err := insertQuery.Query(tx, &dest)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
@ -324,7 +324,7 @@ RETURNING sample_ranges.date_range AS "sample_ranges.date_range",
|
||||||
sample_ranges.num_range AS "sample_ranges.num_range";
|
sample_ranges.num_range AS "sample_ranges.num_range";
|
||||||
`)
|
`)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
var dest []model.SampleRanges
|
var dest []model.SampleRanges
|
||||||
|
|
||||||
err := stmt.Query(tx, &dest)
|
err := stmt.Query(tx, &dest)
|
||||||
|
|
@ -351,7 +351,7 @@ SET int4_range = int4range(-12::integer, 78::integer),
|
||||||
WHERE LOWER(sample_ranges.timestampz_range) > '2024-02-27 00:00:00 UTC'::timestamp with time zone;
|
WHERE LOWER(sample_ranges.timestampz_range) > '2024-02-27 00:00:00 UTC'::timestamp with time zone;
|
||||||
`)
|
`)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
testutils.AssertExec(t, stmt, tx, 1)
|
testutils.AssertExec(t, stmt, tx, 1)
|
||||||
})
|
})
|
||||||
})
|
})
|
||||||
|
|
|
||||||
|
|
@ -2,7 +2,7 @@ package postgres
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"database/sql"
|
"github.com/go-jet/jet/v2/qrm"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
|
@ -128,7 +128,7 @@ RETURNING link.id AS "link.id",
|
||||||
link.description AS "link.description";
|
link.description AS "link.description";
|
||||||
`)
|
`)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
var links []model2.Link
|
var links []model2.Link
|
||||||
err := stmt.Query(tx, &links)
|
err := stmt.Query(tx, &links)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
|
||||||
|
|
@ -1,8 +1,8 @@
|
||||||
package postgres
|
package postgres
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"database/sql"
|
|
||||||
"github.com/go-jet/jet/v2/internal/utils/ptr"
|
"github.com/go-jet/jet/v2/internal/utils/ptr"
|
||||||
|
"github.com/go-jet/jet/v2/qrm"
|
||||||
"github.com/google/uuid"
|
"github.com/google/uuid"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
|
@ -516,7 +516,7 @@ func TestMutableColumnsExcludeGeneratedColumn(t *testing.T) {
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("should insert without generated columns", func(t *testing.T) {
|
t.Run("should insert without generated columns", func(t *testing.T) {
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
insertQuery := People.INSERT(
|
insertQuery := People.INSERT(
|
||||||
People.MutableColumns,
|
People.MutableColumns,
|
||||||
).MODEL(
|
).MODEL(
|
||||||
|
|
|
||||||
|
|
@ -110,7 +110,7 @@ ORDER BY rental.staff_id ASC, rental.customer_id ASC, rental.rental_id ASC;
|
||||||
`)
|
`)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestClassicSelect(t *testing.T) {
|
func TestSelectClassic(t *testing.T) {
|
||||||
expectedSQL := `
|
expectedSQL := `
|
||||||
SELECT payment.payment_id AS "payment.payment_id",
|
SELECT payment.payment_id AS "payment.payment_id",
|
||||||
payment.customer_id AS "payment.customer_id",
|
payment.customer_id AS "payment.customer_id",
|
||||||
|
|
@ -146,7 +146,7 @@ LIMIT 30;
|
||||||
|
|
||||||
testutils.AssertDebugStatementSql(t, query, expectedSQL, int64(30))
|
testutils.AssertDebugStatementSql(t, query, expectedSQL, int64(30))
|
||||||
|
|
||||||
dest := []model.Payment{}
|
var dest []model.Payment
|
||||||
|
|
||||||
err := query.Query(db, &dest)
|
err := query.Query(db, &dest)
|
||||||
|
|
||||||
|
|
@ -232,12 +232,12 @@ LIMIT 12;
|
||||||
|
|
||||||
testutils.AssertDebugStatementSql(t, query, expectedSQL, int64(1), int64(1), int64(10), int64(1), int64(2), int64(1), int64(12))
|
testutils.AssertDebugStatementSql(t, query, expectedSQL, int64(1), int64(1), int64(10), int64(1), int64(2), int64(1), int64(12))
|
||||||
|
|
||||||
dest := []struct{}{}
|
var dest []struct{}
|
||||||
err := query.Query(db, &dest)
|
err := query.Query(db, &dest)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestFetchFirst(t *testing.T) {
|
func TestSelectFetchFirst(t *testing.T) {
|
||||||
|
|
||||||
t.Run("rows only", func(t *testing.T) {
|
t.Run("rows only", func(t *testing.T) {
|
||||||
stmt := SELECT(Actor.AllColumns).
|
stmt := SELECT(Actor.AllColumns).
|
||||||
|
|
@ -321,7 +321,7 @@ FETCH FIRST (
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestOffsetExpression(t *testing.T) {
|
func TestSelectOffsetExpression(t *testing.T) {
|
||||||
|
|
||||||
stmt := SELECT(Actor.AllColumns).
|
stmt := SELECT(Actor.AllColumns).
|
||||||
FROM(Actor).
|
FROM(Actor).
|
||||||
|
|
@ -353,7 +353,7 @@ OFFSET (
|
||||||
require.Equal(t, dest[0].ActorID, int32(3))
|
require.Equal(t, dest[0].ActorID, int32(3))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestJoinQueryStruct(t *testing.T) {
|
func TestSelectJoinQueryStruct(t *testing.T) {
|
||||||
|
|
||||||
expectedSQL := `
|
expectedSQL := `
|
||||||
SELECT film_actor.actor_id AS "film_actor.actor_id",
|
SELECT film_actor.actor_id AS "film_actor.actor_id",
|
||||||
|
|
@ -446,7 +446,7 @@ LIMIT 1000;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestJoinQuerySlice(t *testing.T) {
|
func TestSelectJoinQuerySlice(t *testing.T) {
|
||||||
expectedSQL := `
|
expectedSQL := `
|
||||||
SELECT language.language_id AS "language.language_id",
|
SELECT language.language_id AS "language.language_id",
|
||||||
language.name AS "language.name",
|
language.name AS "language.name",
|
||||||
|
|
@ -475,7 +475,7 @@ LIMIT 15;
|
||||||
Film []model.Film
|
Film []model.Film
|
||||||
}
|
}
|
||||||
|
|
||||||
filmsPerLanguage := []FilmsPerLanguage{}
|
var filmsPerLanguage []FilmsPerLanguage
|
||||||
limit := 15
|
limit := 15
|
||||||
|
|
||||||
query := Film.
|
query := Film.
|
||||||
|
|
@ -505,7 +505,7 @@ LIMIT 15;
|
||||||
}
|
}
|
||||||
|
|
||||||
// https://github.com/go-jet/jet/issues/226
|
// https://github.com/go-jet/jet/issues/226
|
||||||
func TestDuplicateSlicesInDestination(t *testing.T) {
|
func TestSelectDuplicateSlicesInDestination(t *testing.T) {
|
||||||
|
|
||||||
type Staffs struct {
|
type Staffs struct {
|
||||||
StaffList []model.Staff
|
StaffList []model.Staff
|
||||||
|
|
@ -646,7 +646,7 @@ func TestDuplicateSlicesInDestination(t *testing.T) {
|
||||||
`)
|
`)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestExecution1(t *testing.T) {
|
func TestSelectExecution1(t *testing.T) {
|
||||||
stmt := City.
|
stmt := City.
|
||||||
INNER_JOIN(Address, Address.CityID.EQ(City.CityID)).
|
INNER_JOIN(Address, Address.CityID.EQ(City.CityID)).
|
||||||
INNER_JOIN(Customer, Customer.AddressID.EQ(Address.AddressID)).
|
INNER_JOIN(Customer, Customer.AddressID.EQ(Address.AddressID)).
|
||||||
|
|
@ -707,7 +707,7 @@ ORDER BY city.city_id, address.address_id, customer.customer_id;
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestExecution2(t *testing.T) {
|
func TestSelectExecution2(t *testing.T) {
|
||||||
|
|
||||||
type MyAddress struct {
|
type MyAddress struct {
|
||||||
ID int32 `sql:"primary_key"`
|
ID int32 `sql:"primary_key"`
|
||||||
|
|
@ -770,7 +770,7 @@ ORDER BY city.city_id, address.address_id, customer.customer_id;
|
||||||
require.Equal(t, *dest[0].Customers[1].LastName, "Vines")
|
require.Equal(t, *dest[0].Customers[1].LastName, "Vines")
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestExecution3(t *testing.T) {
|
func TestSelectExecution3(t *testing.T) {
|
||||||
|
|
||||||
var dest []struct {
|
var dest []struct {
|
||||||
CityID int32 `sql:"primary_key"`
|
CityID int32 `sql:"primary_key"`
|
||||||
|
|
@ -827,7 +827,7 @@ ORDER BY city.city_id, address.address_id, customer.customer_id;
|
||||||
require.Equal(t, *dest[0].Customers[1].LastName, "Vines")
|
require.Equal(t, *dest[0].Customers[1].LastName, "Vines")
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestExecution4(t *testing.T) {
|
func TestSelectExecution4(t *testing.T) {
|
||||||
|
|
||||||
var dest []struct {
|
var dest []struct {
|
||||||
CityID int32 `sql:"primary_key" alias:"city.city_id"`
|
CityID int32 `sql:"primary_key" alias:"city.city_id"`
|
||||||
|
|
@ -919,7 +919,7 @@ ORDER BY city.city_id, address.address_id, customer.customer_id;
|
||||||
}
|
}
|
||||||
|
|
||||||
// Test join with custom primary keys (sql.NullInt64)
|
// Test join with custom primary keys (sql.NullInt64)
|
||||||
func TestExecutionCustomPKTypes1(t *testing.T) {
|
func TestSelectExecutionCustomPKTypes1(t *testing.T) {
|
||||||
|
|
||||||
var dest []struct {
|
var dest []struct {
|
||||||
CityID sql.NullInt64 `sql:"primary_key" alias:"city.city_id"`
|
CityID sql.NullInt64 `sql:"primary_key" alias:"city.city_id"`
|
||||||
|
|
@ -1040,7 +1040,7 @@ ORDER BY city.city_id, address.address_id, customer.customer_id;
|
||||||
}
|
}
|
||||||
|
|
||||||
// Test join with custom primary keys (null.Int)
|
// Test join with custom primary keys (null.Int)
|
||||||
func TestExecutionCustomPKTypes2(t *testing.T) {
|
func TestSelectExecutionCustomPKTypes2(t *testing.T) {
|
||||||
|
|
||||||
var dest []struct {
|
var dest []struct {
|
||||||
CityID null.Int `sql:"primary_key" alias:"city.city_id"`
|
CityID null.Int `sql:"primary_key" alias:"city.city_id"`
|
||||||
|
|
@ -1136,7 +1136,7 @@ ORDER BY city.city_id, address.address_id, customer.customer_id;
|
||||||
`)
|
`)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestJoinQuerySliceWithPtrs(t *testing.T) {
|
func TestSelectJoinQuerySliceWithPtrs(t *testing.T) {
|
||||||
type FilmsPerLanguage struct {
|
type FilmsPerLanguage struct {
|
||||||
Language model.Language
|
Language model.Language
|
||||||
Film *[]*model.Film
|
Film *[]*model.Film
|
||||||
|
|
@ -1161,7 +1161,7 @@ func TestJoinQuerySliceWithPtrs(t *testing.T) {
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSelect_WithoutUniqueColumnSelected(t *testing.T) {
|
func TestSelect_WithoutUniqueColumnSelected(t *testing.T) {
|
||||||
query := Customer.SELECT(Customer.FirstName, Customer.LastName, Customer.Email)
|
query := Customer.SELECT(Customer.FirstName, Customer.LastName, Customer.Email).ORDER_BY(Customer.Email)
|
||||||
|
|
||||||
var customers []model.Customer
|
var customers []model.Customer
|
||||||
|
|
||||||
|
|
@ -1220,7 +1220,7 @@ func TestSelectOrderByAscDesc(t *testing.T) {
|
||||||
testutils.AssertDeepEqual(t, customerAscDesc327, customersAscDesc[327])
|
testutils.AssertDeepEqual(t, customerAscDesc327, customersAscDesc[327])
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestOrderBy(t *testing.T) {
|
func TestSelectOrderBy(t *testing.T) {
|
||||||
|
|
||||||
t.Run("default", func(t *testing.T) {
|
t.Run("default", func(t *testing.T) {
|
||||||
stmt := SELECT(
|
stmt := SELECT(
|
||||||
|
|
@ -1651,7 +1651,7 @@ LIMIT 1000;
|
||||||
testutils.AssertDeepEqual(t, films[0], thesameLengthFilms{"Alien Center", "Iron Moon", 46})
|
testutils.AssertDeepEqual(t, films[0], thesameLengthFilms{"Alien Center", "Iron Moon", 46})
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSubQuery(t *testing.T) {
|
func TestSelectSubQuery(t *testing.T) {
|
||||||
rRatingFilms :=
|
rRatingFilms :=
|
||||||
SELECT(
|
SELECT(
|
||||||
Film.FilmID,
|
Film.FilmID,
|
||||||
|
|
@ -1817,7 +1817,7 @@ ORDER BY film.film_id ASC;
|
||||||
|
|
||||||
testutils.AssertDebugStatementSql(t, query, expectedSQL)
|
testutils.AssertDebugStatementSql(t, query, expectedSQL)
|
||||||
|
|
||||||
maxRentalRateFilms := []model.Film{}
|
var maxRentalRateFilms []model.Film
|
||||||
err := query.Query(db, &maxRentalRateFilms)
|
err := query.Query(db, &maxRentalRateFilms)
|
||||||
|
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
@ -1921,7 +1921,7 @@ ORDER BY customer.customer_id, SUM(payment.amount) ASC;
|
||||||
testutils.AssertJSONFile(t, dest, "./testdata/results/postgres/customer_payment_sum.json")
|
testutils.AssertJSONFile(t, dest, "./testdata/results/postgres/customer_payment_sum.json")
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestGroupByGroupingSets(t *testing.T) {
|
func TestSelectGroupByGroupingSets(t *testing.T) {
|
||||||
skipForCockroachDB(t)
|
skipForCockroachDB(t)
|
||||||
|
|
||||||
stmt := SELECT(
|
stmt := SELECT(
|
||||||
|
|
@ -2003,7 +2003,7 @@ ORDER BY inventory.film_id, inventory.store_id;
|
||||||
`)
|
`)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestGroupByCube(t *testing.T) {
|
func TestSelectGroupByCube(t *testing.T) {
|
||||||
skipForCockroachDB(t)
|
skipForCockroachDB(t)
|
||||||
|
|
||||||
stmt := SELECT(
|
stmt := SELECT(
|
||||||
|
|
@ -2080,7 +2080,7 @@ ORDER BY country.country, city.city;
|
||||||
`)
|
`)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestGroupByRollup(t *testing.T) {
|
func TestSelectGroupByRollup(t *testing.T) {
|
||||||
skipForCockroachDB(t)
|
skipForCockroachDB(t)
|
||||||
|
|
||||||
stmt := SELECT(
|
stmt := SELECT(
|
||||||
|
|
@ -2167,7 +2167,7 @@ ORDER BY year ASC, EXTRACT(MONTH FROM rental.rental_date) ASC, day ASC;
|
||||||
`)
|
`)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestAggregateFunctionDistinct(t *testing.T) {
|
func TestSelectAggregateFunctionDistinct(t *testing.T) {
|
||||||
stmt := SELECT(
|
stmt := SELECT(
|
||||||
Payment.CustomerID,
|
Payment.CustomerID,
|
||||||
|
|
||||||
|
|
@ -2298,7 +2298,7 @@ ORDER BY customer_payment_sum.amount_sum ASC;
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSelectStaff(t *testing.T) {
|
func TestSelectStaff(t *testing.T) {
|
||||||
staffs := []model.Staff{}
|
var staffs []model.Staff
|
||||||
|
|
||||||
err := Staff.SELECT(Staff.AllColumns).Query(db, &staffs)
|
err := Staff.SELECT(Staff.AllColumns).Query(db, &staffs)
|
||||||
|
|
||||||
|
|
@ -2372,7 +2372,7 @@ ORDER BY payment.payment_date ASC;
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestUnion(t *testing.T) {
|
func TestSelectUnion(t *testing.T) {
|
||||||
expectedQuery := `
|
expectedQuery := `
|
||||||
(
|
(
|
||||||
SELECT payment.payment_id AS "payment.payment_id",
|
SELECT payment.payment_id AS "payment.payment_id",
|
||||||
|
|
@ -2425,7 +2425,7 @@ OFFSET 20;
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestUnionOffsetWithExpression(t *testing.T) {
|
func TestSelectUnionOffsetWithExpression(t *testing.T) {
|
||||||
stmt := UNION(
|
stmt := UNION(
|
||||||
SELECT(Rental.AllColumns).
|
SELECT(Rental.AllColumns).
|
||||||
FROM(Rental).
|
FROM(Rental).
|
||||||
|
|
@ -2475,7 +2475,7 @@ OFFSET (
|
||||||
require.Len(t, dest, 10)
|
require.Len(t, dest, 10)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestAllSetOperators(t *testing.T) {
|
func TestSelectAllSetOperators(t *testing.T) {
|
||||||
var select1 = Payment.SELECT(Payment.AllColumns).WHERE(Payment.PaymentID.GT_EQ(Int(17600)).AND(Payment.PaymentID.LT(Int(17610))))
|
var select1 = Payment.SELECT(Payment.AllColumns).WHERE(Payment.PaymentID.GT_EQ(Int(17600)).AND(Payment.PaymentID.LT(Int(17610))))
|
||||||
var select2 = Payment.SELECT(Payment.AllColumns).WHERE(Payment.PaymentID.GT_EQ(Int(17620)).AND(Payment.PaymentID.LT(Int(17630))))
|
var select2 = Payment.SELECT(Payment.AllColumns).WHERE(Payment.PaymentID.GT_EQ(Int(17620)).AND(Payment.PaymentID.LT(Int(17630))))
|
||||||
|
|
||||||
|
|
@ -2580,46 +2580,30 @@ func getRowLockTestData() map[RowLock]string {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestRowLock(t *testing.T) {
|
func TestSelectRowLock(t *testing.T) {
|
||||||
|
|
||||||
|
query := SELECT(STAR).
|
||||||
|
FROM(Address).
|
||||||
|
WHERE(Address.AddressID.LT(Int(10)))
|
||||||
|
|
||||||
expectedSQL := `
|
expectedSQL := `
|
||||||
SELECT *
|
SELECT *
|
||||||
FROM dvds.address
|
FROM dvds.address
|
||||||
LIMIT 3
|
WHERE address.address_id < 10
|
||||||
FOR`
|
FOR`
|
||||||
query := Address.
|
|
||||||
SELECT(STAR).
|
|
||||||
LIMIT(3)
|
|
||||||
|
|
||||||
for lockType, lockTypeStr := range getRowLockTestData() {
|
for lockType, lockTypeStr := range getRowLockTestData() {
|
||||||
query.FOR(lockType)
|
query.FOR(lockType)
|
||||||
|
|
||||||
testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+";\n", int64(3))
|
testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+";\n", int64(10))
|
||||||
|
testutils.AssertExecAndRollback(t, query, db, 9)
|
||||||
tx, _ := db.Begin()
|
|
||||||
|
|
||||||
res, err := query.Exec(tx)
|
|
||||||
require.NoError(t, err)
|
|
||||||
rowsAffected, _ := res.RowsAffected()
|
|
||||||
require.Equal(t, rowsAffected, int64(3))
|
|
||||||
|
|
||||||
err = tx.Rollback()
|
|
||||||
require.NoError(t, err)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
for lockType, lockTypeStr := range getRowLockTestData() {
|
for lockType, lockTypeStr := range getRowLockTestData() {
|
||||||
query.FOR(lockType.NOWAIT())
|
query.FOR(lockType.NOWAIT())
|
||||||
|
|
||||||
testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+" NOWAIT;\n", int64(3))
|
testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+" NOWAIT;\n", int64(10))
|
||||||
|
testutils.AssertExecAndRollback(t, query, db, 9)
|
||||||
tx, _ := db.Begin()
|
|
||||||
|
|
||||||
res, err := query.Exec(tx)
|
|
||||||
require.NoError(t, err)
|
|
||||||
rowsAffected, _ := res.RowsAffected()
|
|
||||||
require.Equal(t, rowsAffected, int64(3))
|
|
||||||
|
|
||||||
err = tx.Rollback()
|
|
||||||
require.NoError(t, err)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if sourceIsCockroachDB() {
|
if sourceIsCockroachDB() {
|
||||||
|
|
@ -2629,21 +2613,13 @@ FOR`
|
||||||
for lockType, lockTypeStr := range getRowLockTestData() {
|
for lockType, lockTypeStr := range getRowLockTestData() {
|
||||||
query.FOR(lockType.SKIP_LOCKED())
|
query.FOR(lockType.SKIP_LOCKED())
|
||||||
|
|
||||||
testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+" SKIP LOCKED;\n", int64(3))
|
testutils.AssertDebugStatementSql(t, query, expectedSQL+" "+lockTypeStr+" SKIP LOCKED;\n", int64(10))
|
||||||
|
testutils.AssertExecAndRollback(t, query, db, 9)
|
||||||
tx, _ := db.Begin()
|
|
||||||
|
|
||||||
res, err := query.Exec(tx)
|
|
||||||
require.NoError(t, err)
|
|
||||||
rowsAffected, _ := res.RowsAffected()
|
|
||||||
require.Equal(t, rowsAffected, int64(3))
|
|
||||||
|
|
||||||
err = tx.Rollback()
|
|
||||||
require.NoError(t, err)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestRowLockWithUpdateOf(t *testing.T) {
|
func TestSelectRowLockWithUpdateOf(t *testing.T) {
|
||||||
|
|
||||||
stmt := SELECT(
|
stmt := SELECT(
|
||||||
Film.FilmID,
|
Film.FilmID,
|
||||||
Film.Title,
|
Film.Title,
|
||||||
|
|
@ -2673,9 +2649,10 @@ LIMIT 1
|
||||||
FOR UPDATE OF film, actor NOWAIT;
|
FOR UPDATE OF film, actor NOWAIT;
|
||||||
`)
|
`)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
var dest []struct {
|
var dest []struct {
|
||||||
model.Film
|
model.Film
|
||||||
|
|
||||||
CategoryID int
|
CategoryID int
|
||||||
Actor []model.Actor
|
Actor []model.Actor
|
||||||
}
|
}
|
||||||
|
|
@ -2686,7 +2663,7 @@ FOR UPDATE OF film, actor NOWAIT;
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestRowLockWithUpdateOfAliasedTable(t *testing.T) {
|
func TestSelectRowLockWithUpdateOfAliasedTable(t *testing.T) {
|
||||||
|
|
||||||
myFilm := Film.AS("myFilm")
|
myFilm := Film.AS("myFilm")
|
||||||
|
|
||||||
|
|
@ -2709,18 +2686,18 @@ LIMIT 1
|
||||||
FOR UPDATE OF "myFilm";
|
FOR UPDATE OF "myFilm";
|
||||||
`)
|
`)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
var dest []struct {
|
var dest []struct {
|
||||||
model.Film `alias:"myFilm.*"`
|
model.Film `alias:"myFilm.*"`
|
||||||
}
|
}
|
||||||
|
|
||||||
err := stmt.Query(db, &dest)
|
err := stmt.Query(tx, &dest)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Len(t, dest, 1)
|
require.Len(t, dest, 1)
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestQuickStart(t *testing.T) {
|
func TestSelectQuickStart(t *testing.T) {
|
||||||
|
|
||||||
var expectedSQL = `
|
var expectedSQL = `
|
||||||
SELECT actor.actor_id AS "actor.actor_id",
|
SELECT actor.actor_id AS "actor.actor_id",
|
||||||
|
|
@ -2811,7 +2788,7 @@ ORDER BY actor.actor_id ASC, film.film_id ASC;
|
||||||
testutils.AssertJSONFile(t, dest2, "./testdata/results/postgres/quick-start-dest2.json")
|
testutils.AssertJSONFile(t, dest2, "./testdata/results/postgres/quick-start-dest2.json")
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestQuickStartWithSubQueries(t *testing.T) {
|
func TestSelectQuickStartWithSubQueries(t *testing.T) {
|
||||||
|
|
||||||
filmLogerThan180 := Film.
|
filmLogerThan180 := Film.
|
||||||
SELECT(Film.AllColumns).
|
SELECT(Film.AllColumns).
|
||||||
|
|
@ -2876,7 +2853,7 @@ func TestQuickStartWithSubQueries(t *testing.T) {
|
||||||
testutils.AssertJSONFile(t, dest2, "./testdata/results/postgres/quick-start-dest2.json")
|
testutils.AssertJSONFile(t, dest2, "./testdata/results/postgres/quick-start-dest2.json")
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestExpressionWrappers(t *testing.T) {
|
func TestSelectExpressionWrappers(t *testing.T) {
|
||||||
query := SELECT(
|
query := SELECT(
|
||||||
BoolExp(Raw("true")),
|
BoolExp(Raw("true")),
|
||||||
IntExp(Raw("11")),
|
IntExp(Raw("11")),
|
||||||
|
|
@ -2906,7 +2883,7 @@ SELECT true,
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestWindowFunction(t *testing.T) {
|
func TestSelectWindowFunction(t *testing.T) {
|
||||||
var expectedSQL = `
|
var expectedSQL = `
|
||||||
SELECT AVG(payment.amount) OVER (),
|
SELECT AVG(payment.amount) OVER (),
|
||||||
AVG(payment.amount) OVER (PARTITION BY payment.customer_id),
|
AVG(payment.amount) OVER (PARTITION BY payment.customer_id),
|
||||||
|
|
@ -2978,7 +2955,7 @@ GROUP BY payment.amount, payment.customer_id, payment.payment_date;
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestWindowClause(t *testing.T) {
|
func TestSelectWindowClause(t *testing.T) {
|
||||||
var expectedSQL = `
|
var expectedSQL = `
|
||||||
SELECT AVG(payment.amount) OVER (),
|
SELECT AVG(payment.amount) OVER (),
|
||||||
AVG(payment.amount) OVER (w1),
|
AVG(payment.amount) OVER (w1),
|
||||||
|
|
@ -3015,7 +2992,7 @@ ORDER BY payment.customer_id;
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestSimpleView(t *testing.T) {
|
func TestSelectSimpleView(t *testing.T) {
|
||||||
|
|
||||||
query := SELECT(
|
query := SELECT(
|
||||||
view.ActorInfo.AllColumns,
|
view.ActorInfo.AllColumns,
|
||||||
|
|
@ -3054,7 +3031,7 @@ func TestSimpleView(t *testing.T) {
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestJoinViewWithTable(t *testing.T) {
|
func TestSelectJoinViewWithTable(t *testing.T) {
|
||||||
query := SELECT(
|
query := SELECT(
|
||||||
view.CustomerList.AllColumns,
|
view.CustomerList.AllColumns,
|
||||||
Rental.AllColumns,
|
Rental.AllColumns,
|
||||||
|
|
@ -3080,7 +3057,7 @@ func TestJoinViewWithTable(t *testing.T) {
|
||||||
require.Equal(t, len(dest[1].Rentals), 27)
|
require.Equal(t, len(dest[1].Rentals), 27)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestDynamicProjectionList(t *testing.T) {
|
func TestSelectDynamicProjectionList(t *testing.T) {
|
||||||
|
|
||||||
var request struct {
|
var request struct {
|
||||||
ColumnsToSelect []string
|
ColumnsToSelect []string
|
||||||
|
|
@ -3127,7 +3104,7 @@ LIMIT 3;
|
||||||
require.Equal(t, len(dest), 3)
|
require.Equal(t, len(dest), 3)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestDynamicCondition(t *testing.T) {
|
func TestSelectDynamicCondition(t *testing.T) {
|
||||||
var request struct {
|
var request struct {
|
||||||
CustomerID *int64
|
CustomerID *int64
|
||||||
Email *string
|
Email *string
|
||||||
|
|
@ -3177,7 +3154,7 @@ WHERE ($1::boolean AND (customer.customer_id = $2)) AND (customer.activebool = $
|
||||||
testutils.AssertDeepEqual(t, dest[0], customer0)
|
testutils.AssertDeepEqual(t, dest[0], customer0)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestLateral(t *testing.T) {
|
func TestSelectLateral(t *testing.T) {
|
||||||
|
|
||||||
languages := LATERAL(
|
languages := LATERAL(
|
||||||
SELECT(
|
SELECT(
|
||||||
|
|
@ -3349,7 +3326,7 @@ type ActorWrap struct {
|
||||||
Films []FilmWrap
|
Films []FilmWrap
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestRecursionScanNxM(t *testing.T) {
|
func TestSelectRecursionScanNxM(t *testing.T) {
|
||||||
|
|
||||||
stmt := SELECT(
|
stmt := SELECT(
|
||||||
Actor.AllColumns,
|
Actor.AllColumns,
|
||||||
|
|
@ -3492,7 +3469,7 @@ type StaffWrap struct {
|
||||||
Store StoreWrap
|
Store StoreWrap
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestRecursionScanNx1(t *testing.T) {
|
func TestSelectRecursionScanNx1(t *testing.T) {
|
||||||
stmt := SELECT(
|
stmt := SELECT(
|
||||||
Store.AllColumns,
|
Store.AllColumns,
|
||||||
Staff.AllColumns,
|
Staff.AllColumns,
|
||||||
|
|
@ -3639,7 +3616,7 @@ type ManagerInfo struct {
|
||||||
Store *StoreInfo
|
Store *StoreInfo
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestRecursionScan1x1(t *testing.T) {
|
func TestSelectRecursionScan1x1(t *testing.T) {
|
||||||
|
|
||||||
stmt := SELECT(
|
stmt := SELECT(
|
||||||
Store.AllColumns,
|
Store.AllColumns,
|
||||||
|
|
@ -3704,7 +3681,7 @@ func TestRecursionScan1x1(t *testing.T) {
|
||||||
// In parameterized statements integer literals, like Int(num), are replaced with a placeholders. For some expressions,
|
// In parameterized statements integer literals, like Int(num), are replaced with a placeholders. For some expressions,
|
||||||
// postgres interpreter will not have enough information to deduce the type. If this is the case postgres returns an error.
|
// postgres interpreter will not have enough information to deduce the type. If this is the case postgres returns an error.
|
||||||
// Int8, Int16, .... functions will add automatic type cast over placeholder, so type deduction is always possible.
|
// Int8, Int16, .... functions will add automatic type cast over placeholder, so type deduction is always possible.
|
||||||
func TestLiteralTypeDeduction(t *testing.T) {
|
func TestSelectLiteralTypeDeduction(t *testing.T) {
|
||||||
stmt := SELECT(
|
stmt := SELECT(
|
||||||
SUM(
|
SUM(
|
||||||
CASE().WHEN(Staff.Active.IS_TRUE()).
|
CASE().WHEN(Staff.Active.IS_TRUE()).
|
||||||
|
|
@ -3726,7 +3703,7 @@ func GET_FILM_COUNT(lenFrom, lenTo IntegerExpression) IntegerExpression {
|
||||||
return IntExp(Func("dvds.get_film_count", lenFrom, lenTo))
|
return IntExp(Func("dvds.get_film_count", lenFrom, lenTo))
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestCustomFunctionCall(t *testing.T) {
|
func TestSelectCustomFunctionCall(t *testing.T) {
|
||||||
skipForCockroachDB(t)
|
skipForCockroachDB(t)
|
||||||
|
|
||||||
stmt := SELECT(
|
stmt := SELECT(
|
||||||
|
|
@ -3762,7 +3739,7 @@ SELECT dvds.get_film_count(100, 120) AS "film_count";
|
||||||
require.Equal(t, dest.FilmCount, 165)
|
require.Equal(t, dest.FilmCount, 165)
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestScanUsingConn(t *testing.T) {
|
func TestSelectScanUsingConn(t *testing.T) {
|
||||||
conn, err := db.Conn(context.Background())
|
conn, err := db.Conn(context.Background())
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
defer conn.Close()
|
defer conn.Close()
|
||||||
|
|
@ -3804,7 +3781,7 @@ func TestScanUsingConn(t *testing.T) {
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestConditionalFunctions(t *testing.T) {
|
func TestSelectConditionalFunctions(t *testing.T) {
|
||||||
stmt := SELECT(
|
stmt := SELECT(
|
||||||
EXISTS(
|
EXISTS(
|
||||||
Film.SELECT(Film.FilmID).WHERE(Film.RentalDuration.GT(Int(100))),
|
Film.SELECT(Film.FilmID).WHERE(Film.RentalDuration.GT(Int(100))),
|
||||||
|
|
|
||||||
139
tests/postgres/stmtcache_test.go
Normal file
139
tests/postgres/stmtcache_test.go
Normal file
|
|
@ -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)
|
||||||
|
}
|
||||||
|
|
@ -2,9 +2,9 @@ package postgres
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"database/sql"
|
|
||||||
"github.com/go-jet/jet/v2/internal/testutils"
|
"github.com/go-jet/jet/v2/internal/testutils"
|
||||||
. "github.com/go-jet/jet/v2/postgres"
|
. "github.com/go-jet/jet/v2/postgres"
|
||||||
|
"github.com/go-jet/jet/v2/qrm"
|
||||||
model2 "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/model"
|
model2 "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/model"
|
||||||
"github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/table"
|
"github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/table"
|
||||||
"github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/test_sample/model"
|
"github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/test_sample/model"
|
||||||
|
|
@ -27,7 +27,7 @@ SET (name, url) = ('Bong', 'http://bong.com')
|
||||||
WHERE link.name = 'Bing'::text;
|
WHERE link.name = 'Bing'::text;
|
||||||
`, "Bong", "http://bong.com", "Bing")
|
`, "Bong", "http://bong.com", "Bing")
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
|
|
||||||
testutils.AssertExec(t, query, tx, 1)
|
testutils.AssertExec(t, query, tx, 1)
|
||||||
requireLogged(t, query)
|
requireLogged(t, query)
|
||||||
|
|
@ -142,7 +142,7 @@ RETURNING link.id AS "link.id",
|
||||||
link.description AS "link.description";
|
link.description AS "link.description";
|
||||||
`, "DuckDuckGo", "http://www.duckduckgo.com", "Ask")
|
`, "DuckDuckGo", "http://www.duckduckgo.com", "Ask")
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
links := []model.Link{}
|
links := []model.Link{}
|
||||||
|
|
||||||
err := stmt.Query(tx, &links)
|
err := stmt.Query(tx, &links)
|
||||||
|
|
@ -325,8 +325,8 @@ func TestUpdateQueryContext(t *testing.T) {
|
||||||
|
|
||||||
time.Sleep(10 * time.Millisecond)
|
time.Sleep(10 * time.Millisecond)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
dest := []model.Link{}
|
var dest []model.Link
|
||||||
err := updateStmt.QueryContext(ctx, tx, &dest)
|
err := updateStmt.QueryContext(ctx, tx, &dest)
|
||||||
|
|
||||||
require.Error(t, err, "context deadline exceeded")
|
require.Error(t, err, "context deadline exceeded")
|
||||||
|
|
@ -344,8 +344,8 @@ func TestUpdateExecContext(t *testing.T) {
|
||||||
|
|
||||||
time.Sleep(10 * time.Millisecond)
|
time.Sleep(10 * time.Millisecond)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
_, err := updateStmt.ExecContext(ctx, tx)
|
_, err := updateStmt.ExecContext(ctx, db)
|
||||||
require.Error(t, err, "context deadline exceeded")
|
require.Error(t, err, "context deadline exceeded")
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
@ -388,7 +388,7 @@ RETURNING rental.rental_id AS "rental.rental_id",
|
||||||
store.address_id AS "store.address_id";
|
store.address_id AS "store.address_id";
|
||||||
`)
|
`)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
var dest []struct {
|
var dest []struct {
|
||||||
Rental model2.Rental
|
Rental model2.Rental
|
||||||
Store model2.Store
|
Store model2.Store
|
||||||
|
|
|
||||||
|
|
@ -1,9 +1,9 @@
|
||||||
package postgres
|
package postgres
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"database/sql"
|
|
||||||
"github.com/go-jet/jet/v2/internal/testutils"
|
"github.com/go-jet/jet/v2/internal/testutils"
|
||||||
. "github.com/go-jet/jet/v2/postgres"
|
. "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/model"
|
||||||
. "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/table"
|
. "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/table"
|
||||||
"github.com/stretchr/testify/assert"
|
"github.com/stretchr/testify/assert"
|
||||||
|
|
@ -251,7 +251,7 @@ RETURNING payment.payment_id AS "payment.payment_id",
|
||||||
payment.payment_date AS "payment.payment_date";
|
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
|
var payments []model.Payment
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -2,6 +2,7 @@ package sqlite
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"github.com/go-jet/jet/v2/qrm"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
|
@ -60,8 +61,6 @@ LIMIT 1;
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestDeleteContextDeadlineExceeded(t *testing.T) {
|
func TestDeleteContextDeadlineExceeded(t *testing.T) {
|
||||||
tx := beginSampleDBTx(t)
|
|
||||||
defer tx.Rollback()
|
|
||||||
|
|
||||||
deleteStmt := Link.
|
deleteStmt := Link.
|
||||||
DELETE().
|
DELETE().
|
||||||
|
|
@ -70,14 +69,18 @@ func TestDeleteContextDeadlineExceeded(t *testing.T) {
|
||||||
ctx, cancel := context.WithTimeout(context.Background(), 1*time.Microsecond)
|
ctx, cancel := context.WithTimeout(context.Background(), 1*time.Microsecond)
|
||||||
defer cancel()
|
defer cancel()
|
||||||
|
|
||||||
time.Sleep(10 * time.Millisecond)
|
time.Sleep(20 * time.Millisecond)
|
||||||
|
|
||||||
dest := []model.Link{}
|
testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx qrm.DB) {
|
||||||
err := deleteStmt.QueryContext(ctx, tx, &dest)
|
var dest []model.Link
|
||||||
require.Error(t, err, "context deadline exceeded")
|
err := deleteStmt.QueryContext(ctx, tx, &dest)
|
||||||
|
require.Error(t, err, "context deadline exceeded")
|
||||||
|
})
|
||||||
|
|
||||||
_, err = deleteStmt.ExecContext(ctx, tx)
|
testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx qrm.DB) {
|
||||||
require.Error(t, err, "context deadline exceeded")
|
_, err := deleteStmt.ExecContext(ctx, tx)
|
||||||
|
require.Error(t, err, "context deadline exceeded")
|
||||||
|
})
|
||||||
|
|
||||||
requireLogged(t, deleteStmt)
|
requireLogged(t, deleteStmt)
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -2,8 +2,8 @@ package sqlite
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"database/sql"
|
|
||||||
"github.com/go-jet/jet/v2/internal/utils/ptr"
|
"github.com/go-jet/jet/v2/internal/utils/ptr"
|
||||||
|
"github.com/go-jet/jet/v2/qrm"
|
||||||
"math/rand"
|
"math/rand"
|
||||||
|
|
||||||
"testing"
|
"testing"
|
||||||
|
|
@ -31,7 +31,7 @@ VALUES (?, ?, ?, ?),
|
||||||
101, "http://www.google.com", "Google", "Search engine",
|
101, "http://www.google.com", "Google", "Search engine",
|
||||||
102, "http://www.yahoo.com", "Yahoo", nil)
|
102, "http://www.yahoo.com", "Yahoo", nil)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx qrm.DB) {
|
||||||
testutils.AssertExec(t, insertQuery, tx)
|
testutils.AssertExec(t, insertQuery, tx)
|
||||||
requireLogged(t, insertQuery)
|
requireLogged(t, insertQuery)
|
||||||
|
|
||||||
|
|
@ -76,7 +76,7 @@ VALUES (100, 'http://www.postgresqltutorial.com', 'PostgreSQL Tutorial', NULL);
|
||||||
`,
|
`,
|
||||||
100, "http://www.postgresqltutorial.com", "PostgreSQL Tutorial", nil)
|
100, "http://www.postgresqltutorial.com", "PostgreSQL Tutorial", nil)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx qrm.DB) {
|
||||||
_, err := stmt.Exec(tx)
|
_, err := stmt.Exec(tx)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
requireLogged(t, stmt)
|
requireLogged(t, stmt)
|
||||||
|
|
@ -109,7 +109,7 @@ INSERT INTO link (url, name)
|
||||||
VALUES ('http://www.duckduckgo.com', 'Duck Duck go');
|
VALUES ('http://www.duckduckgo.com', 'Duck Duck go');
|
||||||
`, "http://www.duckduckgo.com", "Duck Duck go")
|
`, "http://www.duckduckgo.com", "Duck Duck go")
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx qrm.DB) {
|
||||||
_, err := query.Exec(tx)
|
_, err := query.Exec(tx)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
})
|
})
|
||||||
|
|
@ -131,7 +131,7 @@ INSERT INTO link
|
||||||
VALUES (1000, 'http://www.duckduckgo.com', 'Duck Duck go', NULL);
|
VALUES (1000, 'http://www.duckduckgo.com', 'Duck Duck go', NULL);
|
||||||
`, int32(1000), "http://www.duckduckgo.com", "Duck Duck go", nil)
|
`, int32(1000), "http://www.duckduckgo.com", "Duck Duck go", nil)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx qrm.DB) {
|
||||||
_, err := query.Exec(tx)
|
_, err := query.Exec(tx)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
})
|
})
|
||||||
|
|
@ -220,7 +220,7 @@ WHERE link.id = 24;
|
||||||
)
|
)
|
||||||
|
|
||||||
testutils.AssertDebugStatementSql(t, query, expectedSQL, int64(24))
|
testutils.AssertDebugStatementSql(t, query, expectedSQL, int64(24))
|
||||||
testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx qrm.DB) {
|
||||||
_, err := query.Exec(tx)
|
_, err := query.Exec(tx)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
|
@ -249,7 +249,7 @@ RETURNING link.id AS "link.id",
|
||||||
link.description AS "link.description";
|
link.description AS "link.description";
|
||||||
`)
|
`)
|
||||||
|
|
||||||
testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx qrm.DB) {
|
||||||
var link model.Link
|
var link model.Link
|
||||||
err := stmt.Query(tx, &link)
|
err := stmt.Query(tx, &link)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
@ -387,9 +387,15 @@ func TestInsertContextDeadlineExceeded(t *testing.T) {
|
||||||
time.Sleep(10 * time.Millisecond)
|
time.Sleep(10 * time.Millisecond)
|
||||||
|
|
||||||
var dest []model.Link
|
var dest []model.Link
|
||||||
err := stmt.QueryContext(ctx, sampleDB, &dest)
|
|
||||||
require.Error(t, err, "context deadline exceeded")
|
|
||||||
|
|
||||||
_, err = stmt.ExecContext(ctx, db)
|
testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx qrm.DB) {
|
||||||
require.Error(t, err, "context deadline exceeded")
|
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")
|
||||||
|
})
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -5,8 +5,8 @@ import (
|
||||||
"database/sql"
|
"database/sql"
|
||||||
"fmt"
|
"fmt"
|
||||||
"github.com/go-jet/jet/v2/internal/utils/throw"
|
"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/sqlite"
|
||||||
|
"github.com/go-jet/jet/v2/stmtcache"
|
||||||
"github.com/go-jet/jet/v2/tests/dbconfig"
|
"github.com/go-jet/jet/v2/tests/dbconfig"
|
||||||
"github.com/pkg/profile"
|
"github.com/pkg/profile"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
|
|
@ -17,28 +17,52 @@ import (
|
||||||
_ "github.com/mattn/go-sqlite3"
|
_ "github.com/mattn/go-sqlite3"
|
||||||
)
|
)
|
||||||
|
|
||||||
var db *sql.DB
|
var db *stmtcache.DB
|
||||||
var sampleDB *sql.DB
|
var sampleDB *stmtcache.DB
|
||||||
|
|
||||||
|
var withStatementCaching bool
|
||||||
|
|
||||||
|
func init() {
|
||||||
|
withStatementCaching = os.Getenv("JET_TESTS_WITH_STMT_CACHE") == "true"
|
||||||
|
}
|
||||||
|
|
||||||
func TestMain(m *testing.M) {
|
func TestMain(m *testing.M) {
|
||||||
defer profile.Start().Stop()
|
defer profile.Start().Stop()
|
||||||
|
|
||||||
var err error
|
func() {
|
||||||
db, err = sql.Open("sqlite3", "file:"+dbconfig.SakilaDBPath)
|
fmt.Printf("\nRunning sqlite tests caching enabled: %t \n", withStatementCaching)
|
||||||
throw.OnError(err)
|
|
||||||
defer db.Close()
|
|
||||||
|
|
||||||
_, err = db.Exec(fmt.Sprintf("ATTACH DATABASE '%s' as 'chinook';", dbconfig.ChinookDBPath))
|
sqlDB, err := sql.Open("sqlite3", "file:"+dbconfig.SakilaDBPath)
|
||||||
throw.OnError(err)
|
throw.OnError(err)
|
||||||
|
db = stmtcache.New(sqlDB).SetCaching(withStatementCaching)
|
||||||
|
defer db.Close()
|
||||||
|
|
||||||
sampleDB, err = sql.Open("sqlite3", dbconfig.TestSampleDBPath)
|
_, err = db.Exec(fmt.Sprintf("ATTACH DATABASE '%s' as 'chinook';", dbconfig.ChinookDBPath))
|
||||||
throw.OnError(err)
|
throw.OnError(err)
|
||||||
|
|
||||||
ret := m.Run()
|
sqlSampleDB, err := sql.Open("sqlite3", dbconfig.TestSampleDBPath)
|
||||||
|
throw.OnError(err)
|
||||||
|
sampleDB = stmtcache.New(sqlSampleDB).SetCaching(withStatementCaching)
|
||||||
|
defer sampleDB.Close()
|
||||||
|
|
||||||
if ret != 0 {
|
for i := 0; i < runCount(withStatementCaching); i++ {
|
||||||
os.Exit(ret)
|
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
|
||||||
}
|
}
|
||||||
|
|
||||||
|
return 1
|
||||||
}
|
}
|
||||||
|
|
||||||
var loggedSQL string
|
var loggedSQL string
|
||||||
|
|
@ -62,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()
|
query, args := statement.Sql()
|
||||||
queryLogged, argsLogged := queryInfo.Statement.Sql()
|
queryLogged, argsLogged := queryInfo.Statement.Sql()
|
||||||
|
|
||||||
|
|
@ -84,13 +108,13 @@ func requireLogged(t *testing.T, statement sqlite.Statement) {
|
||||||
require.Equal(t, loggedDebugSQL, statement.DebugSql())
|
require.Equal(t, loggedDebugSQL, statement.DebugSql())
|
||||||
}
|
}
|
||||||
|
|
||||||
func beginSampleDBTx(t *testing.T) *sql.Tx {
|
func beginSampleDBTx(t *testing.T) *stmtcache.Tx {
|
||||||
tx, err := sampleDB.Begin()
|
tx, err := sampleDB.BeginTx(context.Background(), nil)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
return tx
|
return tx
|
||||||
}
|
}
|
||||||
|
|
||||||
func beginDBTx(t *testing.T) *sql.Tx {
|
func beginDBTx(t *testing.T) *stmtcache.Tx {
|
||||||
tx, err := db.Begin()
|
tx, err := db.Begin()
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
return tx
|
return tx
|
||||||
|
|
|
||||||
|
|
@ -1,8 +1,8 @@
|
||||||
package sqlite
|
package sqlite
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"database/sql"
|
|
||||||
"github.com/go-jet/jet/v2/internal/testutils"
|
"github.com/go-jet/jet/v2/internal/testutils"
|
||||||
|
"github.com/go-jet/jet/v2/qrm"
|
||||||
"github.com/go-jet/jet/v2/internal/utils/ptr"
|
"github.com/go-jet/jet/v2/internal/utils/ptr"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
@ -49,7 +49,7 @@ WHERE people.people_id = ?;
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("should insert without generated columns", func(t *testing.T) {
|
t.Run("should insert without generated columns", func(t *testing.T) {
|
||||||
testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx *sql.Tx) {
|
testutils.ExecuteInTxAndRollback(t, sampleDB, func(tx qrm.DB) {
|
||||||
insertQuery := People.INSERT(
|
insertQuery := People.INSERT(
|
||||||
People.MutableColumns,
|
People.MutableColumns,
|
||||||
).MODEL(
|
).MODEL(
|
||||||
|
|
|
||||||
|
|
@ -987,14 +987,14 @@ WHERE artists.''ArtistId'' = 11;
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestRowsScan(t *testing.T) {
|
func TestRowsScan(t *testing.T) {
|
||||||
stmt :=
|
|
||||||
SELECT(
|
stmt := SELECT(
|
||||||
Inventory.AllColumns,
|
Inventory.AllColumns,
|
||||||
).FROM(
|
).FROM(
|
||||||
Inventory,
|
Inventory,
|
||||||
).ORDER_BY(
|
).ORDER_BY(
|
||||||
Inventory.InventoryID.ASC(),
|
Inventory.InventoryID.ASC(),
|
||||||
)
|
)
|
||||||
|
|
||||||
rows, err := stmt.Rows(context.Background(), db)
|
rows, err := stmt.Rows(context.Background(), db)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
|
||||||
131
tests/sqlite/stmtcache_test.go
Normal file
131
tests/sqlite/stmtcache_test.go
Normal file
|
|
@ -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)
|
||||||
|
}
|
||||||
|
|
@ -2,6 +2,7 @@ package sqlite
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"github.com/go-jet/jet/v2/qrm"
|
||||||
model2 "github.com/go-jet/jet/v2/tests/.gentestdata/sqlite/sakila/model"
|
model2 "github.com/go-jet/jet/v2/tests/.gentestdata/sqlite/sakila/model"
|
||||||
"github.com/go-jet/jet/v2/tests/.gentestdata/sqlite/sakila/table"
|
"github.com/go-jet/jet/v2/tests/.gentestdata/sqlite/sakila/table"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
@ -215,7 +216,7 @@ WHERE link.id = 20;
|
||||||
|
|
||||||
testutils.AssertDebugStatementSql(t, stmt, expectedSQL, nil, "DuckDuckGo", "http://www.duckduckgo.com", int32(20))
|
testutils.AssertDebugStatementSql(t, stmt, expectedSQL, nil, "DuckDuckGo", "http://www.duckduckgo.com", int32(20))
|
||||||
|
|
||||||
testutils.AssertExec(t, stmt, tx)
|
testutils.AssertExec(t, stmt, tx, 1)
|
||||||
requireLogged(t, stmt)
|
requireLogged(t, stmt)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -271,8 +272,6 @@ func TestUpdateWithInvalidModelData(t *testing.T) {
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestUpdateContextDeadlineExceeded(t *testing.T) {
|
func TestUpdateContextDeadlineExceeded(t *testing.T) {
|
||||||
tx := beginSampleDBTx(t)
|
|
||||||
defer tx.Rollback()
|
|
||||||
|
|
||||||
updateStmt := Link.UPDATE(Link.Name, Link.URL).
|
updateStmt := Link.UPDATE(Link.Name, Link.URL).
|
||||||
SET("Bong", "http://bong.com").
|
SET("Bong", "http://bong.com").
|
||||||
|
|
@ -283,12 +282,16 @@ func TestUpdateContextDeadlineExceeded(t *testing.T) {
|
||||||
|
|
||||||
time.Sleep(10 * time.Millisecond)
|
time.Sleep(10 * time.Millisecond)
|
||||||
|
|
||||||
var dest []model.Link
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
err := updateStmt.QueryContext(ctx, tx, &dest)
|
var dest []model.Link
|
||||||
require.Error(t, err, "context deadline exceeded")
|
err := updateStmt.QueryContext(ctx, tx, &dest)
|
||||||
|
require.Error(t, err, "context deadline exceeded")
|
||||||
|
})
|
||||||
|
|
||||||
_, err = updateStmt.ExecContext(ctx, tx)
|
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
|
||||||
require.Error(t, err, "context deadline exceeded")
|
_, err := updateStmt.ExecContext(ctx, tx)
|
||||||
|
require.Error(t, err, "context deadline exceeded")
|
||||||
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestUpdateFrom(t *testing.T) {
|
func TestUpdateFrom(t *testing.T) {
|
||||||
|
|
|
||||||
|
|
@ -1,8 +1,8 @@
|
||||||
package sqlite
|
package sqlite
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"database/sql"
|
|
||||||
"github.com/go-jet/jet/v2/internal/testutils"
|
"github.com/go-jet/jet/v2/internal/testutils"
|
||||||
|
"github.com/go-jet/jet/v2/qrm"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
@ -293,7 +293,7 @@ RETURNING payment.payment_id AS "payment.payment_id",
|
||||||
payment.last_update AS "payment.last_update";
|
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
|
var payments []model.Payment
|
||||||
|
|
||||||
err := stmt.Query(tx, &payments)
|
err := stmt.Query(tx, &payments)
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue