Add support for prepared statements caching.

This commit is contained in:
go-jet 2024-10-19 14:06:12 +02:00
parent 4bb9775134
commit 5f220569dd
20 changed files with 591 additions and 134 deletions

View file

@ -1,7 +1,6 @@
package postgres
import (
"database/sql"
"github.com/go-jet/jet/v2/internal/utils/ptr"
"github.com/stretchr/testify/assert"
@ -944,7 +943,7 @@ RETURNING employee.employee_id AS "employee.employee_id",
employee.manager_id AS "employee.manager_id",
employee.pto_accrual AS "employee.pto_accrual";
`
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
var windy model.Employee
windy.PtoAccrual = ptr.Of("3h")
stmt := Employee.UPDATE(Employee.PtoAccrual).SET(
@ -972,7 +971,7 @@ RETURNING employee.employee_id AS "employee.employee_id",
employee.manager_id AS "employee.manager_id",
employee.pto_accrual AS "employee.pto_accrual";
`
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
var employee model.Employee
employee.PtoAccrual = ptr.Of("5h")
stmt := Employee.INSERT(Employee.AllColumns).

View file

@ -4,6 +4,7 @@ import (
"context"
"database/sql"
"fmt"
"github.com/go-jet/jet/v2/stmtcache"
"github.com/go-jet/jet/v2/tests/internal/utils/repo"
"github.com/jackc/pgx/v4/stdlib"
"os"
@ -19,15 +20,17 @@ import (
_ "github.com/jackc/pgx/v4/stdlib"
)
var db *postgres.DB
var db *stmtcache.DB
var testRoot string
var source string
var withStatementCaching bool
const CockroachDB = "COCKROACH_DB"
func init() {
source = os.Getenv("PG_SOURCE")
withStatementCaching = os.Getenv("JET_TESTS_WITH_STMT_CACHE") == "true"
}
func sourceIsCockroachDB() bool {
@ -45,39 +48,50 @@ func TestMain(m *testing.M) {
setTestRoot()
for _, driverName := range []string{"pgx", "postgres"} {
fmt.Printf("\nRunning postgres tests for '%s' driver\n", driverName)
for _, driverName := range []string{"postgres", "pgx"} {
fmt.Printf("\nRunning postgres tests for driver: %s, caching enabled: %t \n", driverName, withStatementCaching)
func() {
connectionString := dbconfig.PostgresConnectString
if sourceIsCockroachDB() {
connectionString = dbconfig.CockroachConnectString
}
sqlDB, err := sql.Open(driverName, connectionString)
sqlDB, err := sql.Open(driverName, getConnectionString())
if err != nil {
fmt.Println(err.Error())
panic("Failed to connect to test db")
}
db = postgres.NewDB(sqlDB).WithStatementsCaching(true)
defer db.Close()
db = stmtcache.New(sqlDB).SetCaching(withStatementCaching)
defer func(db *stmtcache.DB) {
err := db.Close()
if err != nil {
fmt.Printf("ERROR: Failed to close db connection, %v", err)
}
}(db)
for i := 0; i < 2; i++ {
for i := 0; i < runCount(withStatementCaching); i++ {
ret := m.Run()
if ret != 0 {
fmt.Printf("\nFAIL: Running postgres tests failed for driver: %s, caching enabled: %t \n", driverName, withStatementCaching)
os.Exit(ret)
}
}
err = db.Clear()
if err != nil {
os.Exit(-2)
}
}()
}
}
func runCount(stmtCaching bool) int {
if stmtCaching {
return 2
}
return 1
}
func getConnectionString() string {
if sourceIsCockroachDB() {
return dbconfig.CockroachConnectString
}
return dbconfig.PostgresConnectString
}
func setTestRoot() {

View file

@ -1,8 +1,8 @@
package postgres
import (
"github.com/go-jet/jet/v2/qrm"
"github.com/go-jet/jet/v2/internal/utils/ptr"
"github.com/go-jet/jet/v2/qrm"
"github.com/google/uuid"
"testing"

View 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)
}

View file

@ -1,9 +1,9 @@
package postgres
import (
"database/sql"
"github.com/go-jet/jet/v2/internal/testutils"
. "github.com/go-jet/jet/v2/postgres"
"github.com/go-jet/jet/v2/qrm"
"github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/model"
. "github.com/go-jet/jet/v2/tests/.gentestdata/jetdb/dvds/table"
"github.com/stretchr/testify/assert"
@ -251,7 +251,7 @@ RETURNING payment.payment_id AS "payment.payment_id",
payment.payment_date AS "payment.payment_date";
`)
testutils.ExecuteInTxAndRollback(t, db, func(tx *sql.Tx) {
testutils.ExecuteInTxAndRollback(t, db, func(tx qrm.DB) {
var payments []model.Payment