Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 5 additions & 1 deletion client.go
Original file line number Diff line number Diff line change
Expand Up @@ -1110,9 +1110,13 @@ func (c *Client[TTx]) Start(ctx context.Context) error {
// available, the client appears to have started even though it's completely
// non-functional. Here we try to make an initial assessment of health and
// return quickly in case of an apparent problem.
if err := c.driver.GetExecutor().Exec(fetchCtx, "SELECT 1"); err != nil {
executor := c.driver.GetExecutor()
if err := executor.Ping(fetchCtx); err != nil {
return fmt.Errorf("error making initial connection to database: %w", err)
}
if err := executor.InitDriver(fetchCtx); err != nil {
return fmt.Errorf("error initializing driver: %w", err)
}

// Each time we start, we need a fresh completer subscribe channel to
// send job completion events on, because the completer will close it
Expand Down
20 changes: 20 additions & 0 deletions client_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -8316,6 +8316,26 @@ func Test_Client_Start_Error(t *testing.T) {
require.Equal(t, pgerrcode.InvalidCatalogName, pgErr.Code)
})

t.Run("DatabaseErrorAfterSuccessfulStart", func(t *testing.T) {
t.Parallel()

dbPool := riversharedtest.DBPoolClone(ctx, t)
driver := NewDriverPollOnly(dbPool)
schema := riverdbtest.TestSchema(ctx, t, driver, nil)

client, err := NewClient(driver, newTestConfig(t, schema))
require.NoError(t, err)
t.Cleanup(func() { require.NoError(t, client.Stop(ctx)) })

require.NoError(t, client.Start(ctx))
require.NoError(t, client.Stop(ctx))

dbPool.Close()

err = client.Start(ctx)
require.ErrorIs(t, err, riverdriver.ErrClosedPool)
})

t.Run("CanRestartAfterFailure", func(t *testing.T) {
t.Parallel()

Expand Down
5 changes: 0 additions & 5 deletions internal/rivercommon/river_common.go
Original file line number Diff line number Diff line change
Expand Up @@ -45,11 +45,6 @@ const (

// MetadataKeyRescueCount records how many times the job has been rescued.
MetadataKeyRescueCount = "river:rescue_count"

// MetadataKeyUniqueNonce is a special metadata key used by the SQLite driver to
// determine whether an upsert is was skipped or not because the `(xmax != 0)`
// trick we use in Postgres doesn't work in SQLite.
MetadataKeyUniqueNonce = "river:unique_nonce"
)

type ContextKeyClient struct{}
Expand Down
9 changes: 9 additions & 0 deletions riverdriver/river_driver_interface.go
Original file line number Diff line number Diff line change
Expand Up @@ -224,6 +224,11 @@ type Executor interface {
IndexReindex(ctx context.Context, params *IndexReindexParams) error
IndexReindexArtifacts(ctx context.Context, params *IndexReindexArtifactsParams) ([]string, error)

// InitDriver initializes driver-specific state using information read from
// the database. Implementations must be safe to call concurrently and
// repeatedly, and should cache successfully initialized state.
InitDriver(ctx context.Context) error

JobCancel(ctx context.Context, params *JobCancelParams) (*rivertype.JobRow, error)
JobCountByAllStates(ctx context.Context, params *JobCountByAllStatesParams) (map[rivertype.JobState]int, error)
JobCountByQueueAndState(ctx context.Context, params *JobCountByQueueAndStateParams) ([]*JobCountByQueueAndStateResult, error)
Expand Down Expand Up @@ -289,6 +294,10 @@ type Executor interface {
NotificationDeleteBefore(ctx context.Context, params *NotificationDeleteBeforeParams) (int, error)

NotifyMany(ctx context.Context, params *NotifyManyParams) error

// Ping checks that the database is reachable.
Ping(ctx context.Context) error

PGAdvisoryXactLock(ctx context.Context, key int64) (*struct{}, error)

QueueCreateOrSetUpdatedAt(ctx context.Context, params *QueueCreateOrSetUpdatedAtParams) (*rivertype.Queue, error)
Expand Down
18 changes: 18 additions & 0 deletions riverdriver/riverdatabasesql/internal/dbsqlc/pg_misc.sql.go

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

72 changes: 67 additions & 5 deletions riverdriver/riverdatabasesql/river_database_sql_driver.go
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@ import (
"io/fs"
"math"
"strings"
"sync/atomic"
"time"

"github.com/jackc/pgx/v5/pgxpool"
Expand All @@ -28,6 +29,7 @@ import (
"github.com/riverqueue/river/rivershared/uniquestates"
"github.com/riverqueue/river/rivershared/util/dbutil"
"github.com/riverqueue/river/rivershared/util/ptrutil"
"github.com/riverqueue/river/rivershared/util/randutil"
"github.com/riverqueue/river/rivershared/util/savepointutil"
"github.com/riverqueue/river/rivershared/util/sliceutil"
"github.com/riverqueue/river/rivertype"
Expand All @@ -38,9 +40,10 @@ var migrationFS embed.FS

// Driver is an implementation of riverdriver.Driver for database/sql.
type Driver struct {
dbPool *sql.DB
listenerDriver *riverpgxv5.Driver
replacer sqlctemplate.Replacer
dbPool *sql.DB
listenerDriver *riverpgxv5.Driver
replacer sqlctemplate.Replacer
uniqueInsertMode atomic.Uint32
}

// New returns a new database/sql River driver for use with River.
Expand Down Expand Up @@ -247,6 +250,11 @@ func (e *Executor) IndexesExist(ctx context.Context, params *riverdriver.Indexes
return exists, nil
}

func (e *Executor) InitDriver(ctx context.Context) error {
_, err := e.uniqueInsertMode(ctx)
return err
}

func (e *Executor) JobCancel(ctx context.Context, params *riverdriver.JobCancelParams) (*rivertype.JobRow, error) {
cancelledAt, err := params.CancelAttemptedAt.MarshalJSON()
if err != nil {
Expand Down Expand Up @@ -402,6 +410,16 @@ func (e *Executor) JobGetStuck(ctx context.Context, params *riverdriver.JobGetSt
}

func (e *Executor) JobInsertFastMany(ctx context.Context, params *riverdriver.JobInsertFastManyParams) ([]*riverdriver.JobInsertFastResult, error) {
uniqueInsertMode, err := e.uniqueInsertMode(ctx)
if err != nil {
return nil, err
}

var uniqueNonce string
if uniqueInsertMode == riverdriver.UniqueInsertModeMetadataNonce {
uniqueNonce = randutil.Hex(8)
}

insertJobsParams := &dbsqlc.JobInsertFastManyParams{
ID: make([]int64, len(params.Jobs)),
Args: make([]string, len(params.Jobs)),
Expand Down Expand Up @@ -442,7 +460,16 @@ func (e *Executor) JobInsertFastMany(ctx context.Context, params *riverdriver.Jo
insertJobsParams.CreatedAt[i] = createdAt
insertJobsParams.Kind[i] = params.Kind
insertJobsParams.MaxAttempts[i] = int16(min(params.MaxAttempts, math.MaxInt16)) //nolint:gosec
insertJobsParams.Metadata[i] = cmp.Or(string(params.Metadata), "{}")
metadata := []byte(cmp.Or(string(params.Metadata), "{}"))
if uniqueNonce != "" {
var err error
metadata, err = riverdriver.UniqueInsertMetadataWithNonce(metadata, uniqueNonce)
if err != nil {
return nil, err
}
}

insertJobsParams.Metadata[i] = string(metadata)
insertJobsParams.Priority[i] = int16(min(params.Priority, math.MaxInt16)) //nolint:gosec
insertJobsParams.Queue[i] = params.Queue
insertJobsParams.ScheduledAt[i] = scheduledAt
Expand All @@ -452,6 +479,9 @@ func (e *Executor) JobInsertFastMany(ctx context.Context, params *riverdriver.Jo
insertJobsParams.UniqueStates[i] = int32(params.UniqueStates)
}

ctx = sqlctemplate.WithReplacements(ctx, map[string]sqlctemplate.Replacement{
"unique_skipped_as_duplicate": {Value: uniqueInsertMode.SQL(), Stable: true},
}, nil)
items, err := dbsqlc.New().JobInsertFastMany(schemaTemplateParam(ctx, params.Schema), e.dbtx, insertJobsParams)
if err != nil {
return nil, interpretError(err)
Expand All @@ -462,7 +492,13 @@ func (e *Executor) JobInsertFastMany(ctx context.Context, params *riverdriver.Jo
if err != nil {
return nil, err
}
return &riverdriver.JobInsertFastResult{Job: job, UniqueSkippedAsDuplicate: row.UniqueSkippedAsDuplicate}, nil

uniqueSkippedAsDuplicate := row.UniqueSkippedAsDuplicate
if uniqueInsertMode == riverdriver.UniqueInsertModeMetadataNonce {
uniqueSkippedAsDuplicate = riverdriver.UniqueInsertMetadataIsDuplicate(job.Metadata, uniqueNonce)
}

return &riverdriver.JobInsertFastResult{Job: job, UniqueSkippedAsDuplicate: uniqueSkippedAsDuplicate}, nil
})
}

Expand Down Expand Up @@ -932,6 +968,10 @@ func (e *Executor) NotifyMany(ctx context.Context, params *riverdriver.NotifyMan
})
}

func (e *Executor) Ping(ctx context.Context) error {
return e.Exec(ctx, "SELECT 1")
}

func (e *Executor) PGAdvisoryXactLock(ctx context.Context, key int64) (*struct{}, error) {
err := dbsqlc.New().PGAdvisoryXactLock(ctx, e.dbtx, key)
return &struct{}{}, interpretError(err)
Expand Down Expand Up @@ -1090,6 +1130,28 @@ func (e *Executor) TableTruncate(ctx context.Context, params *riverdriver.TableT
return interpretError(err)
}

func (e *Executor) uniqueInsertMode(ctx context.Context) (riverdriver.UniqueInsertMode, error) {
if e.driver != nil {
if mode := riverdriver.UniqueInsertMode(e.driver.uniqueInsertMode.Load()); mode != riverdriver.UniqueInsertModeUnknown {
return mode, nil
}
}

productAndVersion, err := dbsqlc.New().PGGetProductAndVersion(ctx, e.dbtx)
if err != nil {
return riverdriver.UniqueInsertModeUnknown, interpretError(err)
}

mode := riverdriver.UniqueInsertModeFromProductAndVersion(productAndVersion.Product, productAndVersion.VersionNum)
if e.driver != nil {
// Concurrent callers may both detect, but the first successful result
// becomes the driver's cached mode.
e.driver.uniqueInsertMode.CompareAndSwap(uint32(riverdriver.UniqueInsertModeUnknown), uint32(mode))
mode = riverdriver.UniqueInsertMode(e.driver.uniqueInsertMode.Load())
}
return mode, nil
}

type ExecutorTx struct {
Executor

Expand Down
61 changes: 61 additions & 0 deletions riverdriver/riverdatabasesql/river_database_sql_driver_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,18 +5,34 @@ import (
"database/sql"
"errors"
"testing"
"time"

"github.com/jackc/pgx/v5/pgxpool"
_ "github.com/jackc/pgx/v5/stdlib"
"github.com/stretchr/testify/require"

"github.com/riverqueue/river/riverdriver"
"github.com/riverqueue/river/rivershared/riversharedtest"
"github.com/riverqueue/river/rivershared/sqlctemplate"
"github.com/riverqueue/river/rivershared/testsignal"
"github.com/riverqueue/river/rivershared/util/urlutil"
"github.com/riverqueue/river/rivertype"
)

// Verify interface compliance.
var _ riverdriver.Driver[*sql.Tx] = New(nil)

type executorInitDriverTestDBTX struct {
*sql.DB

QueryRowStarted testsignal.TestSignal[struct{}]
}

func (d *executorInitDriverTestDBTX) QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row {
d.QueryRowStarted.Signal(struct{}{})
return d.DB.QueryRowContext(ctx, query, args...)
}

func TestNew(t *testing.T) {
t.Parallel()

Expand Down Expand Up @@ -48,6 +64,51 @@ func TestNew(t *testing.T) {
})
}

func TestExecutor_InitDriverDoesNotBlockTransaction(t *testing.T) {
t.Parallel()

ctx := context.Background()
dbPool, err := sql.Open("pgx", urlutil.DatabaseSQLCompatibleURL(riversharedtest.TestDatabaseURL()))
require.NoError(t, err)
dbPool.SetMaxOpenConns(1)
t.Cleanup(func() { require.NoError(t, dbPool.Close()) })

driver := New(dbPool)
tx, err := dbPool.BeginTx(ctx, nil)
require.NoError(t, err)
t.Cleanup(func() { _ = tx.Rollback() })

poolDBTX := &executorInitDriverTestDBTX{DB: dbPool}
poolDBTX.QueryRowStarted.Init(t)
poolExecutor := &Executor{
dbPool: dbPool,
dbtx: templateReplaceWrapper{dbtx: poolDBTX, replacer: &driver.replacer},
driver: driver,
}

initCtx, initCancel := context.WithTimeout(ctx, 10*time.Second)
t.Cleanup(initCancel)

var poolInitFinished testsignal.TestSignal[error]
poolInitFinished.Init(t)
go func() { poolInitFinished.Signal(poolExecutor.InitDriver(initCtx)) }()
poolDBTX.QueryRowStarted.WaitOrTimeout()

var txInitFinished testsignal.TestSignal[error]
txInitFinished.Init(t)
go func() { txInitFinished.Signal(driver.UnwrapExecutor(tx).InitDriver(ctx)) }()

select {
case err := <-txInitFinished.WaitC():
require.NoError(t, err)
case <-time.After(2 * time.Second):
require.FailNow(t, "transactional driver initialization blocked behind pool initialization")
}

require.NoError(t, tx.Rollback())
require.NoError(t, poolInitFinished.WaitOrTimeout())
}

func TestNewWithPgxListener(t *testing.T) {
t.Parallel()

Expand Down
57 changes: 57 additions & 0 deletions riverdriver/riverdatabasesql/yugabyte_compatibility_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,57 @@
package riverdatabasesql

import (
"fmt"
"io/fs"
"os"
"regexp"
"strings"
"testing"

"github.com/stretchr/testify/require"
)

func TestYugabyteCompatibility(t *testing.T) {
t.Parallel()

// YugabyteDB doesn't expose PostgreSQL's transaction-related system columns:
// https://docs.yugabyte.com/stable/yugabyte-voyager/known-issues/postgresql/#system-columns-is-not-yet-supported
// The unique insert query may use xmax only because the entire expression is
// replaced when the driver detects YugabyteDB.
var (
uniqueInsertModeTemplateRE = regexp.MustCompile(`(?s)/\*\s*TEMPLATE_BEGIN: unique_skipped_as_duplicate\s*\*/.*?/\*\s*TEMPLATE_END\s*\*/`)
unsupportedSystemColumnRE = regexp.MustCompile(`(?i)\b(?:cmax|cmin|ctid|xmax|xmin)\b`)
)

sourceRoot, err := os.OpenRoot(".")
require.NoError(t, err)
t.Cleanup(func() { require.NoError(t, sourceRoot.Close()) })

var violations []string
err = fs.WalkDir(sourceRoot.FS(), ".", func(path string, entry fs.DirEntry, err error) error {
if err != nil {
return err
}
if entry.IsDir() || (!strings.HasSuffix(path, ".go") && !strings.HasSuffix(path, ".sql")) || strings.HasSuffix(path, "_test.go") {
return nil
}

contents, err := sourceRoot.ReadFile(path)
if err != nil {
return err
}
contents = uniqueInsertModeTemplateRE.ReplaceAll(contents, nil)

for lineNum, line := range strings.Split(string(contents), "\n") {
for _, column := range unsupportedSystemColumnRE.FindAllString(line, -1) {
violations = append(violations, fmt.Sprintf("%s:%d: %s", path, lineNum+1, column))
}
}
return nil
})
require.NoError(t, err)

require.Empty(t, violations,
"YugabyteDB-incompatible PostgreSQL system columns must only appear inside SQL templates that replace them for YugabyteDB",
)
}
Loading
Loading