diff --git a/.github/workflows/ci.yaml b/.github/workflows/ci.yaml index d62f2dec..1c7e6d65 100644 --- a/.github/workflows/ci.yaml +++ b/.github/workflows/ci.yaml @@ -77,6 +77,8 @@ jobs: # Run the latest Go version against all supported Postgres versions: - go-version: "1.26" postgres-version: 18 + # MySQL tests are opt-in and only need to run in one matrix job. + mysql-enabled: true - go-version: "1.26" postgres-version: 17 - go-version: "1.26" @@ -91,6 +93,7 @@ jobs: # latest Postgres version: - go-version: "1.25" postgres-version: 18 + fail-fast: false timeout-minutes: 5 @@ -128,7 +131,17 @@ jobs: - name: Set up database run: psql -c "CREATE DATABASE river_test" $ADMIN_DATABASE_URL + # MySQL is pre-installed on GitHub-hosted Ubuntu runners. Start the + # service and configure passwordless root access for MySQL-enabled runs. + - name: Start MySQL + if: matrix.mysql-enabled + run: | + sudo systemctl start mysql + mysql -u root -proot -e "ALTER USER 'root'@'localhost' IDENTIFIED BY ''; FLUSH PRIVILEGES;" + - name: Test + env: + RIVER_MYSQL_TESTS_ENABLED: ${{ matrix.mysql-enabled && 'true' || '' }} run: make test/race cli: diff --git a/Makefile b/Makefile index ce82cd4f..bbe9b86a 100644 --- a/Makefile +++ b/Makefile @@ -27,6 +27,7 @@ generate/migrations: ## Sync changes of pgxv5 migrations to database/sql .PHONY: generate/sqlc generate/sqlc: ## Generate sqlc cd riverdriver/riverdatabasesql/internal/dbsqlc && sqlc generate + cd riverdriver/rivermysql/internal/dbsqlc && sqlc generate cd riverdriver/riverpgxv5/internal/dbsqlc && sqlc generate cd riverdriver/riversqlite/internal/dbsqlc && sqlc generate @@ -101,5 +102,6 @@ verify/migrations: ## Verify synced migrations .PHONY: verify/sqlc verify/sqlc: ## Verify generated sqlc cd riverdriver/riverdatabasesql/internal/dbsqlc && sqlc diff + cd riverdriver/rivermysql/internal/dbsqlc && sqlc diff cd riverdriver/riverpgxv5/internal/dbsqlc && sqlc diff cd riverdriver/riversqlite/internal/dbsqlc && sqlc diff diff --git a/client.go b/client.go index 422eb4d7..ce85cd28 100644 --- a/client.go +++ b/client.go @@ -1011,11 +1011,11 @@ func NewClient[TTx any](driver riverdriver.Driver[TTx], config *Config) (*Client client.testSignals.queueCleaner = &queueCleaner.TestSignals } - if driver.DatabaseName() == riverdriver.DatabaseNameSQLite { - sqliteNotificationCleaner := maintenance.NewSQLiteNotificationCleaner(archetype, &maintenance.SQLiteNotificationCleanerConfig{ + if driver.DatabaseName() == riverdriver.DatabaseNameMySQL || driver.DatabaseName() == riverdriver.DatabaseNameSQLite { + notificationCleaner := maintenance.NewSQLiteNotificationCleaner(archetype, &maintenance.SQLiteNotificationCleanerConfig{ Schema: config.Schema, }, driver.GetExecutor()) - maintenanceServices = append(maintenanceServices, sqliteNotificationCleaner) + maintenanceServices = append(maintenanceServices, notificationCleaner) } { @@ -2425,11 +2425,13 @@ func (c *Client[TTx]) JobList(ctx context.Context, params *JobListParams) (*JobL if params == nil { params = NewJobListParams() } - params.schema = c.config.Schema - if c.driver.DatabaseName() == riverdriver.DatabaseNameSQLite && params.metadataCalled { return nil, errJobListParamsMetadataNotSupportedSQLite } + if c.driver.DatabaseName() == riverdriver.DatabaseNameMySQL && params.metadataCalled { + params = params.withMetadataPredicateSQL("JSON_CONTAINS(metadata, CAST(@metadata_fragment AS JSON))") + } + params.schema = c.config.Schema dbParams, err := params.toDBParams() if err != nil { @@ -2466,11 +2468,13 @@ func (c *Client[TTx]) JobListTx(ctx context.Context, tx TTx, params *JobListPara if params == nil { params = NewJobListParams() } - params.schema = c.config.Schema - if c.driver.DatabaseName() == riverdriver.DatabaseNameSQLite && params.metadataCalled { return nil, errJobListParamsMetadataNotSupportedSQLite } + if c.driver.DatabaseName() == riverdriver.DatabaseNameMySQL && params.metadataCalled { + params = params.withMetadataPredicateSQL("JSON_CONTAINS(metadata, CAST(@metadata_fragment AS JSON))") + } + params.schema = c.config.Schema dbParams, err := params.toDBParams() if err != nil { diff --git a/client_test.go b/client_test.go index 0274f5ef..09dcb9a2 100644 --- a/client_test.go +++ b/client_test.go @@ -8299,6 +8299,32 @@ func Test_NewClient_Overrides(t *testing.T) { require.Len(t, client.config.WorkerMiddleware, 1) } +func Test_NewClient_MySQLNotificationCleaner(t *testing.T) { + t.Parallel() + + ctx := context.Background() + + var ( + dbPool = riversharedtest.DBPool(ctx, t) + pgxDriver = riverpgxv5.New(dbPool) + schema = riverdbtest.TestSchema(ctx, t, pgxDriver, nil) + ) + + workers := NewWorkers() + AddWorker(workers, &noOpWorker{}) + + client, err := NewClient(NewDriverMySQLDatabaseName(dbPool), &Config{ + Queues: map[string]QueueConfig{QueueDefault: {MaxWorkers: 1}}, + Schema: schema, + TestOnly: true, + Workers: workers, + }) + require.NoError(t, err) + + notificationCleaner := maintenance.GetService[*maintenance.SQLiteNotificationCleaner](client.queueMaintainer) + require.Equal(t, schema, notificationCleaner.Config.Schema) +} + func Test_NewClient_PluginsAndHybrids(t *testing.T) { t.Parallel() @@ -9640,6 +9666,23 @@ func TestDefaultClientIDWithHost(t *testing.T) { require.Equal(t, strings.Repeat("a", 60)+"_2024_03_07T04_39_12_123456", defaultClientIDWithHost(startedAt, strings.Repeat("a", 61))) } +// DriverMySQLDatabaseName is a pgx driver that reports MySQL as its database +// name so MySQL-specific client wiring can be tested without importing the +// separate MySQL driver module. +type DriverMySQLDatabaseName struct { + riverpgxv5.Driver +} + +// NewDriverMySQLDatabaseName returns a new test driver that reports MySQL as +// its database name. +func NewDriverMySQLDatabaseName(dbPool *pgxpool.Pool) *DriverMySQLDatabaseName { + return &DriverMySQLDatabaseName{ + Driver: *riverpgxv5.New(dbPool), + } +} + +func (d *DriverMySQLDatabaseName) DatabaseName() string { return riverdriver.DatabaseNameMySQL } + // DriverPollOnly simulates a driver without a listener. An example of this is // Postgres through `riverdatabasesql`, which is Postgres (so it can notify), // but where `database/sql` provides no listener mechanism. We could use the diff --git a/cmd/river/go.sum b/cmd/river/go.sum index b5e5d990..0f09a7aa 100644 --- a/cmd/river/go.sum +++ b/cmd/river/go.sum @@ -36,18 +36,18 @@ github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZb github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE= github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= -github.com/riverqueue/river v0.40.0 h1:4dynKqqU1P22iPmwWDfDj/YZXnuUTZysXXF3wNHekNw= -github.com/riverqueue/river v0.40.0/go.mod h1:auvB4kHqM97tnshEQxzy2E7aFvJFhl010NB6N29DXXc= -github.com/riverqueue/river/riverdriver v0.40.0 h1:QjBHHh+kaxgUgK9tPumrfx7W14vhAHLIt+0dK40ikI8= -github.com/riverqueue/river/riverdriver v0.40.0/go.mod h1:7wEqxsqvtjGk3hBKGKK3IvdnULlPZltpAfkbMr9M2VY= -github.com/riverqueue/river/riverdriver/riverpgxv5 v0.40.0 h1:WbzXgGukvOqYpsORX4rr03sPZfasjoMckLr4On+Je+4= -github.com/riverqueue/river/riverdriver/riverpgxv5 v0.40.0/go.mod h1:HtgFIcn/eJ9O3f371S1cSb0ttS7xJMMkkTD/K6gQKQg= -github.com/riverqueue/river/riverdriver/riversqlite v0.40.0 h1:DE/22nz7BKgQ+0XpPFg/3hoO4mP8Ij+gGcIOS9PusFg= -github.com/riverqueue/river/riverdriver/riversqlite v0.40.0/go.mod h1:/Xuw1OOo+X4Pj5DCKu4MLbF5BhGIBl/OWp93kN2v5ec= -github.com/riverqueue/river/rivershared v0.40.0 h1:6ZX1Ok94Nkcx/WpHPhexpfdkRB1BqRWFEWrbDszg+JY= -github.com/riverqueue/river/rivershared v0.40.0/go.mod h1:Z77wB2/ctD+zSvlNDKjxC1SW9c1Y2jyE+PXR2CxBlzQ= -github.com/riverqueue/river/rivertype v0.40.0 h1:6sFLhs0OtkxDZ/iyZIUqaQJb/G1krnb2bh9lfXdCimg= -github.com/riverqueue/river/rivertype v0.40.0/go.mod h1:D1Ad+EaZiaXbQbJcJcfeicXJMBKno0n6UcfKI5Q7DIQ= +github.com/riverqueue/river v0.41.0 h1:E7Yfyhn74IgVaCBKsBpTHoC0LxWJvoG0jKt7Tb+lJZg= +github.com/riverqueue/river v0.41.0/go.mod h1:WXiAF1/2gfUPj3H+WqQG3Q5NW5w9JmRXRTe7U9zybNc= +github.com/riverqueue/river/riverdriver v0.41.0 h1:A5g80n6fCGu9Bp4hSq7gys7mvfOiuWkbQHuFx3iPp7w= +github.com/riverqueue/river/riverdriver v0.41.0/go.mod h1:OeGOhZYoH7Ac0CrVanNAAwD9C0WAXy4LzUEl65Ex0B4= +github.com/riverqueue/river/riverdriver/riverpgxv5 v0.41.0 h1:uHGluToMMWvCnpNkAbEWEmP4YAjdseuriEXGu1dwOp4= +github.com/riverqueue/river/riverdriver/riverpgxv5 v0.41.0/go.mod h1:0ueUi+3eW5fIsibS5VkD2HzQmwLKTAiM/uiiVMX/SME= +github.com/riverqueue/river/riverdriver/riversqlite v0.41.0 h1:4IBOD4Uzil620AT/WT5xPHVoQSk2OiQ5UvG7OQ2piOc= +github.com/riverqueue/river/riverdriver/riversqlite v0.41.0/go.mod h1:k16A8Tt6Ke49rCU4k33T5EKXAhAXiExyU2qV+6hpkNs= +github.com/riverqueue/river/rivershared v0.41.0 h1:ax3uz5KiqfmcvRQ3kKB6iYs9Nt8J5L1RNfOKoVMWkqw= +github.com/riverqueue/river/rivershared v0.41.0/go.mod h1:FaZ7bxC2DORhyFDVyHKRiZzgQPf2b/IELwPBXnHZVLA= +github.com/riverqueue/river/rivertype v0.41.0 h1:dfscvt1asf1PpeeHTMFxQdV1ZfMoReiiMNFCPPfR0as= +github.com/riverqueue/river/rivertype v0.41.0/go.mod h1:D1Ad+EaZiaXbQbJcJcfeicXJMBKno0n6UcfKI5Q7DIQ= github.com/robfig/cron/v3 v3.0.1 h1:WdRxkvbJztn8LMz/QEvLN5sBU+xKpSqwwUO1Pjr4qDs= github.com/robfig/cron/v3 v3.0.1/go.mod h1:eQICP3HwyT7UooqI/z+Ov+PtYAWygg1TEWWzGIFLtro= github.com/rogpeppe/go-internal v1.12.0 h1:exVL4IDcn6na9z1rAb56Vxr+CgyK3nn3O+epU5NdKM8= @@ -76,8 +76,8 @@ github.com/tidwall/sjson v1.2.5/go.mod h1:Fvgq9kS/6ociJEDnK0Fk1cpYF4FIW6ZF7LAe+6 go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto= go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE= go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg= -golang.org/x/mod v0.37.0 h1:vF1DjpVEshcIqoEaauuHebaLk1O1forxjxBaVn884JQ= -golang.org/x/mod v0.37.0/go.mod h1:m8S8VeM9r4dzDwjrKO0a1sZP3YjeMamRRlD+fmR2Q/0= +golang.org/x/mod v0.38.0 h1:MECBjubtXD7yj4HrhIUcywNaGeNVUdfVnxmPajOk4yk= +golang.org/x/mod v0.38.0/go.mod h1:V6Xz0pq8TQ3dGqVQ1FVHuelZpAL0uNhSkk9ogYP3c40= golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek= golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= diff --git a/go.sum b/go.sum index 4f7c9572..db5863fa 100644 --- a/go.sum +++ b/go.sum @@ -17,14 +17,14 @@ github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= -github.com/riverqueue/river/riverdriver v0.40.0 h1:QjBHHh+kaxgUgK9tPumrfx7W14vhAHLIt+0dK40ikI8= -github.com/riverqueue/river/riverdriver v0.40.0/go.mod h1:7wEqxsqvtjGk3hBKGKK3IvdnULlPZltpAfkbMr9M2VY= -github.com/riverqueue/river/riverdriver/riverpgxv5 v0.40.0 h1:WbzXgGukvOqYpsORX4rr03sPZfasjoMckLr4On+Je+4= -github.com/riverqueue/river/riverdriver/riverpgxv5 v0.40.0/go.mod h1:HtgFIcn/eJ9O3f371S1cSb0ttS7xJMMkkTD/K6gQKQg= -github.com/riverqueue/river/rivershared v0.40.0 h1:6ZX1Ok94Nkcx/WpHPhexpfdkRB1BqRWFEWrbDszg+JY= -github.com/riverqueue/river/rivershared v0.40.0/go.mod h1:Z77wB2/ctD+zSvlNDKjxC1SW9c1Y2jyE+PXR2CxBlzQ= -github.com/riverqueue/river/rivertype v0.40.0 h1:6sFLhs0OtkxDZ/iyZIUqaQJb/G1krnb2bh9lfXdCimg= -github.com/riverqueue/river/rivertype v0.40.0/go.mod h1:D1Ad+EaZiaXbQbJcJcfeicXJMBKno0n6UcfKI5Q7DIQ= +github.com/riverqueue/river/riverdriver v0.41.0 h1:A5g80n6fCGu9Bp4hSq7gys7mvfOiuWkbQHuFx3iPp7w= +github.com/riverqueue/river/riverdriver v0.41.0/go.mod h1:OeGOhZYoH7Ac0CrVanNAAwD9C0WAXy4LzUEl65Ex0B4= +github.com/riverqueue/river/riverdriver/riverpgxv5 v0.41.0 h1:uHGluToMMWvCnpNkAbEWEmP4YAjdseuriEXGu1dwOp4= +github.com/riverqueue/river/riverdriver/riverpgxv5 v0.41.0/go.mod h1:0ueUi+3eW5fIsibS5VkD2HzQmwLKTAiM/uiiVMX/SME= +github.com/riverqueue/river/rivershared v0.41.0 h1:ax3uz5KiqfmcvRQ3kKB6iYs9Nt8J5L1RNfOKoVMWkqw= +github.com/riverqueue/river/rivershared v0.41.0/go.mod h1:FaZ7bxC2DORhyFDVyHKRiZzgQPf2b/IELwPBXnHZVLA= +github.com/riverqueue/river/rivertype v0.41.0 h1:dfscvt1asf1PpeeHTMFxQdV1ZfMoReiiMNFCPPfR0as= +github.com/riverqueue/river/rivertype v0.41.0/go.mod h1:D1Ad+EaZiaXbQbJcJcfeicXJMBKno0n6UcfKI5Q7DIQ= github.com/robfig/cron/v3 v3.0.1 h1:WdRxkvbJztn8LMz/QEvLN5sBU+xKpSqwwUO1Pjr4qDs= github.com/robfig/cron/v3 v3.0.1/go.mod h1:eQICP3HwyT7UooqI/z+Ov+PtYAWygg1TEWWzGIFLtro= github.com/rogpeppe/go-internal v1.12.0 h1:exVL4IDcn6na9z1rAb56Vxr+CgyK3nn3O+epU5NdKM8= diff --git a/go.work b/go.work index 9a9927fe..6bbee011 100644 --- a/go.work +++ b/go.work @@ -9,6 +9,7 @@ use ( ./riverdriver/riverdatabasesql ./riverdriver/riverdrivertest ./riverdriver/riverpgxv5 + ./riverdriver/rivermysql ./riverdriver/riversqlite ./rivershared ./rivertype diff --git a/internal/maintenance/sqlite_notification_cleaner.go b/internal/maintenance/sqlite_notification_cleaner.go index de9e9526..328a8e0d 100644 --- a/internal/maintenance/sqlite_notification_cleaner.go +++ b/internal/maintenance/sqlite_notification_cleaner.go @@ -60,9 +60,8 @@ func (c *SQLiteNotificationCleanerConfig) mustValidate() *SQLiteNotificationClea return c } -// SQLiteNotificationCleaner periodically removes old rows from SQLite's -// notification outbox. It is only needed for the SQLite driver's emulated -// listen/notify support. +// SQLiteNotificationCleaner periodically removes old rows from the durable +// notification outbox used to emulate listen/notify for SQLite and MySQL. type SQLiteNotificationCleaner struct { riversharedmaintenance.QueueMaintainerServiceBase startstop.BaseStartStop @@ -74,7 +73,7 @@ type SQLiteNotificationCleaner struct { exec riverdriver.Executor } -// NewSQLiteNotificationCleaner returns a SQLite notification cleaner. +// NewSQLiteNotificationCleaner returns a notification outbox cleaner. func NewSQLiteNotificationCleaner(archetype *baseservice.Archetype, config *SQLiteNotificationCleanerConfig, exec riverdriver.Executor) *SQLiteNotificationCleaner { return baseservice.Init(archetype, &SQLiteNotificationCleaner{ Config: (&SQLiteNotificationCleanerConfig{ @@ -113,7 +112,7 @@ func (s *SQLiteNotificationCleaner) Start(ctx context.Context) error { //nolint: res, err := s.runOnce(ctx) if err != nil { if !errors.Is(err, context.Canceled) { - s.Logger.ErrorContext(ctx, s.Name+": Error cleaning SQLite notifications", slog.String("error", err.Error())) + s.Logger.ErrorContext(ctx, s.Name+": Error cleaning notification outbox", slog.String("error", err.Error())) } continue } diff --git a/internal/rivercommon/river_common.go b/internal/rivercommon/river_common.go index 4f769b39..81d322b5 100644 --- a/internal/rivercommon/river_common.go +++ b/internal/rivercommon/river_common.go @@ -46,9 +46,9 @@ 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 is a special metadata key used by the MySQL and + // SQLite drivers to determine whether an upsert was skipped because the + // `(xmax != 0)` trick used in Postgres isn't available. MetadataKeyUniqueNonce = "river:unique_nonce" ) diff --git a/job_list_params.go b/job_list_params.go index c9f18016..595d4571 100644 --- a/job_list_params.go +++ b/job_list_params.go @@ -180,6 +180,8 @@ type JobListParams struct { where []dblist.WherePredicate } +const jobListMetadataPredicatePostgres = `metadata @> @metadata_fragment::jsonb` + // NewJobListParams creates a new JobListParams to return available jobs sorted // by time in ascending order, returning 100 jobs at most. func NewJobListParams() *JobListParams { @@ -278,9 +280,9 @@ func (p *JobListParams) toDBParams() (*dblist.JobListParams, error) { } else { namedArgs["cursor_time"] = p.after.time if sortOrder == dblist.SortOrderAsc { - p.where = append(p.where, dblist.WherePredicate{NamedArgs: namedArgs, SQL: fmt.Sprintf(`("%s" > @cursor_time OR ("%s" = @cursor_time AND "id" > @after_id))`, timeField, timeField)}) + p.where = append(p.where, dblist.WherePredicate{NamedArgs: namedArgs, SQL: fmt.Sprintf(`(%s > @cursor_time OR (%s = @cursor_time AND id > @after_id))`, timeField, timeField)}) } else { - p.where = append(p.where, dblist.WherePredicate{NamedArgs: namedArgs, SQL: fmt.Sprintf(`("%s" < @cursor_time OR ("%s" = @cursor_time AND "id" < @after_id))`, timeField, timeField)}) + p.where = append(p.where, dblist.WherePredicate{NamedArgs: namedArgs, SQL: fmt.Sprintf(`(%s < @cursor_time OR (%s = @cursor_time AND id < @after_id))`, timeField, timeField)}) } } } @@ -298,6 +300,16 @@ func (p *JobListParams) toDBParams() (*dblist.JobListParams, error) { }, nil } +func (p *JobListParams) withMetadataPredicateSQL(predicateSQL string) *JobListParams { + paramsCopy := p.copy() + for i := range paramsCopy.where { + if paramsCopy.where[i].SQL == jobListMetadataPredicatePostgres { + paramsCopy.where[i].SQL = predicateSQL + } + } + return paramsCopy +} + // After returns an updated filter set that will only return jobs // after the given cursor. func (p *JobListParams) After(cursor *JobListCursor) *JobListParams { @@ -359,7 +371,7 @@ func (p *JobListParams) Metadata(json string) *JobListParams { paramsCopy.metadataCalled = true paramsCopy.where = append(paramsCopy.where, dblist.WherePredicate{ NamedArgs: map[string]any{"metadata_fragment": json}, - SQL: `metadata @> @metadata_fragment::jsonb`, + SQL: jobListMetadataPredicatePostgres, }) return paramsCopy } diff --git a/riverdbtest/riverdbtest.go b/riverdbtest/riverdbtest.go index 64468ca9..652e4d54 100644 --- a/riverdbtest/riverdbtest.go +++ b/riverdbtest/riverdbtest.go @@ -481,7 +481,8 @@ type TestTxOpts struct { // run using TestSchema. This is meant for environments where parallelism // doesn't work as well, like SQLite, which will emit "busy" errors when // multiple clients try to share a schema, even when they're in separate - // transactions. + // transactions. Also applies to MySQL, where InnoDB deadlocks are common + // when multiple transactions are sharing a database. DisableSchemaSharing bool // IsTestTxHelper should be set to true for if TestTx is being called from diff --git a/riverdriver/go.sum b/riverdriver/go.sum index 3a02b5b2..ab809626 100644 --- a/riverdriver/go.sum +++ b/riverdriver/go.sum @@ -7,8 +7,8 @@ github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= -github.com/riverqueue/river/rivertype v0.40.0 h1:6sFLhs0OtkxDZ/iyZIUqaQJb/G1krnb2bh9lfXdCimg= -github.com/riverqueue/river/rivertype v0.40.0/go.mod h1:D1Ad+EaZiaXbQbJcJcfeicXJMBKno0n6UcfKI5Q7DIQ= +github.com/riverqueue/river/rivertype v0.41.0 h1:dfscvt1asf1PpeeHTMFxQdV1ZfMoReiiMNFCPPfR0as= +github.com/riverqueue/river/rivertype v0.41.0/go.mod h1:D1Ad+EaZiaXbQbJcJcfeicXJMBKno0n6UcfKI5Q7DIQ= github.com/rogpeppe/go-internal v1.12.0 h1:exVL4IDcn6na9z1rAb56Vxr+CgyK3nn3O+epU5NdKM8= github.com/rogpeppe/go-internal v1.12.0/go.mod h1:E+RYuTGaKKdloAfM02xzb0FW3Paa99yedzYV+kq4uf4= github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= diff --git a/riverdriver/river_driver_interface.go b/riverdriver/river_driver_interface.go index 64412aad..8ce7dd11 100644 --- a/riverdriver/river_driver_interface.go +++ b/riverdriver/river_driver_interface.go @@ -25,6 +25,7 @@ import ( const AllQueuesString = "*" const ( + DatabaseNameMySQL = "mysql" DatabaseNamePostgres = "postgres" DatabaseNameSQLite = "sqlite" ) @@ -53,7 +54,7 @@ var ( type Driver[TTx any] interface { // ArgPlaceholder is the placeholder character used in query positional // arguments, so "$" for "$1", "$2", "$3", etc. This is a "$" for Postgres - // and "?" for SQLite. + // and "?" for MySQL and SQLite. // // API is not stable. DO NOT USE. ArgPlaceholder() string @@ -128,6 +129,13 @@ type Driver[TTx any] interface { // API is not stable. DO NOT USE. PoolSet(dbPool any) error + // SafeIdentifier returns a safely quoted identifier (e.g. a table or + // schema name) for use in SQL queries. Each driver quotes using its + // native syntax: double quotes for Postgres/SQLite, backticks for MySQL. + // + // API is not stable. DO NOT USE. + SafeIdentifier(ident string) string + // SQLFragmentColumnIn generates an SQL fragment to be included as a // predicate in a `WHERE` query for the existence of a set of values in a // column like `id IN (...)`. The actual implementation depends on support diff --git a/riverdriver/riverdatabasesql/go.sum b/riverdriver/riverdatabasesql/go.sum index ace09258..801ce52d 100644 --- a/riverdriver/riverdatabasesql/go.sum +++ b/riverdriver/riverdatabasesql/go.sum @@ -19,16 +19,16 @@ github.com/lib/pq v1.12.3 h1:tTWxr2YLKwIvK90ZXEw8GP7UFHtcbTtty8zsI+YjrfQ= github.com/lib/pq v1.12.3/go.mod h1:/p+8NSbOcwzAEI7wiMXFlgydTwcgTr3OSKMsD2BitpA= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= -github.com/riverqueue/river v0.40.0 h1:4dynKqqU1P22iPmwWDfDj/YZXnuUTZysXXF3wNHekNw= -github.com/riverqueue/river v0.40.0/go.mod h1:auvB4kHqM97tnshEQxzy2E7aFvJFhl010NB6N29DXXc= -github.com/riverqueue/river/riverdriver v0.40.0 h1:QjBHHh+kaxgUgK9tPumrfx7W14vhAHLIt+0dK40ikI8= -github.com/riverqueue/river/riverdriver v0.40.0/go.mod h1:7wEqxsqvtjGk3hBKGKK3IvdnULlPZltpAfkbMr9M2VY= -github.com/riverqueue/river/riverdriver/riverpgxv5 v0.40.0 h1:WbzXgGukvOqYpsORX4rr03sPZfasjoMckLr4On+Je+4= -github.com/riverqueue/river/riverdriver/riverpgxv5 v0.40.0/go.mod h1:HtgFIcn/eJ9O3f371S1cSb0ttS7xJMMkkTD/K6gQKQg= -github.com/riverqueue/river/rivershared v0.40.0 h1:6ZX1Ok94Nkcx/WpHPhexpfdkRB1BqRWFEWrbDszg+JY= -github.com/riverqueue/river/rivershared v0.40.0/go.mod h1:Z77wB2/ctD+zSvlNDKjxC1SW9c1Y2jyE+PXR2CxBlzQ= -github.com/riverqueue/river/rivertype v0.40.0 h1:6sFLhs0OtkxDZ/iyZIUqaQJb/G1krnb2bh9lfXdCimg= -github.com/riverqueue/river/rivertype v0.40.0/go.mod h1:D1Ad+EaZiaXbQbJcJcfeicXJMBKno0n6UcfKI5Q7DIQ= +github.com/riverqueue/river v0.41.0 h1:E7Yfyhn74IgVaCBKsBpTHoC0LxWJvoG0jKt7Tb+lJZg= +github.com/riverqueue/river v0.41.0/go.mod h1:WXiAF1/2gfUPj3H+WqQG3Q5NW5w9JmRXRTe7U9zybNc= +github.com/riverqueue/river/riverdriver v0.41.0 h1:A5g80n6fCGu9Bp4hSq7gys7mvfOiuWkbQHuFx3iPp7w= +github.com/riverqueue/river/riverdriver v0.41.0/go.mod h1:OeGOhZYoH7Ac0CrVanNAAwD9C0WAXy4LzUEl65Ex0B4= +github.com/riverqueue/river/riverdriver/riverpgxv5 v0.41.0 h1:uHGluToMMWvCnpNkAbEWEmP4YAjdseuriEXGu1dwOp4= +github.com/riverqueue/river/riverdriver/riverpgxv5 v0.41.0/go.mod h1:0ueUi+3eW5fIsibS5VkD2HzQmwLKTAiM/uiiVMX/SME= +github.com/riverqueue/river/rivershared v0.41.0 h1:ax3uz5KiqfmcvRQ3kKB6iYs9Nt8J5L1RNfOKoVMWkqw= +github.com/riverqueue/river/rivershared v0.41.0/go.mod h1:FaZ7bxC2DORhyFDVyHKRiZzgQPf2b/IELwPBXnHZVLA= +github.com/riverqueue/river/rivertype v0.41.0 h1:dfscvt1asf1PpeeHTMFxQdV1ZfMoReiiMNFCPPfR0as= +github.com/riverqueue/river/rivertype v0.41.0/go.mod h1:D1Ad+EaZiaXbQbJcJcfeicXJMBKno0n6UcfKI5Q7DIQ= github.com/robfig/cron/v3 v3.0.1 h1:WdRxkvbJztn8LMz/QEvLN5sBU+xKpSqwwUO1Pjr4qDs= github.com/robfig/cron/v3 v3.0.1/go.mod h1:eQICP3HwyT7UooqI/z+Ov+PtYAWygg1TEWWzGIFLtro= github.com/rogpeppe/go-internal v1.12.0 h1:exVL4IDcn6na9z1rAb56Vxr+CgyK3nn3O+epU5NdKM8= diff --git a/riverdriver/riverdatabasesql/river_database_sql_driver.go b/riverdriver/riverdatabasesql/river_database_sql_driver.go index 8320b17a..ad6906bc 100644 --- a/riverdriver/riverdatabasesql/river_database_sql_driver.go +++ b/riverdriver/riverdatabasesql/river_database_sql_driver.go @@ -53,8 +53,9 @@ func New(dbPool *sql.DB) *Driver { const argPlaceholder = "$" -func (d *Driver) ArgPlaceholder() string { return argPlaceholder } -func (d *Driver) DatabaseName() string { return riverdriver.DatabaseNamePostgres } +func (d *Driver) ArgPlaceholder() string { return argPlaceholder } +func (d *Driver) DatabaseName() string { return riverdriver.DatabaseNamePostgres } +func (d *Driver) SafeIdentifier(ident string) string { return dbutil.SafeIdentifier(ident) } func (d *Driver) GetExecutor() riverdriver.Executor { return &Executor{d.dbPool, templateReplaceWrapper{d.dbPool, &d.replacer}, d} diff --git a/riverdriver/riverdrivertest/driver_client_test.go b/riverdriver/riverdrivertest/driver_client_test.go index 923c7a2b..77963f4f 100644 --- a/riverdriver/riverdrivertest/driver_client_test.go +++ b/riverdriver/riverdrivertest/driver_client_test.go @@ -7,6 +7,7 @@ import ( "testing" "time" + _ "github.com/go-sql-driver/mysql" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/stdlib" "github.com/lib/pq" @@ -19,10 +20,12 @@ import ( "github.com/riverqueue/river/riverdbtest" "github.com/riverqueue/river/riverdriver" "github.com/riverqueue/river/riverdriver/riverdatabasesql" + "github.com/riverqueue/river/riverdriver/rivermysql" "github.com/riverqueue/river/riverdriver/riverpgxv5" "github.com/riverqueue/river/riverdriver/riversqlite" "github.com/riverqueue/river/rivershared/riversharedtest" "github.com/riverqueue/river/rivershared/testfactory" + "github.com/riverqueue/river/rivershared/util/ptrutil" "github.com/riverqueue/river/rivershared/util/testutil" "github.com/riverqueue/river/rivershared/util/urlutil" "github.com/riverqueue/river/rivertype" @@ -110,6 +113,26 @@ func TestClientWithDriverRiverLibSQL(t *testing.T) { ) } +func TestClientWithDriverRiverMySQL(t *testing.T) { + t.Parallel() + + riversharedtest.SkipIfMySQLNotEnabled(t) + + var ( + ctx = context.Background() + dbPool = riversharedtest.DBPoolMySQL(ctx, t) + driver = rivermysql.New(dbPool) + ) + + ExerciseClient(ctx, t, + func(ctx context.Context, t *testing.T) (riverdriver.Driver[*sql.Tx], string) { + t.Helper() + + return driver, riverdbtest.TestSchema(ctx, t, driver, nil) + }, + ) +} + func TestClientWithDriverRiverSQLiteModernC(t *testing.T) { t.Parallel() @@ -552,6 +575,28 @@ func ExerciseClient[TTx any](ctx context.Context, t *testing.T, require.Equal(t, job.ID, listRes.Jobs[0].ID) }) + t.Run("JobListPaginatesByScheduledAt", func(t *testing.T) { + t.Parallel() + + client, bundle := setup(t) + if bundle.driver.DatabaseName() == riverdriver.DatabaseNameSQLite { + t.Skip("SQLite cursor pagination is covered by the main client suite") + } + + now := time.Now().UTC() + job1 := testfactory.Job(ctx, t, bundle.exec, &testfactory.JobOpts{ScheduledAt: &now, Schema: bundle.schema}) + job2 := testfactory.Job(ctx, t, bundle.exec, &testfactory.JobOpts{ScheduledAt: ptrutil.Ptr(now.Add(time.Second)), Schema: bundle.schema}) + + listRes, err := client.JobList(ctx, + river.NewJobListParams(). + OrderBy(river.JobListOrderByScheduledAt, river.SortOrderAsc). + After(river.JobListCursorFromJob(job1)), + ) + require.NoError(t, err) + require.Len(t, listRes.Jobs, 1) + require.Equal(t, job2.ID, listRes.Jobs[0].ID) + }) + t.Run("JobListTx", func(t *testing.T) { t.Parallel() @@ -630,9 +675,12 @@ func ExerciseClient[TTx any](ctx context.Context, t *testing.T, listParams := river.NewJobListParams() - if bundle.driver.DatabaseName() == riverdriver.DatabaseNameSQLite { + switch bundle.driver.DatabaseName() { + case riverdriver.DatabaseNameSQLite: listParams = listParams.Where("metadata ->> @json_path = @json_val", river.NamedArgs{"json_path": "$.foo", "json_val": "bar"}) - } else { + case riverdriver.DatabaseNameMySQL: + listParams = listParams.Where("JSON_UNQUOTE(JSON_EXTRACT(metadata, @json_path)) = @json_val", river.NamedArgs{"json_path": "$.foo", "json_val": "bar"}) + default: // "bar" is quoted in this branch because `jsonb_path_query_first` needs to be compared to a JSON value listParams = listParams.Where("jsonb_path_query_first(metadata, @json_path) = @json_val", river.NamedArgs{"json_path": "$.foo", "json_val": `"bar"`}) } diff --git a/riverdriver/riverdrivertest/driver_test.go b/riverdriver/riverdrivertest/driver_test.go index bbf9db89..ef5b490f 100644 --- a/riverdriver/riverdrivertest/driver_test.go +++ b/riverdriver/riverdrivertest/driver_test.go @@ -8,6 +8,7 @@ import ( "testing" "time" + _ "github.com/go-sql-driver/mysql" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" "github.com/jackc/pgx/v5/stdlib" @@ -22,6 +23,7 @@ import ( "github.com/riverqueue/river/riverdriver" "github.com/riverqueue/river/riverdriver/riverdatabasesql" "github.com/riverqueue/river/riverdriver/riverdrivertest" + "github.com/riverqueue/river/riverdriver/rivermysql" "github.com/riverqueue/river/riverdriver/riverpgxv5" "github.com/riverqueue/river/riverdriver/riversqlite" "github.com/riverqueue/river/rivershared/riversharedtest" @@ -204,6 +206,43 @@ func TestDriverRiverLiteLibSQL(t *testing.T) { //nolint:dupl }) } +func TestDriverRiverMySQL(t *testing.T) { + t.Parallel() + + riversharedtest.SkipIfMySQLNotEnabled(t) + + var ( + ctx = context.Background() + dbPool = riversharedtest.DBPoolMySQL(ctx, t) + driver = rivermysql.New(dbPool) + ) + + riverdrivertest.Exercise(ctx, t, + func(ctx context.Context, t *testing.T, opts *riverdbtest.TestSchemaOpts) (riverdriver.Driver[*sql.Tx], string) { + t.Helper() + + return driver, riverdbtest.TestSchema(ctx, t, driver, opts) + }, + func(ctx context.Context, t *testing.T) (riverdriver.Executor, riverdriver.Driver[*sql.Tx]) { + t.Helper() + + tx, schema := riverdbtest.TestTx(ctx, t, driver, &riverdbtest.TestTxOpts{ + // Disable schema sharing to reduce InnoDB deadlocks from + // parallel tests contending on the same database. + DisableSchemaSharing: true, + }) + + // MySQL has no search_path equivalent, so USE the test schema + // database so that unqualified queries resolve correctly. + if schema != "" { + _, err := tx.ExecContext(ctx, "USE "+schema) + require.NoError(t, err) + } + + return driver.UnwrapExecutor(tx), driver + }) +} + func TestDriverRiverSQLiteModernC(t *testing.T) { //nolint:dupl t.Parallel() diff --git a/riverdriver/riverdrivertest/executor_tx.go b/riverdriver/riverdrivertest/executor_tx.go index e4313932..e540423c 100644 --- a/riverdriver/riverdrivertest/executor_tx.go +++ b/riverdriver/riverdrivertest/executor_tx.go @@ -157,7 +157,13 @@ func exerciseExecutorTx[TTx any](ctx context.Context, t *testing.T, exec := setup(ctx, t) - require.NoError(t, exec.Exec(ctx, "SELECT $1 || $2", "foo", "bar")) + _, driver := executorWithTx(ctx, t) + switch driver.DatabaseName() { + case riverdriver.DatabaseNameMySQL: + require.NoError(t, exec.Exec(ctx, "SELECT CONCAT(?, ?)", "foo", "bar")) + default: + require.NoError(t, exec.Exec(ctx, "SELECT $1 || $2", "foo", "bar")) + } }) }) @@ -166,9 +172,8 @@ func exerciseExecutorTx[TTx any](ctx context.Context, t *testing.T, { driver, _ := driverWithSchema(ctx, t, nil) - if driver.DatabaseName() == riverdriver.DatabaseNameSQLite { - t.Logf("Skipping PGAdvisoryXactLock test for SQLite") - return + if driver.DatabaseName() == riverdriver.DatabaseNameSQLite || driver.DatabaseName() == riverdriver.DatabaseNameMySQL { + t.Skipf("Skipping PGAdvisoryXactLock test for %s", driver.DatabaseName()) } } diff --git a/riverdriver/riverdrivertest/go.mod b/riverdriver/riverdrivertest/go.mod index 34fa1abf..8ccdb3fb 100644 --- a/riverdriver/riverdrivertest/go.mod +++ b/riverdriver/riverdrivertest/go.mod @@ -6,12 +6,14 @@ toolchain go1.25.7 require ( github.com/davecgh/go-spew v1.1.1 + github.com/go-sql-driver/mysql v1.10.0 github.com/jackc/pgerrcode v0.0.0-20240316143900-6e2875d9b438 github.com/jackc/pgx/v5 v5.10.0 github.com/lib/pq v1.12.3 github.com/riverqueue/river v0.41.0 github.com/riverqueue/river/riverdriver v0.41.0 github.com/riverqueue/river/riverdriver/riverdatabasesql v0.41.0 + github.com/riverqueue/river/riverdriver/rivermysql v0.0.0-00010101000000-000000000000 github.com/riverqueue/river/riverdriver/riverpgxv5 v0.41.0 github.com/riverqueue/river/riverdriver/riversqlite v0.41.0 github.com/riverqueue/river/rivershared v0.41.0 @@ -25,7 +27,12 @@ require ( turso.tech/database/tursogo v0.7.0 ) +// Remove this replacement after rivermysql has had its first coordinated +// release. It lets go mod tidy resolve the new sibling module in the meantime. +replace github.com/riverqueue/river/riverdriver/rivermysql => ../rivermysql + require ( + filippo.io/edwards25519 v1.2.0 // indirect github.com/antlr4-go/antlr/v4 v4.13.0 // indirect github.com/coder/websocket v1.8.12 // indirect github.com/dustin/go-humanize v1.0.1 // indirect diff --git a/riverdriver/riverdrivertest/go.sum b/riverdriver/riverdrivertest/go.sum index e946393b..39bd26e2 100644 --- a/riverdriver/riverdrivertest/go.sum +++ b/riverdriver/riverdrivertest/go.sum @@ -1,3 +1,5 @@ +filippo.io/edwards25519 v1.2.0 h1:crnVqOiS4jqYleHd9vaKZ+HKtHfllngJIiOpNpoJsjo= +filippo.io/edwards25519 v1.2.0/go.mod h1:xzAOLCNug/yB62zG1bQ8uziwrIqIuxhctzJT18Q77mc= github.com/antlr4-go/antlr/v4 v4.13.0 h1:lxCg3LAv+EUK6t1i0y1V6/SLeUi0eKEKdhQAlS8TVTI= github.com/antlr4-go/antlr/v4 v4.13.0/go.mod h1:pfChB/xh/Unjila75QW7+VU4TSnWnnk9UTnmpPaOR2g= github.com/coder/websocket v1.8.12 h1:5bUXkEPPIbewrnkU8LTCLVaxi4N4J8ahufH2vlo4NAo= @@ -9,6 +11,8 @@ github.com/dustin/go-humanize v1.0.1 h1:GzkhY7T5VNhEkwH0PVJgjz+fX1rhBrR7pRT3mDkp github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto= github.com/ebitengine/purego v0.9.1 h1:a/k2f2HQU3Pi399RPW1MOaZyhKJL9w/xFpKAg4q1s0A= github.com/ebitengine/purego v0.9.1/go.mod h1:iIjxzd6CiRiOG0UyXP+V1+jWqUXVjPKLAI0mRfJZTmQ= +github.com/go-sql-driver/mysql v1.10.0 h1:Q+1LV8DkHJvSYAdR83XzuhDaTykuDx0l6fkXxoWCWfw= +github.com/go-sql-driver/mysql v1.10.0/go.mod h1:M+cqaI7+xxXGG9swrdeUIoPG3Y3KCkF0pZej+SK+nWk= github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs= github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= @@ -41,20 +45,20 @@ github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZb github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94icq4NjY3clb7Lk8O1qJ8BdBEF8z0ibU0rE= github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= -github.com/riverqueue/river v0.40.0 h1:4dynKqqU1P22iPmwWDfDj/YZXnuUTZysXXF3wNHekNw= -github.com/riverqueue/river v0.40.0/go.mod h1:auvB4kHqM97tnshEQxzy2E7aFvJFhl010NB6N29DXXc= -github.com/riverqueue/river/riverdriver v0.40.0 h1:QjBHHh+kaxgUgK9tPumrfx7W14vhAHLIt+0dK40ikI8= -github.com/riverqueue/river/riverdriver v0.40.0/go.mod h1:7wEqxsqvtjGk3hBKGKK3IvdnULlPZltpAfkbMr9M2VY= -github.com/riverqueue/river/riverdriver/riverdatabasesql v0.40.0 h1:EtKlKjaKzkNckyGrFnrpK5mcc7+xnp7Z3IV9XVRf+7E= -github.com/riverqueue/river/riverdriver/riverdatabasesql v0.40.0/go.mod h1:Si7twoX1cou3msM9G6rYCKelaSHLqTgJ65H/m3Zi8rk= -github.com/riverqueue/river/riverdriver/riverpgxv5 v0.40.0 h1:WbzXgGukvOqYpsORX4rr03sPZfasjoMckLr4On+Je+4= -github.com/riverqueue/river/riverdriver/riverpgxv5 v0.40.0/go.mod h1:HtgFIcn/eJ9O3f371S1cSb0ttS7xJMMkkTD/K6gQKQg= -github.com/riverqueue/river/riverdriver/riversqlite v0.40.0 h1:DE/22nz7BKgQ+0XpPFg/3hoO4mP8Ij+gGcIOS9PusFg= -github.com/riverqueue/river/riverdriver/riversqlite v0.40.0/go.mod h1:/Xuw1OOo+X4Pj5DCKu4MLbF5BhGIBl/OWp93kN2v5ec= -github.com/riverqueue/river/rivershared v0.40.0 h1:6ZX1Ok94Nkcx/WpHPhexpfdkRB1BqRWFEWrbDszg+JY= -github.com/riverqueue/river/rivershared v0.40.0/go.mod h1:Z77wB2/ctD+zSvlNDKjxC1SW9c1Y2jyE+PXR2CxBlzQ= -github.com/riverqueue/river/rivertype v0.40.0 h1:6sFLhs0OtkxDZ/iyZIUqaQJb/G1krnb2bh9lfXdCimg= -github.com/riverqueue/river/rivertype v0.40.0/go.mod h1:D1Ad+EaZiaXbQbJcJcfeicXJMBKno0n6UcfKI5Q7DIQ= +github.com/riverqueue/river v0.41.0 h1:E7Yfyhn74IgVaCBKsBpTHoC0LxWJvoG0jKt7Tb+lJZg= +github.com/riverqueue/river v0.41.0/go.mod h1:WXiAF1/2gfUPj3H+WqQG3Q5NW5w9JmRXRTe7U9zybNc= +github.com/riverqueue/river/riverdriver v0.41.0 h1:A5g80n6fCGu9Bp4hSq7gys7mvfOiuWkbQHuFx3iPp7w= +github.com/riverqueue/river/riverdriver v0.41.0/go.mod h1:OeGOhZYoH7Ac0CrVanNAAwD9C0WAXy4LzUEl65Ex0B4= +github.com/riverqueue/river/riverdriver/riverdatabasesql v0.41.0 h1:HEcEzKx8IMVnoGT5hnLEBfq+qiJDLE303C07qyGVHFY= +github.com/riverqueue/river/riverdriver/riverdatabasesql v0.41.0/go.mod h1:YTvDkpQzN4cVM3BIAIo8dK3ifkzN65BrbaluC/LBu0w= +github.com/riverqueue/river/riverdriver/riverpgxv5 v0.41.0 h1:uHGluToMMWvCnpNkAbEWEmP4YAjdseuriEXGu1dwOp4= +github.com/riverqueue/river/riverdriver/riverpgxv5 v0.41.0/go.mod h1:0ueUi+3eW5fIsibS5VkD2HzQmwLKTAiM/uiiVMX/SME= +github.com/riverqueue/river/riverdriver/riversqlite v0.41.0 h1:4IBOD4Uzil620AT/WT5xPHVoQSk2OiQ5UvG7OQ2piOc= +github.com/riverqueue/river/riverdriver/riversqlite v0.41.0/go.mod h1:k16A8Tt6Ke49rCU4k33T5EKXAhAXiExyU2qV+6hpkNs= +github.com/riverqueue/river/rivershared v0.41.0 h1:ax3uz5KiqfmcvRQ3kKB6iYs9Nt8J5L1RNfOKoVMWkqw= +github.com/riverqueue/river/rivershared v0.41.0/go.mod h1:FaZ7bxC2DORhyFDVyHKRiZzgQPf2b/IELwPBXnHZVLA= +github.com/riverqueue/river/rivertype v0.41.0 h1:dfscvt1asf1PpeeHTMFxQdV1ZfMoReiiMNFCPPfR0as= +github.com/riverqueue/river/rivertype v0.41.0/go.mod h1:D1Ad+EaZiaXbQbJcJcfeicXJMBKno0n6UcfKI5Q7DIQ= github.com/robfig/cron/v3 v3.0.1 h1:WdRxkvbJztn8LMz/QEvLN5sBU+xKpSqwwUO1Pjr4qDs= github.com/robfig/cron/v3 v3.0.1/go.mod h1:eQICP3HwyT7UooqI/z+Ov+PtYAWygg1TEWWzGIFLtro= github.com/rogpeppe/go-internal v1.12.0 h1:exVL4IDcn6na9z1rAb56Vxr+CgyK3nn3O+epU5NdKM8= @@ -83,8 +87,8 @@ go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto= go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE= golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 h1:mgKeJMpvi0yx/sU5GsxQ7p6s2wtOnGAHZWCHUM4KGzY= golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546/go.mod h1:j/pmGrbnkbPtQfxEe5D0VQhZC6qKbfKifgD0oM7sR70= -golang.org/x/mod v0.37.0 h1:vF1DjpVEshcIqoEaauuHebaLk1O1forxjxBaVn884JQ= -golang.org/x/mod v0.37.0/go.mod h1:m8S8VeM9r4dzDwjrKO0a1sZP3YjeMamRRlD+fmR2Q/0= +golang.org/x/mod v0.38.0 h1:MECBjubtXD7yj4HrhIUcywNaGeNVUdfVnxmPajOk4yk= +golang.org/x/mod v0.38.0/go.mod h1:V6Xz0pq8TQ3dGqVQ1FVHuelZpAL0uNhSkk9ogYP3c40= golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek= golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= diff --git a/riverdriver/riverdrivertest/job_delete.go b/riverdriver/riverdrivertest/job_delete.go index 2b1919dc..9f9368a8 100644 --- a/riverdriver/riverdrivertest/job_delete.go +++ b/riverdriver/riverdrivertest/job_delete.go @@ -138,6 +138,34 @@ func exerciseJobDelete[TTx any](ctx context.Context, t *testing.T, executorWithT afterHorizon = horizon.Add(1 * time.Minute) ) + t.Run("RespectsDisabledStates", func(t *testing.T) { + t.Parallel() + + exec, bundle := setup(ctx, t) + if bundle.driver.DatabaseName() == riverdriver.DatabaseNameSQLite { + t.Skip("SQLite does not currently support disabling individual retention periods") + } + + cancelledJob := testfactory.Job(ctx, t, exec, &testfactory.JobOpts{FinalizedAt: &beforeHorizon, State: ptrutil.Ptr(rivertype.JobStateCancelled)}) + completedJob := testfactory.Job(ctx, t, exec, &testfactory.JobOpts{FinalizedAt: &beforeHorizon, State: ptrutil.Ptr(rivertype.JobStateCompleted)}) + discardedJob := testfactory.Job(ctx, t, exec, &testfactory.JobOpts{FinalizedAt: &beforeHorizon, State: ptrutil.Ptr(rivertype.JobStateDiscarded)}) + + numDeleted, err := exec.JobDeleteBefore(ctx, &riverdriver.JobDeleteBeforeParams{ + CompletedDoDelete: true, + CompletedFinalizedAtHorizon: horizon, + Max: 1_000, + }) + require.NoError(t, err) + require.Equal(t, 1, numDeleted) + + _, err = exec.JobGetByID(ctx, &riverdriver.JobGetByIDParams{ID: cancelledJob.ID}) + require.NoError(t, err) + _, err = exec.JobGetByID(ctx, &riverdriver.JobGetByIDParams{ID: completedJob.ID}) + require.ErrorIs(t, err, rivertype.ErrNotFound) + _, err = exec.JobGetByID(ctx, &riverdriver.JobGetByIDParams{ID: discardedJob.ID}) + require.NoError(t, err) + }) + t.Run("Success", func(t *testing.T) { t.Parallel() @@ -200,21 +228,30 @@ func exerciseJobDelete[TTx any](ctx context.Context, t *testing.T, executorWithT t.Run("QueuesExcluded", func(t *testing.T) { t.Parallel() - exec, _ := setup(ctx, t) - - var ( //nolint:dupl - cancelledJob = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{FinalizedAt: &beforeHorizon, State: ptrutil.Ptr(rivertype.JobStateCancelled)}) - completedJob = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{FinalizedAt: &beforeHorizon, State: ptrutil.Ptr(rivertype.JobStateCompleted)}) - discardedJob = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{FinalizedAt: &beforeHorizon, State: ptrutil.Ptr(rivertype.JobStateDiscarded)}) + exec, bundle := setup(ctx, t) + var ( excludedQueue1 = "excluded1" excludedQueue2 = "excluded2" - // Not deleted because in an omitted queue. + // Insert excluded jobs first to verify that they don't consume the + // limited candidate batch and starve later eligible jobs. notDeletedJob1 = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{FinalizedAt: &beforeHorizon, Queue: &excludedQueue1, State: ptrutil.Ptr(rivertype.JobStateCompleted)}) notDeletedJob2 = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{FinalizedAt: &beforeHorizon, Queue: &excludedQueue2, State: ptrutil.Ptr(rivertype.JobStateCompleted)}) + + cancelledJob = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{FinalizedAt: &beforeHorizon, State: ptrutil.Ptr(rivertype.JobStateCancelled)}) + completedJob = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{FinalizedAt: &beforeHorizon, State: ptrutil.Ptr(rivertype.JobStateCompleted)}) + discardedJob = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{FinalizedAt: &beforeHorizon, State: ptrutil.Ptr(rivertype.JobStateDiscarded)}) ) + maxJobs := 3 + if bundle.driver.DatabaseName() == riverdriver.DatabaseNameSQLite { + // SQLite currently applies queue exclusions after selecting the + // limited candidate batch, so retain coverage of its existing weaker + // behavior while checking starvation resistance on other drivers. + maxJobs = 1_000 + } + numDeleted, err := exec.JobDeleteBefore(ctx, &riverdriver.JobDeleteBeforeParams{ CancelledDoDelete: true, CancelledFinalizedAtHorizon: horizon, @@ -222,7 +259,7 @@ func exerciseJobDelete[TTx any](ctx context.Context, t *testing.T, executorWithT CompletedFinalizedAtHorizon: horizon, DiscardedDoDelete: true, DiscardedFinalizedAtHorizon: horizon, - Max: 1_000, + Max: maxJobs, QueuesExcluded: []string{excludedQueue1, excludedQueue2}, }) require.NoError(t, err) @@ -262,11 +299,11 @@ func exerciseJobDelete[TTx any](ctx context.Context, t *testing.T, executorWithT // `queues_included` for the foreseeable future), I've just set // SQLite to not support `queues_included` for the time being. if bundle.driver.DatabaseName() == riverdriver.DatabaseNameSQLite { - t.Logf("Skipping JobDeleteBefore with QueuesIncluded test for SQLite") + t.Skipf("Skipping JobDeleteBefore with QueuesIncluded test for %s", bundle.driver.DatabaseName()) return } - var ( //nolint:dupl + var ( cancelledJob = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{FinalizedAt: &beforeHorizon, State: ptrutil.Ptr(rivertype.JobStateCancelled)}) completedJob = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{FinalizedAt: &beforeHorizon, State: ptrutil.Ptr(rivertype.JobStateCompleted)}) discardedJob = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{FinalizedAt: &beforeHorizon, State: ptrutil.Ptr(rivertype.JobStateDiscarded)}) @@ -274,11 +311,12 @@ func exerciseJobDelete[TTx any](ctx context.Context, t *testing.T, executorWithT includedQueue1 = "included1" includedQueue2 = "included2" - // Not deleted because in an omitted queue. + // Deleted because they're in included queues. deletedJob1 = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{FinalizedAt: &beforeHorizon, Queue: &includedQueue1, State: ptrutil.Ptr(rivertype.JobStateCompleted)}) deletedJob2 = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{FinalizedAt: &beforeHorizon, Queue: &includedQueue2, State: ptrutil.Ptr(rivertype.JobStateCompleted)}) ) + // The three older non-included jobs must not consume this batch. numDeleted, err := exec.JobDeleteBefore(ctx, &riverdriver.JobDeleteBeforeParams{ CancelledDoDelete: true, CancelledFinalizedAtHorizon: horizon, @@ -286,7 +324,7 @@ func exerciseJobDelete[TTx any](ctx context.Context, t *testing.T, executorWithT CompletedFinalizedAtHorizon: horizon, DiscardedDoDelete: true, DiscardedFinalizedAtHorizon: horizon, - Max: 1_000, + Max: 2, QueuesIncluded: []string{includedQueue1, includedQueue2}, }) require.NoError(t, err) @@ -456,5 +494,30 @@ func exerciseJobDelete[TTx any](ctx context.Context, t *testing.T, executorWithT require.NoError(t, err) require.Equal(t, []int64{job1.ID, job2.ID, job3.ID, job4.ID, job5.ID}, sliceutil.Map(deletedJobs, func(j *rivertype.JobRow) int64 { return j.ID })) }) + + t.Run("SortedResultsDescending", func(t *testing.T) { + t.Parallel() + + exec, bundle := setup(ctx, t) + if bundle.driver.DatabaseName() == riverdriver.DatabaseNameSQLite { + t.Skip("SQLite always returns JobDeleteMany results ordered by ID") + } + + var ( + job1 = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{}) + job2 = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{}) + job3 = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{}) + job4 = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{}) + job5 = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{}) + ) + + deletedJobs, err := exec.JobDeleteMany(ctx, &riverdriver.JobDeleteManyParams{ + Max: 100, + OrderByClause: "id DESC", + WhereClause: "true", + }) + require.NoError(t, err) + require.Equal(t, []int64{job5.ID, job4.ID, job3.ID, job2.ID, job1.ID}, sliceutil.Map(deletedJobs, func(j *rivertype.JobRow) int64 { return j.ID })) + }) }) } diff --git a/riverdriver/riverdrivertest/job_insert.go b/riverdriver/riverdrivertest/job_insert.go index ee4fa745..35dbdef5 100644 --- a/riverdriver/riverdrivertest/job_insert.go +++ b/riverdriver/riverdrivertest/job_insert.go @@ -86,9 +86,8 @@ func exerciseJobInsert[TTx any](ctx context.Context, t *testing.T, require.False(t, result.UniqueSkippedAsDuplicate) job := result.Job - // SQLite needs to set a special metadata key to be able to - // check for duplicates. Remove this for purposes of comparing - // inserted metadata. + // Some drivers use a special metadata key to detect duplicates. + // Remove it for purposes of comparing inserted metadata. job.Metadata, err = sjson.DeleteBytes(job.Metadata, rivercommon.MetadataKeyUniqueNonce) require.NoError(t, err) @@ -113,6 +112,33 @@ func exerciseJobInsert[TTx any](ctx context.Context, t *testing.T, } }) + t.Run("DuplicateIDWithoutUniqueKeyErrors", func(t *testing.T) { + t.Parallel() + + exec, _ := setup(ctx, t) + id := rand.Int64() + params := &riverdriver.JobInsertFastManyParams{ + Jobs: []*riverdriver.JobInsertFastParams{ + { + EncodedArgs: []byte(`{}`), + ID: &id, + Kind: "test_kind", + MaxAttempts: rivercommon.MaxAttemptsDefault, + Priority: rivercommon.PriorityDefault, + Queue: rivercommon.QueueDefault, + State: rivertype.JobStateAvailable, + Tags: []string{}, + }, + }, + } + + _, err := exec.JobInsertFastMany(ctx, params) + require.NoError(t, err) + + _, err = exec.JobInsertFastMany(ctx, params) + require.Error(t, err) + }) + t.Run("MissingValuesDefaultAsExpected", func(t *testing.T) { t.Parallel() @@ -699,7 +725,7 @@ func exerciseJobInsert[TTx any](ctx context.Context, t *testing.T, _, err := exec.JobInsertFull(ctx, params) require.Error(t, err) // two separate error messages here for Postgres and SQLite - require.Regexp(t, `(CHECK constraint failed: finalized_or_finalized_at_null|violates check constraint "finalized_or_finalized_at_null")`, err.Error()) + require.Regexp(t, `(CHECK constraint failed: finalized_or_finalized_at_null|violates check constraint "finalized_or_finalized_at_null"|Check constraint 'finalized_or_finalized_at_null' is violated)`, err.Error()) }) t.Run(fmt.Sprintf("CanSetState%sWithFinalizedAt", capitalizeJobState(state)), func(t *testing.T) { @@ -749,7 +775,7 @@ func exerciseJobInsert[TTx any](ctx context.Context, t *testing.T, })) require.Error(t, err) // two separate error messages here for Postgres and SQLite - require.Regexp(t, `(CHECK constraint failed: finalized_or_finalized_at_null|violates check constraint "finalized_or_finalized_at_null")`, err.Error()) + require.Regexp(t, `(CHECK constraint failed: finalized_or_finalized_at_null|violates check constraint "finalized_or_finalized_at_null"|Check constraint 'finalized_or_finalized_at_null' is violated)`, err.Error()) }) } }) diff --git a/riverdriver/riverdrivertest/job_read.go b/riverdriver/riverdrivertest/job_read.go index 942ca446..6434d6c3 100644 --- a/riverdriver/riverdrivertest/job_read.go +++ b/riverdriver/riverdrivertest/job_read.go @@ -498,6 +498,7 @@ func exerciseJobRead[TTx any](ctx context.Context, t *testing.T, executorWithTx job2 := testfactory.Job(ctx, t, exec, &testfactory.JobOpts{Kind: ptrutil.Ptr("kind2")}) // Not returned. + _ = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{Kind: ptrutil.Ptr("KIND1")}) _ = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{Kind: ptrutil.Ptr("kind3")}) jobs, err := exec.JobGetByKindMany(ctx, &riverdriver.JobGetByKindManyParams{ diff --git a/riverdriver/riverdrivertest/listener.go b/riverdriver/riverdrivertest/listener.go index 7a096075..cd895d0d 100644 --- a/riverdriver/riverdrivertest/listener.go +++ b/riverdriver/riverdrivertest/listener.go @@ -128,8 +128,8 @@ func exerciseListener[TTx any](ctx context.Context, t *testing.T, driverWithPool listener = driver.GetListener(&riverdriver.GetListenenerParams{Schema: ""}) ) - if driver.DatabaseName() == riverdriver.DatabaseNameSQLite { - t.Skip("SQLite has no search_path") + if driver.DatabaseName() == riverdriver.DatabaseNameMySQL || driver.DatabaseName() == riverdriver.DatabaseNameSQLite { + t.Skipf("%s has no search_path", driver.DatabaseName()) } listener.SetAfterConnectExec("SET search_path TO 'public'") @@ -146,6 +146,10 @@ func exerciseListener[TTx any](ctx context.Context, t *testing.T, driverWithPool listener = driver.GetListener(&riverdriver.GetListenenerParams{Schema: ""}) ) + if driver.DatabaseName() == riverdriver.DatabaseNameMySQL { + t.Skip("MySQL has no search_path") + } + connectListener(ctx, t, listener) require.Empty(t, listener.Schema()) }) diff --git a/riverdriver/riverdrivertest/notification.go b/riverdriver/riverdrivertest/notification.go index 311f77dd..e3ff4585 100644 --- a/riverdriver/riverdrivertest/notification.go +++ b/riverdriver/riverdrivertest/notification.go @@ -25,7 +25,7 @@ func exerciseNotification[TTx any](ctx context.Context, t *testing.T, executorWi ($4, $5, $6), ($7, $8, $9) ` - if driver.DatabaseName() == riverdriver.DatabaseNameSQLite { + if driver.DatabaseName() == riverdriver.DatabaseNameMySQL || driver.DatabaseName() == riverdriver.DatabaseNameSQLite { insertQuery = ` INSERT INTO river_notification (created_at, payload, topic) VALUES diff --git a/riverdriver/riverdrivertest/riverdrivertest.go b/riverdriver/riverdrivertest/riverdrivertest.go index b970f80d..b670fc1a 100644 --- a/riverdriver/riverdrivertest/riverdrivertest.go +++ b/riverdriver/riverdrivertest/riverdrivertest.go @@ -80,6 +80,25 @@ func exerciseDriverPool[TTx any](ctx context.Context, t *testing.T, }) }) + t.Run("SafeIdentifier", func(t *testing.T) { + t.Parallel() + + _, driver := executorWithTx(ctx, t) + + switch driver.DatabaseName() { + case riverdriver.DatabaseNamePostgres, riverdriver.DatabaseNameSQLite: + require.Equal(t, `"my_schema"`, driver.SafeIdentifier("my_schema")) + require.Equal(t, `"has space"`, driver.SafeIdentifier("has space")) + require.Equal(t, `"has""quote"`, driver.SafeIdentifier(`has"quote`)) + case riverdriver.DatabaseNameMySQL: + require.Equal(t, "`my_schema`", driver.SafeIdentifier("my_schema")) + require.Equal(t, "`has space`", driver.SafeIdentifier("has space")) + require.Equal(t, "`has``backtick`", driver.SafeIdentifier("has`backtick")) + default: + require.FailNow(t, "Don't know how to check SafeIdentifier for: "+driver.DatabaseName()) + } + }) + t.Run("SupportsListenNotify", func(t *testing.T) { t.Parallel() @@ -90,6 +109,8 @@ func exerciseDriverPool[TTx any](ctx context.Context, t *testing.T, require.True(t, driver.SupportsListenNotify()) case riverdriver.DatabaseNameSQLite: require.True(t, driver.SupportsListenNotify()) + case riverdriver.DatabaseNameMySQL: + require.True(t, driver.SupportsListenNotify()) default: require.FailNow(t, "Don't know how to check SupportsListenNotify for: "+driver.DatabaseName()) } @@ -107,6 +128,7 @@ func requireMissingRelation(t *testing.T, err error, schema, missingRelation str // lib/pq: pq: relation %s.%s does not exist // SQLite: no such table: %s.%s // Turso: turso: error: Invalid argument supplied: no such database: %s - require.Regexp(t, fmt.Sprintf(`(pq: relation "%s\.%s" does not exist|no such table: %s\.%s|no such database: %s)`, schema, missingRelation, schema, missingRelation, schema), err.Error()) + // MySQL: Unknown database '%s' + require.Regexp(t, fmt.Sprintf(`(pq: relation "%s\.%s" does not exist|no such table: %s\.%s|no such database: %s|Unknown database '%s')`, schema, missingRelation, schema, missingRelation, schema, schema), err.Error()) } } diff --git a/riverdriver/rivermysql/example_mysql_test.go b/riverdriver/rivermysql/example_mysql_test.go new file mode 100644 index 00000000..74aa0177 --- /dev/null +++ b/riverdriver/rivermysql/example_mysql_test.go @@ -0,0 +1,125 @@ +package rivermysql_test + +import ( + "cmp" + "context" + "database/sql" + "fmt" + "log/slog" + "os" + "sort" + + _ "github.com/go-sql-driver/mysql" + + "github.com/riverqueue/river" + "github.com/riverqueue/river/riverdriver/rivermysql" + "github.com/riverqueue/river/rivermigrate" + "github.com/riverqueue/river/rivershared/riversharedtest" + "github.com/riverqueue/river/rivershared/util/slogutil" + "github.com/riverqueue/river/rivershared/util/testutil" +) + +type MySQLSortArgs struct { + // Strings is a slice of strings to sort. + Strings []string `json:"strings"` +} + +func (MySQLSortArgs) Kind() string { return "sort" } + +type MySQLSortWorker struct { + river.WorkerDefaults[MySQLSortArgs] +} + +func (w *MySQLSortWorker) Work(ctx context.Context, job *river.Job[MySQLSortArgs]) error { + sort.Strings(job.Args.Strings) + fmt.Printf("Sorted strings: %+v\n", job.Args.Strings) + return nil +} + +// Example_mysql demonstrates use of River's MySQL driver. +func Example_mysql() { + // MySQL tests are opt-in because they require a running server. When + // disabled, print the expected output so the example test always passes. + val := os.Getenv("RIVER_MYSQL_TESTS_ENABLED") + if val != "1" && val != "true" { + fmt.Println("Sorted strings: [bear tiger whale]") + return + } + + ctx := context.Background() + + dsn := cmp.Or( + os.Getenv("TEST_MYSQL_URL"), + "root@tcp(localhost:3306)/?parseTime=true&multiStatements=true&loc=UTC&time_zone=%27%2B00%3A00%27", + ) + + dbPool, err := sql.Open("mysql", dsn) + if err != nil { + panic(err) + } + defer dbPool.Close() + + // Create a temporary database for the example. + const exampleDB = "river_example_mysql" + if _, err := dbPool.ExecContext(ctx, "CREATE DATABASE IF NOT EXISTS "+exampleDB); err != nil { + panic(err) + } + defer func() { + _, _ = dbPool.ExecContext(ctx, "DROP DATABASE IF EXISTS "+exampleDB) + }() + + // Run River's migrations to prepare the schema. + migrator, err := rivermigrate.New(rivermysql.New(dbPool), &rivermigrate.Config{ + Logger: slog.New(slog.NewTextHandler(os.Stdout, &slog.HandlerOptions{Level: slog.LevelWarn, ReplaceAttr: slogutil.NoLevelTime})), + Schema: exampleDB, + }) + if err != nil { + panic(err) + } + if _, err := migrator.Migrate(ctx, rivermigrate.DirectionUp, nil); err != nil { + panic(err) + } + + workers := river.NewWorkers() + river.AddWorker(workers, &MySQLSortWorker{}) + + riverClient, err := river.NewClient(rivermysql.New(dbPool), &river.Config{ + Logger: slog.New(slog.NewTextHandler(os.Stdout, &slog.HandlerOptions{Level: slog.LevelWarn, ReplaceAttr: slogutil.NoLevelTime})), + Queues: map[string]river.QueueConfig{ + river.QueueDefault: {MaxWorkers: 100}, + }, + Schema: exampleDB, + TestOnly: true, // suitable only for use in tests; remove for live environments + Workers: workers, + }) + if err != nil { + panic(err) + } + + // Out of example scope, but used to wait until a job is worked. + subscribeChan, subscribeCancel := riverClient.Subscribe(river.EventKindJobCompleted) + defer subscribeCancel() + + if err := riverClient.Start(ctx); err != nil { + panic(err) + } + + _, err = riverClient.Insert(ctx, MySQLSortArgs{ + Strings: []string{ + "whale", "tiger", "bear", + }, + }, nil) + if err != nil { + panic(err) + } + + // Wait for jobs to complete. Only needed for purposes of the example test. + riversharedtest.WaitOrTimeoutN(testutil.PanicTB(), subscribeChan, 1) + + if err := riverClient.Stop(ctx); err != nil { + panic(err) + } + + // Output: + // Sorted strings: [bear tiger whale] +} diff --git a/riverdriver/rivermysql/go.mod b/riverdriver/rivermysql/go.mod new file mode 100644 index 00000000..584dbd21 --- /dev/null +++ b/riverdriver/rivermysql/go.mod @@ -0,0 +1,33 @@ +module github.com/riverqueue/river/riverdriver/rivermysql + +go 1.25.0 + +toolchain go1.25.7 + +require ( + github.com/go-sql-driver/mysql v1.10.0 + github.com/riverqueue/river v0.41.0 + github.com/riverqueue/river/riverdriver v0.41.0 + github.com/riverqueue/river/rivershared v0.41.0 + github.com/riverqueue/river/rivertype v0.41.0 + github.com/stretchr/testify v1.11.1 + github.com/tidwall/gjson v1.19.0 + github.com/tidwall/sjson v1.2.5 +) + +require ( + filippo.io/edwards25519 v1.2.0 // indirect + github.com/davecgh/go-spew v1.1.1 // indirect + github.com/jackc/pgpassfile v1.0.0 // indirect + github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect + github.com/jackc/pgx/v5 v5.10.0 // indirect + github.com/jackc/puddle/v2 v2.2.2 // indirect + github.com/pmezard/go-difflib v1.0.0 // indirect + github.com/riverqueue/river/riverdriver/riverpgxv5 v0.41.0 // indirect + github.com/tidwall/match v1.2.0 // indirect + github.com/tidwall/pretty v1.2.1 // indirect + go.uber.org/goleak v1.3.0 // indirect + golang.org/x/sync v0.22.0 // indirect + golang.org/x/text v0.40.0 // indirect + gopkg.in/yaml.v3 v3.0.1 // indirect +) diff --git a/riverdriver/rivermysql/go.sum b/riverdriver/rivermysql/go.sum new file mode 100644 index 00000000..e695dd94 --- /dev/null +++ b/riverdriver/rivermysql/go.sum @@ -0,0 +1,65 @@ +filippo.io/edwards25519 v1.2.0 h1:crnVqOiS4jqYleHd9vaKZ+HKtHfllngJIiOpNpoJsjo= +filippo.io/edwards25519 v1.2.0/go.mod h1:xzAOLCNug/yB62zG1bQ8uziwrIqIuxhctzJT18Q77mc= +github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= +github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/go-sql-driver/mysql v1.10.0 h1:Q+1LV8DkHJvSYAdR83XzuhDaTykuDx0l6fkXxoWCWfw= +github.com/go-sql-driver/mysql v1.10.0/go.mod h1:M+cqaI7+xxXGG9swrdeUIoPG3Y3KCkF0pZej+SK+nWk= +github.com/jackc/pgerrcode v0.0.0-20240316143900-6e2875d9b438 h1:Dj0L5fhJ9F82ZJyVOmBx6msDp/kfd1t9GRfny/mfJA0= +github.com/jackc/pgerrcode v0.0.0-20240316143900-6e2875d9b438/go.mod h1:a/s9Lp5W7n/DD0VrVoyJ00FbP2ytTPDVOivvn2bMlds= +github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM= +github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg= +github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo= +github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM= +github.com/jackc/pgx/v5 v5.10.0 h1:VhSvgU2jSli8o3AqIEOTJr7rZwAEUVo4E4XhR94Zfr0= +github.com/jackc/pgx/v5 v5.10.0/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4= +github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo= +github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4= +github.com/kr/pretty v0.3.0 h1:WgNl7dwNpEZ6jJ9k1snq4pZsg7DOEN8hP9Xw0Tsjwk0= +github.com/kr/pretty v0.3.0/go.mod h1:640gp4NfQd8pI5XOwp5fnNeVWj67G7CFk/SaSQn7NBk= +github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= +github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= +github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/riverqueue/river v0.41.0 h1:E7Yfyhn74IgVaCBKsBpTHoC0LxWJvoG0jKt7Tb+lJZg= +github.com/riverqueue/river v0.41.0/go.mod h1:WXiAF1/2gfUPj3H+WqQG3Q5NW5w9JmRXRTe7U9zybNc= +github.com/riverqueue/river/riverdriver v0.41.0 h1:A5g80n6fCGu9Bp4hSq7gys7mvfOiuWkbQHuFx3iPp7w= +github.com/riverqueue/river/riverdriver v0.41.0/go.mod h1:OeGOhZYoH7Ac0CrVanNAAwD9C0WAXy4LzUEl65Ex0B4= +github.com/riverqueue/river/riverdriver/riverpgxv5 v0.41.0 h1:uHGluToMMWvCnpNkAbEWEmP4YAjdseuriEXGu1dwOp4= +github.com/riverqueue/river/riverdriver/riverpgxv5 v0.41.0/go.mod h1:0ueUi+3eW5fIsibS5VkD2HzQmwLKTAiM/uiiVMX/SME= +github.com/riverqueue/river/rivershared v0.41.0 h1:ax3uz5KiqfmcvRQ3kKB6iYs9Nt8J5L1RNfOKoVMWkqw= +github.com/riverqueue/river/rivershared v0.41.0/go.mod h1:FaZ7bxC2DORhyFDVyHKRiZzgQPf2b/IELwPBXnHZVLA= +github.com/riverqueue/river/rivertype v0.41.0 h1:dfscvt1asf1PpeeHTMFxQdV1ZfMoReiiMNFCPPfR0as= +github.com/riverqueue/river/rivertype v0.41.0/go.mod h1:D1Ad+EaZiaXbQbJcJcfeicXJMBKno0n6UcfKI5Q7DIQ= +github.com/robfig/cron/v3 v3.0.1 h1:WdRxkvbJztn8LMz/QEvLN5sBU+xKpSqwwUO1Pjr4qDs= +github.com/robfig/cron/v3 v3.0.1/go.mod h1:eQICP3HwyT7UooqI/z+Ov+PtYAWygg1TEWWzGIFLtro= +github.com/rogpeppe/go-internal v1.12.0 h1:exVL4IDcn6na9z1rAb56Vxr+CgyK3nn3O+epU5NdKM8= +github.com/rogpeppe/go-internal v1.12.0/go.mod h1:E+RYuTGaKKdloAfM02xzb0FW3Paa99yedzYV+kq4uf4= +github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= +github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI= +github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= +github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= +github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= +github.com/tidwall/gjson v1.14.2/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk= +github.com/tidwall/gjson v1.19.0 h1:xwxm7n691Uf3u5OFjzngavjGTh55KX5q/9w9xHW88JU= +github.com/tidwall/gjson v1.19.0/go.mod h1:V37/opeE/JbLUOfH0QTXiNez2l0RUjYUhpT4szFQAfc= +github.com/tidwall/match v1.1.1/go.mod h1:eRSPERbgtNPcGhD8UCthc6PmLEQXEWd3PRB5JTxsfmM= +github.com/tidwall/match v1.2.0 h1:0pt8FlkOwjN2fPt4bIl4BoNxb98gGHN2ObFEDkrfZnM= +github.com/tidwall/match v1.2.0/go.mod h1:eRSPERbgtNPcGhD8UCthc6PmLEQXEWd3PRB5JTxsfmM= +github.com/tidwall/pretty v1.2.0/go.mod h1:ITEVvHYasfjBbM0u2Pg8T2nJnzm8xPwvNhhsoaGGjNU= +github.com/tidwall/pretty v1.2.1 h1:qjsOFOWWQl+N3RsoF5/ssm1pHmJJwhjlSbZ51I6wMl4= +github.com/tidwall/pretty v1.2.1/go.mod h1:ITEVvHYasfjBbM0u2Pg8T2nJnzm8xPwvNhhsoaGGjNU= +github.com/tidwall/sjson v1.2.5 h1:kLy8mja+1c9jlljvWTlSazM7cKDRfJuR/bOJhcY5NcY= +github.com/tidwall/sjson v1.2.5/go.mod h1:Fvgq9kS/6ociJEDnK0Fk1cpYF4FIW6ZF7LAe+6jwd28= +go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto= +go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE= +golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek= +golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= +golang.org/x/text v0.40.0 h1:Ub2Z6/xjgF1WrYQz2nuITOEegKFtiIy+rieRJ5lHZKs= +golang.org/x/text v0.40.0/go.mod h1:hpnzDAfGV753zIKo+wk3u1bVKCGPbrnF7+7LBF/UHVY= +gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q= +gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/riverdriver/rivermysql/internal/dbsqlc/db.go b/riverdriver/rivermysql/internal/dbsqlc/db.go new file mode 100644 index 00000000..3bebd3a3 --- /dev/null +++ b/riverdriver/rivermysql/internal/dbsqlc/db.go @@ -0,0 +1,24 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.0 + +package dbsqlc + +import ( + "context" + "database/sql" +) + +type DBTX interface { + ExecContext(context.Context, string, ...interface{}) (sql.Result, error) + PrepareContext(context.Context, string) (*sql.Stmt, error) + QueryContext(context.Context, string, ...interface{}) (*sql.Rows, error) + QueryRowContext(context.Context, string, ...interface{}) *sql.Row +} + +func New() *Queries { + return &Queries{} +} + +type Queries struct { +} diff --git a/riverdriver/rivermysql/internal/dbsqlc/models.go b/riverdriver/rivermysql/internal/dbsqlc/models.go new file mode 100644 index 00000000..270911e6 --- /dev/null +++ b/riverdriver/rivermysql/internal/dbsqlc/models.go @@ -0,0 +1,80 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.0 + +package dbsqlc + +import ( + "database/sql" + "time" +) + +type Columns struct { + ColumnName string + TableName string + TableSchema string +} + +type RiverJob struct { + ID int64 + Args []byte + Attempt int64 + AttemptedAt sql.NullTime + AttemptedBy []byte + CreatedAt time.Time + Errors []byte + FinalizedAt sql.NullTime + Kind string + MaxAttempts int64 + Metadata []byte + Priority int16 + Queue string + State string + ScheduledAt time.Time + Tags []byte + UniqueKey sql.NullString + UniqueStates sql.NullInt16 +} + +type RiverLeader struct { + ElectedAt time.Time + ExpiresAt time.Time + LeaderID string + Name string +} + +type RiverMigration struct { + Line string + Version int64 + CreatedAt time.Time +} + +type RiverNotification struct { + ID int64 + CreatedAt time.Time + Payload string + Topic string +} + +type RiverQueue struct { + Name string + CreatedAt time.Time + Metadata []byte + PausedAt sql.NullTime + UpdatedAt time.Time +} + +type Schemata struct { + SchemaName string +} + +type Statistics struct { + IndexName string + TableName string + TableSchema string +} + +type Tables struct { + TableName string + TableSchema string +} diff --git a/riverdriver/rivermysql/internal/dbsqlc/river_job.sql b/riverdriver/rivermysql/internal/dbsqlc/river_job.sql new file mode 100644 index 00000000..aef27534 --- /dev/null +++ b/riverdriver/rivermysql/internal/dbsqlc/river_job.sql @@ -0,0 +1,452 @@ +CREATE TABLE river_job ( + id BIGINT AUTO_INCREMENT PRIMARY KEY, + args JSON NOT NULL DEFAULT (JSON_OBJECT()), + attempt INT NOT NULL DEFAULT 0, + attempted_at DATETIME(6) NULL, + attempted_by JSON NULL, + -- sqlc's MySQL parser doesn't accept UTC_TIMESTAMP(6) in a DEFAULT + -- expression. Runtime migrations use UTC_TIMESTAMP(6); this codegen-only + -- schema declaration uses NOW(6). + created_at DATETIME(6) NOT NULL DEFAULT (NOW(6)), + errors JSON NULL, + finalized_at DATETIME(6) NULL, + kind VARCHAR(128) CHARACTER SET utf8mb4 COLLATE utf8mb4_bin NOT NULL, + max_attempts INT NOT NULL DEFAULT 25, + metadata JSON NOT NULL DEFAULT (JSON_OBJECT()), + priority SMALLINT NOT NULL DEFAULT 1, + queue VARCHAR(128) CHARACTER SET utf8mb4 COLLATE utf8mb4_bin NOT NULL DEFAULT 'default', + state VARCHAR(20) CHARACTER SET utf8mb4 COLLATE utf8mb4_bin NOT NULL DEFAULT 'available', + scheduled_at DATETIME(6) NOT NULL DEFAULT (NOW(6)), + tags JSON NOT NULL DEFAULT (JSON_ARRAY()), + unique_key VARBINARY(255) NULL, + unique_states SMALLINT NULL +); + +-- name: JobGetByID :one +SELECT * +FROM /* TEMPLATE: schema */river_job +WHERE id = sqlc.arg('id') +LIMIT 1; + +-- name: JobGetByIDMany :many +SELECT * +FROM /* TEMPLATE: schema */river_job +WHERE id IN (sqlc.slice('id')) +ORDER BY id; + +-- name: JobGetByKindMany :many +SELECT * +FROM /* TEMPLATE: schema */river_job +WHERE kind IN (sqlc.slice('kind')) +ORDER BY id; + +-- name: JobCancelExec :execresult +UPDATE /* TEMPLATE: schema */river_job +SET + state = CASE WHEN state = 'running' THEN state ELSE 'cancelled' END, + finalized_at = CASE WHEN state = 'running' THEN finalized_at ELSE COALESCE(sqlc.narg('now'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) END, + metadata = JSON_SET(metadata, '$.cancel_attempted_at', CAST(sqlc.arg('cancel_attempted_at') AS CHAR)) +WHERE id = sqlc.arg('id') + AND state NOT IN ('cancelled', 'completed', 'discarded') + AND finalized_at IS NULL; + +-- name: JobClearUniqueNonce :exec +UPDATE /* TEMPLATE: schema */river_job +SET metadata = JSON_REMOVE(metadata, '$."river:unique_nonce"') +WHERE id IN (sqlc.slice('id')); + +-- name: JobCountByAllStates :many +SELECT state, count(*) AS count +FROM /* TEMPLATE: schema */river_job +GROUP BY state; + +-- name: JobCountByQueueAndState :many +WITH all_queues AS ( + SELECT DISTINCT river_job.queue + FROM /* TEMPLATE: schema */river_job + WHERE river_job.queue IN (sqlc.slice('queue_names')) +), + +running_job_counts AS ( + SELECT river_job.queue, COUNT(*) AS count + FROM /* TEMPLATE: schema */river_job + WHERE river_job.queue IN (sqlc.slice('queue_names')) + AND river_job.state = 'running' + GROUP BY river_job.queue +), + +available_job_counts AS ( + SELECT river_job.queue, COUNT(*) AS count + FROM /* TEMPLATE: schema */river_job + WHERE river_job.queue IN (sqlc.slice('queue_names')) + AND river_job.state = 'available' + GROUP BY river_job.queue +) + +SELECT + all_queues.queue, + COALESCE(available_job_counts.count, 0) AS count_available, + COALESCE(running_job_counts.count, 0) AS count_running +FROM all_queues +LEFT JOIN running_job_counts ON all_queues.queue = running_job_counts.queue +LEFT JOIN available_job_counts ON all_queues.queue = available_job_counts.queue +ORDER BY all_queues.queue ASC; + +-- name: JobCountByState :one +SELECT count(*) AS count +FROM /* TEMPLATE: schema */river_job +WHERE state = sqlc.arg('state'); + +-- name: JobDeleteExec :execresult +DELETE FROM /* TEMPLATE: schema */river_job +WHERE id = sqlc.arg('id') + AND river_job.state != 'running'; + +-- name: JobDeleteSelect :one +SELECT * +FROM /* TEMPLATE: schema */river_job +WHERE id = sqlc.arg('id') +LIMIT 1 +FOR UPDATE; + +-- name: JobDeleteBefore :execresult +DELETE FROM /* TEMPLATE: schema */river_job +WHERE + id IN ( + SELECT id FROM ( + SELECT rj2.id + FROM /* TEMPLATE: schema */river_job rj2 + WHERE ( + ( + CAST(sqlc.arg('cancelled_do_delete') AS UNSIGNED) + AND rj2.state = 'cancelled' + AND rj2.finalized_at < sqlc.arg('cancelled_finalized_at_horizon') + ) OR ( + CAST(sqlc.arg('completed_do_delete') AS UNSIGNED) + AND rj2.state = 'completed' + AND rj2.finalized_at < sqlc.arg('completed_finalized_at_horizon') + ) OR ( + CAST(sqlc.arg('discarded_do_delete') AS UNSIGNED) + AND rj2.state = 'discarded' + AND rj2.finalized_at < sqlc.arg('discarded_finalized_at_horizon') + ) + ) + AND ( + CAST(sqlc.arg('queues_excluded_empty') AS SIGNED) + OR rj2.queue NOT IN (sqlc.slice('queues_excluded')) + ) + AND ( + CAST(sqlc.arg('queues_included_empty') AS SIGNED) + OR rj2.queue IN (sqlc.slice('queues_included')) + ) + ORDER BY rj2.id + LIMIT ? + ) AS tmp + ); + +-- name: JobDeleteManySelect :many +SELECT * +FROM /* TEMPLATE: schema */river_job +WHERE id IN ( + SELECT id FROM ( + SELECT id + FROM /* TEMPLATE: schema */river_job + WHERE /* TEMPLATE_BEGIN: where_clause */ true /* TEMPLATE_END */ + AND state != 'running' + ORDER BY /* TEMPLATE_BEGIN: order_by_clause */ id /* TEMPLATE_END */ + LIMIT ? + ) AS tmp +) +ORDER BY /* TEMPLATE_BEGIN: order_by_clause */ id /* TEMPLATE_END */ +FOR UPDATE; + +-- name: JobDeleteManyExec :exec +DELETE FROM /* TEMPLATE: schema */river_job +WHERE id IN (sqlc.slice('id')); + +-- name: JobGetAvailableIDs :many +SELECT id +FROM /* TEMPLATE: schema */river_job +WHERE + priority >= 0 + AND queue = sqlc.arg('queue') + AND scheduled_at <= COALESCE(sqlc.narg('now'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) + AND state = 'available' +ORDER BY priority ASC, scheduled_at ASC, id ASC +LIMIT ? +FOR UPDATE SKIP LOCKED; + +-- name: JobGetAvailableUpdate :exec +UPDATE /* TEMPLATE: schema */river_job +SET + attempt = attempt + 1, + attempted_at = COALESCE(sqlc.narg('now'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')), + attempted_by = JSON_ARRAY_APPEND( + CASE + WHEN JSON_LENGTH(COALESCE(attempted_by, JSON_ARRAY())) < CAST(sqlc.arg('max_attempted_by') AS SIGNED) + THEN COALESCE(attempted_by, JSON_ARRAY()) + WHEN CAST(sqlc.arg('max_attempted_by') AS SIGNED) <= 1 + THEN JSON_ARRAY() + ELSE COALESCE( + JSON_EXTRACT( + attempted_by, + CONCAT('$[last-', CAST(sqlc.arg('max_attempted_by') AS SIGNED) - 2, ' to last]') + ), + JSON_ARRAY() + ) + END, + '$', + CAST(sqlc.arg('attempted_by') AS CHAR) + ), + state = 'running' +WHERE id IN (sqlc.slice('id')); + +-- name: JobGetByIDManyOrdered :many +SELECT * +FROM /* TEMPLATE: schema */river_job +WHERE id IN (sqlc.slice('id')) +ORDER BY priority ASC, scheduled_at ASC, id ASC; + +-- name: JobGetStuck :many +SELECT * +FROM /* TEMPLATE: schema */river_job +WHERE state = 'running' + AND id > sqlc.arg('after_id') + AND attempted_at < sqlc.arg('stuck_horizon') +ORDER BY id +LIMIT ?; + +-- name: JobInsertFast :execresult +INSERT INTO /* TEMPLATE: schema */river_job( + id, + args, + created_at, + kind, + max_attempts, + metadata, + priority, + queue, + scheduled_at, + state, + tags, + unique_key, + unique_states +) VALUES ( + sqlc.narg('id'), + sqlc.arg('args'), + COALESCE(sqlc.narg('created_at'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')), + sqlc.arg('kind'), + sqlc.arg('max_attempts'), + CAST(sqlc.arg('metadata') AS JSON), + sqlc.arg('priority'), + sqlc.arg('queue'), + COALESCE(sqlc.narg('scheduled_at'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')), + sqlc.arg('state'), + CAST(sqlc.arg('tags') AS JSON), + sqlc.narg('unique_key'), + sqlc.narg('unique_states') +) /* TEMPLATE_BEGIN: on_duplicate_key */ ON DUPLICATE KEY UPDATE id = LAST_INSERT_ID(id) /* TEMPLATE_END */; + +-- name: JobInsertFullExec :execlastid +INSERT INTO /* TEMPLATE: schema */river_job( + args, + attempt, + attempted_at, + attempted_by, + created_at, + errors, + finalized_at, + kind, + max_attempts, + metadata, + priority, + queue, + scheduled_at, + state, + tags, + unique_key, + unique_states +) VALUES ( + sqlc.arg('args'), + sqlc.arg('attempt'), + sqlc.narg('attempted_at'), + CAST(sqlc.narg('attempted_by') AS JSON), + COALESCE(sqlc.narg('created_at'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')), + CAST(sqlc.narg('errors') AS JSON), + sqlc.narg('finalized_at'), + sqlc.arg('kind'), + sqlc.arg('max_attempts'), + CAST(sqlc.arg('metadata') AS JSON), + sqlc.arg('priority'), + sqlc.arg('queue'), + COALESCE(sqlc.narg('scheduled_at'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')), + sqlc.arg('state'), + CAST(sqlc.arg('tags') AS JSON), + sqlc.narg('unique_key'), + sqlc.narg('unique_states') +); + +-- name: JobKindList :many +SELECT DISTINCT kind +FROM /* TEMPLATE: schema */river_job +WHERE (sqlc.arg('match') = '' OR LOWER(kind) COLLATE utf8mb4_general_ci LIKE CONCAT('%', LOWER(sqlc.arg('match')), '%')) + AND (sqlc.arg('after') = '' OR kind > sqlc.arg('after')) + AND kind NOT IN (sqlc.slice('exclude')) +ORDER BY kind ASC +LIMIT ?; + +-- name: JobList :many +SELECT * +FROM /* TEMPLATE: schema */river_job +WHERE /* TEMPLATE_BEGIN: where_clause */ true /* TEMPLATE_END */ +ORDER BY /* TEMPLATE_BEGIN: order_by_clause */ id /* TEMPLATE_END */ +LIMIT ?; + +-- name: JobRescue :exec +UPDATE /* TEMPLATE: schema */river_job +SET + errors = JSON_ARRAY_APPEND(COALESCE(errors, JSON_ARRAY()), '$', CAST(sqlc.arg('error') AS JSON)), + finalized_at = sqlc.narg('finalized_at'), + scheduled_at = sqlc.arg('scheduled_at'), + metadata = JSON_SET( + metadata, + '$."river:rescue_count"', + COALESCE( + CASE JSON_TYPE(JSON_EXTRACT(metadata, '$."river:rescue_count"')) + WHEN 'INTEGER' THEN JSON_EXTRACT(metadata, '$."river:rescue_count"') + WHEN 'DOUBLE' THEN JSON_EXTRACT(metadata, '$."river:rescue_count"') + ELSE NULL + END, + 0 + ) + 1 + ), + state = sqlc.arg('state') +WHERE id = sqlc.arg('id'); + +-- name: JobRetryExec :execresult +UPDATE /* TEMPLATE: schema */river_job +SET + state = 'available', + max_attempts = CASE WHEN attempt = max_attempts THEN max_attempts + 1 ELSE max_attempts END, + finalized_at = NULL, + scheduled_at = COALESCE(sqlc.narg('now'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) +WHERE id = sqlc.arg('id') + AND state != 'running' + AND ( + state <> 'available' + OR scheduled_at > COALESCE(sqlc.narg('now'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) + ); + +-- name: JobSchedule :many +WITH eligible AS ( + SELECT river_job.id, river_job.unique_key, river_job.unique_states, river_job.priority, river_job.scheduled_at, + CASE + WHEN river_job.unique_key IS NOT NULL AND river_job.unique_states IS NOT NULL THEN + ROW_NUMBER() OVER (PARTITION BY river_job.unique_key ORDER BY river_job.priority, river_job.scheduled_at, river_job.id) + ELSE NULL + END AS row_num + FROM /* TEMPLATE: schema */river_job + WHERE + river_job.state IN ('retryable', 'scheduled') + AND river_job.scheduled_at <= COALESCE(sqlc.narg('now'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) + ORDER BY + river_job.priority, + river_job.scheduled_at, + river_job.id + LIMIT ? + FOR UPDATE SKIP LOCKED +), +unique_conflicts AS ( + SELECT DISTINCT eligible.unique_key + FROM /* TEMPLATE: schema */river_job + JOIN eligible + ON river_job.unique_key = eligible.unique_key + AND river_job.id != eligible.id + WHERE + river_job.unique_key IS NOT NULL + AND river_job.unique_states IS NOT NULL + AND CASE river_job.state + WHEN 'available' THEN river_job.unique_states & (1 << 0) + WHEN 'cancelled' THEN river_job.unique_states & (1 << 1) + WHEN 'completed' THEN river_job.unique_states & (1 << 2) + WHEN 'discarded' THEN river_job.unique_states & (1 << 3) + WHEN 'pending' THEN river_job.unique_states & (1 << 4) + WHEN 'retryable' THEN river_job.unique_states & (1 << 5) + WHEN 'running' THEN river_job.unique_states & (1 << 6) + WHEN 'scheduled' THEN river_job.unique_states & (1 << 7) + ELSE 0 + END >= 1 +) +SELECT eligible.id, + CASE + WHEN eligible.unique_key IS NULL OR eligible.unique_states IS NULL THEN FALSE + WHEN uc.unique_key IS NOT NULL THEN TRUE + WHEN eligible.row_num > 1 THEN TRUE + ELSE FALSE + END AS conflict_discarded +FROM eligible +LEFT JOIN unique_conflicts uc ON eligible.unique_key = uc.unique_key +ORDER BY eligible.priority, eligible.scheduled_at, eligible.id; + +-- name: JobScheduleSetAvailableExec :exec +UPDATE /* TEMPLATE: schema */river_job +SET state = 'available' +WHERE id IN (sqlc.slice('id')); + +-- name: JobScheduleSetDiscardedExec :exec +UPDATE /* TEMPLATE: schema */river_job +SET metadata = JSON_MERGE_PATCH(metadata, '{"unique_key_conflict": "scheduler_discarded"}'), + finalized_at = COALESCE(sqlc.narg('now'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')), + state = 'discarded' +WHERE id IN (sqlc.slice('id')); + +-- name: JobSetMetadataIfNotRunningExec :execresult +UPDATE /* TEMPLATE: schema */river_job +SET metadata = JSON_MERGE_PATCH(metadata, CAST(sqlc.arg('metadata_updates') AS JSON)) +WHERE id = sqlc.arg('id') + AND state != 'running'; + +-- name: JobSetStateIfRunningExec :exec +UPDATE /* TEMPLATE: schema */river_job +SET + attempt = CASE WHEN (CAST(sqlc.arg('state') AS CHAR) <> 'retryable' AND sqlc.arg('state') <> 'scheduled' OR JSON_EXTRACT(metadata, '$.cancel_attempted_at') IS NULL) AND CAST(sqlc.arg('attempt_do_update') AS SIGNED) + THEN sqlc.arg('attempt') + ELSE attempt END, + errors = CASE WHEN CAST(sqlc.arg('errors_do_update') AS SIGNED) + THEN JSON_ARRAY_APPEND(COALESCE(errors, JSON_ARRAY()), '$', CAST(sqlc.arg('error') AS JSON)) + ELSE errors END, + finalized_at = CASE WHEN ((sqlc.arg('state') = 'retryable' OR sqlc.arg('state') = 'scheduled') AND JSON_EXTRACT(metadata, '$.cancel_attempted_at') IS NOT NULL) + THEN COALESCE(sqlc.narg('now'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) + WHEN CAST(sqlc.arg('finalized_at_do_update') AS SIGNED) + THEN sqlc.narg('finalized_at') + ELSE finalized_at END, + metadata = CASE WHEN CAST(sqlc.arg('metadata_do_merge') AS SIGNED) + THEN JSON_MERGE_PATCH(metadata, CAST(sqlc.arg('metadata_updates') AS JSON)) + ELSE metadata END, + scheduled_at = CASE WHEN (CAST(sqlc.arg('state') AS CHAR) <> 'retryable' AND sqlc.arg('state') <> 'scheduled' OR JSON_EXTRACT(metadata, '$.cancel_attempted_at') IS NULL) AND CAST(sqlc.arg('scheduled_at_do_update') AS SIGNED) + THEN sqlc.arg('scheduled_at') + ELSE scheduled_at END, + state = CASE WHEN ((sqlc.arg('state') = 'retryable' OR sqlc.arg('state') = 'scheduled') AND JSON_EXTRACT(metadata, '$.cancel_attempted_at') IS NOT NULL) + THEN 'cancelled' + ELSE sqlc.arg('state') END +WHERE id = sqlc.arg('id') + AND state = 'running'; + +-- name: JobUpdateExec :exec +UPDATE /* TEMPLATE: schema */river_job +SET + metadata = CASE WHEN CAST(sqlc.arg('metadata_do_merge') AS SIGNED) THEN JSON_MERGE_PATCH(metadata, CAST(sqlc.arg('metadata') AS JSON)) ELSE metadata END +WHERE id = sqlc.arg('id'); + +-- name: JobUpdateFullExec :exec +UPDATE /* TEMPLATE: schema */river_job +SET + attempt = CASE WHEN CAST(sqlc.arg('attempt_do_update') AS SIGNED) THEN sqlc.arg('attempt') ELSE attempt END, + attempted_at = CASE WHEN CAST(sqlc.arg('attempted_at_do_update') AS SIGNED) THEN sqlc.narg('attempted_at') ELSE attempted_at END, + attempted_by = CASE WHEN CAST(sqlc.arg('attempted_by_do_update') AS SIGNED) THEN CAST(sqlc.arg('attempted_by') AS JSON) ELSE attempted_by END, + errors = CASE WHEN CAST(sqlc.arg('errors_do_update') AS SIGNED) THEN CAST(sqlc.arg('errors') AS JSON) ELSE errors END, + finalized_at = CASE WHEN CAST(sqlc.arg('finalized_at_do_update') AS SIGNED) THEN sqlc.narg('finalized_at') ELSE finalized_at END, + max_attempts = CASE WHEN CAST(sqlc.arg('max_attempts_do_update') AS SIGNED) THEN sqlc.arg('max_attempts') ELSE max_attempts END, + metadata = CASE WHEN CAST(sqlc.arg('metadata_do_update') AS SIGNED) THEN CAST(sqlc.arg('metadata') AS JSON) ELSE metadata END, + state = CASE WHEN CAST(sqlc.arg('state_do_update') AS SIGNED) THEN sqlc.arg('state') ELSE state END +WHERE id = sqlc.arg('id'); diff --git a/riverdriver/rivermysql/internal/dbsqlc/river_job.sql.go b/riverdriver/rivermysql/internal/dbsqlc/river_job.sql.go new file mode 100644 index 00000000..d2916a70 --- /dev/null +++ b/riverdriver/rivermysql/internal/dbsqlc/river_job.sql.go @@ -0,0 +1,1379 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.0 +// source: river_job.sql + +package dbsqlc + +import ( + "context" + "database/sql" + "strings" + "time" +) + +const jobCancelExec = `-- name: JobCancelExec :execresult +UPDATE /* TEMPLATE: schema */river_job +SET + state = CASE WHEN state = 'running' THEN state ELSE 'cancelled' END, + finalized_at = CASE WHEN state = 'running' THEN finalized_at ELSE COALESCE(?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) END, + metadata = JSON_SET(metadata, '$.cancel_attempted_at', CAST(? AS CHAR)) +WHERE id = ? + AND state NOT IN ('cancelled', 'completed', 'discarded') + AND finalized_at IS NULL +` + +type JobCancelExecParams struct { + Now sql.NullTime + CancelAttemptedAt interface{} + ID int64 +} + +func (q *Queries) JobCancelExec(ctx context.Context, db DBTX, arg *JobCancelExecParams) (sql.Result, error) { + return db.ExecContext(ctx, jobCancelExec, arg.Now, arg.CancelAttemptedAt, arg.ID) +} + +const jobClearUniqueNonce = `-- name: JobClearUniqueNonce :exec +UPDATE /* TEMPLATE: schema */river_job +SET metadata = JSON_REMOVE(metadata, '$."river:unique_nonce"') +WHERE id IN (/*SLICE:id*/?) +` + +func (q *Queries) JobClearUniqueNonce(ctx context.Context, db DBTX, id []int64) error { + query := jobClearUniqueNonce + var queryParams []interface{} + if len(id) > 0 { + for _, v := range id { + queryParams = append(queryParams, v) + } + query = strings.Replace(query, "/*SLICE:id*/?", strings.Repeat(",?", len(id))[1:], 1) + } else { + query = strings.Replace(query, "/*SLICE:id*/?", "NULL", 1) + } + _, err := db.ExecContext(ctx, query, queryParams...) + return err +} + +const jobCountByAllStates = `-- name: JobCountByAllStates :many +SELECT state, count(*) AS count +FROM /* TEMPLATE: schema */river_job +GROUP BY state +` + +type JobCountByAllStatesRow struct { + State string + Count int64 +} + +func (q *Queries) JobCountByAllStates(ctx context.Context, db DBTX) ([]*JobCountByAllStatesRow, error) { + rows, err := db.QueryContext(ctx, jobCountByAllStates) + if err != nil { + return nil, err + } + defer rows.Close() + var items []*JobCountByAllStatesRow + for rows.Next() { + var i JobCountByAllStatesRow + if err := rows.Scan(&i.State, &i.Count); err != nil { + return nil, err + } + items = append(items, &i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const jobCountByQueueAndState = `-- name: JobCountByQueueAndState :many +WITH all_queues AS ( + SELECT DISTINCT river_job.queue + FROM /* TEMPLATE: schema */river_job + WHERE river_job.queue IN (/*SLICE:queue_names*/?) +), + +running_job_counts AS ( + SELECT river_job.queue, COUNT(*) AS count + FROM /* TEMPLATE: schema */river_job + WHERE river_job.queue IN (/*SLICE:queue_names*/?) + AND river_job.state = 'running' + GROUP BY river_job.queue +), + +available_job_counts AS ( + SELECT river_job.queue, COUNT(*) AS count + FROM /* TEMPLATE: schema */river_job + WHERE river_job.queue IN (/*SLICE:queue_names*/?) + AND river_job.state = 'available' + GROUP BY river_job.queue +) + +SELECT + all_queues.queue, + COALESCE(available_job_counts.count, 0) AS count_available, + COALESCE(running_job_counts.count, 0) AS count_running +FROM all_queues +LEFT JOIN running_job_counts ON all_queues.queue = running_job_counts.queue +LEFT JOIN available_job_counts ON all_queues.queue = available_job_counts.queue +ORDER BY all_queues.queue ASC +` + +type JobCountByQueueAndStateParams struct { + QueueNames []string +} + +type JobCountByQueueAndStateRow struct { + Queue string + CountAvailable int64 + CountRunning int64 +} + +func (q *Queries) JobCountByQueueAndState(ctx context.Context, db DBTX, arg *JobCountByQueueAndStateParams) ([]*JobCountByQueueAndStateRow, error) { + query := jobCountByQueueAndState + var queryParams []interface{} + if len(arg.QueueNames) > 0 { + for _, v := range arg.QueueNames { + queryParams = append(queryParams, v) + } + query = strings.Replace(query, "/*SLICE:queue_names*/?", strings.Repeat(",?", len(arg.QueueNames))[1:], 1) + } else { + query = strings.Replace(query, "/*SLICE:queue_names*/?", "NULL", 1) + } + if len(arg.QueueNames) > 0 { + for _, v := range arg.QueueNames { + queryParams = append(queryParams, v) + } + query = strings.Replace(query, "/*SLICE:queue_names*/?", strings.Repeat(",?", len(arg.QueueNames))[1:], 1) + } else { + query = strings.Replace(query, "/*SLICE:queue_names*/?", "NULL", 1) + } + if len(arg.QueueNames) > 0 { + for _, v := range arg.QueueNames { + queryParams = append(queryParams, v) + } + query = strings.Replace(query, "/*SLICE:queue_names*/?", strings.Repeat(",?", len(arg.QueueNames))[1:], 1) + } else { + query = strings.Replace(query, "/*SLICE:queue_names*/?", "NULL", 1) + } + rows, err := db.QueryContext(ctx, query, queryParams...) + if err != nil { + return nil, err + } + defer rows.Close() + var items []*JobCountByQueueAndStateRow + for rows.Next() { + var i JobCountByQueueAndStateRow + if err := rows.Scan(&i.Queue, &i.CountAvailable, &i.CountRunning); err != nil { + return nil, err + } + items = append(items, &i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const jobCountByState = `-- name: JobCountByState :one +SELECT count(*) AS count +FROM /* TEMPLATE: schema */river_job +WHERE state = ? +` + +func (q *Queries) JobCountByState(ctx context.Context, db DBTX, state string) (int64, error) { + row := db.QueryRowContext(ctx, jobCountByState, state) + var count int64 + err := row.Scan(&count) + return count, err +} + +const jobDeleteBefore = `-- name: JobDeleteBefore :execresult +DELETE FROM /* TEMPLATE: schema */river_job +WHERE + id IN ( + SELECT id FROM ( + SELECT rj2.id + FROM /* TEMPLATE: schema */river_job rj2 + WHERE ( + ( + CAST(? AS UNSIGNED) + AND rj2.state = 'cancelled' + AND rj2.finalized_at < ? + ) OR ( + CAST(? AS UNSIGNED) + AND rj2.state = 'completed' + AND rj2.finalized_at < ? + ) OR ( + CAST(? AS UNSIGNED) + AND rj2.state = 'discarded' + AND rj2.finalized_at < ? + ) + ) + AND ( + CAST(? AS SIGNED) + OR rj2.queue NOT IN (/*SLICE:queues_excluded*/?) + ) + AND ( + CAST(? AS SIGNED) + OR rj2.queue IN (/*SLICE:queues_included*/?) + ) + ORDER BY rj2.id + LIMIT ? + ) AS tmp + ) +` + +type JobDeleteBeforeParams struct { + CancelledDoDelete int64 + CancelledFinalizedAtHorizon sql.NullTime + CompletedDoDelete int64 + CompletedFinalizedAtHorizon sql.NullTime + DiscardedDoDelete int64 + DiscardedFinalizedAtHorizon sql.NullTime + QueuesExcludedEmpty int64 + QueuesExcluded []string + QueuesIncludedEmpty int64 + QueuesIncluded []string + Limit int32 +} + +func (q *Queries) JobDeleteBefore(ctx context.Context, db DBTX, arg *JobDeleteBeforeParams) (sql.Result, error) { + query := jobDeleteBefore + var queryParams []interface{} + queryParams = append(queryParams, arg.CancelledDoDelete) + queryParams = append(queryParams, arg.CancelledFinalizedAtHorizon) + queryParams = append(queryParams, arg.CompletedDoDelete) + queryParams = append(queryParams, arg.CompletedFinalizedAtHorizon) + queryParams = append(queryParams, arg.DiscardedDoDelete) + queryParams = append(queryParams, arg.DiscardedFinalizedAtHorizon) + queryParams = append(queryParams, arg.QueuesExcludedEmpty) + if len(arg.QueuesExcluded) > 0 { + for _, v := range arg.QueuesExcluded { + queryParams = append(queryParams, v) + } + query = strings.Replace(query, "/*SLICE:queues_excluded*/?", strings.Repeat(",?", len(arg.QueuesExcluded))[1:], 1) + } else { + query = strings.Replace(query, "/*SLICE:queues_excluded*/?", "NULL", 1) + } + queryParams = append(queryParams, arg.QueuesIncludedEmpty) + if len(arg.QueuesIncluded) > 0 { + for _, v := range arg.QueuesIncluded { + queryParams = append(queryParams, v) + } + query = strings.Replace(query, "/*SLICE:queues_included*/?", strings.Repeat(",?", len(arg.QueuesIncluded))[1:], 1) + } else { + query = strings.Replace(query, "/*SLICE:queues_included*/?", "NULL", 1) + } + queryParams = append(queryParams, arg.Limit) + return db.ExecContext(ctx, query, queryParams...) +} + +const jobDeleteExec = `-- name: JobDeleteExec :execresult +DELETE FROM /* TEMPLATE: schema */river_job +WHERE id = ? + AND river_job.state != 'running' +` + +func (q *Queries) JobDeleteExec(ctx context.Context, db DBTX, id int64) (sql.Result, error) { + return db.ExecContext(ctx, jobDeleteExec, id) +} + +const jobDeleteManyExec = `-- name: JobDeleteManyExec :exec +DELETE FROM /* TEMPLATE: schema */river_job +WHERE id IN (/*SLICE:id*/?) +` + +func (q *Queries) JobDeleteManyExec(ctx context.Context, db DBTX, id []int64) error { + query := jobDeleteManyExec + var queryParams []interface{} + if len(id) > 0 { + for _, v := range id { + queryParams = append(queryParams, v) + } + query = strings.Replace(query, "/*SLICE:id*/?", strings.Repeat(",?", len(id))[1:], 1) + } else { + query = strings.Replace(query, "/*SLICE:id*/?", "NULL", 1) + } + _, err := db.ExecContext(ctx, query, queryParams...) + return err +} + +const jobDeleteManySelect = `-- name: JobDeleteManySelect :many +SELECT id, args, attempt, attempted_at, attempted_by, created_at, errors, finalized_at, kind, max_attempts, metadata, priority, queue, state, scheduled_at, tags, unique_key, unique_states +FROM /* TEMPLATE: schema */river_job +WHERE id IN ( + SELECT id FROM ( + SELECT id + FROM /* TEMPLATE: schema */river_job + WHERE /* TEMPLATE_BEGIN: where_clause */ true /* TEMPLATE_END */ + AND state != 'running' + ORDER BY /* TEMPLATE_BEGIN: order_by_clause */ id /* TEMPLATE_END */ + LIMIT ? + ) AS tmp +) +ORDER BY /* TEMPLATE_BEGIN: order_by_clause */ id /* TEMPLATE_END */ +FOR UPDATE +` + +func (q *Queries) JobDeleteManySelect(ctx context.Context, db DBTX, limit int32) ([]*RiverJob, error) { + rows, err := db.QueryContext(ctx, jobDeleteManySelect, limit) + if err != nil { + return nil, err + } + defer rows.Close() + var items []*RiverJob + for rows.Next() { + var i RiverJob + if err := rows.Scan( + &i.ID, + &i.Args, + &i.Attempt, + &i.AttemptedAt, + &i.AttemptedBy, + &i.CreatedAt, + &i.Errors, + &i.FinalizedAt, + &i.Kind, + &i.MaxAttempts, + &i.Metadata, + &i.Priority, + &i.Queue, + &i.State, + &i.ScheduledAt, + &i.Tags, + &i.UniqueKey, + &i.UniqueStates, + ); err != nil { + return nil, err + } + items = append(items, &i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const jobDeleteSelect = `-- name: JobDeleteSelect :one +SELECT id, args, attempt, attempted_at, attempted_by, created_at, errors, finalized_at, kind, max_attempts, metadata, priority, queue, state, scheduled_at, tags, unique_key, unique_states +FROM /* TEMPLATE: schema */river_job +WHERE id = ? +LIMIT 1 +FOR UPDATE +` + +func (q *Queries) JobDeleteSelect(ctx context.Context, db DBTX, id int64) (*RiverJob, error) { + row := db.QueryRowContext(ctx, jobDeleteSelect, id) + var i RiverJob + err := row.Scan( + &i.ID, + &i.Args, + &i.Attempt, + &i.AttemptedAt, + &i.AttemptedBy, + &i.CreatedAt, + &i.Errors, + &i.FinalizedAt, + &i.Kind, + &i.MaxAttempts, + &i.Metadata, + &i.Priority, + &i.Queue, + &i.State, + &i.ScheduledAt, + &i.Tags, + &i.UniqueKey, + &i.UniqueStates, + ) + return &i, err +} + +const jobGetAvailableIDs = `-- name: JobGetAvailableIDs :many +SELECT id +FROM /* TEMPLATE: schema */river_job +WHERE + priority >= 0 + AND queue = ? + AND scheduled_at <= COALESCE(?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) + AND state = 'available' +ORDER BY priority ASC, scheduled_at ASC, id ASC +LIMIT ? +FOR UPDATE SKIP LOCKED +` + +type JobGetAvailableIDsParams struct { + Queue string + Now sql.NullTime + Limit int32 +} + +func (q *Queries) JobGetAvailableIDs(ctx context.Context, db DBTX, arg *JobGetAvailableIDsParams) ([]int64, error) { + rows, err := db.QueryContext(ctx, jobGetAvailableIDs, arg.Queue, arg.Now, arg.Limit) + if err != nil { + return nil, err + } + defer rows.Close() + var items []int64 + for rows.Next() { + var id int64 + if err := rows.Scan(&id); err != nil { + return nil, err + } + items = append(items, id) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const jobGetAvailableUpdate = `-- name: JobGetAvailableUpdate :exec +UPDATE /* TEMPLATE: schema */river_job +SET + attempt = attempt + 1, + attempted_at = COALESCE(?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')), + attempted_by = JSON_ARRAY_APPEND( + CASE + WHEN JSON_LENGTH(COALESCE(attempted_by, JSON_ARRAY())) < CAST(? AS SIGNED) + THEN COALESCE(attempted_by, JSON_ARRAY()) + WHEN CAST(? AS SIGNED) <= 1 + THEN JSON_ARRAY() + ELSE COALESCE( + JSON_EXTRACT( + attempted_by, + CONCAT('$[last-', CAST(? AS SIGNED) - 2, ' to last]') + ), + JSON_ARRAY() + ) + END, + '$', + CAST(? AS CHAR) + ), + state = 'running' +WHERE id IN (/*SLICE:id*/?) +` + +type JobGetAvailableUpdateParams struct { + Now sql.NullTime + MaxAttemptedBy int64 + AttemptedBy interface{} + ID []int64 +} + +func (q *Queries) JobGetAvailableUpdate(ctx context.Context, db DBTX, arg *JobGetAvailableUpdateParams) error { + query := jobGetAvailableUpdate + var queryParams []interface{} + queryParams = append(queryParams, arg.Now) + queryParams = append(queryParams, arg.MaxAttemptedBy) + queryParams = append(queryParams, arg.MaxAttemptedBy) + queryParams = append(queryParams, arg.MaxAttemptedBy) + queryParams = append(queryParams, arg.AttemptedBy) + if len(arg.ID) > 0 { + for _, v := range arg.ID { + queryParams = append(queryParams, v) + } + query = strings.Replace(query, "/*SLICE:id*/?", strings.Repeat(",?", len(arg.ID))[1:], 1) + } else { + query = strings.Replace(query, "/*SLICE:id*/?", "NULL", 1) + } + _, err := db.ExecContext(ctx, query, queryParams...) + return err +} + +const jobGetByID = `-- name: JobGetByID :one +SELECT id, args, attempt, attempted_at, attempted_by, created_at, errors, finalized_at, kind, max_attempts, metadata, priority, queue, state, scheduled_at, tags, unique_key, unique_states +FROM /* TEMPLATE: schema */river_job +WHERE id = ? +LIMIT 1 +` + +func (q *Queries) JobGetByID(ctx context.Context, db DBTX, id int64) (*RiverJob, error) { + row := db.QueryRowContext(ctx, jobGetByID, id) + var i RiverJob + err := row.Scan( + &i.ID, + &i.Args, + &i.Attempt, + &i.AttemptedAt, + &i.AttemptedBy, + &i.CreatedAt, + &i.Errors, + &i.FinalizedAt, + &i.Kind, + &i.MaxAttempts, + &i.Metadata, + &i.Priority, + &i.Queue, + &i.State, + &i.ScheduledAt, + &i.Tags, + &i.UniqueKey, + &i.UniqueStates, + ) + return &i, err +} + +const jobGetByIDMany = `-- name: JobGetByIDMany :many +SELECT id, args, attempt, attempted_at, attempted_by, created_at, errors, finalized_at, kind, max_attempts, metadata, priority, queue, state, scheduled_at, tags, unique_key, unique_states +FROM /* TEMPLATE: schema */river_job +WHERE id IN (/*SLICE:id*/?) +ORDER BY id +` + +func (q *Queries) JobGetByIDMany(ctx context.Context, db DBTX, id []int64) ([]*RiverJob, error) { + query := jobGetByIDMany + var queryParams []interface{} + if len(id) > 0 { + for _, v := range id { + queryParams = append(queryParams, v) + } + query = strings.Replace(query, "/*SLICE:id*/?", strings.Repeat(",?", len(id))[1:], 1) + } else { + query = strings.Replace(query, "/*SLICE:id*/?", "NULL", 1) + } + rows, err := db.QueryContext(ctx, query, queryParams...) + if err != nil { + return nil, err + } + defer rows.Close() + var items []*RiverJob + for rows.Next() { + var i RiverJob + if err := rows.Scan( + &i.ID, + &i.Args, + &i.Attempt, + &i.AttemptedAt, + &i.AttemptedBy, + &i.CreatedAt, + &i.Errors, + &i.FinalizedAt, + &i.Kind, + &i.MaxAttempts, + &i.Metadata, + &i.Priority, + &i.Queue, + &i.State, + &i.ScheduledAt, + &i.Tags, + &i.UniqueKey, + &i.UniqueStates, + ); err != nil { + return nil, err + } + items = append(items, &i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const jobGetByIDManyOrdered = `-- name: JobGetByIDManyOrdered :many +SELECT id, args, attempt, attempted_at, attempted_by, created_at, errors, finalized_at, kind, max_attempts, metadata, priority, queue, state, scheduled_at, tags, unique_key, unique_states +FROM /* TEMPLATE: schema */river_job +WHERE id IN (/*SLICE:id*/?) +ORDER BY priority ASC, scheduled_at ASC, id ASC +` + +func (q *Queries) JobGetByIDManyOrdered(ctx context.Context, db DBTX, id []int64) ([]*RiverJob, error) { + query := jobGetByIDManyOrdered + var queryParams []interface{} + if len(id) > 0 { + for _, v := range id { + queryParams = append(queryParams, v) + } + query = strings.Replace(query, "/*SLICE:id*/?", strings.Repeat(",?", len(id))[1:], 1) + } else { + query = strings.Replace(query, "/*SLICE:id*/?", "NULL", 1) + } + rows, err := db.QueryContext(ctx, query, queryParams...) + if err != nil { + return nil, err + } + defer rows.Close() + var items []*RiverJob + for rows.Next() { + var i RiverJob + if err := rows.Scan( + &i.ID, + &i.Args, + &i.Attempt, + &i.AttemptedAt, + &i.AttemptedBy, + &i.CreatedAt, + &i.Errors, + &i.FinalizedAt, + &i.Kind, + &i.MaxAttempts, + &i.Metadata, + &i.Priority, + &i.Queue, + &i.State, + &i.ScheduledAt, + &i.Tags, + &i.UniqueKey, + &i.UniqueStates, + ); err != nil { + return nil, err + } + items = append(items, &i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const jobGetByKindMany = `-- name: JobGetByKindMany :many +SELECT id, args, attempt, attempted_at, attempted_by, created_at, errors, finalized_at, kind, max_attempts, metadata, priority, queue, state, scheduled_at, tags, unique_key, unique_states +FROM /* TEMPLATE: schema */river_job +WHERE kind IN (/*SLICE:kind*/?) +ORDER BY id +` + +func (q *Queries) JobGetByKindMany(ctx context.Context, db DBTX, kind []string) ([]*RiverJob, error) { + query := jobGetByKindMany + var queryParams []interface{} + if len(kind) > 0 { + for _, v := range kind { + queryParams = append(queryParams, v) + } + query = strings.Replace(query, "/*SLICE:kind*/?", strings.Repeat(",?", len(kind))[1:], 1) + } else { + query = strings.Replace(query, "/*SLICE:kind*/?", "NULL", 1) + } + rows, err := db.QueryContext(ctx, query, queryParams...) + if err != nil { + return nil, err + } + defer rows.Close() + var items []*RiverJob + for rows.Next() { + var i RiverJob + if err := rows.Scan( + &i.ID, + &i.Args, + &i.Attempt, + &i.AttemptedAt, + &i.AttemptedBy, + &i.CreatedAt, + &i.Errors, + &i.FinalizedAt, + &i.Kind, + &i.MaxAttempts, + &i.Metadata, + &i.Priority, + &i.Queue, + &i.State, + &i.ScheduledAt, + &i.Tags, + &i.UniqueKey, + &i.UniqueStates, + ); err != nil { + return nil, err + } + items = append(items, &i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const jobGetStuck = `-- name: JobGetStuck :many +SELECT id, args, attempt, attempted_at, attempted_by, created_at, errors, finalized_at, kind, max_attempts, metadata, priority, queue, state, scheduled_at, tags, unique_key, unique_states +FROM /* TEMPLATE: schema */river_job +WHERE state = 'running' + AND id > ? + AND attempted_at < ? +ORDER BY id +LIMIT ? +` + +type JobGetStuckParams struct { + AfterID int64 + StuckHorizon sql.NullTime + Limit int32 +} + +func (q *Queries) JobGetStuck(ctx context.Context, db DBTX, arg *JobGetStuckParams) ([]*RiverJob, error) { + rows, err := db.QueryContext(ctx, jobGetStuck, arg.AfterID, arg.StuckHorizon, arg.Limit) + if err != nil { + return nil, err + } + defer rows.Close() + var items []*RiverJob + for rows.Next() { + var i RiverJob + if err := rows.Scan( + &i.ID, + &i.Args, + &i.Attempt, + &i.AttemptedAt, + &i.AttemptedBy, + &i.CreatedAt, + &i.Errors, + &i.FinalizedAt, + &i.Kind, + &i.MaxAttempts, + &i.Metadata, + &i.Priority, + &i.Queue, + &i.State, + &i.ScheduledAt, + &i.Tags, + &i.UniqueKey, + &i.UniqueStates, + ); err != nil { + return nil, err + } + items = append(items, &i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const jobInsertFast = `-- name: JobInsertFast :execresult +INSERT INTO /* TEMPLATE: schema */river_job( + id, + args, + created_at, + kind, + max_attempts, + metadata, + priority, + queue, + scheduled_at, + state, + tags, + unique_key, + unique_states +) VALUES ( + ?, + ?, + COALESCE(?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')), + ?, + ?, + CAST(? AS JSON), + ?, + ?, + COALESCE(?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')), + ?, + CAST(? AS JSON), + ?, + ? +) /* TEMPLATE_BEGIN: on_duplicate_key */ ON DUPLICATE KEY UPDATE id = LAST_INSERT_ID(id) /* TEMPLATE_END */ +` + +type JobInsertFastParams struct { + ID sql.NullInt64 + Args []byte + CreatedAt interface{} + Kind string + MaxAttempts int64 + Metadata []byte + Priority int16 + Queue string + ScheduledAt interface{} + State string + Tags []byte + UniqueKey sql.NullString + UniqueStates sql.NullInt16 +} + +func (q *Queries) JobInsertFast(ctx context.Context, db DBTX, arg *JobInsertFastParams) (sql.Result, error) { + return db.ExecContext(ctx, jobInsertFast, + arg.ID, + arg.Args, + arg.CreatedAt, + arg.Kind, + arg.MaxAttempts, + arg.Metadata, + arg.Priority, + arg.Queue, + arg.ScheduledAt, + arg.State, + arg.Tags, + arg.UniqueKey, + arg.UniqueStates, + ) +} + +const jobInsertFullExec = `-- name: JobInsertFullExec :execlastid +INSERT INTO /* TEMPLATE: schema */river_job( + args, + attempt, + attempted_at, + attempted_by, + created_at, + errors, + finalized_at, + kind, + max_attempts, + metadata, + priority, + queue, + scheduled_at, + state, + tags, + unique_key, + unique_states +) VALUES ( + ?, + ?, + ?, + CAST(? AS JSON), + COALESCE(?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')), + CAST(? AS JSON), + ?, + ?, + ?, + CAST(? AS JSON), + ?, + ?, + COALESCE(?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')), + ?, + CAST(? AS JSON), + ?, + ? +) +` + +type JobInsertFullExecParams struct { + Args []byte + Attempt int64 + AttemptedAt sql.NullTime + AttemptedBy []byte + CreatedAt interface{} + Errors []byte + FinalizedAt sql.NullTime + Kind string + MaxAttempts int64 + Metadata []byte + Priority int16 + Queue string + ScheduledAt interface{} + State string + Tags []byte + UniqueKey sql.NullString + UniqueStates sql.NullInt16 +} + +func (q *Queries) JobInsertFullExec(ctx context.Context, db DBTX, arg *JobInsertFullExecParams) (int64, error) { + result, err := db.ExecContext(ctx, jobInsertFullExec, + arg.Args, + arg.Attempt, + arg.AttemptedAt, + arg.AttemptedBy, + arg.CreatedAt, + arg.Errors, + arg.FinalizedAt, + arg.Kind, + arg.MaxAttempts, + arg.Metadata, + arg.Priority, + arg.Queue, + arg.ScheduledAt, + arg.State, + arg.Tags, + arg.UniqueKey, + arg.UniqueStates, + ) + if err != nil { + return 0, err + } + return result.LastInsertId() +} + +const jobKindList = `-- name: JobKindList :many +SELECT DISTINCT kind +FROM /* TEMPLATE: schema */river_job +WHERE (? = '' OR LOWER(kind) COLLATE utf8mb4_general_ci LIKE CONCAT('%', LOWER(?), '%')) + AND (? = '' OR kind > ?) + AND kind NOT IN (/*SLICE:exclude*/?) +ORDER BY kind ASC +LIMIT ? +` + +type JobKindListParams struct { + Match string + After string + Exclude []string + Limit int32 +} + +func (q *Queries) JobKindList(ctx context.Context, db DBTX, arg *JobKindListParams) ([]string, error) { + query := jobKindList + var queryParams []interface{} + queryParams = append(queryParams, arg.Match) + queryParams = append(queryParams, arg.Match) + queryParams = append(queryParams, arg.After) + queryParams = append(queryParams, arg.After) + if len(arg.Exclude) > 0 { + for _, v := range arg.Exclude { + queryParams = append(queryParams, v) + } + query = strings.Replace(query, "/*SLICE:exclude*/?", strings.Repeat(",?", len(arg.Exclude))[1:], 1) + } else { + query = strings.Replace(query, "/*SLICE:exclude*/?", "NULL", 1) + } + queryParams = append(queryParams, arg.Limit) + rows, err := db.QueryContext(ctx, query, queryParams...) + if err != nil { + return nil, err + } + defer rows.Close() + var items []string + for rows.Next() { + var kind string + if err := rows.Scan(&kind); err != nil { + return nil, err + } + items = append(items, kind) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const jobList = `-- name: JobList :many +SELECT id, args, attempt, attempted_at, attempted_by, created_at, errors, finalized_at, kind, max_attempts, metadata, priority, queue, state, scheduled_at, tags, unique_key, unique_states +FROM /* TEMPLATE: schema */river_job +WHERE /* TEMPLATE_BEGIN: where_clause */ true /* TEMPLATE_END */ +ORDER BY /* TEMPLATE_BEGIN: order_by_clause */ id /* TEMPLATE_END */ +LIMIT ? +` + +func (q *Queries) JobList(ctx context.Context, db DBTX, limit int32) ([]*RiverJob, error) { + rows, err := db.QueryContext(ctx, jobList, limit) + if err != nil { + return nil, err + } + defer rows.Close() + var items []*RiverJob + for rows.Next() { + var i RiverJob + if err := rows.Scan( + &i.ID, + &i.Args, + &i.Attempt, + &i.AttemptedAt, + &i.AttemptedBy, + &i.CreatedAt, + &i.Errors, + &i.FinalizedAt, + &i.Kind, + &i.MaxAttempts, + &i.Metadata, + &i.Priority, + &i.Queue, + &i.State, + &i.ScheduledAt, + &i.Tags, + &i.UniqueKey, + &i.UniqueStates, + ); err != nil { + return nil, err + } + items = append(items, &i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const jobRescue = `-- name: JobRescue :exec +UPDATE /* TEMPLATE: schema */river_job +SET + errors = JSON_ARRAY_APPEND(COALESCE(errors, JSON_ARRAY()), '$', CAST(? AS JSON)), + finalized_at = ?, + scheduled_at = ?, + metadata = JSON_SET( + metadata, + '$."river:rescue_count"', + COALESCE( + CASE JSON_TYPE(JSON_EXTRACT(metadata, '$."river:rescue_count"')) + WHEN 'INTEGER' THEN JSON_EXTRACT(metadata, '$."river:rescue_count"') + WHEN 'DOUBLE' THEN JSON_EXTRACT(metadata, '$."river:rescue_count"') + ELSE NULL + END, + 0 + ) + 1 + ), + state = ? +WHERE id = ? +` + +type JobRescueParams struct { + Error []byte + FinalizedAt sql.NullTime + ScheduledAt time.Time + State string + ID int64 +} + +func (q *Queries) JobRescue(ctx context.Context, db DBTX, arg *JobRescueParams) error { + _, err := db.ExecContext(ctx, jobRescue, + arg.Error, + arg.FinalizedAt, + arg.ScheduledAt, + arg.State, + arg.ID, + ) + return err +} + +const jobRetryExec = `-- name: JobRetryExec :execresult +UPDATE /* TEMPLATE: schema */river_job +SET + state = 'available', + max_attempts = CASE WHEN attempt = max_attempts THEN max_attempts + 1 ELSE max_attempts END, + finalized_at = NULL, + scheduled_at = COALESCE(?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) +WHERE id = ? + AND state != 'running' + AND ( + state <> 'available' + OR scheduled_at > COALESCE(?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) + ) +` + +type JobRetryExecParams struct { + Now sql.NullTime + ID int64 +} + +func (q *Queries) JobRetryExec(ctx context.Context, db DBTX, arg *JobRetryExecParams) (sql.Result, error) { + return db.ExecContext(ctx, jobRetryExec, arg.Now, arg.ID, arg.Now) +} + +const jobSchedule = `-- name: JobSchedule :many +WITH eligible AS ( + SELECT river_job.id, river_job.unique_key, river_job.unique_states, river_job.priority, river_job.scheduled_at, + CASE + WHEN river_job.unique_key IS NOT NULL AND river_job.unique_states IS NOT NULL THEN + ROW_NUMBER() OVER (PARTITION BY river_job.unique_key ORDER BY river_job.priority, river_job.scheduled_at, river_job.id) + ELSE NULL + END AS row_num + FROM /* TEMPLATE: schema */river_job + WHERE + river_job.state IN ('retryable', 'scheduled') + AND river_job.scheduled_at <= COALESCE(?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) + ORDER BY + river_job.priority, + river_job.scheduled_at, + river_job.id + LIMIT ? + FOR UPDATE SKIP LOCKED +), +unique_conflicts AS ( + SELECT DISTINCT eligible.unique_key + FROM /* TEMPLATE: schema */river_job + JOIN eligible + ON river_job.unique_key = eligible.unique_key + AND river_job.id != eligible.id + WHERE + river_job.unique_key IS NOT NULL + AND river_job.unique_states IS NOT NULL + AND CASE river_job.state + WHEN 'available' THEN river_job.unique_states & (1 << 0) + WHEN 'cancelled' THEN river_job.unique_states & (1 << 1) + WHEN 'completed' THEN river_job.unique_states & (1 << 2) + WHEN 'discarded' THEN river_job.unique_states & (1 << 3) + WHEN 'pending' THEN river_job.unique_states & (1 << 4) + WHEN 'retryable' THEN river_job.unique_states & (1 << 5) + WHEN 'running' THEN river_job.unique_states & (1 << 6) + WHEN 'scheduled' THEN river_job.unique_states & (1 << 7) + ELSE 0 + END >= 1 +) +SELECT eligible.id, + CASE + WHEN eligible.unique_key IS NULL OR eligible.unique_states IS NULL THEN FALSE + WHEN uc.unique_key IS NOT NULL THEN TRUE + WHEN eligible.row_num > 1 THEN TRUE + ELSE FALSE + END AS conflict_discarded +FROM eligible +LEFT JOIN unique_conflicts uc ON eligible.unique_key = uc.unique_key +ORDER BY eligible.priority, eligible.scheduled_at, eligible.id +` + +type JobScheduleParams struct { + Now sql.NullTime + Limit int32 +} + +type JobScheduleRow struct { + ID int64 + ConflictDiscarded int64 +} + +func (q *Queries) JobSchedule(ctx context.Context, db DBTX, arg *JobScheduleParams) ([]*JobScheduleRow, error) { + rows, err := db.QueryContext(ctx, jobSchedule, arg.Now, arg.Limit) + if err != nil { + return nil, err + } + defer rows.Close() + var items []*JobScheduleRow + for rows.Next() { + var i JobScheduleRow + if err := rows.Scan(&i.ID, &i.ConflictDiscarded); err != nil { + return nil, err + } + items = append(items, &i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const jobScheduleSetAvailableExec = `-- name: JobScheduleSetAvailableExec :exec +UPDATE /* TEMPLATE: schema */river_job +SET state = 'available' +WHERE id IN (/*SLICE:id*/?) +` + +func (q *Queries) JobScheduleSetAvailableExec(ctx context.Context, db DBTX, id []int64) error { + query := jobScheduleSetAvailableExec + var queryParams []interface{} + if len(id) > 0 { + for _, v := range id { + queryParams = append(queryParams, v) + } + query = strings.Replace(query, "/*SLICE:id*/?", strings.Repeat(",?", len(id))[1:], 1) + } else { + query = strings.Replace(query, "/*SLICE:id*/?", "NULL", 1) + } + _, err := db.ExecContext(ctx, query, queryParams...) + return err +} + +const jobScheduleSetDiscardedExec = `-- name: JobScheduleSetDiscardedExec :exec +UPDATE /* TEMPLATE: schema */river_job +SET metadata = JSON_MERGE_PATCH(metadata, '{"unique_key_conflict": "scheduler_discarded"}'), + finalized_at = COALESCE(?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')), + state = 'discarded' +WHERE id IN (/*SLICE:id*/?) +` + +type JobScheduleSetDiscardedExecParams struct { + Now sql.NullTime + ID []int64 +} + +func (q *Queries) JobScheduleSetDiscardedExec(ctx context.Context, db DBTX, arg *JobScheduleSetDiscardedExecParams) error { + query := jobScheduleSetDiscardedExec + var queryParams []interface{} + queryParams = append(queryParams, arg.Now) + if len(arg.ID) > 0 { + for _, v := range arg.ID { + queryParams = append(queryParams, v) + } + query = strings.Replace(query, "/*SLICE:id*/?", strings.Repeat(",?", len(arg.ID))[1:], 1) + } else { + query = strings.Replace(query, "/*SLICE:id*/?", "NULL", 1) + } + _, err := db.ExecContext(ctx, query, queryParams...) + return err +} + +const jobSetMetadataIfNotRunningExec = `-- name: JobSetMetadataIfNotRunningExec :execresult +UPDATE /* TEMPLATE: schema */river_job +SET metadata = JSON_MERGE_PATCH(metadata, CAST(? AS JSON)) +WHERE id = ? + AND state != 'running' +` + +type JobSetMetadataIfNotRunningExecParams struct { + MetadataUpdates []byte + ID int64 +} + +func (q *Queries) JobSetMetadataIfNotRunningExec(ctx context.Context, db DBTX, arg *JobSetMetadataIfNotRunningExecParams) (sql.Result, error) { + return db.ExecContext(ctx, jobSetMetadataIfNotRunningExec, arg.MetadataUpdates, arg.ID) +} + +const jobSetStateIfRunningExec = `-- name: JobSetStateIfRunningExec :exec +UPDATE /* TEMPLATE: schema */river_job +SET + attempt = CASE WHEN (CAST(? AS CHAR) <> 'retryable' AND ? <> 'scheduled' OR JSON_EXTRACT(metadata, '$.cancel_attempted_at') IS NULL) AND CAST(? AS SIGNED) + THEN ? + ELSE attempt END, + errors = CASE WHEN CAST(? AS SIGNED) + THEN JSON_ARRAY_APPEND(COALESCE(errors, JSON_ARRAY()), '$', CAST(? AS JSON)) + ELSE errors END, + finalized_at = CASE WHEN ((? = 'retryable' OR ? = 'scheduled') AND JSON_EXTRACT(metadata, '$.cancel_attempted_at') IS NOT NULL) + THEN COALESCE(?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) + WHEN CAST(? AS SIGNED) + THEN ? + ELSE finalized_at END, + metadata = CASE WHEN CAST(? AS SIGNED) + THEN JSON_MERGE_PATCH(metadata, CAST(? AS JSON)) + ELSE metadata END, + scheduled_at = CASE WHEN (CAST(? AS CHAR) <> 'retryable' AND ? <> 'scheduled' OR JSON_EXTRACT(metadata, '$.cancel_attempted_at') IS NULL) AND CAST(? AS SIGNED) + THEN ? + ELSE scheduled_at END, + state = CASE WHEN ((? = 'retryable' OR ? = 'scheduled') AND JSON_EXTRACT(metadata, '$.cancel_attempted_at') IS NOT NULL) + THEN 'cancelled' + ELSE ? END +WHERE id = ? + AND state = 'running' +` + +type JobSetStateIfRunningExecParams struct { + State string + AttemptDoUpdate int64 + Attempt int64 + ErrorsDoUpdate int64 + Error []byte + Now sql.NullTime + FinalizedAtDoUpdate int64 + FinalizedAt sql.NullTime + MetadataDoMerge int64 + MetadataUpdates []byte + ScheduledAtDoUpdate int64 + ScheduledAt time.Time + ID int64 +} + +func (q *Queries) JobSetStateIfRunningExec(ctx context.Context, db DBTX, arg *JobSetStateIfRunningExecParams) error { + _, err := db.ExecContext(ctx, jobSetStateIfRunningExec, + arg.State, + arg.State, + arg.AttemptDoUpdate, + arg.Attempt, + arg.ErrorsDoUpdate, + arg.Error, + arg.State, + arg.State, + arg.Now, + arg.FinalizedAtDoUpdate, + arg.FinalizedAt, + arg.MetadataDoMerge, + arg.MetadataUpdates, + arg.State, + arg.State, + arg.ScheduledAtDoUpdate, + arg.ScheduledAt, + arg.State, + arg.State, + arg.State, + arg.ID, + ) + return err +} + +const jobUpdateExec = `-- name: JobUpdateExec :exec +UPDATE /* TEMPLATE: schema */river_job +SET + metadata = CASE WHEN CAST(? AS SIGNED) THEN JSON_MERGE_PATCH(metadata, CAST(? AS JSON)) ELSE metadata END +WHERE id = ? +` + +type JobUpdateExecParams struct { + MetadataDoMerge int64 + Metadata []byte + ID int64 +} + +func (q *Queries) JobUpdateExec(ctx context.Context, db DBTX, arg *JobUpdateExecParams) error { + _, err := db.ExecContext(ctx, jobUpdateExec, arg.MetadataDoMerge, arg.Metadata, arg.ID) + return err +} + +const jobUpdateFullExec = `-- name: JobUpdateFullExec :exec +UPDATE /* TEMPLATE: schema */river_job +SET + attempt = CASE WHEN CAST(? AS SIGNED) THEN ? ELSE attempt END, + attempted_at = CASE WHEN CAST(? AS SIGNED) THEN ? ELSE attempted_at END, + attempted_by = CASE WHEN CAST(? AS SIGNED) THEN CAST(? AS JSON) ELSE attempted_by END, + errors = CASE WHEN CAST(? AS SIGNED) THEN CAST(? AS JSON) ELSE errors END, + finalized_at = CASE WHEN CAST(? AS SIGNED) THEN ? ELSE finalized_at END, + max_attempts = CASE WHEN CAST(? AS SIGNED) THEN ? ELSE max_attempts END, + metadata = CASE WHEN CAST(? AS SIGNED) THEN CAST(? AS JSON) ELSE metadata END, + state = CASE WHEN CAST(? AS SIGNED) THEN ? ELSE state END +WHERE id = ? +` + +type JobUpdateFullExecParams struct { + AttemptDoUpdate int64 + Attempt int64 + AttemptedAtDoUpdate int64 + AttemptedAt sql.NullTime + AttemptedByDoUpdate int64 + AttemptedBy []byte + ErrorsDoUpdate int64 + Errors []byte + FinalizedAtDoUpdate int64 + FinalizedAt sql.NullTime + MaxAttemptsDoUpdate int64 + MaxAttempts int64 + MetadataDoUpdate int64 + Metadata []byte + StateDoUpdate int64 + State string + ID int64 +} + +func (q *Queries) JobUpdateFullExec(ctx context.Context, db DBTX, arg *JobUpdateFullExecParams) error { + _, err := db.ExecContext(ctx, jobUpdateFullExec, + arg.AttemptDoUpdate, + arg.Attempt, + arg.AttemptedAtDoUpdate, + arg.AttemptedAt, + arg.AttemptedByDoUpdate, + arg.AttemptedBy, + arg.ErrorsDoUpdate, + arg.Errors, + arg.FinalizedAtDoUpdate, + arg.FinalizedAt, + arg.MaxAttemptsDoUpdate, + arg.MaxAttempts, + arg.MetadataDoUpdate, + arg.Metadata, + arg.StateDoUpdate, + arg.State, + arg.ID, + ) + return err +} diff --git a/riverdriver/rivermysql/internal/dbsqlc/river_leader.sql b/riverdriver/rivermysql/internal/dbsqlc/river_leader.sql new file mode 100644 index 00000000..4163c08b --- /dev/null +++ b/riverdriver/rivermysql/internal/dbsqlc/river_leader.sql @@ -0,0 +1,50 @@ +CREATE TABLE river_leader ( + elected_at DATETIME(6) NOT NULL, + expires_at DATETIME(6) NOT NULL, + leader_id VARCHAR(128) CHARACTER SET utf8mb4 COLLATE utf8mb4_bin NOT NULL, + name VARCHAR(128) CHARACTER SET utf8mb4 COLLATE utf8mb4_bin NOT NULL DEFAULT 'default' PRIMARY KEY +); + +-- name: LeaderAttemptElectExec :exec +INSERT INTO /* TEMPLATE: schema */river_leader ( + leader_id, + elected_at, + expires_at +) VALUES ( + sqlc.arg('leader_id'), + COALESCE(sqlc.narg('now'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')), + TIMESTAMPADD(MICROSECOND, sqlc.arg('ttl'), COALESCE(sqlc.narg('now'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00'))) +); + +-- name: LeaderAttemptReelectExec :execresult +UPDATE /* TEMPLATE: schema */river_leader +SET expires_at = TIMESTAMPADD(MICROSECOND, sqlc.arg('ttl'), COALESCE(sqlc.narg('now'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00'))) +WHERE + elected_at = sqlc.arg('elected_at') + AND expires_at >= COALESCE(sqlc.narg('now'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) + AND leader_id = sqlc.arg('leader_id'); + +-- name: LeaderDeleteExpired :execrows +DELETE FROM /* TEMPLATE: schema */river_leader +WHERE expires_at < COALESCE(sqlc.narg('now'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')); + +-- name: LeaderGetElectedLeader :one +SELECT elected_at, expires_at, leader_id, name +FROM /* TEMPLATE: schema */river_leader; + +-- name: LeaderInsertExec :exec +INSERT INTO /* TEMPLATE: schema */river_leader ( + elected_at, + expires_at, + leader_id +) VALUES ( + COALESCE(sqlc.narg('elected_at'), sqlc.narg('now'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')), + COALESCE(sqlc.narg('expires_at'), TIMESTAMPADD(MICROSECOND, sqlc.arg('ttl'), COALESCE(sqlc.narg('now'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')))), + sqlc.arg('leader_id') +); + +-- name: LeaderResign :execrows +DELETE FROM /* TEMPLATE: schema */river_leader +WHERE + elected_at = sqlc.arg('elected_at') + AND leader_id = sqlc.arg('leader_id'); diff --git a/riverdriver/rivermysql/internal/dbsqlc/river_leader.sql.go b/riverdriver/rivermysql/internal/dbsqlc/river_leader.sql.go new file mode 100644 index 00000000..b9c08e23 --- /dev/null +++ b/riverdriver/rivermysql/internal/dbsqlc/river_leader.sql.go @@ -0,0 +1,148 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.0 +// source: river_leader.sql + +package dbsqlc + +import ( + "context" + "database/sql" + "time" +) + +const leaderAttemptElectExec = `-- name: LeaderAttemptElectExec :exec +INSERT INTO /* TEMPLATE: schema */river_leader ( + leader_id, + elected_at, + expires_at +) VALUES ( + ?, + COALESCE(?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')), + TIMESTAMPADD(MICROSECOND, ?, COALESCE(?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00'))) +) +` + +type LeaderAttemptElectExecParams struct { + LeaderID string + Now interface{} + TTL int64 +} + +func (q *Queries) LeaderAttemptElectExec(ctx context.Context, db DBTX, arg *LeaderAttemptElectExecParams) error { + _, err := db.ExecContext(ctx, leaderAttemptElectExec, + arg.LeaderID, + arg.Now, + arg.TTL, + arg.Now, + ) + return err +} + +const leaderAttemptReelectExec = `-- name: LeaderAttemptReelectExec :execresult +UPDATE /* TEMPLATE: schema */river_leader +SET expires_at = TIMESTAMPADD(MICROSECOND, ?, COALESCE(?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00'))) +WHERE + elected_at = ? + AND expires_at >= COALESCE(?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) + AND leader_id = ? +` + +type LeaderAttemptReelectExecParams struct { + TTL int64 + Now sql.NullTime + ElectedAt time.Time + LeaderID string +} + +func (q *Queries) LeaderAttemptReelectExec(ctx context.Context, db DBTX, arg *LeaderAttemptReelectExecParams) (sql.Result, error) { + return db.ExecContext(ctx, leaderAttemptReelectExec, + arg.TTL, + arg.Now, + arg.ElectedAt, + arg.Now, + arg.LeaderID, + ) +} + +const leaderDeleteExpired = `-- name: LeaderDeleteExpired :execrows +DELETE FROM /* TEMPLATE: schema */river_leader +WHERE expires_at < COALESCE(?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) +` + +func (q *Queries) LeaderDeleteExpired(ctx context.Context, db DBTX, now sql.NullTime) (int64, error) { + result, err := db.ExecContext(ctx, leaderDeleteExpired, now) + if err != nil { + return 0, err + } + return result.RowsAffected() +} + +const leaderGetElectedLeader = `-- name: LeaderGetElectedLeader :one +SELECT elected_at, expires_at, leader_id, name +FROM /* TEMPLATE: schema */river_leader +` + +func (q *Queries) LeaderGetElectedLeader(ctx context.Context, db DBTX) (*RiverLeader, error) { + row := db.QueryRowContext(ctx, leaderGetElectedLeader) + var i RiverLeader + err := row.Scan( + &i.ElectedAt, + &i.ExpiresAt, + &i.LeaderID, + &i.Name, + ) + return &i, err +} + +const leaderInsertExec = `-- name: LeaderInsertExec :exec +INSERT INTO /* TEMPLATE: schema */river_leader ( + elected_at, + expires_at, + leader_id +) VALUES ( + COALESCE(?, ?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')), + COALESCE(?, TIMESTAMPADD(MICROSECOND, ?, COALESCE(?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')))), + ? +) +` + +type LeaderInsertExecParams struct { + ElectedAt interface{} + Now interface{} + ExpiresAt interface{} + TTL int64 + LeaderID string +} + +func (q *Queries) LeaderInsertExec(ctx context.Context, db DBTX, arg *LeaderInsertExecParams) error { + _, err := db.ExecContext(ctx, leaderInsertExec, + arg.ElectedAt, + arg.Now, + arg.ExpiresAt, + arg.TTL, + arg.Now, + arg.LeaderID, + ) + return err +} + +const leaderResign = `-- name: LeaderResign :execrows +DELETE FROM /* TEMPLATE: schema */river_leader +WHERE + elected_at = ? + AND leader_id = ? +` + +type LeaderResignParams struct { + ElectedAt time.Time + LeaderID string +} + +func (q *Queries) LeaderResign(ctx context.Context, db DBTX, arg *LeaderResignParams) (int64, error) { + result, err := db.ExecContext(ctx, leaderResign, arg.ElectedAt, arg.LeaderID) + if err != nil { + return 0, err + } + return result.RowsAffected() +} diff --git a/riverdriver/rivermysql/internal/dbsqlc/river_migration.sql b/riverdriver/rivermysql/internal/dbsqlc/river_migration.sql new file mode 100644 index 00000000..30849a13 --- /dev/null +++ b/riverdriver/rivermysql/internal/dbsqlc/river_migration.sql @@ -0,0 +1,80 @@ +CREATE TABLE river_migration ( + line VARCHAR(128) CHARACTER SET utf8mb4 COLLATE utf8mb4_bin NOT NULL, + version BIGINT NOT NULL, + -- sqlc's MySQL parser doesn't accept UTC_TIMESTAMP(6) in a DEFAULT + -- expression. Runtime migrations use UTC_TIMESTAMP(6); this codegen-only + -- schema declaration uses NOW(6). + created_at DATETIME(6) NOT NULL DEFAULT (NOW(6)), + PRIMARY KEY (line, version) +); + +-- name: RiverMigrationDeleteAssumingMainMany :many +SELECT created_at, version +FROM /* TEMPLATE: schema */river_migration +WHERE version IN (sqlc.slice('version')); + +-- name: RiverMigrationDeleteAssumingMainManyExec :exec +DELETE FROM /* TEMPLATE: schema */river_migration +WHERE version IN (sqlc.slice('version')); + +-- name: RiverMigrationDeleteByLineAndVersionMany :many +SELECT line, version, created_at +FROM /* TEMPLATE: schema */river_migration +WHERE line = sqlc.arg('line') + AND version IN (sqlc.slice('version')); + +-- name: RiverMigrationDeleteByLineAndVersionManyExec :exec +DELETE FROM /* TEMPLATE: schema */river_migration +WHERE line = sqlc.arg('line') + AND version IN (sqlc.slice('version')); + +-- name: RiverMigrationGetAllAssumingMain :many +SELECT + created_at, + version +FROM /* TEMPLATE: schema */river_migration +ORDER BY version; + +-- name: RiverMigrationGetByLine :many +SELECT line, version, created_at +FROM /* TEMPLATE: schema */river_migration +WHERE line = sqlc.arg('line') +ORDER BY version; + +-- name: RiverMigrationInsertExec :exec +INSERT INTO /* TEMPLATE: schema */river_migration ( + line, + version +) VALUES ( + sqlc.arg('line'), + sqlc.arg('version') +); + +-- name: RiverMigrationGetByLineAndVersion :one +SELECT line, version, created_at +FROM /* TEMPLATE: schema */river_migration +WHERE line = sqlc.arg('line') AND version = sqlc.arg('version'); + +-- name: RiverMigrationGetByLineAndVersionMany :many +SELECT line, version, created_at +FROM /* TEMPLATE: schema */river_migration +WHERE line = sqlc.arg('line') AND version IN (sqlc.slice('version')) +ORDER BY version; + +-- name: RiverMigrationInsertAssumingMainExec :exec +INSERT INTO /* TEMPLATE: schema */river_migration ( + version +) VALUES ( + sqlc.arg('version') +); + +-- name: RiverMigrationGetByVersion :one +SELECT created_at, version +FROM /* TEMPLATE: schema */river_migration +WHERE version = sqlc.arg('version'); + +-- name: RiverMigrationGetByVersionMany :many +SELECT created_at, version +FROM /* TEMPLATE: schema */river_migration +WHERE version IN (sqlc.slice('version')) +ORDER BY version; diff --git a/riverdriver/rivermysql/internal/dbsqlc/river_migration.sql.go b/riverdriver/rivermysql/internal/dbsqlc/river_migration.sql.go new file mode 100644 index 00000000..2e3cbb89 --- /dev/null +++ b/riverdriver/rivermysql/internal/dbsqlc/river_migration.sql.go @@ -0,0 +1,375 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.0 +// source: river_migration.sql + +package dbsqlc + +import ( + "context" + "strings" + "time" +) + +const riverMigrationDeleteAssumingMainMany = `-- name: RiverMigrationDeleteAssumingMainMany :many +SELECT created_at, version +FROM /* TEMPLATE: schema */river_migration +WHERE version IN (/*SLICE:version*/?) +` + +type RiverMigrationDeleteAssumingMainManyRow struct { + CreatedAt time.Time + Version int64 +} + +func (q *Queries) RiverMigrationDeleteAssumingMainMany(ctx context.Context, db DBTX, version []int64) ([]*RiverMigrationDeleteAssumingMainManyRow, error) { + query := riverMigrationDeleteAssumingMainMany + var queryParams []interface{} + if len(version) > 0 { + for _, v := range version { + queryParams = append(queryParams, v) + } + query = strings.Replace(query, "/*SLICE:version*/?", strings.Repeat(",?", len(version))[1:], 1) + } else { + query = strings.Replace(query, "/*SLICE:version*/?", "NULL", 1) + } + rows, err := db.QueryContext(ctx, query, queryParams...) + if err != nil { + return nil, err + } + defer rows.Close() + var items []*RiverMigrationDeleteAssumingMainManyRow + for rows.Next() { + var i RiverMigrationDeleteAssumingMainManyRow + if err := rows.Scan(&i.CreatedAt, &i.Version); err != nil { + return nil, err + } + items = append(items, &i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const riverMigrationDeleteAssumingMainManyExec = `-- name: RiverMigrationDeleteAssumingMainManyExec :exec +DELETE FROM /* TEMPLATE: schema */river_migration +WHERE version IN (/*SLICE:version*/?) +` + +func (q *Queries) RiverMigrationDeleteAssumingMainManyExec(ctx context.Context, db DBTX, version []int64) error { + query := riverMigrationDeleteAssumingMainManyExec + var queryParams []interface{} + if len(version) > 0 { + for _, v := range version { + queryParams = append(queryParams, v) + } + query = strings.Replace(query, "/*SLICE:version*/?", strings.Repeat(",?", len(version))[1:], 1) + } else { + query = strings.Replace(query, "/*SLICE:version*/?", "NULL", 1) + } + _, err := db.ExecContext(ctx, query, queryParams...) + return err +} + +const riverMigrationDeleteByLineAndVersionMany = `-- name: RiverMigrationDeleteByLineAndVersionMany :many +SELECT line, version, created_at +FROM /* TEMPLATE: schema */river_migration +WHERE line = ? + AND version IN (/*SLICE:version*/?) +` + +type RiverMigrationDeleteByLineAndVersionManyParams struct { + Line string + Version []int64 +} + +func (q *Queries) RiverMigrationDeleteByLineAndVersionMany(ctx context.Context, db DBTX, arg *RiverMigrationDeleteByLineAndVersionManyParams) ([]*RiverMigration, error) { + query := riverMigrationDeleteByLineAndVersionMany + var queryParams []interface{} + queryParams = append(queryParams, arg.Line) + if len(arg.Version) > 0 { + for _, v := range arg.Version { + queryParams = append(queryParams, v) + } + query = strings.Replace(query, "/*SLICE:version*/?", strings.Repeat(",?", len(arg.Version))[1:], 1) + } else { + query = strings.Replace(query, "/*SLICE:version*/?", "NULL", 1) + } + rows, err := db.QueryContext(ctx, query, queryParams...) + if err != nil { + return nil, err + } + defer rows.Close() + var items []*RiverMigration + for rows.Next() { + var i RiverMigration + if err := rows.Scan(&i.Line, &i.Version, &i.CreatedAt); err != nil { + return nil, err + } + items = append(items, &i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const riverMigrationDeleteByLineAndVersionManyExec = `-- name: RiverMigrationDeleteByLineAndVersionManyExec :exec +DELETE FROM /* TEMPLATE: schema */river_migration +WHERE line = ? + AND version IN (/*SLICE:version*/?) +` + +type RiverMigrationDeleteByLineAndVersionManyExecParams struct { + Line string + Version []int64 +} + +func (q *Queries) RiverMigrationDeleteByLineAndVersionManyExec(ctx context.Context, db DBTX, arg *RiverMigrationDeleteByLineAndVersionManyExecParams) error { + query := riverMigrationDeleteByLineAndVersionManyExec + var queryParams []interface{} + queryParams = append(queryParams, arg.Line) + if len(arg.Version) > 0 { + for _, v := range arg.Version { + queryParams = append(queryParams, v) + } + query = strings.Replace(query, "/*SLICE:version*/?", strings.Repeat(",?", len(arg.Version))[1:], 1) + } else { + query = strings.Replace(query, "/*SLICE:version*/?", "NULL", 1) + } + _, err := db.ExecContext(ctx, query, queryParams...) + return err +} + +const riverMigrationGetAllAssumingMain = `-- name: RiverMigrationGetAllAssumingMain :many +SELECT + created_at, + version +FROM /* TEMPLATE: schema */river_migration +ORDER BY version +` + +type RiverMigrationGetAllAssumingMainRow struct { + CreatedAt time.Time + Version int64 +} + +func (q *Queries) RiverMigrationGetAllAssumingMain(ctx context.Context, db DBTX) ([]*RiverMigrationGetAllAssumingMainRow, error) { + rows, err := db.QueryContext(ctx, riverMigrationGetAllAssumingMain) + if err != nil { + return nil, err + } + defer rows.Close() + var items []*RiverMigrationGetAllAssumingMainRow + for rows.Next() { + var i RiverMigrationGetAllAssumingMainRow + if err := rows.Scan(&i.CreatedAt, &i.Version); err != nil { + return nil, err + } + items = append(items, &i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const riverMigrationGetByLine = `-- name: RiverMigrationGetByLine :many +SELECT line, version, created_at +FROM /* TEMPLATE: schema */river_migration +WHERE line = ? +ORDER BY version +` + +func (q *Queries) RiverMigrationGetByLine(ctx context.Context, db DBTX, line string) ([]*RiverMigration, error) { + rows, err := db.QueryContext(ctx, riverMigrationGetByLine, line) + if err != nil { + return nil, err + } + defer rows.Close() + var items []*RiverMigration + for rows.Next() { + var i RiverMigration + if err := rows.Scan(&i.Line, &i.Version, &i.CreatedAt); err != nil { + return nil, err + } + items = append(items, &i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const riverMigrationGetByLineAndVersion = `-- name: RiverMigrationGetByLineAndVersion :one +SELECT line, version, created_at +FROM /* TEMPLATE: schema */river_migration +WHERE line = ? AND version = ? +` + +type RiverMigrationGetByLineAndVersionParams struct { + Line string + Version int64 +} + +func (q *Queries) RiverMigrationGetByLineAndVersion(ctx context.Context, db DBTX, arg *RiverMigrationGetByLineAndVersionParams) (*RiverMigration, error) { + row := db.QueryRowContext(ctx, riverMigrationGetByLineAndVersion, arg.Line, arg.Version) + var i RiverMigration + err := row.Scan(&i.Line, &i.Version, &i.CreatedAt) + return &i, err +} + +const riverMigrationGetByLineAndVersionMany = `-- name: RiverMigrationGetByLineAndVersionMany :many +SELECT line, version, created_at +FROM /* TEMPLATE: schema */river_migration +WHERE line = ? AND version IN (/*SLICE:version*/?) +ORDER BY version +` + +type RiverMigrationGetByLineAndVersionManyParams struct { + Line string + Version []int64 +} + +func (q *Queries) RiverMigrationGetByLineAndVersionMany(ctx context.Context, db DBTX, arg *RiverMigrationGetByLineAndVersionManyParams) ([]*RiverMigration, error) { + query := riverMigrationGetByLineAndVersionMany + var queryParams []interface{} + queryParams = append(queryParams, arg.Line) + if len(arg.Version) > 0 { + for _, v := range arg.Version { + queryParams = append(queryParams, v) + } + query = strings.Replace(query, "/*SLICE:version*/?", strings.Repeat(",?", len(arg.Version))[1:], 1) + } else { + query = strings.Replace(query, "/*SLICE:version*/?", "NULL", 1) + } + rows, err := db.QueryContext(ctx, query, queryParams...) + if err != nil { + return nil, err + } + defer rows.Close() + var items []*RiverMigration + for rows.Next() { + var i RiverMigration + if err := rows.Scan(&i.Line, &i.Version, &i.CreatedAt); err != nil { + return nil, err + } + items = append(items, &i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const riverMigrationGetByVersion = `-- name: RiverMigrationGetByVersion :one +SELECT created_at, version +FROM /* TEMPLATE: schema */river_migration +WHERE version = ? +` + +type RiverMigrationGetByVersionRow struct { + CreatedAt time.Time + Version int64 +} + +func (q *Queries) RiverMigrationGetByVersion(ctx context.Context, db DBTX, version int64) (*RiverMigrationGetByVersionRow, error) { + row := db.QueryRowContext(ctx, riverMigrationGetByVersion, version) + var i RiverMigrationGetByVersionRow + err := row.Scan(&i.CreatedAt, &i.Version) + return &i, err +} + +const riverMigrationGetByVersionMany = `-- name: RiverMigrationGetByVersionMany :many +SELECT created_at, version +FROM /* TEMPLATE: schema */river_migration +WHERE version IN (/*SLICE:version*/?) +ORDER BY version +` + +type RiverMigrationGetByVersionManyRow struct { + CreatedAt time.Time + Version int64 +} + +func (q *Queries) RiverMigrationGetByVersionMany(ctx context.Context, db DBTX, version []int64) ([]*RiverMigrationGetByVersionManyRow, error) { + query := riverMigrationGetByVersionMany + var queryParams []interface{} + if len(version) > 0 { + for _, v := range version { + queryParams = append(queryParams, v) + } + query = strings.Replace(query, "/*SLICE:version*/?", strings.Repeat(",?", len(version))[1:], 1) + } else { + query = strings.Replace(query, "/*SLICE:version*/?", "NULL", 1) + } + rows, err := db.QueryContext(ctx, query, queryParams...) + if err != nil { + return nil, err + } + defer rows.Close() + var items []*RiverMigrationGetByVersionManyRow + for rows.Next() { + var i RiverMigrationGetByVersionManyRow + if err := rows.Scan(&i.CreatedAt, &i.Version); err != nil { + return nil, err + } + items = append(items, &i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const riverMigrationInsertAssumingMainExec = `-- name: RiverMigrationInsertAssumingMainExec :exec +INSERT INTO /* TEMPLATE: schema */river_migration ( + version +) VALUES ( + ? +) +` + +func (q *Queries) RiverMigrationInsertAssumingMainExec(ctx context.Context, db DBTX, version int64) error { + _, err := db.ExecContext(ctx, riverMigrationInsertAssumingMainExec, version) + return err +} + +const riverMigrationInsertExec = `-- name: RiverMigrationInsertExec :exec +INSERT INTO /* TEMPLATE: schema */river_migration ( + line, + version +) VALUES ( + ?, + ? +) +` + +type RiverMigrationInsertExecParams struct { + Line string + Version int64 +} + +func (q *Queries) RiverMigrationInsertExec(ctx context.Context, db DBTX, arg *RiverMigrationInsertExecParams) error { + _, err := db.ExecContext(ctx, riverMigrationInsertExec, arg.Line, arg.Version) + return err +} diff --git a/riverdriver/rivermysql/internal/dbsqlc/river_notification.sql b/riverdriver/rivermysql/internal/dbsqlc/river_notification.sql new file mode 100644 index 00000000..be508b9e --- /dev/null +++ b/riverdriver/rivermysql/internal/dbsqlc/river_notification.sql @@ -0,0 +1,43 @@ +CREATE TABLE river_notification ( + id BIGINT AUTO_INCREMENT PRIMARY KEY, + -- sqlc's MySQL parser doesn't accept UTC_TIMESTAMP(6) in a DEFAULT + -- expression. Runtime migrations use UTC_TIMESTAMP(6); this codegen-only + -- schema declaration uses NOW(6). + created_at DATETIME(6) NOT NULL DEFAULT (NOW(6)), + payload TEXT NOT NULL, + topic VARCHAR(127) CHARACTER SET utf8mb4 COLLATE utf8mb4_bin NOT NULL, + CONSTRAINT topic_length CHECK (CHAR_LENGTH(topic) > 0 AND CHAR_LENGTH(topic) < 128) +); + +-- name: NotificationDeleteBefore :execrows +DELETE FROM /* TEMPLATE: schema */river_notification +WHERE created_at < sqlc.arg('created_at_horizon'); + +-- name: NotificationGetAfterForUpdate :one +-- InnoDB allocates auto-increment IDs before commit, so IDs may commit out of +-- order. The ascending locking read waits for an earlier uncommitted row rather +-- than returning a later committed row and causing the listener to skip it. +SELECT * +FROM /* TEMPLATE: schema */river_notification +WHERE id > sqlc.arg('after') +ORDER BY id ASC +LIMIT 1 +FOR UPDATE; + +-- name: NotificationGetIDsForUpdate :many +-- Used to establish a listener's initial high-water mark. This intentionally +-- scans in ascending order instead of using MAX(id) or a descending LIMIT so +-- that it encounters and waits for any lower uncommitted auto-increment IDs. +SELECT id +FROM /* TEMPLATE: schema */river_notification +ORDER BY id +FOR UPDATE; + +-- name: NotificationInsert :exec +INSERT INTO /* TEMPLATE: schema */river_notification ( + payload, + topic +) VALUES ( + sqlc.arg('payload'), + sqlc.arg('topic') +); diff --git a/riverdriver/rivermysql/internal/dbsqlc/river_notification.sql.go b/riverdriver/rivermysql/internal/dbsqlc/river_notification.sql.go new file mode 100644 index 00000000..df59f4df --- /dev/null +++ b/riverdriver/rivermysql/internal/dbsqlc/river_notification.sql.go @@ -0,0 +1,101 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.0 +// source: river_notification.sql + +package dbsqlc + +import ( + "context" + "time" +) + +const notificationDeleteBefore = `-- name: NotificationDeleteBefore :execrows +DELETE FROM /* TEMPLATE: schema */river_notification +WHERE created_at < ? +` + +func (q *Queries) NotificationDeleteBefore(ctx context.Context, db DBTX, createdAtHorizon time.Time) (int64, error) { + result, err := db.ExecContext(ctx, notificationDeleteBefore, createdAtHorizon) + if err != nil { + return 0, err + } + return result.RowsAffected() +} + +const notificationGetAfterForUpdate = `-- name: NotificationGetAfterForUpdate :one +SELECT id, created_at, payload, topic +FROM /* TEMPLATE: schema */river_notification +WHERE id > ? +ORDER BY id ASC +LIMIT 1 +FOR UPDATE +` + +// InnoDB allocates auto-increment IDs before commit, so IDs may commit out of +// order. The ascending locking read waits for an earlier uncommitted row rather +// than returning a later committed row and causing the listener to skip it. +func (q *Queries) NotificationGetAfterForUpdate(ctx context.Context, db DBTX, after int64) (*RiverNotification, error) { + row := db.QueryRowContext(ctx, notificationGetAfterForUpdate, after) + var i RiverNotification + err := row.Scan( + &i.ID, + &i.CreatedAt, + &i.Payload, + &i.Topic, + ) + return &i, err +} + +const notificationGetIDsForUpdate = `-- name: NotificationGetIDsForUpdate :many +SELECT id +FROM /* TEMPLATE: schema */river_notification +ORDER BY id +FOR UPDATE +` + +// Used to establish a listener's initial high-water mark. This intentionally +// scans in ascending order instead of using MAX(id) or a descending LIMIT so +// that it encounters and waits for any lower uncommitted auto-increment IDs. +func (q *Queries) NotificationGetIDsForUpdate(ctx context.Context, db DBTX) ([]int64, error) { + rows, err := db.QueryContext(ctx, notificationGetIDsForUpdate) + if err != nil { + return nil, err + } + defer rows.Close() + var items []int64 + for rows.Next() { + var id int64 + if err := rows.Scan(&id); err != nil { + return nil, err + } + items = append(items, id) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const notificationInsert = `-- name: NotificationInsert :exec +INSERT INTO /* TEMPLATE: schema */river_notification ( + payload, + topic +) VALUES ( + ?, + ? +) +` + +type NotificationInsertParams struct { + Payload string + Topic string +} + +func (q *Queries) NotificationInsert(ctx context.Context, db DBTX, arg *NotificationInsertParams) error { + _, err := db.ExecContext(ctx, notificationInsert, arg.Payload, arg.Topic) + return err +} diff --git a/riverdriver/rivermysql/internal/dbsqlc/river_queue.sql b/riverdriver/rivermysql/internal/dbsqlc/river_queue.sql new file mode 100644 index 00000000..aebd2ef4 --- /dev/null +++ b/riverdriver/rivermysql/internal/dbsqlc/river_queue.sql @@ -0,0 +1,95 @@ +CREATE TABLE river_queue ( + name VARCHAR(128) CHARACTER SET utf8mb4 COLLATE utf8mb4_bin NOT NULL PRIMARY KEY, + -- sqlc's MySQL parser doesn't accept UTC_TIMESTAMP(6) in a DEFAULT + -- expression. Runtime migrations use UTC_TIMESTAMP(6); this codegen-only + -- schema declaration uses NOW(6). + created_at DATETIME(6) NOT NULL DEFAULT (NOW(6)), + metadata JSON NOT NULL DEFAULT (JSON_OBJECT()), + paused_at DATETIME(6) NULL, + updated_at DATETIME(6) NOT NULL DEFAULT (NOW(6)) +); + +-- name: QueueCreateOrSetUpdatedAtExec :exec +INSERT INTO /* TEMPLATE: schema */river_queue ( + created_at, + metadata, + name, + paused_at, + updated_at +) VALUES ( + COALESCE(sqlc.narg('now'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')), + CAST(sqlc.arg('metadata') AS JSON), + sqlc.arg('name'), + sqlc.narg('paused_at'), + COALESCE(sqlc.narg('updated_at'), sqlc.narg('now'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) +) ON DUPLICATE KEY UPDATE + updated_at = COALESCE(sqlc.narg('updated_at'), sqlc.narg('now'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')); + +-- name: QueueGet :one +SELECT name, created_at, metadata, paused_at, updated_at +FROM /* TEMPLATE: schema */river_queue +WHERE name = sqlc.arg('name'); + +-- name: QueueDeleteExpiredSelect :many +SELECT name +FROM /* TEMPLATE: schema */river_queue +WHERE updated_at < sqlc.arg('updated_at_horizon') +ORDER BY name ASC +LIMIT ? +FOR UPDATE; + +-- name: QueueDeleteExpiredExec :exec +DELETE FROM /* TEMPLATE: schema */river_queue +WHERE name IN (sqlc.slice('names')); + +-- name: QueueList :many +SELECT name, created_at, metadata, paused_at, updated_at +FROM /* TEMPLATE: schema */river_queue +ORDER BY name ASC +LIMIT ?; + +-- name: QueueNameList :many +SELECT name +FROM /* TEMPLATE: schema */river_queue +WHERE + name > sqlc.arg('after') + AND (sqlc.arg('match') = '' OR LOWER(name) COLLATE utf8mb4_general_ci LIKE CONCAT('%', LOWER(sqlc.arg('match')), '%')) + AND name NOT IN (sqlc.slice('exclude')) +ORDER BY name ASC +LIMIT ?; + +-- MySQL evaluates SET clauses left-to-right using already-updated values, so +-- updated_at must be set BEFORE paused_at to see the original paused_at value. + +-- name: QueuePauseAll :execresult +UPDATE /* TEMPLATE: schema */river_queue +SET + updated_at = CASE WHEN paused_at IS NULL THEN COALESCE(sqlc.narg('now'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) ELSE updated_at END, + paused_at = CASE WHEN paused_at IS NULL THEN COALESCE(sqlc.narg('now'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) ELSE paused_at END; + +-- name: QueuePauseByName :execresult +UPDATE /* TEMPLATE: schema */river_queue +SET + updated_at = CASE WHEN paused_at IS NULL THEN COALESCE(sqlc.narg('now'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) ELSE updated_at END, + paused_at = CASE WHEN paused_at IS NULL THEN COALESCE(sqlc.narg('now'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) ELSE paused_at END +WHERE name = sqlc.arg('name'); + +-- name: QueueResumeAll :execresult +UPDATE /* TEMPLATE: schema */river_queue +SET + updated_at = CASE WHEN paused_at IS NOT NULL THEN COALESCE(sqlc.narg('now'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) ELSE updated_at END, + paused_at = NULL; + +-- name: QueueResumeByName :execresult +UPDATE /* TEMPLATE: schema */river_queue +SET + updated_at = CASE WHEN paused_at IS NOT NULL THEN COALESCE(sqlc.narg('now'), CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) ELSE updated_at END, + paused_at = NULL +WHERE name = sqlc.arg('name'); + +-- name: QueueUpdateExec :exec +UPDATE /* TEMPLATE: schema */river_queue +SET + metadata = CASE WHEN CAST(sqlc.arg('metadata_do_update') AS SIGNED) THEN CAST(sqlc.arg('metadata') AS JSON) ELSE metadata END, + updated_at = CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00') +WHERE name = sqlc.arg('name'); diff --git a/riverdriver/rivermysql/internal/dbsqlc/river_queue.sql.go b/riverdriver/rivermysql/internal/dbsqlc/river_queue.sql.go new file mode 100644 index 00000000..8690eb3f --- /dev/null +++ b/riverdriver/rivermysql/internal/dbsqlc/river_queue.sql.go @@ -0,0 +1,301 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.0 +// source: river_queue.sql + +package dbsqlc + +import ( + "context" + "database/sql" + "strings" + "time" +) + +const queueCreateOrSetUpdatedAtExec = `-- name: QueueCreateOrSetUpdatedAtExec :exec +INSERT INTO /* TEMPLATE: schema */river_queue ( + created_at, + metadata, + name, + paused_at, + updated_at +) VALUES ( + COALESCE(?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')), + CAST(? AS JSON), + ?, + ?, + COALESCE(?, ?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) +) ON DUPLICATE KEY UPDATE + updated_at = COALESCE(?, ?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) +` + +type QueueCreateOrSetUpdatedAtExecParams struct { + Now sql.NullTime + Metadata []byte + Name string + PausedAt sql.NullTime + UpdatedAt sql.NullTime +} + +func (q *Queries) QueueCreateOrSetUpdatedAtExec(ctx context.Context, db DBTX, arg *QueueCreateOrSetUpdatedAtExecParams) error { + _, err := db.ExecContext(ctx, queueCreateOrSetUpdatedAtExec, + arg.Now, + arg.Metadata, + arg.Name, + arg.PausedAt, + arg.UpdatedAt, + arg.Now, + arg.UpdatedAt, + arg.Now, + ) + return err +} + +const queueDeleteExpiredExec = `-- name: QueueDeleteExpiredExec :exec +DELETE FROM /* TEMPLATE: schema */river_queue +WHERE name IN (/*SLICE:names*/?) +` + +func (q *Queries) QueueDeleteExpiredExec(ctx context.Context, db DBTX, names []string) error { + query := queueDeleteExpiredExec + var queryParams []interface{} + if len(names) > 0 { + for _, v := range names { + queryParams = append(queryParams, v) + } + query = strings.Replace(query, "/*SLICE:names*/?", strings.Repeat(",?", len(names))[1:], 1) + } else { + query = strings.Replace(query, "/*SLICE:names*/?", "NULL", 1) + } + _, err := db.ExecContext(ctx, query, queryParams...) + return err +} + +const queueDeleteExpiredSelect = `-- name: QueueDeleteExpiredSelect :many +SELECT name +FROM /* TEMPLATE: schema */river_queue +WHERE updated_at < ? +ORDER BY name ASC +LIMIT ? +FOR UPDATE +` + +type QueueDeleteExpiredSelectParams struct { + UpdatedAtHorizon time.Time + Limit int32 +} + +func (q *Queries) QueueDeleteExpiredSelect(ctx context.Context, db DBTX, arg *QueueDeleteExpiredSelectParams) ([]string, error) { + rows, err := db.QueryContext(ctx, queueDeleteExpiredSelect, arg.UpdatedAtHorizon, arg.Limit) + if err != nil { + return nil, err + } + defer rows.Close() + var items []string + for rows.Next() { + var name string + if err := rows.Scan(&name); err != nil { + return nil, err + } + items = append(items, name) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const queueGet = `-- name: QueueGet :one +SELECT name, created_at, metadata, paused_at, updated_at +FROM /* TEMPLATE: schema */river_queue +WHERE name = ? +` + +func (q *Queries) QueueGet(ctx context.Context, db DBTX, name string) (*RiverQueue, error) { + row := db.QueryRowContext(ctx, queueGet, name) + var i RiverQueue + err := row.Scan( + &i.Name, + &i.CreatedAt, + &i.Metadata, + &i.PausedAt, + &i.UpdatedAt, + ) + return &i, err +} + +const queueList = `-- name: QueueList :many +SELECT name, created_at, metadata, paused_at, updated_at +FROM /* TEMPLATE: schema */river_queue +ORDER BY name ASC +LIMIT ? +` + +func (q *Queries) QueueList(ctx context.Context, db DBTX, limit int32) ([]*RiverQueue, error) { + rows, err := db.QueryContext(ctx, queueList, limit) + if err != nil { + return nil, err + } + defer rows.Close() + var items []*RiverQueue + for rows.Next() { + var i RiverQueue + if err := rows.Scan( + &i.Name, + &i.CreatedAt, + &i.Metadata, + &i.PausedAt, + &i.UpdatedAt, + ); err != nil { + return nil, err + } + items = append(items, &i) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const queueNameList = `-- name: QueueNameList :many +SELECT name +FROM /* TEMPLATE: schema */river_queue +WHERE + name > ? + AND (? = '' OR LOWER(name) COLLATE utf8mb4_general_ci LIKE CONCAT('%', LOWER(?), '%')) + AND name NOT IN (/*SLICE:exclude*/?) +ORDER BY name ASC +LIMIT ? +` + +type QueueNameListParams struct { + After string + Match string + Exclude []string + Limit int32 +} + +func (q *Queries) QueueNameList(ctx context.Context, db DBTX, arg *QueueNameListParams) ([]string, error) { + query := queueNameList + var queryParams []interface{} + queryParams = append(queryParams, arg.After) + queryParams = append(queryParams, arg.Match) + queryParams = append(queryParams, arg.Match) + if len(arg.Exclude) > 0 { + for _, v := range arg.Exclude { + queryParams = append(queryParams, v) + } + query = strings.Replace(query, "/*SLICE:exclude*/?", strings.Repeat(",?", len(arg.Exclude))[1:], 1) + } else { + query = strings.Replace(query, "/*SLICE:exclude*/?", "NULL", 1) + } + queryParams = append(queryParams, arg.Limit) + rows, err := db.QueryContext(ctx, query, queryParams...) + if err != nil { + return nil, err + } + defer rows.Close() + var items []string + for rows.Next() { + var name string + if err := rows.Scan(&name); err != nil { + return nil, err + } + items = append(items, name) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const queuePauseAll = `-- name: QueuePauseAll :execresult + +UPDATE /* TEMPLATE: schema */river_queue +SET + updated_at = CASE WHEN paused_at IS NULL THEN COALESCE(?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) ELSE updated_at END, + paused_at = CASE WHEN paused_at IS NULL THEN COALESCE(?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) ELSE paused_at END +` + +type QueuePauseAllParams struct { + Now sql.NullTime +} + +// MySQL evaluates SET clauses left-to-right using already-updated values, so +// updated_at must be set BEFORE paused_at to see the original paused_at value. +func (q *Queries) QueuePauseAll(ctx context.Context, db DBTX, arg *QueuePauseAllParams) (sql.Result, error) { + return db.ExecContext(ctx, queuePauseAll, arg.Now, arg.Now) +} + +const queuePauseByName = `-- name: QueuePauseByName :execresult +UPDATE /* TEMPLATE: schema */river_queue +SET + updated_at = CASE WHEN paused_at IS NULL THEN COALESCE(?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) ELSE updated_at END, + paused_at = CASE WHEN paused_at IS NULL THEN COALESCE(?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) ELSE paused_at END +WHERE name = ? +` + +type QueuePauseByNameParams struct { + Now sql.NullTime + Name string +} + +func (q *Queries) QueuePauseByName(ctx context.Context, db DBTX, arg *QueuePauseByNameParams) (sql.Result, error) { + return db.ExecContext(ctx, queuePauseByName, arg.Now, arg.Now, arg.Name) +} + +const queueResumeAll = `-- name: QueueResumeAll :execresult +UPDATE /* TEMPLATE: schema */river_queue +SET + updated_at = CASE WHEN paused_at IS NOT NULL THEN COALESCE(?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) ELSE updated_at END, + paused_at = NULL +` + +func (q *Queries) QueueResumeAll(ctx context.Context, db DBTX, now sql.NullTime) (sql.Result, error) { + return db.ExecContext(ctx, queueResumeAll, now) +} + +const queueResumeByName = `-- name: QueueResumeByName :execresult +UPDATE /* TEMPLATE: schema */river_queue +SET + updated_at = CASE WHEN paused_at IS NOT NULL THEN COALESCE(?, CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00')) ELSE updated_at END, + paused_at = NULL +WHERE name = ? +` + +type QueueResumeByNameParams struct { + Now sql.NullTime + Name string +} + +func (q *Queries) QueueResumeByName(ctx context.Context, db DBTX, arg *QueueResumeByNameParams) (sql.Result, error) { + return db.ExecContext(ctx, queueResumeByName, arg.Now, arg.Name) +} + +const queueUpdateExec = `-- name: QueueUpdateExec :exec +UPDATE /* TEMPLATE: schema */river_queue +SET + metadata = CASE WHEN CAST(? AS SIGNED) THEN CAST(? AS JSON) ELSE metadata END, + updated_at = CONVERT_TZ(NOW(6), @@session.time_zone, '+00:00') +WHERE name = ? +` + +type QueueUpdateExecParams struct { + MetadataDoUpdate int64 + Metadata []byte + Name string +} + +func (q *Queries) QueueUpdateExec(ctx context.Context, db DBTX, arg *QueueUpdateExecParams) error { + _, err := db.ExecContext(ctx, queueUpdateExec, arg.MetadataDoUpdate, arg.Metadata, arg.Name) + return err +} diff --git a/riverdriver/rivermysql/internal/dbsqlc/schema.sql b/riverdriver/rivermysql/internal/dbsqlc/schema.sql new file mode 100644 index 00000000..7c0838c8 --- /dev/null +++ b/riverdriver/rivermysql/internal/dbsqlc/schema.sql @@ -0,0 +1,76 @@ +-- Dummy table definitions for INFORMATION_SCHEMA system tables so sqlc can +-- resolve column types. At runtime the template prefix replaces the empty +-- default with "INFORMATION_SCHEMA.", making queries target the real system +-- tables. +CREATE TABLE COLUMNS ( + COLUMN_NAME VARCHAR(128) NOT NULL, + TABLE_NAME VARCHAR(128) NOT NULL, + TABLE_SCHEMA VARCHAR(128) NOT NULL +); + +CREATE TABLE SCHEMATA ( + SCHEMA_NAME VARCHAR(128) NOT NULL +); + +CREATE TABLE STATISTICS ( + INDEX_NAME VARCHAR(128) NOT NULL, + TABLE_NAME VARCHAR(128) NOT NULL, + TABLE_SCHEMA VARCHAR(128) NOT NULL +); + +CREATE TABLE TABLES ( + TABLE_NAME VARCHAR(128) NOT NULL, + TABLE_SCHEMA VARCHAR(128) NOT NULL +); + +-- name: ColumnExists :one +SELECT EXISTS ( + SELECT 1 + FROM /* TEMPLATE: information_schema */COLUMNS + WHERE COLUMN_NAME = sqlc.arg('column_name') + AND TABLE_NAME = sqlc.arg('table_name') + AND TABLE_SCHEMA = COALESCE(sqlc.narg('schema'), DATABASE()) +); + +-- name: IndexExists :one +SELECT EXISTS ( + SELECT 1 + FROM /* TEMPLATE: information_schema */STATISTICS + WHERE INDEX_NAME = sqlc.arg('index_name') + AND TABLE_SCHEMA = COALESCE(sqlc.narg('schema'), DATABASE()) +); + +-- name: IndexGetTableName :one +SELECT TABLE_NAME +FROM /* TEMPLATE: information_schema */STATISTICS +WHERE INDEX_NAME = sqlc.arg('index_name') + AND TABLE_SCHEMA = COALESCE(sqlc.narg('schema'), DATABASE()) +LIMIT 1; + +-- name: IndexReindexArtifacts :many +SELECT DISTINCT INDEX_NAME AS index_name +FROM /* TEMPLATE: information_schema */STATISTICS +WHERE TABLE_SCHEMA = COALESCE(sqlc.narg('schema'), DATABASE()) + AND REGEXP_LIKE(INDEX_NAME, sqlc.arg('artifact_pattern'), 'c') +ORDER BY INDEX_NAME; + +-- name: IndexesExist :many +SELECT DISTINCT INDEX_NAME AS index_name +FROM /* TEMPLATE: information_schema */STATISTICS +WHERE INDEX_NAME IN (sqlc.slice('index_names')) + AND TABLE_SCHEMA = COALESCE(sqlc.narg('schema'), DATABASE()); + +-- name: SchemaGetExpired :many +SELECT SCHEMA_NAME +FROM /* TEMPLATE: information_schema */SCHEMATA +WHERE SCHEMA_NAME LIKE CONCAT(sqlc.arg('prefix'), '%') + AND SCHEMA_NAME < sqlc.arg('before_name') +ORDER BY SCHEMA_NAME; + +-- name: TableExists :one +SELECT EXISTS ( + SELECT 1 + FROM /* TEMPLATE: information_schema */TABLES + WHERE TABLE_NAME = sqlc.arg('table_name') + AND TABLE_SCHEMA = COALESCE(sqlc.narg('schema'), DATABASE()) +); diff --git a/riverdriver/rivermysql/internal/dbsqlc/schema.sql.go b/riverdriver/rivermysql/internal/dbsqlc/schema.sql.go new file mode 100644 index 00000000..50d84a62 --- /dev/null +++ b/riverdriver/rivermysql/internal/dbsqlc/schema.sql.go @@ -0,0 +1,215 @@ +// Code generated by sqlc. DO NOT EDIT. +// versions: +// sqlc v1.31.0 +// source: schema.sql + +package dbsqlc + +import ( + "context" + "database/sql" + "strings" +) + +const columnExists = `-- name: ColumnExists :one +SELECT EXISTS ( + SELECT 1 + FROM /* TEMPLATE: information_schema */COLUMNS + WHERE COLUMN_NAME = ? + AND TABLE_NAME = ? + AND TABLE_SCHEMA = COALESCE(?, DATABASE()) +) +` + +type ColumnExistsParams struct { + ColumnName string + TableName string + Schema sql.NullString +} + +func (q *Queries) ColumnExists(ctx context.Context, db DBTX, arg *ColumnExistsParams) (bool, error) { + row := db.QueryRowContext(ctx, columnExists, arg.ColumnName, arg.TableName, arg.Schema) + var exists bool + err := row.Scan(&exists) + return exists, err +} + +const indexExists = `-- name: IndexExists :one +SELECT EXISTS ( + SELECT 1 + FROM /* TEMPLATE: information_schema */STATISTICS + WHERE INDEX_NAME = ? + AND TABLE_SCHEMA = COALESCE(?, DATABASE()) +) +` + +type IndexExistsParams struct { + IndexName string + Schema sql.NullString +} + +func (q *Queries) IndexExists(ctx context.Context, db DBTX, arg *IndexExistsParams) (bool, error) { + row := db.QueryRowContext(ctx, indexExists, arg.IndexName, arg.Schema) + var exists bool + err := row.Scan(&exists) + return exists, err +} + +const indexGetTableName = `-- name: IndexGetTableName :one +SELECT TABLE_NAME +FROM /* TEMPLATE: information_schema */STATISTICS +WHERE INDEX_NAME = ? + AND TABLE_SCHEMA = COALESCE(?, DATABASE()) +LIMIT 1 +` + +type IndexGetTableNameParams struct { + IndexName string + Schema sql.NullString +} + +func (q *Queries) IndexGetTableName(ctx context.Context, db DBTX, arg *IndexGetTableNameParams) (string, error) { + row := db.QueryRowContext(ctx, indexGetTableName, arg.IndexName, arg.Schema) + var table_name string + err := row.Scan(&table_name) + return table_name, err +} + +const indexReindexArtifacts = `-- name: IndexReindexArtifacts :many +SELECT DISTINCT INDEX_NAME AS index_name +FROM /* TEMPLATE: information_schema */STATISTICS +WHERE TABLE_SCHEMA = COALESCE(?, DATABASE()) + AND REGEXP_LIKE(INDEX_NAME, ?, 'c') +ORDER BY INDEX_NAME +` + +type IndexReindexArtifactsParams struct { + Schema sql.NullString + ArtifactPattern string +} + +func (q *Queries) IndexReindexArtifacts(ctx context.Context, db DBTX, arg *IndexReindexArtifactsParams) ([]string, error) { + rows, err := db.QueryContext(ctx, indexReindexArtifacts, arg.Schema, arg.ArtifactPattern) + if err != nil { + return nil, err + } + defer rows.Close() + var items []string + for rows.Next() { + var index_name string + if err := rows.Scan(&index_name); err != nil { + return nil, err + } + items = append(items, index_name) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const indexesExist = `-- name: IndexesExist :many +SELECT DISTINCT INDEX_NAME AS index_name +FROM /* TEMPLATE: information_schema */STATISTICS +WHERE INDEX_NAME IN (/*SLICE:index_names*/?) + AND TABLE_SCHEMA = COALESCE(?, DATABASE()) +` + +type IndexesExistParams struct { + IndexNames []string + Schema sql.NullString +} + +func (q *Queries) IndexesExist(ctx context.Context, db DBTX, arg *IndexesExistParams) ([]string, error) { + query := indexesExist + var queryParams []interface{} + if len(arg.IndexNames) > 0 { + for _, v := range arg.IndexNames { + queryParams = append(queryParams, v) + } + query = strings.Replace(query, "/*SLICE:index_names*/?", strings.Repeat(",?", len(arg.IndexNames))[1:], 1) + } else { + query = strings.Replace(query, "/*SLICE:index_names*/?", "NULL", 1) + } + queryParams = append(queryParams, arg.Schema) + rows, err := db.QueryContext(ctx, query, queryParams...) + if err != nil { + return nil, err + } + defer rows.Close() + var items []string + for rows.Next() { + var index_name string + if err := rows.Scan(&index_name); err != nil { + return nil, err + } + items = append(items, index_name) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const schemaGetExpired = `-- name: SchemaGetExpired :many +SELECT SCHEMA_NAME +FROM /* TEMPLATE: information_schema */SCHEMATA +WHERE SCHEMA_NAME LIKE CONCAT(?, '%') + AND SCHEMA_NAME < ? +ORDER BY SCHEMA_NAME +` + +type SchemaGetExpiredParams struct { + Prefix interface{} + BeforeName string +} + +func (q *Queries) SchemaGetExpired(ctx context.Context, db DBTX, arg *SchemaGetExpiredParams) ([]string, error) { + rows, err := db.QueryContext(ctx, schemaGetExpired, arg.Prefix, arg.BeforeName) + if err != nil { + return nil, err + } + defer rows.Close() + var items []string + for rows.Next() { + var schema_name string + if err := rows.Scan(&schema_name); err != nil { + return nil, err + } + items = append(items, schema_name) + } + if err := rows.Close(); err != nil { + return nil, err + } + if err := rows.Err(); err != nil { + return nil, err + } + return items, nil +} + +const tableExists = `-- name: TableExists :one +SELECT EXISTS ( + SELECT 1 + FROM /* TEMPLATE: information_schema */TABLES + WHERE TABLE_NAME = ? + AND TABLE_SCHEMA = COALESCE(?, DATABASE()) +) +` + +type TableExistsParams struct { + TableName string + Schema sql.NullString +} + +func (q *Queries) TableExists(ctx context.Context, db DBTX, arg *TableExistsParams) (bool, error) { + row := db.QueryRowContext(ctx, tableExists, arg.TableName, arg.Schema) + var exists bool + err := row.Scan(&exists) + return exists, err +} diff --git a/riverdriver/rivermysql/internal/dbsqlc/sqlc.yaml b/riverdriver/rivermysql/internal/dbsqlc/sqlc.yaml new file mode 100644 index 00000000..9ddeb227 --- /dev/null +++ b/riverdriver/rivermysql/internal/dbsqlc/sqlc.yaml @@ -0,0 +1,41 @@ +version: "2" +sql: + - engine: "mysql" + queries: + - river_job.sql + - river_leader.sql + - river_migration.sql + - river_notification.sql + - river_queue.sql + - schema.sql + schema: + - river_job.sql + - river_leader.sql + - river_migration.sql + - river_notification.sql + - river_queue.sql + - schema.sql + gen: + go: + package: "dbsqlc" + out: "." + emit_exact_table_names: true + emit_methods_with_db_argument: true + emit_params_struct_pointers: true + emit_pointers_for_null_types: true + emit_result_struct_pointers: true + + rename: + ids: "IDs" + ttl: "TTL" + + overrides: + - db_type: "json" + go_type: + type: "[]byte" + - db_type: "json" + go_type: + type: "[]byte" + nullable: true + - db_type: "int" + go_type: "int64" diff --git a/riverdriver/rivermysql/migration/main/001_create_river_migration.down.sql b/riverdriver/rivermysql/migration/main/001_create_river_migration.down.sql new file mode 100644 index 00000000..d2f6b3fa --- /dev/null +++ b/riverdriver/rivermysql/migration/main/001_create_river_migration.down.sql @@ -0,0 +1 @@ +DROP TABLE /* TEMPLATE: schema */river_migration; diff --git a/riverdriver/rivermysql/migration/main/001_create_river_migration.up.sql b/riverdriver/rivermysql/migration/main/001_create_river_migration.up.sql new file mode 100644 index 00000000..9aee5ff2 --- /dev/null +++ b/riverdriver/rivermysql/migration/main/001_create_river_migration.up.sql @@ -0,0 +1,8 @@ +CREATE TABLE /* TEMPLATE: schema */river_migration ( + id BIGINT AUTO_INCREMENT PRIMARY KEY, + created_at DATETIME(6) NOT NULL DEFAULT (UTC_TIMESTAMP(6)), + version BIGINT NOT NULL, + CONSTRAINT version CHECK (version >= 1) +) ENGINE=InnoDB; + +CREATE UNIQUE INDEX river_migration_version_idx ON /* TEMPLATE: schema */river_migration (version); diff --git a/riverdriver/rivermysql/migration/main/002_initial_schema.down.sql b/riverdriver/rivermysql/migration/main/002_initial_schema.down.sql new file mode 100644 index 00000000..095aa1dd --- /dev/null +++ b/riverdriver/rivermysql/migration/main/002_initial_schema.down.sql @@ -0,0 +1,2 @@ +DROP TABLE /* TEMPLATE: schema */river_job; +DROP TABLE /* TEMPLATE: schema */river_leader; diff --git a/riverdriver/rivermysql/migration/main/002_initial_schema.up.sql b/riverdriver/rivermysql/migration/main/002_initial_schema.up.sql new file mode 100644 index 00000000..c680cb3d --- /dev/null +++ b/riverdriver/rivermysql/migration/main/002_initial_schema.up.sql @@ -0,0 +1,43 @@ +CREATE TABLE /* TEMPLATE: schema */river_job( + id BIGINT AUTO_INCREMENT PRIMARY KEY, + state VARCHAR(20) CHARACTER SET utf8mb4 COLLATE utf8mb4_bin NOT NULL DEFAULT 'available', + attempt INT NOT NULL DEFAULT 0, + max_attempts INT NOT NULL, + attempted_at DATETIME(6) NULL, + created_at DATETIME(6) NOT NULL DEFAULT (UTC_TIMESTAMP(6)), + finalized_at DATETIME(6) NULL, + scheduled_at DATETIME(6) NOT NULL DEFAULT (UTC_TIMESTAMP(6)), + priority SMALLINT NOT NULL DEFAULT 1, + args JSON NULL, + attempted_by JSON NULL, -- JSON array of strings (no native text[] in MySQL) + errors JSON NULL, -- JSON array of error objects (no native jsonb[] in MySQL) + kind VARCHAR(128) CHARACTER SET utf8mb4 COLLATE utf8mb4_bin NOT NULL, + metadata JSON NOT NULL DEFAULT (JSON_OBJECT()), + queue VARCHAR(128) CHARACTER SET utf8mb4 COLLATE utf8mb4_bin NOT NULL DEFAULT 'default', + tags JSON NULL, -- JSON array of strings (no native varchar[] in MySQL) + CONSTRAINT finalized_or_finalized_at_null CHECK ( + (state IN ('cancelled', 'completed', 'discarded') AND finalized_at IS NOT NULL) OR + finalized_at IS NULL + ), + CONSTRAINT max_attempts_is_positive CHECK (max_attempts > 0), + CONSTRAINT priority_in_range CHECK (priority >= 1 AND priority <= 4), + CONSTRAINT queue_length CHECK (CHAR_LENGTH(queue) > 0 AND CHAR_LENGTH(queue) < 128), + CONSTRAINT kind_length CHECK (CHAR_LENGTH(kind) > 0 AND CHAR_LENGTH(kind) < 128), + CONSTRAINT state_valid CHECK (state IN ('available', 'cancelled', 'completed', 'discarded', 'retryable', 'running', 'scheduled')) +) ENGINE=InnoDB; + +CREATE INDEX river_job_kind ON /* TEMPLATE: schema */river_job (kind); +CREATE INDEX river_job_state_and_finalized_at_index ON /* TEMPLATE: schema */river_job (state, finalized_at); +CREATE INDEX river_job_prioritized_fetching_index ON /* TEMPLATE: schema */river_job (state, queue, priority, scheduled_at, id); + +-- MySQL does not support triggers for LISTEN/NOTIFY, so river_job_notify is +-- omitted. + +CREATE TABLE /* TEMPLATE: schema */river_leader ( + elected_at DATETIME(6) NOT NULL, + expires_at DATETIME(6) NOT NULL, + leader_id VARCHAR(128) CHARACTER SET utf8mb4 COLLATE utf8mb4_bin NOT NULL, + name VARCHAR(128) CHARACTER SET utf8mb4 COLLATE utf8mb4_bin NOT NULL PRIMARY KEY, + CONSTRAINT name_length CHECK (CHAR_LENGTH(name) > 0 AND CHAR_LENGTH(name) < 128), + CONSTRAINT leader_id_length CHECK (CHAR_LENGTH(leader_id) > 0 AND CHAR_LENGTH(leader_id) < 128) +) ENGINE=InnoDB; diff --git a/riverdriver/rivermysql/migration/main/003_river_job_tags_non_null.down.sql b/riverdriver/rivermysql/migration/main/003_river_job_tags_non_null.down.sql new file mode 100644 index 00000000..93c33638 --- /dev/null +++ b/riverdriver/rivermysql/migration/main/003_river_job_tags_non_null.down.sql @@ -0,0 +1 @@ +ALTER TABLE /* TEMPLATE: schema */river_job MODIFY COLUMN tags JSON NULL; diff --git a/riverdriver/rivermysql/migration/main/003_river_job_tags_non_null.up.sql b/riverdriver/rivermysql/migration/main/003_river_job_tags_non_null.up.sql new file mode 100644 index 00000000..61b2f430 --- /dev/null +++ b/riverdriver/rivermysql/migration/main/003_river_job_tags_non_null.up.sql @@ -0,0 +1,2 @@ +UPDATE /* TEMPLATE: schema */river_job SET tags = JSON_ARRAY() WHERE tags IS NULL; +ALTER TABLE /* TEMPLATE: schema */river_job MODIFY COLUMN tags JSON NOT NULL DEFAULT (JSON_ARRAY()); diff --git a/riverdriver/rivermysql/migration/main/004_pending_and_more.down.sql b/riverdriver/rivermysql/migration/main/004_pending_and_more.down.sql new file mode 100644 index 00000000..891a0e81 --- /dev/null +++ b/riverdriver/rivermysql/migration/main/004_pending_and_more.down.sql @@ -0,0 +1,21 @@ +ALTER TABLE /* TEMPLATE: schema */river_job MODIFY COLUMN args JSON NULL; + +ALTER TABLE /* TEMPLATE: schema */river_job MODIFY COLUMN metadata JSON NOT NULL DEFAULT (JSON_OBJECT()); + +-- Cannot safely remove 'pending' from the CHECK constraint if rows reference +-- it, but we restore the original constraint form. +ALTER TABLE /* TEMPLATE: schema */river_job DROP CONSTRAINT finalized_or_finalized_at_null; +ALTER TABLE /* TEMPLATE: schema */river_job ADD CONSTRAINT finalized_or_finalized_at_null CHECK ( + (state IN ('cancelled', 'completed', 'discarded') AND finalized_at IS NOT NULL) OR + finalized_at IS NULL +); + +-- MySQL does not support triggers for LISTEN/NOTIFY, so no trigger changes. + +DROP TABLE /* TEMPLATE: schema */river_queue; + +ALTER TABLE /* TEMPLATE: schema */river_leader + ALTER COLUMN name DROP DEFAULT; + +ALTER TABLE /* TEMPLATE: schema */river_leader DROP CONSTRAINT name_length; +ALTER TABLE /* TEMPLATE: schema */river_leader ADD CONSTRAINT name_length CHECK (CHAR_LENGTH(name) > 0 AND CHAR_LENGTH(name) < 128); diff --git a/riverdriver/rivermysql/migration/main/004_pending_and_more.up.sql b/riverdriver/rivermysql/migration/main/004_pending_and_more.up.sql new file mode 100644 index 00000000..7725c9f1 --- /dev/null +++ b/riverdriver/rivermysql/migration/main/004_pending_and_more.up.sql @@ -0,0 +1,48 @@ +-- Make args NOT NULL with a default. +UPDATE /* TEMPLATE: schema */river_job SET args = JSON_OBJECT() WHERE args IS NULL; +ALTER TABLE /* TEMPLATE: schema */river_job MODIFY COLUMN args JSON NOT NULL DEFAULT (JSON_OBJECT()); + +-- Make metadata NOT NULL (it already had a default). +UPDATE /* TEMPLATE: schema */river_job SET metadata = JSON_OBJECT() WHERE metadata IS NULL; +ALTER TABLE /* TEMPLATE: schema */river_job MODIFY COLUMN metadata JSON NOT NULL DEFAULT (JSON_OBJECT()); + +-- Add 'pending' to the set of valid states. MySQL doesn't have enum types to +-- alter, so we update the CHECK constraint instead. +ALTER TABLE /* TEMPLATE: schema */river_job DROP CONSTRAINT state_valid; +ALTER TABLE /* TEMPLATE: schema */river_job ADD CONSTRAINT state_valid CHECK ( + state IN ('available', 'cancelled', 'completed', 'discarded', 'pending', 'retryable', 'running', 'scheduled') +); + +-- Update the finalized_at constraint to use the inverted form (matching +-- Postgres migration 004). +ALTER TABLE /* TEMPLATE: schema */river_job DROP CONSTRAINT finalized_or_finalized_at_null; +ALTER TABLE /* TEMPLATE: schema */river_job ADD CONSTRAINT finalized_or_finalized_at_null CHECK ( + (finalized_at IS NULL AND state NOT IN ('cancelled', 'completed', 'discarded')) OR + (finalized_at IS NOT NULL AND state IN ('cancelled', 'completed', 'discarded')) +); + +-- MySQL does not support triggers for LISTEN/NOTIFY, so river_job_notify +-- changes from Postgres are omitted. + +-- +-- Create table `river_queue`. +-- + +CREATE TABLE /* TEMPLATE: schema */river_queue ( + name VARCHAR(128) CHARACTER SET utf8mb4 COLLATE utf8mb4_bin NOT NULL PRIMARY KEY, + created_at DATETIME(6) NOT NULL DEFAULT (UTC_TIMESTAMP(6)), + metadata JSON NOT NULL DEFAULT (JSON_OBJECT()), + paused_at DATETIME(6) NULL, + updated_at DATETIME(6) NOT NULL +) ENGINE=InnoDB; + +-- +-- Alter `river_leader` to add a default value of 'default' to `name`. +-- MySQL supports ALTER TABLE for changing defaults and constraints. +-- + +ALTER TABLE /* TEMPLATE: schema */river_leader + ALTER COLUMN name SET DEFAULT 'default'; + +ALTER TABLE /* TEMPLATE: schema */river_leader DROP CONSTRAINT name_length; +ALTER TABLE /* TEMPLATE: schema */river_leader ADD CONSTRAINT name_length CHECK (name = 'default'); diff --git a/riverdriver/rivermysql/migration/main/005_migration_unique_client.down.sql b/riverdriver/rivermysql/migration/main/005_migration_unique_client.down.sql new file mode 100644 index 00000000..108c807a --- /dev/null +++ b/riverdriver/rivermysql/migration/main/005_migration_unique_client.down.sql @@ -0,0 +1,35 @@ +-- +-- Revert to migration table based only on `(version)`. +-- + +CREATE TABLE /* TEMPLATE: schema */river_migration_old ( + id BIGINT AUTO_INCREMENT PRIMARY KEY, + created_at DATETIME(6) NOT NULL DEFAULT (UTC_TIMESTAMP(6)), + version BIGINT NOT NULL, + CONSTRAINT version CHECK (version >= 1) +) ENGINE=InnoDB; + +CREATE UNIQUE INDEX river_migration_version_idx ON /* TEMPLATE: schema */river_migration_old (version); + +INSERT INTO /* TEMPLATE: schema */river_migration_old + (created_at, version) +SELECT created_at, version +FROM /* TEMPLATE: schema */river_migration; + +DROP TABLE /* TEMPLATE: schema */river_migration; + +ALTER TABLE /* TEMPLATE: schema */river_migration_old RENAME TO /* TEMPLATE: schema */river_migration; + +-- +-- Drop `river_job.unique_key` and its index. +-- + +DROP INDEX river_job_kind_unique_key_idx ON /* TEMPLATE: schema */river_job; +ALTER TABLE /* TEMPLATE: schema */river_job DROP COLUMN unique_key; + +-- +-- Drop `river_client` and derivative. +-- + +DROP TABLE /* TEMPLATE: schema */river_client_queue; +DROP TABLE /* TEMPLATE: schema */river_client; diff --git a/riverdriver/rivermysql/migration/main/005_migration_unique_client.up.sql b/riverdriver/rivermysql/migration/main/005_migration_unique_client.up.sql new file mode 100644 index 00000000..fc5889c4 --- /dev/null +++ b/riverdriver/rivermysql/migration/main/005_migration_unique_client.up.sql @@ -0,0 +1,60 @@ +-- +-- Rebuild the migration table so it's based on `(line, version)`. +-- + +DROP INDEX river_migration_version_idx ON /* TEMPLATE: schema */river_migration; + +CREATE TABLE /* TEMPLATE: schema */river_migration_new ( + line VARCHAR(128) CHARACTER SET utf8mb4 COLLATE utf8mb4_bin NOT NULL, + version BIGINT NOT NULL, + created_at DATETIME(6) NOT NULL DEFAULT (UTC_TIMESTAMP(6)), + CONSTRAINT line_length CHECK (CHAR_LENGTH(line) > 0 AND CHAR_LENGTH(line) < 128), + CONSTRAINT version_gte_1 CHECK (version >= 1), + PRIMARY KEY (line, version) +) ENGINE=InnoDB; + +INSERT INTO /* TEMPLATE: schema */river_migration_new + (created_at, line, version) +SELECT created_at, 'main', version +FROM /* TEMPLATE: schema */river_migration; + +DROP TABLE /* TEMPLATE: schema */river_migration; + +ALTER TABLE /* TEMPLATE: schema */river_migration_new RENAME TO /* TEMPLATE: schema */river_migration; + +-- +-- Add `river_job.unique_key` and bring up an index on it. +-- + +ALTER TABLE /* TEMPLATE: schema */river_job ADD COLUMN unique_key VARBINARY(255) NULL; + +CREATE UNIQUE INDEX river_job_kind_unique_key_idx ON /* TEMPLATE: schema */river_job (kind, unique_key); + +-- +-- Create `river_client` and derivative. +-- + +CREATE TABLE /* TEMPLATE: schema */river_client ( + id VARCHAR(128) CHARACTER SET utf8mb4 COLLATE utf8mb4_bin NOT NULL PRIMARY KEY, + created_at DATETIME(6) NOT NULL DEFAULT (UTC_TIMESTAMP(6)), + metadata JSON NOT NULL DEFAULT (JSON_OBJECT()), + paused_at DATETIME(6) NULL, + updated_at DATETIME(6) NOT NULL, + CONSTRAINT client_name_length CHECK (CHAR_LENGTH(id) > 0 AND CHAR_LENGTH(id) < 128) +) ENGINE=InnoDB; + +CREATE TABLE /* TEMPLATE: schema */river_client_queue ( + river_client_id VARCHAR(128) CHARACTER SET utf8mb4 COLLATE utf8mb4_bin NOT NULL, + name VARCHAR(128) CHARACTER SET utf8mb4 COLLATE utf8mb4_bin NOT NULL, + created_at DATETIME(6) NOT NULL DEFAULT (UTC_TIMESTAMP(6)), + max_workers INT NOT NULL DEFAULT 0, + metadata JSON NOT NULL DEFAULT (JSON_OBJECT()), + num_jobs_completed BIGINT NOT NULL DEFAULT 0, + num_jobs_running BIGINT NOT NULL DEFAULT 0, + updated_at DATETIME(6) NOT NULL, + PRIMARY KEY (river_client_id, name), + CONSTRAINT fk_river_client FOREIGN KEY (river_client_id) REFERENCES river_client (id) ON DELETE CASCADE, + CONSTRAINT cq_name_length CHECK (CHAR_LENGTH(name) > 0 AND CHAR_LENGTH(name) < 128), + CONSTRAINT num_jobs_completed_zero_or_positive CHECK (num_jobs_completed >= 0), + CONSTRAINT num_jobs_running_zero_or_positive CHECK (num_jobs_running >= 0) +) ENGINE=InnoDB; diff --git a/riverdriver/rivermysql/migration/main/006_bulk_unique.down.sql b/riverdriver/rivermysql/migration/main/006_bulk_unique.down.sql new file mode 100644 index 00000000..0e4bf09f --- /dev/null +++ b/riverdriver/rivermysql/migration/main/006_bulk_unique.down.sql @@ -0,0 +1,9 @@ +-- +-- Drop the functional unique index and unique_states column. +-- + +DROP INDEX river_job_unique_idx ON /* TEMPLATE: schema */river_job; +ALTER TABLE /* TEMPLATE: schema */river_job DROP COLUMN unique_states; + +-- Recreate the old unique index from migration 005. +CREATE UNIQUE INDEX river_job_kind_unique_key_idx ON /* TEMPLATE: schema */river_job (kind, unique_key); diff --git a/riverdriver/rivermysql/migration/main/006_bulk_unique.up.sql b/riverdriver/rivermysql/migration/main/006_bulk_unique.up.sql new file mode 100644 index 00000000..2f4883cc --- /dev/null +++ b/riverdriver/rivermysql/migration/main/006_bulk_unique.up.sql @@ -0,0 +1,45 @@ +-- +-- Add `river_job.unique_states` and bring up an index on it. +-- +-- MySQL 8.0+ supports functional indexes (index on an expression). We use one +-- to index `unique_key` only when the job's state matches the bitmask, which +-- is equivalent to Postgres's partial index with `river_job_state_in_bitmask`. +-- The expression evaluates to NULL when the constraint shouldn't be active, +-- and MySQL's UNIQUE allows multiple NULLs. +-- + +ALTER TABLE /* TEMPLATE: schema */river_job ADD COLUMN unique_states SMALLINT NULL; + +-- A "functional index" (feature specific to MySQL 8.0+) over `unique_key`, +-- `unique_states`, and `state`. The expression evaluates to `unique_key` when +-- the job's current state has its corresponding bit set in `unique_states`, and +-- to NULL otherwise. MySQL's UNIQUE indexes permit multiple NULL values, so +-- rows where the constraint is inactive (NULL result) never conflict with each +-- other, while rows where it's active are checked for uniqueness on +-- `unique_key`. This is the MySQL equivalent of Postgres's partial index with +-- `WHERE ... AND river_job_state_in_bitmask(unique_states, state)`. +-- +-- Unlike Postgres's partial indexes which exclude non-matching rows from the +-- index entirely, MySQL's functional index still stores an entry for every row +-- (NULLs included), so the storage savings aren't equivalent. The uniqueness +-- semantics are the same though. +CREATE UNIQUE INDEX river_job_unique_idx ON /* TEMPLATE: schema */river_job (( + CASE WHEN unique_key IS NOT NULL AND unique_states IS NOT NULL AND + CASE state + WHEN 'available' THEN unique_states & (1 << 0) + WHEN 'cancelled' THEN unique_states & (1 << 1) + WHEN 'completed' THEN unique_states & (1 << 2) + WHEN 'discarded' THEN unique_states & (1 << 3) + WHEN 'pending' THEN unique_states & (1 << 4) + WHEN 'retryable' THEN unique_states & (1 << 5) + WHEN 'running' THEN unique_states & (1 << 6) + WHEN 'scheduled' THEN unique_states & (1 << 7) + ELSE 0 + END >= 1 + THEN unique_key + ELSE NULL + END +)); + +-- Remove the old unique index from migration 005. +DROP INDEX river_job_kind_unique_key_idx ON /* TEMPLATE: schema */river_job; diff --git a/riverdriver/rivermysql/migration/main/007_notification_outbox_sqlite_jsonb_and_sql_cleanup.down.sql b/riverdriver/rivermysql/migration/main/007_notification_outbox_sqlite_jsonb_and_sql_cleanup.down.sql new file mode 100644 index 00000000..c4c740ea --- /dev/null +++ b/riverdriver/rivermysql/migration/main/007_notification_outbox_sqlite_jsonb_and_sql_cleanup.down.sql @@ -0,0 +1,57 @@ +-- +-- SQL cleanup rollback. +-- + +-- +-- Add back unused tables `river_client` and `river_client_queue`. +-- + +CREATE TABLE /* TEMPLATE: schema */river_client ( + id VARCHAR(128) CHARACTER SET utf8mb4 COLLATE utf8mb4_bin NOT NULL PRIMARY KEY, + created_at DATETIME(6) NOT NULL DEFAULT (UTC_TIMESTAMP(6)), + metadata JSON NOT NULL DEFAULT (JSON_OBJECT()), + paused_at DATETIME(6) NULL, + updated_at DATETIME(6) NOT NULL, + CONSTRAINT client_name_length CHECK (CHAR_LENGTH(id) > 0 AND CHAR_LENGTH(id) < 128) +) ENGINE=InnoDB; + +CREATE TABLE /* TEMPLATE: schema */river_client_queue ( + river_client_id VARCHAR(128) CHARACTER SET utf8mb4 COLLATE utf8mb4_bin NOT NULL, + name VARCHAR(128) CHARACTER SET utf8mb4 COLLATE utf8mb4_bin NOT NULL, + created_at DATETIME(6) NOT NULL DEFAULT (UTC_TIMESTAMP(6)), + max_workers INT NOT NULL DEFAULT 0, + metadata JSON NOT NULL DEFAULT (JSON_OBJECT()), + num_jobs_completed BIGINT NOT NULL DEFAULT 0, + num_jobs_running BIGINT NOT NULL DEFAULT 0, + updated_at DATETIME(6) NOT NULL, + PRIMARY KEY (river_client_id, name), + CONSTRAINT fk_river_client FOREIGN KEY (river_client_id) REFERENCES river_client (id) ON DELETE CASCADE, + CONSTRAINT cq_name_length CHECK (CHAR_LENGTH(name) > 0 AND CHAR_LENGTH(name) < 128), + CONSTRAINT num_jobs_completed_zero_or_positive CHECK (num_jobs_completed >= 0), + CONSTRAINT num_jobs_running_zero_or_positive CHECK (num_jobs_running >= 0) +) ENGINE=InnoDB; + +-- +-- Revert addition of `DEFAULT 25` to `river_job.max_attempts`. +-- + +ALTER TABLE /* TEMPLATE: schema */river_job + MODIFY COLUMN max_attempts INT NOT NULL; + +-- +-- Changes `river_queue.updated_at` to revert the default of `CURRENT_TIMESTAMP`. +-- + +ALTER TABLE /* TEMPLATE: schema */river_queue + MODIFY COLUMN updated_at DATETIME(6) NOT NULL; + +-- +-- SQLite JSONB conversion rollback. +-- +-- No-op. MySQL stores River JSON columns in JSON columns. + +-- +-- Notification outbox rollback. +-- + +DROP TABLE /* TEMPLATE: schema */river_notification; diff --git a/riverdriver/rivermysql/migration/main/007_notification_outbox_sqlite_jsonb_and_sql_cleanup.up.sql b/riverdriver/rivermysql/migration/main/007_notification_outbox_sqlite_jsonb_and_sql_cleanup.up.sql new file mode 100644 index 00000000..d5dc9a6b --- /dev/null +++ b/riverdriver/rivermysql/migration/main/007_notification_outbox_sqlite_jsonb_and_sql_cleanup.up.sql @@ -0,0 +1,44 @@ +-- +-- Notification outbox. +-- + +CREATE TABLE /* TEMPLATE: schema */river_notification ( + id BIGINT AUTO_INCREMENT PRIMARY KEY, + created_at DATETIME(6) NOT NULL DEFAULT (UTC_TIMESTAMP(6)), + payload TEXT NOT NULL, + topic VARCHAR(127) CHARACTER SET utf8mb4 COLLATE utf8mb4_bin NOT NULL, + CONSTRAINT topic_length CHECK (CHAR_LENGTH(topic) > 0 AND CHAR_LENGTH(topic) < 128) +) ENGINE=InnoDB; + +CREATE INDEX river_notification_created_at_idx ON /* TEMPLATE: schema */river_notification (created_at); +CREATE INDEX river_notification_topic_id_idx ON /* TEMPLATE: schema */river_notification (topic, id); + +-- +-- SQLite JSONB conversion. +-- +-- No-op. MySQL stores River JSON columns in JSON columns. + +-- +-- SQL cleanup. +-- + +-- +-- Drop unused tables `river_client` and `river_client_queue`. +-- + +DROP TABLE /* TEMPLATE: schema */river_client_queue; +DROP TABLE /* TEMPLATE: schema */river_client; + +-- +-- Adds `DEFAULT 25` to `river_job.max_attempts`. +-- + +ALTER TABLE /* TEMPLATE: schema */river_job + ALTER COLUMN max_attempts SET DEFAULT 25; + +-- +-- Changes `river_queue.updated_at` to have a default of `CURRENT_TIMESTAMP`. +-- + +ALTER TABLE /* TEMPLATE: schema */river_queue + MODIFY COLUMN updated_at DATETIME(6) NOT NULL DEFAULT (UTC_TIMESTAMP(6)); diff --git a/riverdriver/rivermysql/river_mysql_driver.go b/riverdriver/rivermysql/river_mysql_driver.go new file mode 100644 index 00000000..e4891020 --- /dev/null +++ b/riverdriver/rivermysql/river_mysql_driver.go @@ -0,0 +1,2084 @@ +// Package rivermysql provides a River driver implementation for MySQL. +// +// This driver targets MySQL 8.0+ and requires `go-sql-driver/mysql` or a +// compatible driver registered with `database/sql`. DSNs used with River +// should include `parseTime=true`, `loc=UTC`, and `multiStatements=true`. +// `parseTime` and `loc` ensure `DATETIME` columns are scanned as UTC +// [time.Time] values, while `multiStatements` is needed to run River's +// migrations. +// +// MySQL does not support LISTEN/NOTIFY, so this driver uses a +// `river_notification` outbox table for notifications. It also does not support +// `RETURNING` clauses, so most write operations are carried out as two-step +// operations (write + read). +// +// This driver is currently in early development. It's exercised in the test +// suite, but has minimal real world use as of yet. +package rivermysql + +import ( + "bytes" + "context" + "database/sql" + "embed" + "encoding/json" + "errors" + "fmt" + "io/fs" + "math" + "regexp" + "slices" + "strings" + "time" + + "github.com/go-sql-driver/mysql" + "github.com/tidwall/gjson" + "github.com/tidwall/sjson" + + "github.com/riverqueue/river/internal/rivercommon" + "github.com/riverqueue/river/riverdriver" + "github.com/riverqueue/river/riverdriver/rivermysql/internal/dbsqlc" + "github.com/riverqueue/river/rivershared/sqlctemplate" + "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" +) + +//go:embed migration/*/*.sql +var migrationFS embed.FS + +// Driver is an implementation of riverdriver.Driver for MySQL. +type Driver struct { + dbPool *sql.DB + replacer sqlctemplate.Replacer +} + +// New returns a new MySQL driver for use with River. +// +// It takes an sql.DB to use with River. The DSN should include +// `parseTime=true&loc=UTC&multiStatements=true`; see the package documentation +// for details. The pool must not be closed while associated River objects are +// running. +func New(dbPool *sql.DB) *Driver { + return &Driver{ + dbPool: dbPool, + replacer: sqlctemplate.Replacer{UnnumberedPlaceholders: true}, + } +} + +const argPlaceholder = "?" + +func (d *Driver) ArgPlaceholder() string { return argPlaceholder } +func (d *Driver) DatabaseName() string { return riverdriver.DatabaseNameMySQL } +func (d *Driver) SafeIdentifier(ident string) string { return mysqlIdentifier(ident) } + +func (d *Driver) GetExecutor() riverdriver.Executor { + return &Executor{d.dbPool, templateReplaceWrapper{d.dbPool, &d.replacer}, d, nil} +} + +func (d *Driver) GetListener(params *riverdriver.GetListenenerParams) riverdriver.Listener { + return &Listener{ + dbPool: d.dbPool, + pollInterval: notificationPollIntervalDefault, + replacer: &d.replacer, + schema: params.Schema, + topics: make(map[string]struct{}), + } +} + +func (d *Driver) GetMigrationDefaultLines() []string { return []string{riverdriver.MigrationLineMain} } +func (d *Driver) GetMigrationFS(line string) fs.FS { + if line == riverdriver.MigrationLineMain { + return migrationFS + } + panic("migration line does not exist: " + line) +} +func (d *Driver) GetMigrationLines() []string { return []string{riverdriver.MigrationLineMain} } +func (d *Driver) GetMigrationTruncateTables(line string, version int) []string { + if line == riverdriver.MigrationLineMain { + return riverdriver.MigrationLineMainTruncateTables(version) + } + panic("migration line does not exist: " + line) +} + +func (d *Driver) PoolIsSet() bool { return d.dbPool != nil } +func (d *Driver) PoolSet(dbPool any) error { + if d.dbPool != nil { + return errors.New("cannot PoolSet when internal pool is already non-nil") + } + d.dbPool = dbPool.(*sql.DB) //nolint:forcetypeassert + return nil +} + +func (d *Driver) SQLFragmentColumnIn(column string, values any) (string, any, error) { + arg, err := json.Marshal(values) + if err != nil { + return "", nil, err + } + + // Use JSON_TABLE to expand the JSON array into rows for an IN clause. + // The arg is passed through the template system's NamedArgs as @column, + // which the template replacer turns into a positional placeholder. The + // templateReplaceWrapper strips the number suffix to produce plain `?` + // that MySQL expects. Use VARCHAR(255) as the column type since it + // handles both integer and string comparisons via implicit conversion. Use + // River's binary collation explicitly so string comparisons remain + // case-sensitive and don't depend on the database's default collation. + return fmt.Sprintf("%s IN (SELECT jt.val FROM JSON_TABLE(CAST(@%s AS JSON), '$[*]' COLUMNS(val VARCHAR(255) CHARACTER SET utf8mb4 COLLATE utf8mb4_bin PATH '$')) AS jt)", column, column), arg, nil +} + +func (d *Driver) SupportsListener() bool { return true } +func (d *Driver) SupportsListenNotify() bool { return true } +func (d *Driver) TimePrecision() time.Duration { return time.Microsecond } + +func (d *Driver) UnwrapExecutor(tx *sql.Tx) riverdriver.ExecutorTx { + // Allows UnwrapExecutor to be invoked even if driver is nil. + var replacer *sqlctemplate.Replacer + if d == nil { + replacer = &sqlctemplate.Replacer{UnnumberedPlaceholders: true} + } else { + replacer = &d.replacer + } + + executorTx := &ExecutorTx{tx: tx} + executorTx.Executor = Executor{nil, templateReplaceWrapper{tx, replacer}, d, executorTx} + + return executorTx +} + +func (d *Driver) UnwrapTx(execTx riverdriver.ExecutorTx) *sql.Tx { + switch execTx := execTx.(type) { + case *ExecutorSubTx: + return execTx.tx + case *ExecutorTx: + return execTx.tx + } + panic("unhandled executor type") +} + +type Executor struct { + dbPool *sql.DB + dbtx templateReplaceWrapper + driver *Driver + execTx riverdriver.ExecutorTx +} + +func (e *Executor) Begin(ctx context.Context) (riverdriver.ExecutorTx, error) { + if e.execTx != nil { + return e.execTx.Begin(ctx) + } + + tx, err := e.dbPool.BeginTx(ctx, nil) + if err != nil { + return nil, err + } + + executorTx := &ExecutorTx{tx: tx} + executorTx.Executor = Executor{nil, templateReplaceWrapper{tx, &e.driver.replacer}, e.driver, executorTx} + return executorTx, nil +} + +func (e *Executor) ColumnExists(ctx context.Context, params *riverdriver.ColumnExistsParams) (bool, error) { + exists, err := dbsqlc.New().ColumnExists(informationSchemaParam(ctx), e.dbtx, &dbsqlc.ColumnExistsParams{ + ColumnName: params.Column, + Schema: sql.NullString{String: params.Schema, Valid: params.Schema != ""}, + TableName: params.Table, + }) + if err != nil { + return false, interpretError(err) + } + return exists, nil +} + +func (e *Executor) Exec(ctx context.Context, sql string, args ...any) error { + _, err := e.dbtx.ExecContext(ctx, sql, args...) + return interpretError(err) +} + +func (e *Executor) IndexDropIfExists(ctx context.Context, params *riverdriver.IndexDropIfExistsParams) error { + indexName := strings.TrimSpace(params.Index) + + exists, err := e.IndexExists(ctx, &riverdriver.IndexExistsParams{ + Index: indexName, + Schema: params.Schema, + }) + if err != nil { + return err + } + if !exists { + return nil + } + + // MySQL's DROP INDEX requires the table name. Look it up from the index. + tableName, err := dbsqlc.New().IndexGetTableName(informationSchemaParam(ctx), e.dbtx, &dbsqlc.IndexGetTableNameParams{ + IndexName: indexName, + Schema: sql.NullString{String: params.Schema, Valid: params.Schema != ""}, + }) + if err != nil { + return interpretError(err) + } + + var maybeSchema string + if params.Schema != "" { + maybeSchema = mysqlIdentifier(params.Schema) + "." + } + + _, err = e.dbtx.ExecContext(ctx, "DROP INDEX "+mysqlIdentifier(indexName)+" ON "+maybeSchema+mysqlIdentifier(tableName)) + return interpretError(err) +} + +func (e *Executor) IndexExists(ctx context.Context, params *riverdriver.IndexExistsParams) (bool, error) { + ctx = informationSchemaParam(ctx) + + exists, err := dbsqlc.New().IndexExists(ctx, e.dbtx, &dbsqlc.IndexExistsParams{ + IndexName: params.Index, + Schema: sql.NullString{String: params.Schema, Valid: params.Schema != ""}, + }) + if err != nil { + return false, interpretError(err) + } + return exists, nil +} + +func (e *Executor) IndexReindex(ctx context.Context, params *riverdriver.IndexReindexParams) error { + // MySQL has no per-index REINDEX operation. ANALYZE TABLE is the safe + // online approximation: it refreshes the selected index's table statistics + // without rebuilding and locking the entire table as OPTIMIZE TABLE would. + tableName, err := dbsqlc.New().IndexGetTableName(informationSchemaParam(ctx), e.dbtx, &dbsqlc.IndexGetTableNameParams{ + IndexName: params.Index, + Schema: sql.NullString{String: params.Schema, Valid: params.Schema != ""}, + }) + if err != nil { + return interpretError(err) + } + + var maybeSchema string + if params.Schema != "" { + maybeSchema = mysqlIdentifier(params.Schema) + "." + } + + _, err = e.dbtx.ExecContext(ctx, "ANALYZE TABLE "+maybeSchema+mysqlIdentifier(tableName)) + return interpretError(err) +} + +func (e *Executor) IndexReindexArtifacts(ctx context.Context, params *riverdriver.IndexReindexArtifactsParams) ([]string, error) { + artifacts, err := dbsqlc.New().IndexReindexArtifacts(informationSchemaParam(ctx), e.dbtx, &dbsqlc.IndexReindexArtifactsParams{ + ArtifactPattern: "^" + regexp.QuoteMeta(params.Index) + `_(ccnew|ccold)[0-9]*$`, + Schema: sql.NullString{String: params.Schema, Valid: params.Schema != ""}, + }) + return artifacts, interpretError(err) +} + +func (e *Executor) IndexesExist(ctx context.Context, params *riverdriver.IndexesExistParams) (map[string]bool, error) { + foundNames, err := dbsqlc.New().IndexesExist(informationSchemaParam(ctx), e.dbtx, &dbsqlc.IndexesExistParams{ + IndexNames: params.IndexNames, + Schema: sql.NullString{String: params.Schema, Valid: params.Schema != ""}, + }) + if err != nil { + return nil, interpretError(err) + } + + exists := make(map[string]bool, len(params.IndexNames)) + for _, name := range foundNames { + exists[name] = true + } + return exists, nil +} + +func (e *Executor) JobCancel(ctx context.Context, params *riverdriver.JobCancelParams) (*rivertype.JobRow, error) { + return dbutil.WithTxV(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) (*rivertype.JobRow, error) { + ctx = schemaTemplateParam(ctx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + + res, err := dbsqlc.New().JobCancelExec(ctx, dbtx, &dbsqlc.JobCancelExecParams{ + ID: params.ID, + CancelAttemptedAt: params.CancelAttemptedAt.UTC().Format(time.RFC3339Nano), + Now: nullTimeFromPtr(params.Now), + }) + if err != nil { + return nil, interpretError(err) + } + + rowsAffected, err := res.RowsAffected() + if err != nil { + return nil, interpretError(err) + } + + job, err := dbsqlc.New().JobGetByID(ctx, dbtx, params.ID) + if err != nil { + return nil, interpretError(err) + } + + if rowsAffected < 1 { + // No rows were updated, return the job as-is (already finalized) + return jobRowFromInternal(job) + } + + return jobRowFromInternal(job) + }) +} + +func (e *Executor) JobCountByAllStates(ctx context.Context, params *riverdriver.JobCountByAllStatesParams) (map[rivertype.JobState]int, error) { + counts, err := dbsqlc.New().JobCountByAllStates(schemaTemplateParam(ctx, params.Schema), e.dbtx) + if err != nil { + return nil, interpretError(err) + } + countsMap := make(map[rivertype.JobState]int) + for _, state := range rivertype.JobStates() { + countsMap[state] = 0 + } + for _, count := range counts { + countsMap[rivertype.JobState(count.State)] = int(count.Count) + } + return countsMap, nil +} + +func (e *Executor) JobCountByQueueAndState(ctx context.Context, params *riverdriver.JobCountByQueueAndStateParams) ([]*riverdriver.JobCountByQueueAndStateResult, error) { + rows, err := dbsqlc.New().JobCountByQueueAndState(schemaTemplateParam(ctx, params.Schema), e.dbtx, &dbsqlc.JobCountByQueueAndStateParams{ + QueueNames: params.QueueNames, + }) + if err != nil { + return nil, interpretError(err) + } + + // MySQL's GROUP BY only returns queues that have jobs, so fill in zero + // counts for any requested queues not in the result. + countsByQueue := make(map[string]*riverdriver.JobCountByQueueAndStateResult, len(rows)) + for _, row := range rows { + countsByQueue[row.Queue] = &riverdriver.JobCountByQueueAndStateResult{ + CountAvailable: row.CountAvailable, + CountRunning: row.CountRunning, + Queue: row.Queue, + } + } + + var ( + queueNames = slices.Compact(slices.Sorted(slices.Values(params.QueueNames))) + results = make([]*riverdriver.JobCountByQueueAndStateResult, len(queueNames)) + ) + for i, name := range queueNames { + if result, ok := countsByQueue[name]; ok { + results[i] = result + } else { + results[i] = &riverdriver.JobCountByQueueAndStateResult{Queue: name} + } + } + + return results, nil +} + +func (e *Executor) JobCountByState(ctx context.Context, params *riverdriver.JobCountByStateParams) (int, error) { + numJobs, err := dbsqlc.New().JobCountByState(schemaTemplateParam(ctx, params.Schema), e.dbtx, string(params.State)) + if err != nil { + return 0, err + } + return int(numJobs), nil +} + +func (e *Executor) JobDelete(ctx context.Context, params *riverdriver.JobDeleteParams) (*rivertype.JobRow, error) { + return dbutil.WithTxV(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) (*rivertype.JobRow, error) { + ctx = schemaTemplateParam(ctx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + + // Lock the job while reading it so a concurrent producer can't change it + // to running between the read and delete statements. + job, err := dbsqlc.New().JobDeleteSelect(ctx, dbtx, params.ID) + if err != nil { + return nil, interpretError(err) + } + + res, err := dbsqlc.New().JobDeleteExec(ctx, dbtx, params.ID) + if err != nil { + return nil, interpretError(err) + } + + rowsAffected, err := res.RowsAffected() + if err != nil { + return nil, interpretError(err) + } + + if rowsAffected < 1 { + if rivertype.JobState(job.State) == rivertype.JobStateRunning { + return nil, rivertype.ErrJobRunning + } + return nil, fmt.Errorf("bug; expected only to fetch a job with state %q, but was: %q", rivertype.JobStateRunning, job.State) + } + + return jobRowFromInternal(job) + }) +} + +func (e *Executor) JobDeleteBefore(ctx context.Context, params *riverdriver.JobDeleteBeforeParams) (int, error) { + queuesExcluded := params.QueuesExcluded + if len(queuesExcluded) < 1 { + queuesExcluded = []string{""} + } + + queuesIncluded := params.QueuesIncluded + if len(queuesIncluded) < 1 { + queuesIncluded = []string{""} + } + + res, err := dbsqlc.New().JobDeleteBefore(schemaTemplateParam(ctx, params.Schema), e.dbtx, &dbsqlc.JobDeleteBeforeParams{ + CancelledDoDelete: boolToInt64(params.CancelledDoDelete), + CancelledFinalizedAtHorizon: sql.NullTime{Time: params.CancelledFinalizedAtHorizon.UTC(), Valid: true}, + CompletedDoDelete: boolToInt64(params.CompletedDoDelete), + CompletedFinalizedAtHorizon: sql.NullTime{Time: params.CompletedFinalizedAtHorizon.UTC(), Valid: true}, + DiscardedDoDelete: boolToInt64(params.DiscardedDoDelete), + DiscardedFinalizedAtHorizon: sql.NullTime{Time: params.DiscardedFinalizedAtHorizon.UTC(), Valid: true}, + Limit: int32(params.Max), //nolint:gosec + QueuesExcluded: queuesExcluded, + QueuesExcludedEmpty: boolToInt64(len(params.QueuesExcluded) < 1), + QueuesIncluded: queuesIncluded, + QueuesIncludedEmpty: boolToInt64(len(params.QueuesIncluded) < 1), + }) + if err != nil { + return 0, interpretError(err) + } + rowsAffected, err := res.RowsAffected() + if err != nil { + return 0, interpretError(err) + } + return int(rowsAffected), nil +} + +func (e *Executor) JobDeleteMany(ctx context.Context, params *riverdriver.JobDeleteManyParams) ([]*rivertype.JobRow, error) { + return dbutil.WithTxV(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) ([]*rivertype.JobRow, error) { + baseCtx := ctx + selectCtx := sqlctemplate.WithReplacements(ctx, map[string]sqlctemplate.Replacement{ + "order_by_clause": {Value: params.OrderByClause}, + "where_clause": {Value: params.WhereClause}, + }, params.NamedArgs) + selectCtx = schemaTemplateParam(selectCtx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + + // Select and lock the rows before deleting them so that a producer can't + // start one of the selected jobs between the two statements. + jobs, err := dbsqlc.New().JobDeleteManySelect(selectCtx, dbtx, params.Max) + if err != nil { + return nil, interpretError(err) + } + + if len(jobs) > 0 { + // Use a fresh context with only the schema replacement. The select + // context has templates that aren't present in the DELETE query. + deleteCtx := schemaTemplateParam(baseCtx, params.Schema) + ids := sliceutil.Map(jobs, func(j *dbsqlc.RiverJob) int64 { return j.ID }) + if err := dbsqlc.New().JobDeleteManyExec(deleteCtx, dbtx, ids); err != nil { + return nil, interpretError(err) + } + } + + return sliceutil.MapError(jobs, jobRowFromInternal) + }) +} + +func (e *Executor) JobGetAvailable(ctx context.Context, params *riverdriver.JobGetAvailableParams) ([]*rivertype.JobRow, error) { + return dbutil.WithTxV(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) ([]*rivertype.JobRow, error) { + ctx = schemaTemplateParam(ctx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + + ids, err := dbsqlc.New().JobGetAvailableIDs(ctx, dbtx, &dbsqlc.JobGetAvailableIDsParams{ + Queue: params.Queue, + Now: nullTimeFromPtr(params.Now), + Limit: int32(params.MaxToLock), //nolint:gosec + }) + if err != nil { + return nil, interpretError(err) + } + + if len(ids) == 0 { + return nil, nil + } + + if err := dbsqlc.New().JobGetAvailableUpdate(ctx, dbtx, &dbsqlc.JobGetAvailableUpdateParams{ + Now: nullTimeFromPtr(params.Now), + MaxAttemptedBy: int64(params.MaxAttemptedBy), + AttemptedBy: params.ClientID, + ID: ids, + }); err != nil { + return nil, interpretError(err) + } + + jobs, err := dbsqlc.New().JobGetByIDManyOrdered(ctx, dbtx, ids) + if err != nil { + return nil, interpretError(err) + } + + return sliceutil.MapError(jobs, jobRowFromInternal) + }) +} + +func (e *Executor) JobGetByID(ctx context.Context, params *riverdriver.JobGetByIDParams) (*rivertype.JobRow, error) { + job, err := dbsqlc.New().JobGetByID(schemaTemplateParam(ctx, params.Schema), e.dbtx, params.ID) + if err != nil { + return nil, interpretError(err) + } + return jobRowFromInternal(job) +} + +func (e *Executor) JobGetByIDMany(ctx context.Context, params *riverdriver.JobGetByIDManyParams) ([]*rivertype.JobRow, error) { + jobs, err := dbsqlc.New().JobGetByIDMany(schemaTemplateParam(ctx, params.Schema), e.dbtx, params.ID) + if err != nil { + return nil, interpretError(err) + } + return sliceutil.MapError(jobs, jobRowFromInternal) +} + +func (e *Executor) JobGetByKindMany(ctx context.Context, params *riverdriver.JobGetByKindManyParams) ([]*rivertype.JobRow, error) { + jobs, err := dbsqlc.New().JobGetByKindMany(schemaTemplateParam(ctx, params.Schema), e.dbtx, params.Kind) + if err != nil { + return nil, interpretError(err) + } + return sliceutil.MapError(jobs, jobRowFromInternal) +} + +func (e *Executor) JobGetStuck(ctx context.Context, params *riverdriver.JobGetStuckParams) ([]*rivertype.JobRow, error) { + jobs, err := dbsqlc.New().JobGetStuck(schemaTemplateParam(ctx, params.Schema), e.dbtx, &dbsqlc.JobGetStuckParams{ + AfterID: params.AfterID, + Limit: int32(params.Max), //nolint:gosec + StuckHorizon: sql.NullTime{Time: params.StuckHorizon.UTC(), Valid: true}, + }) + if err != nil { + return nil, interpretError(err) + } + return sliceutil.MapError(jobs, jobRowFromInternal) +} + +func (e *Executor) JobInsertFastMany(ctx context.Context, params *riverdriver.JobInsertFastManyParams) ([]*riverdriver.JobInsertFastResult, error) { + var ( + insertRes = make([]*riverdriver.JobInsertFastResult, len(params.Jobs)) + uniqueNonceByIndex = make([]string, len(params.Jobs)) + ) + + if err := dbutil.WithTx(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) error { + baseCtx := ctx + ctx = schemaTemplateParam(ctx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + plainInsertCtx := schemaTemplateParam(baseCtx, params.Schema) + plainInsertCtx = sqlctemplate.WithReplacements(plainInsertCtx, map[string]sqlctemplate.Replacement{ + "on_duplicate_key": {Value: ""}, + }, nil) + uniqueInsertCtx := schemaTemplateParam(baseCtx, params.Schema) + uniqueInsertCtx = sqlctemplate.WithReplacements(uniqueInsertCtx, map[string]sqlctemplate.Replacement{ + "on_duplicate_key": {Value: "ON DUPLICATE KEY UPDATE id = LAST_INSERT_ID(id)"}, + }, nil) + + var idsWithNonce []int64 + ids := make([]int64, len(params.Jobs)) + for i, jobParams := range params.Jobs { + // Only use MySQL's broad ON DUPLICATE KEY behavior when a River unique + // key is present. Otherwise, a primary key conflict or a user-added + // unique constraint must be returned as an error. + insertCtx := plainInsertCtx + if len(jobParams.UniqueKey) > 0 { + insertCtx = uniqueInsertCtx + + // Postgres can inspect xmax to determine whether an upsert inserted a + // row. MySQL has no equivalent, so add a per-row nonce to metadata and + // compare it with the selected result. Each input needs its own nonce + // so duplicates within this batch are classified correctly. + uniqueNonceByIndex[i] = randutil.Hex(8) + } + + insertParams, err := jobInsertFastParams(jobParams, uniqueNonceByIndex[i]) + if err != nil { + return err + } + + res, err := dbsqlc.New().JobInsertFast(insertCtx, dbtx, insertParams) + if err != nil { + return interpretError(err) + } + ids[i], err = res.LastInsertId() + if err != nil { + return err + } + if uniqueNonceByIndex[i] != "" { + idsWithNonce = append(idsWithNonce, ids[i]) + } + } + + jobs, err := dbsqlc.New().JobGetByIDMany(ctx, dbtx, ids) + if err != nil { + return interpretError(err) + } + + jobsByID := make(map[int64]*dbsqlc.RiverJob, len(jobs)) + for _, j := range jobs { + jobsByID[j.ID] = j + } + + for i, id := range ids { + job, err := jobRowFromInternal(jobsByID[id]) + if err != nil { + return err + } + + insertRes[i] = &riverdriver.JobInsertFastResult{Job: job} + if uniqueNonceByIndex[i] != "" { + uniqueSkippedAsDuplicate := gjson.GetBytes(job.Metadata, rivercommon.MetadataKeyUniqueNonce).Str != uniqueNonceByIndex[i] + riverUniqueConstraintConflict := bytes.Equal(job.UniqueKey, params.Jobs[i].UniqueKey) && + slices.Contains(job.UniqueStates, job.State) && + slices.Contains(uniquestates.UniqueBitmaskToStates(params.Jobs[i].UniqueStates), params.Jobs[i].State) + if uniqueSkippedAsDuplicate && !riverUniqueConstraintConflict { + // Unlike PostgreSQL and SQLite, MySQL's ON DUPLICATE KEY + // cannot target River's unique-job index. Reject a conflict + // on any other unique constraint instead of silently treating + // the unrelated row as the requested unique job. + return fmt.Errorf("job insert at index %d conflicted on a non-River unique constraint", i) + } + insertRes[i].UniqueSkippedAsDuplicate = uniqueSkippedAsDuplicate + job.Metadata, err = sjson.DeleteBytes(job.Metadata, rivercommon.MetadataKeyUniqueNonce) + if err != nil { + return fmt.Errorf("error removing unique nonce from returned job metadata: %w", err) + } + } + } + + if len(idsWithNonce) > 0 { + if err := dbsqlc.New().JobClearUniqueNonce(ctx, dbtx, idsWithNonce); err != nil { + return interpretError(err) + } + } + + return nil + }); err != nil { + return nil, err + } + + return insertRes, nil +} + +func (e *Executor) JobInsertFastManyNoReturning(ctx context.Context, params *riverdriver.JobInsertFastManyParams) (int, error) { + // RowsAffected can't distinguish a no-op duplicate from an insert when the + // DSN enables clientFoundRows. Fall back to nonce-based duplicate detection + // for any batch where a conflict is possible. + if slices.ContainsFunc(params.Jobs, func(job *riverdriver.JobInsertFastParams) bool { + return len(job.UniqueKey) > 0 + }) { + results, err := e.JobInsertFastMany(ctx, params) + if err != nil { + return 0, err + } + + numInserted := 0 + for _, result := range results { + if !result.UniqueSkippedAsDuplicate { + numInserted++ + } + } + return numInserted, nil + } + + if err := dbutil.WithTx(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) error { + baseCtx := ctx + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + insertCtx := schemaTemplateParam(baseCtx, params.Schema) + insertCtx = sqlctemplate.WithReplacements(insertCtx, map[string]sqlctemplate.Replacement{ + "on_duplicate_key": {Value: ""}, + }, nil) + + for _, jobParams := range params.Jobs { + insertParams, err := jobInsertFastParams(jobParams, "") + if err != nil { + return err + } + + if _, err := dbsqlc.New().JobInsertFast(insertCtx, dbtx, insertParams); err != nil { + return interpretError(err) + } + } + + return nil + }); err != nil { + return 0, err + } + + return len(params.Jobs), nil +} + +// jobInsertFastParams builds the common insert parameters for a single job. +// If uniqueNonce is non-empty, it's set in the job's metadata for duplicate +// detection. +func jobInsertFastParams(params *riverdriver.JobInsertFastParams, uniqueNonce string) (*dbsqlc.JobInsertFastParams, error) { + metadata := sliceutil.FirstNonEmpty(params.Metadata, []byte("{}")) + if uniqueNonce != "" { + var err error + metadata, err = sjson.SetBytes(metadata, rivercommon.MetadataKeyUniqueNonce, uniqueNonce) + if err != nil { + return nil, fmt.Errorf("error adding unique nonce to job metadata: %w", err) + } + } + + tags, err := json.Marshal(params.Tags) + if err != nil { + return nil, err + } + + var uniqueStates sql.NullInt16 + if params.UniqueStates != 0 { + uniqueStates = sql.NullInt16{Int16: int16(params.UniqueStates), Valid: true} + } + + var id sql.NullInt64 + if params.ID != nil { + id = sql.NullInt64{Int64: *params.ID, Valid: true} + } + + return &dbsqlc.JobInsertFastParams{ + ID: id, + Args: params.EncodedArgs, + CreatedAt: utcTimePtr(params.CreatedAt), + Kind: params.Kind, + MaxAttempts: int64(params.MaxAttempts), + Metadata: metadata, + Priority: int16(params.Priority), //nolint:gosec + Queue: params.Queue, + ScheduledAt: utcTimePtr(params.ScheduledAt), + State: string(params.State), + Tags: tags, + UniqueKey: nullStringFromBytes(params.UniqueKey), + UniqueStates: uniqueStates, + }, nil +} + +func (e *Executor) JobInsertFull(ctx context.Context, params *riverdriver.JobInsertFullParams) (*rivertype.JobRow, error) { + var attemptedBy []byte + if params.AttemptedBy != nil { + var err error + attemptedBy, err = json.Marshal(params.AttemptedBy) + if err != nil { + return nil, err + } + } + + var errorsData []byte + if len(params.Errors) > 0 { + var err error + errorsData, err = json.Marshal(sliceutil.Map(params.Errors, func(e []byte) json.RawMessage { return json.RawMessage(e) })) + if err != nil { + return nil, err + } + } + + tags, err := json.Marshal(params.Tags) + if err != nil { + return nil, err + } + + var uniqueStates sql.NullInt16 + if params.UniqueStates != 0 { + uniqueStates = sql.NullInt16{Int16: int16(params.UniqueStates), Valid: true} + } + + return dbutil.WithTxV(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) (*rivertype.JobRow, error) { + ctx = schemaTemplateParam(ctx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + + lastInsertID, err := dbsqlc.New().JobInsertFullExec(ctx, dbtx, &dbsqlc.JobInsertFullExecParams{ + Attempt: int64(params.Attempt), + AttemptedAt: nullTimeFromPtr(params.AttemptedAt), + AttemptedBy: attemptedBy, + Args: params.EncodedArgs, + CreatedAt: utcTimePtr(params.CreatedAt), + Errors: errorsData, + FinalizedAt: nullTimeFromPtr(params.FinalizedAt), + Kind: params.Kind, + MaxAttempts: int64(params.MaxAttempts), + Metadata: sliceutil.FirstNonEmpty(params.Metadata, []byte("{}")), + Priority: int16(params.Priority), //nolint:gosec + Queue: params.Queue, + ScheduledAt: utcTimePtr(params.ScheduledAt), + State: string(params.State), + Tags: tags, + UniqueKey: nullStringFromBytes(params.UniqueKey), + UniqueStates: uniqueStates, + }) + if err != nil { + return nil, interpretError(err) + } + + job, err := dbsqlc.New().JobGetByID(ctx, dbtx, lastInsertID) + if err != nil { + return nil, interpretError(err) + } + return jobRowFromInternal(job) + }) +} + +func (e *Executor) JobInsertFullMany(ctx context.Context, params *riverdriver.JobInsertFullManyParams) ([]*rivertype.JobRow, error) { + insertRes := make([]*rivertype.JobRow, len(params.Jobs)) + + if err := dbutil.WithTx(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) error { + ctx = schemaTemplateParam(ctx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + + ids := make([]int64, len(params.Jobs)) + + for i, jobParams := range params.Jobs { + var attemptedBy []byte + if jobParams.AttemptedBy != nil { + var err error + attemptedBy, err = json.Marshal(jobParams.AttemptedBy) + if err != nil { + return err + } + } + + var errorsData []byte + if len(jobParams.Errors) > 0 { + var err error + errorsData, err = json.Marshal(sliceutil.Map(jobParams.Errors, func(e []byte) json.RawMessage { return json.RawMessage(e) })) + if err != nil { + return err + } + } + + tags, err := json.Marshal(jobParams.Tags) + if err != nil { + return err + } + + var uniqueStates sql.NullInt16 + if jobParams.UniqueStates != 0 { + uniqueStates = sql.NullInt16{Int16: int16(jobParams.UniqueStates), Valid: true} + } + + ids[i], err = dbsqlc.New().JobInsertFullExec(ctx, dbtx, &dbsqlc.JobInsertFullExecParams{ + Attempt: int64(jobParams.Attempt), + AttemptedAt: nullTimeFromPtr(jobParams.AttemptedAt), + AttemptedBy: attemptedBy, + Args: jobParams.EncodedArgs, + CreatedAt: utcTimePtr(jobParams.CreatedAt), + Errors: errorsData, + FinalizedAt: nullTimeFromPtr(jobParams.FinalizedAt), + Kind: jobParams.Kind, + MaxAttempts: int64(jobParams.MaxAttempts), + Metadata: sliceutil.FirstNonEmpty(jobParams.Metadata, []byte("{}")), + Priority: int16(jobParams.Priority), //nolint:gosec + Queue: jobParams.Queue, + ScheduledAt: utcTimePtr(jobParams.ScheduledAt), + State: string(jobParams.State), + Tags: tags, + UniqueKey: nullStringFromBytes(jobParams.UniqueKey), + UniqueStates: uniqueStates, + }) + if err != nil { + return interpretError(err) + } + } + + jobs, err := dbsqlc.New().JobGetByIDMany(ctx, dbtx, ids) + if err != nil { + return interpretError(err) + } + + jobsByID := make(map[int64]*dbsqlc.RiverJob, len(jobs)) + for _, j := range jobs { + jobsByID[j.ID] = j + } + + for i, id := range ids { + insertRes[i], err = jobRowFromInternal(jobsByID[id]) + if err != nil { + return err + } + } + + return nil + }); err != nil { + return nil, err + } + + return insertRes, nil +} + +func (e *Executor) JobKindList(ctx context.Context, params *riverdriver.JobKindListParams) ([]string, error) { + exclude := params.Exclude + if len(exclude) == 0 { + exclude = []string{""} + } + + kinds, err := dbsqlc.New().JobKindList(schemaTemplateParam(ctx, params.Schema), e.dbtx, &dbsqlc.JobKindListParams{ + After: params.After, + Exclude: exclude, + Match: params.Match, + Limit: int32(min(params.Max, math.MaxInt32)), //nolint:gosec + }) + if err != nil { + return nil, interpretError(err) + } + return kinds, nil +} + +func (e *Executor) JobList(ctx context.Context, params *riverdriver.JobListParams) ([]*rivertype.JobRow, error) { + ctx = sqlctemplate.WithReplacements(ctx, map[string]sqlctemplate.Replacement{ + "order_by_clause": {Value: params.OrderByClause}, + "where_clause": {Value: params.WhereClause}, + }, params.NamedArgs) + + jobs, err := dbsqlc.New().JobList(schemaTemplateParam(ctx, params.Schema), e.dbtx, params.Max) + if err != nil { + return nil, interpretError(err) + } + return sliceutil.MapError(jobs, jobRowFromInternal) +} + +func (e *Executor) JobRescueMany(ctx context.Context, params *riverdriver.JobRescueManyParams) (*struct{}, error) { + if err := dbutil.WithTx(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) error { + ctx = schemaTemplateParam(ctx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + + for i := range params.ID { + if err := dbsqlc.New().JobRescue(ctx, dbtx, &dbsqlc.JobRescueParams{ + ID: params.ID[i], + Error: params.Error[i], + FinalizedAt: nullTimeFromPtr(params.FinalizedAt[i]), + ScheduledAt: params.ScheduledAt[i].UTC(), + State: params.State[i], + }); err != nil { + return interpretError(err) + } + } + + return nil + }); err != nil { + return nil, err + } + + return &struct{}{}, nil +} + +func (e *Executor) JobRetry(ctx context.Context, params *riverdriver.JobRetryParams) (*rivertype.JobRow, error) { + return dbutil.WithTxV(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) (*rivertype.JobRow, error) { + ctx = schemaTemplateParam(ctx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + + _, err := dbsqlc.New().JobRetryExec(ctx, dbtx, &dbsqlc.JobRetryExecParams{ + ID: params.ID, + Now: nullTimeFromPtr(params.Now), + }) + if err != nil { + return nil, interpretError(err) + } + + job, err := dbsqlc.New().JobGetByID(ctx, dbtx, params.ID) + if err != nil { + return nil, interpretError(err) + } + return jobRowFromInternal(job) + }) +} + +func (e *Executor) JobSchedule(ctx context.Context, params *riverdriver.JobScheduleParams) ([]*riverdriver.JobScheduleResult, error) { + return dbutil.WithTxV(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) ([]*riverdriver.JobScheduleResult, error) { + ctx = schemaTemplateParam(ctx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + + scheduleResults, err := dbsqlc.New().JobSchedule(ctx, dbtx, &dbsqlc.JobScheduleParams{ + Limit: int32(params.Max), //nolint:gosec + Now: nullTimeFromPtr(params.Now), + }) + if err != nil { + return nil, interpretError(err) + } + + var ( + allIDs []int64 + availIDs []int64 + discardIDs []int64 + discardSet = make(map[int64]bool) + ) + + for _, result := range scheduleResults { + allIDs = append(allIDs, result.ID) + if result.ConflictDiscarded != 0 { + discardIDs = append(discardIDs, result.ID) + discardSet[result.ID] = true + } else { + availIDs = append(availIDs, result.ID) + } + } + + if len(availIDs) > 0 { + if err := dbsqlc.New().JobScheduleSetAvailableExec(ctx, dbtx, availIDs); err != nil { + return nil, interpretError(err) + } + } + + if len(discardIDs) > 0 { + if err := dbsqlc.New().JobScheduleSetDiscardedExec(ctx, dbtx, &dbsqlc.JobScheduleSetDiscardedExecParams{ + ID: discardIDs, + Now: nullTimeFromPtr(params.Now), + }); err != nil { + return nil, interpretError(err) + } + } + + if len(allIDs) == 0 { + return nil, nil + } + + updatedJobs, err := dbsqlc.New().JobGetByIDMany(ctx, dbtx, allIDs) + if err != nil { + return nil, interpretError(err) + } + + jobsByID := make(map[int64]*dbsqlc.RiverJob, len(updatedJobs)) + for _, j := range updatedJobs { + jobsByID[j.ID] = j + } + + // Return results in the same order as scheduleResults. + results := make([]*riverdriver.JobScheduleResult, len(scheduleResults)) + for i, sr := range scheduleResults { + job, err := jobRowFromInternal(jobsByID[sr.ID]) + if err != nil { + return nil, err + } + results[i] = &riverdriver.JobScheduleResult{ConflictDiscarded: discardSet[sr.ID], Job: *job} + } + + return results, nil + }) +} + +func (e *Executor) JobSetStateIfRunningMany(ctx context.Context, params *riverdriver.JobSetStateIfRunningManyParams) ([]*rivertype.JobRow, error) { + setRes := make([]*rivertype.JobRow, 0, len(params.ID)) + + if err := dbutil.WithTx(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) error { + ctx = schemaTemplateParam(ctx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + + // Step 1: Execute all state changes. + for i := range params.ID { + setStateParams := &dbsqlc.JobSetStateIfRunningExecParams{ + ID: params.ID[i], + Error: []byte("{}"), + MetadataUpdates: []byte("{}"), + Now: nullTimeFromPtr(params.Now), + State: string(params.State[i]), + } + + if params.Attempt[i] != nil { + setStateParams.AttemptDoUpdate = 1 + setStateParams.Attempt = int64(*params.Attempt[i]) + } + if params.ErrData[i] != nil { + setStateParams.ErrorsDoUpdate = 1 + setStateParams.Error = params.ErrData[i] + } + if params.FinalizedAt[i] != nil { + setStateParams.FinalizedAtDoUpdate = 1 + setStateParams.FinalizedAt = nullTimeFromPtr(params.FinalizedAt[i]) + } + if params.MetadataDoMerge[i] { + setStateParams.MetadataDoMerge = 1 + setStateParams.MetadataUpdates = params.MetadataUpdates[i] + } + if params.ScheduledAt[i] != nil { + setStateParams.ScheduledAtDoUpdate = 1 + setStateParams.ScheduledAt = params.ScheduledAt[i].UTC() + } + + if err := dbsqlc.New().JobSetStateIfRunningExec(ctx, dbtx, setStateParams); err != nil { + return fmt.Errorf("error setting job state: %w", err) + } + } + + // Step 2: Batch fetch all jobs. + jobs, err := dbsqlc.New().JobGetByIDMany(ctx, dbtx, params.ID) + if err != nil { + return interpretError(err) + } + + jobsByID := make(map[int64]*dbsqlc.RiverJob, len(jobs)) + for _, j := range jobs { + jobsByID[j.ID] = j + } + + // Step 3: For jobs that weren't running, merge metadata if requested. + var metadataMergedIDs []int64 + for i := range params.ID { + job := jobsByID[params.ID[i]] + if job == nil { + continue + } + + if rivertype.JobState(job.State) != rivertype.JobStateRunning && params.MetadataDoMerge[i] { + res, err := dbsqlc.New().JobSetMetadataIfNotRunningExec(ctx, dbtx, &dbsqlc.JobSetMetadataIfNotRunningExecParams{ + ID: params.ID[i], + MetadataUpdates: sliceutil.FirstNonEmpty(params.MetadataUpdates[i], []byte("{}")), + }) + if err != nil { + return fmt.Errorf("error setting job metadata: %w", err) + } + + rowsAffected, err := res.RowsAffected() + if err != nil { + return err + } + + if rowsAffected > 0 { + metadataMergedIDs = append(metadataMergedIDs, params.ID[i]) + } + } + } + + // Step 4: Re-fetch jobs that had metadata merged. + if len(metadataMergedIDs) > 0 { + refreshed, err := dbsqlc.New().JobGetByIDMany(ctx, dbtx, metadataMergedIDs) + if err != nil { + return interpretError(err) + } + for _, j := range refreshed { + jobsByID[j.ID] = j + } + } + + // Step 5: Build results in original order. + for _, id := range params.ID { + job := jobsByID[id] + if job == nil { + continue + } + + jobRow, err := jobRowFromInternal(job) + if err != nil { + return err + } + setRes = append(setRes, jobRow) + } + + return nil + }); err != nil { + return nil, err + } + + return setRes, nil +} + +func (e *Executor) JobUpdate(ctx context.Context, params *riverdriver.JobUpdateParams) (*rivertype.JobRow, error) { + metadata := params.Metadata + if metadata == nil { + metadata = []byte("{}") + } + + var metadataDoMerge int64 + if params.MetadataDoMerge { + metadataDoMerge = 1 + } + + return dbutil.WithTxV(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) (*rivertype.JobRow, error) { + ctx = schemaTemplateParam(ctx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + + if err := dbsqlc.New().JobUpdateExec(ctx, dbtx, &dbsqlc.JobUpdateExecParams{ + ID: params.ID, + MetadataDoMerge: metadataDoMerge, + Metadata: metadata, + }); err != nil { + return nil, interpretError(err) + } + + job, err := dbsqlc.New().JobGetByID(ctx, dbtx, params.ID) + if err != nil { + return nil, interpretError(err) + } + + return jobRowFromInternal(job) + }) +} + +func (e *Executor) JobUpdateFull(ctx context.Context, params *riverdriver.JobUpdateFullParams) (*rivertype.JobRow, error) { + attemptedAt := params.AttemptedAt + if attemptedAt != nil { + attemptedAt = ptrutil.Ptr(attemptedAt.UTC()) + } + + attemptedBy, err := json.Marshal(params.AttemptedBy) + if err != nil { + return nil, err + } + + errorsData, err := json.Marshal(sliceutil.Map(params.Errors, func(e []byte) json.RawMessage { return json.RawMessage(e) })) + if err != nil { + return nil, err + } + + finalizedAt := params.FinalizedAt + if finalizedAt != nil { + finalizedAt = ptrutil.Ptr(finalizedAt.UTC()) + } + + metadata := params.Metadata + if metadata == nil { + metadata = []byte("{}") + } + + return dbutil.WithTxV(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) (*rivertype.JobRow, error) { + ctx = schemaTemplateParam(ctx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + + if err := dbsqlc.New().JobUpdateFullExec(ctx, dbtx, &dbsqlc.JobUpdateFullExecParams{ + ID: params.ID, + Attempt: int64(params.Attempt), + AttemptDoUpdate: boolToInt64(params.AttemptDoUpdate), + AttemptedAt: nullTimeFromPtr(attemptedAt), + AttemptedAtDoUpdate: boolToInt64(params.AttemptedAtDoUpdate), + AttemptedBy: attemptedBy, + AttemptedByDoUpdate: boolToInt64(params.AttemptedByDoUpdate), + ErrorsDoUpdate: boolToInt64(params.ErrorsDoUpdate), + Errors: errorsData, + FinalizedAtDoUpdate: boolToInt64(params.FinalizedAtDoUpdate), + FinalizedAt: nullTimeFromPtr(finalizedAt), + MaxAttemptsDoUpdate: boolToInt64(params.MaxAttemptsDoUpdate), + MaxAttempts: int64(min(params.MaxAttempts, math.MaxInt64)), + MetadataDoUpdate: boolToInt64(params.MetadataDoUpdate), + Metadata: metadata, + StateDoUpdate: boolToInt64(params.StateDoUpdate), + State: string(params.State), + }); err != nil { + return nil, interpretError(err) + } + + job, err := dbsqlc.New().JobGetByID(ctx, dbtx, params.ID) + if err != nil { + return nil, interpretError(err) + } + + return jobRowFromInternal(job) + }) +} + +func (e *Executor) LeaderAttemptElect(ctx context.Context, params *riverdriver.LeaderElectParams) (*riverdriver.Leader, error) { + return dbutil.WithTxV(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) (*riverdriver.Leader, error) { + ctx = schemaTemplateParam(ctx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + + err := dbsqlc.New().LeaderAttemptElectExec(ctx, dbtx, &dbsqlc.LeaderAttemptElectExecParams{ + LeaderID: params.LeaderID, + Now: nullTimeFromPtr(params.Now), + TTL: params.TTL.Microseconds(), + }) + if err != nil { + if isDuplicateEntry(err) { + return nil, rivertype.ErrNotFound + } + return nil, interpretError(err) + } + + leader, err := dbsqlc.New().LeaderGetElectedLeader(ctx, dbtx) + if err != nil { + return nil, interpretError(err) + } + return leaderFromInternal(leader), nil + }) +} + +func (e *Executor) LeaderAttemptReelect(ctx context.Context, params *riverdriver.LeaderReelectParams) (*riverdriver.Leader, error) { + return dbutil.WithTxV(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) (*riverdriver.Leader, error) { + ctx = schemaTemplateParam(ctx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + + res, err := dbsqlc.New().LeaderAttemptReelectExec(ctx, dbtx, &dbsqlc.LeaderAttemptReelectExecParams{ + ElectedAt: params.ElectedAt.UTC(), + LeaderID: params.LeaderID, + Now: nullTimeFromPtr(params.Now), + TTL: params.TTL.Microseconds(), + }) + if err != nil { + return nil, interpretError(err) + } + + affected, err := res.RowsAffected() + if err != nil { + return nil, err + } + if affected == 0 { + return nil, rivertype.ErrNotFound + } + + leader, err := dbsqlc.New().LeaderGetElectedLeader(ctx, dbtx) + if err != nil { + return nil, interpretError(err) + } + return leaderFromInternal(leader), nil + }) +} + +func (e *Executor) LeaderDeleteExpired(ctx context.Context, params *riverdriver.LeaderDeleteExpiredParams) (int, error) { + numDeleted, err := dbsqlc.New().LeaderDeleteExpired(schemaTemplateParam(ctx, params.Schema), e.dbtx, nullTimeFromPtr(params.Now)) + if err != nil { + return 0, interpretError(err) + } + return int(numDeleted), nil +} + +func (e *Executor) LeaderGetElectedLeader(ctx context.Context, params *riverdriver.LeaderGetElectedLeaderParams) (*riverdriver.Leader, error) { + leader, err := dbsqlc.New().LeaderGetElectedLeader(schemaTemplateParam(ctx, params.Schema), e.dbtx) + if err != nil { + return nil, interpretError(err) + } + return leaderFromInternal(leader), nil +} + +func (e *Executor) LeaderInsert(ctx context.Context, params *riverdriver.LeaderInsertParams) (*riverdriver.Leader, error) { + return dbutil.WithTxV(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) (*riverdriver.Leader, error) { + ctx = schemaTemplateParam(ctx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + + if err := dbsqlc.New().LeaderInsertExec(ctx, dbtx, &dbsqlc.LeaderInsertExecParams{ + ElectedAt: utcTimePtr(params.ElectedAt), + ExpiresAt: utcTimePtr(params.ExpiresAt), + Now: nullTimeFromPtr(params.Now), + LeaderID: params.LeaderID, + TTL: params.TTL.Microseconds(), + }); err != nil { + return nil, interpretError(err) + } + + leader, err := dbsqlc.New().LeaderGetElectedLeader(ctx, dbtx) + if err != nil { + return nil, interpretError(err) + } + return leaderFromInternal(leader), nil + }) +} + +func (e *Executor) LeaderResign(ctx context.Context, params *riverdriver.LeaderResignParams) (bool, error) { + numResigned, err := dbsqlc.New().LeaderResign(schemaTemplateParam(ctx, params.Schema), e.dbtx, &dbsqlc.LeaderResignParams{ + ElectedAt: params.ElectedAt.UTC(), + LeaderID: params.LeaderID, + }) + if err != nil { + return false, interpretError(err) + } + return numResigned > 0, nil +} + +func (e *Executor) MigrationDeleteAssumingMainMany(ctx context.Context, params *riverdriver.MigrationDeleteAssumingMainManyParams) ([]*riverdriver.Migration, error) { + return dbutil.WithTxV(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) ([]*riverdriver.Migration, error) { + ctx = schemaTemplateParam(ctx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + versions := sliceutil.Map(params.Versions, func(v int) int64 { return int64(v) }) + + migrations, err := dbsqlc.New().RiverMigrationDeleteAssumingMainMany(ctx, dbtx, versions) + if err != nil { + return nil, interpretError(err) + } + + if len(versions) > 0 { + if err := dbsqlc.New().RiverMigrationDeleteAssumingMainManyExec(ctx, dbtx, versions); err != nil { + return nil, interpretError(err) + } + } + + return sliceutil.Map(migrations, func(internal *dbsqlc.RiverMigrationDeleteAssumingMainManyRow) *riverdriver.Migration { + return &riverdriver.Migration{ + CreatedAt: internal.CreatedAt.UTC(), + Line: riverdriver.MigrationLineMain, + Version: int(internal.Version), + } + }), nil + }) +} + +func (e *Executor) MigrationDeleteByLineAndVersionMany(ctx context.Context, params *riverdriver.MigrationDeleteByLineAndVersionManyParams) ([]*riverdriver.Migration, error) { + return dbutil.WithTxV(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) ([]*riverdriver.Migration, error) { + ctx = schemaTemplateParam(ctx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + versions := sliceutil.Map(params.Versions, func(v int) int64 { return int64(v) }) + + migrations, err := dbsqlc.New().RiverMigrationDeleteByLineAndVersionMany(ctx, dbtx, &dbsqlc.RiverMigrationDeleteByLineAndVersionManyParams{ + Line: params.Line, + Version: versions, + }) + if err != nil { + return nil, interpretError(err) + } + + if len(versions) > 0 { + if err := dbsqlc.New().RiverMigrationDeleteByLineAndVersionManyExec(ctx, dbtx, &dbsqlc.RiverMigrationDeleteByLineAndVersionManyExecParams{ + Line: params.Line, + Version: versions, + }); err != nil { + return nil, interpretError(err) + } + } + + return sliceutil.Map(migrations, migrationFromInternal), nil + }) +} + +func (e *Executor) MigrationGetAllAssumingMain(ctx context.Context, params *riverdriver.MigrationGetAllAssumingMainParams) ([]*riverdriver.Migration, error) { + migrations, err := dbsqlc.New().RiverMigrationGetAllAssumingMain(schemaTemplateParam(ctx, params.Schema), e.dbtx) + if err != nil { + return nil, interpretError(err) + } + return sliceutil.Map(migrations, func(internal *dbsqlc.RiverMigrationGetAllAssumingMainRow) *riverdriver.Migration { + return &riverdriver.Migration{ + CreatedAt: internal.CreatedAt.UTC(), + Line: riverdriver.MigrationLineMain, + Version: int(internal.Version), + } + }), nil +} + +func (e *Executor) MigrationGetByLine(ctx context.Context, params *riverdriver.MigrationGetByLineParams) ([]*riverdriver.Migration, error) { + migrations, err := dbsqlc.New().RiverMigrationGetByLine(schemaTemplateParam(ctx, params.Schema), e.dbtx, params.Line) + if err != nil { + return nil, interpretError(err) + } + return sliceutil.Map(migrations, migrationFromInternal), nil +} + +func (e *Executor) MigrationInsertMany(ctx context.Context, params *riverdriver.MigrationInsertManyParams) ([]*riverdriver.Migration, error) { + var migrations []*riverdriver.Migration + + if err := dbutil.WithTx(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) error { + ctx = schemaTemplateParam(ctx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + + for _, version := range params.Versions { + if err := dbsqlc.New().RiverMigrationInsertExec(ctx, dbtx, &dbsqlc.RiverMigrationInsertExecParams{ + Line: params.Line, + Version: int64(version), + }); err != nil { + return interpretError(err) + } + } + + versions := sliceutil.Map(params.Versions, func(v int) int64 { return int64(v) }) + + internals, err := dbsqlc.New().RiverMigrationGetByLineAndVersionMany(ctx, dbtx, &dbsqlc.RiverMigrationGetByLineAndVersionManyParams{ + Line: params.Line, + Version: versions, + }) + if err != nil { + return interpretError(err) + } + + migrations = sliceutil.Map(internals, migrationFromInternal) + return nil + }); err != nil { + return nil, err + } + + return migrations, nil +} + +func (e *Executor) MigrationInsertManyAssumingMain(ctx context.Context, params *riverdriver.MigrationInsertManyAssumingMainParams) ([]*riverdriver.Migration, error) { + var migrations []*riverdriver.Migration + + if err := dbutil.WithTx(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) error { + ctx = schemaTemplateParam(ctx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + + for _, version := range params.Versions { + if err := dbsqlc.New().RiverMigrationInsertAssumingMainExec(ctx, dbtx, int64(version)); err != nil { + return interpretError(err) + } + } + + versions := sliceutil.Map(params.Versions, func(v int) int64 { return int64(v) }) + + internals, err := dbsqlc.New().RiverMigrationGetByVersionMany(ctx, dbtx, versions) + if err != nil { + return interpretError(err) + } + + migrations = sliceutil.Map(internals, func(internal *dbsqlc.RiverMigrationGetByVersionManyRow) *riverdriver.Migration { + return &riverdriver.Migration{ + CreatedAt: internal.CreatedAt.UTC(), + Line: riverdriver.MigrationLineMain, + Version: int(internal.Version), + } + }) + return nil + }); err != nil { + return nil, err + } + + return migrations, nil +} + +func (e *Executor) NotificationDeleteBefore(ctx context.Context, params *riverdriver.NotificationDeleteBeforeParams) (int, error) { + numDeleted, err := dbsqlc.New().NotificationDeleteBefore(schemaTemplateParam(ctx, params.Schema), e.dbtx, params.CreatedAtHorizon.UTC()) + return int(numDeleted), interpretError(err) +} + +func (e *Executor) NotifyMany(ctx context.Context, params *riverdriver.NotifyManyParams) error { + if len(params.Payload) < 1 { + return nil + } + + return dbutil.WithTx(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) error { + ctx = schemaTemplateParam(ctx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + + for _, payload := range params.Payload { + if err := dbsqlc.New().NotificationInsert(ctx, dbtx, &dbsqlc.NotificationInsertParams{ + Payload: payload, + Topic: params.Topic, + }); err != nil { + return interpretError(err) + } + } + + return nil + }) +} + +func (e *Executor) PGAdvisoryXactLock(ctx context.Context, key int64) (*struct{}, error) { + // MySQL has GET_LOCK which is similar to PostgreSQL advisory locks, but + // it's session-scoped rather than transaction-scoped. For now, return + // not implemented. + return nil, riverdriver.ErrNotImplemented +} + +func (e *Executor) QueueCreateOrSetUpdatedAt(ctx context.Context, params *riverdriver.QueueCreateOrSetUpdatedAtParams) (*rivertype.Queue, error) { + return dbutil.WithTxV(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) (*rivertype.Queue, error) { + ctx = schemaTemplateParam(ctx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + + if err := dbsqlc.New().QueueCreateOrSetUpdatedAtExec(ctx, dbtx, &dbsqlc.QueueCreateOrSetUpdatedAtExecParams{ + Metadata: sliceutil.FirstNonEmpty(params.Metadata, []byte("{}")), + Name: params.Name, + Now: nullTimeFromPtr(params.Now), + PausedAt: nullTimeFromPtr(params.PausedAt), + UpdatedAt: nullTimeFromPtr(params.UpdatedAt), + }); err != nil { + return nil, interpretError(err) + } + + queue, err := dbsqlc.New().QueueGet(ctx, dbtx, params.Name) + if err != nil { + return nil, interpretError(err) + } + return queueFromInternal(queue), nil + }) +} + +func (e *Executor) QueueDeleteExpired(ctx context.Context, params *riverdriver.QueueDeleteExpiredParams) ([]string, error) { + return dbutil.WithTxV(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) ([]string, error) { + ctx = schemaTemplateParam(ctx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + + queueNames, err := dbsqlc.New().QueueDeleteExpiredSelect(ctx, dbtx, &dbsqlc.QueueDeleteExpiredSelectParams{ + Limit: int32(params.Max), //nolint:gosec + UpdatedAtHorizon: params.UpdatedAtHorizon.UTC(), + }) + if err != nil { + return nil, interpretError(err) + } + + if len(queueNames) > 0 { + if err := dbsqlc.New().QueueDeleteExpiredExec(ctx, dbtx, queueNames); err != nil { + return nil, interpretError(err) + } + } + + return queueNames, nil + }) +} + +func (e *Executor) QueueGet(ctx context.Context, params *riverdriver.QueueGetParams) (*rivertype.Queue, error) { + queue, err := dbsqlc.New().QueueGet(schemaTemplateParam(ctx, params.Schema), e.dbtx, params.Name) + if err != nil { + return nil, interpretError(err) + } + return queueFromInternal(queue), nil +} + +func (e *Executor) QueueList(ctx context.Context, params *riverdriver.QueueListParams) ([]*rivertype.Queue, error) { + queues, err := dbsqlc.New().QueueList(schemaTemplateParam(ctx, params.Schema), e.dbtx, int32(params.Max)) //nolint:gosec + if err != nil { + return nil, interpretError(err) + } + return sliceutil.Map(queues, queueFromInternal), nil +} + +func (e *Executor) QueueNameList(ctx context.Context, params *riverdriver.QueueNameListParams) ([]string, error) { + exclude := params.Exclude + if len(exclude) == 0 { + exclude = []string{""} + } + queueNames, err := dbsqlc.New().QueueNameList(schemaTemplateParam(ctx, params.Schema), e.dbtx, &dbsqlc.QueueNameListParams{ + After: params.After, + Exclude: exclude, + Match: params.Match, + Limit: int32(min(params.Max, math.MaxInt32)), //nolint:gosec + }) + if err != nil { + return nil, interpretError(err) + } + return queueNames, nil +} + +func (e *Executor) QueuePause(ctx context.Context, params *riverdriver.QueuePauseParams) error { + return dbutil.WithTx(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) error { + ctx = schemaTemplateParam(ctx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + now := nullTimeFromPtr(params.Now) + + var ( + res sql.Result + err error + ) + + if params.Name == riverdriver.AllQueuesString { + res, err = dbsqlc.New().QueuePauseAll(ctx, dbtx, &dbsqlc.QueuePauseAllParams{Now: now}) + } else { + res, err = dbsqlc.New().QueuePauseByName(ctx, dbtx, &dbsqlc.QueuePauseByNameParams{ + Name: params.Name, + Now: now, + }) + } + if err != nil { + return interpretError(err) + } + + rowsAffected, err := res.RowsAffected() + if err != nil { + return interpretError(err) + } + + // MySQL only reports rows that actually changed values, not rows + // matched. If no rows were affected for a named queue, check if the + // queue exists before returning ErrNotFound (it may already be paused). + if rowsAffected < 1 && params.Name != riverdriver.AllQueuesString { + if _, err := dbsqlc.New().QueueGet(ctx, dbtx, params.Name); err != nil { + return interpretError(err) + } + } + return nil + }) +} + +func (e *Executor) QueueResume(ctx context.Context, params *riverdriver.QueueResumeParams) error { + return dbutil.WithTx(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) error { + ctx = schemaTemplateParam(ctx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + now := nullTimeFromPtr(params.Now) + + var ( + res sql.Result + err error + ) + + if params.Name == riverdriver.AllQueuesString { + res, err = dbsqlc.New().QueueResumeAll(ctx, dbtx, now) + } else { + res, err = dbsqlc.New().QueueResumeByName(ctx, dbtx, &dbsqlc.QueueResumeByNameParams{ + Name: params.Name, + Now: now, + }) + } + if err != nil { + return interpretError(err) + } + + rowsAffected, err := res.RowsAffected() + if err != nil { + return interpretError(err) + } + + // MySQL only reports rows that actually changed values, not rows + // matched. If no rows were affected for a named queue, check if the + // queue exists before returning ErrNotFound (it may already be resumed). + if rowsAffected < 1 && params.Name != riverdriver.AllQueuesString { + if _, err := dbsqlc.New().QueueGet(ctx, dbtx, params.Name); err != nil { + return interpretError(err) + } + } + return nil + }) +} + +func (e *Executor) QueueUpdate(ctx context.Context, params *riverdriver.QueueUpdateParams) (*rivertype.Queue, error) { + var metadataDoUpdate int64 + if params.MetadataDoUpdate { + metadataDoUpdate = 1 + } + + return dbutil.WithTxV(ctx, e, func(ctx context.Context, execTx riverdriver.ExecutorTx) (*rivertype.Queue, error) { + ctx = schemaTemplateParam(ctx, params.Schema) + dbtx := templateReplaceWrapper{dbtx: e.driver.UnwrapTx(execTx), replacer: &e.driver.replacer} + + if err := dbsqlc.New().QueueUpdateExec(ctx, dbtx, &dbsqlc.QueueUpdateExecParams{ + Metadata: sliceutil.FirstNonEmpty(params.Metadata, []byte("{}")), + MetadataDoUpdate: metadataDoUpdate, + Name: params.Name, + }); err != nil { + return nil, interpretError(err) + } + + queue, err := dbsqlc.New().QueueGet(ctx, dbtx, params.Name) + if err != nil { + return nil, interpretError(err) + } + return queueFromInternal(queue), nil + }) +} + +func (e *Executor) QueryRow(ctx context.Context, sql string, args ...any) riverdriver.Row { + return e.dbtx.QueryRowContext(ctx, sql, args...) +} + +func (e *Executor) SchemaCreate(ctx context.Context, params *riverdriver.SchemaCreateParams) error { + // In MySQL, schemas are databases. Create one if it doesn't exist. + if params.Schema != "" { + _, err := e.dbtx.ExecContext(ctx, "CREATE DATABASE IF NOT EXISTS "+mysqlIdentifier(params.Schema)) + return interpretError(err) + } + return nil +} + +func (e *Executor) SchemaDrop(ctx context.Context, params *riverdriver.SchemaDropParams) error { + if params.Schema != "" { + _, err := e.dbtx.ExecContext(ctx, "DROP DATABASE IF EXISTS "+mysqlIdentifier(params.Schema)) + return interpretError(err) + } + return nil +} + +func (e *Executor) SchemaGetExpired(ctx context.Context, params *riverdriver.SchemaGetExpiredParams) ([]string, error) { + schemas, err := dbsqlc.New().SchemaGetExpired(informationSchemaParam(ctx), e.dbtx, &dbsqlc.SchemaGetExpiredParams{ + BeforeName: params.BeforeName, + Prefix: params.Prefix, + }) + return schemas, interpretError(err) +} + +func (e *Executor) TableExists(ctx context.Context, params *riverdriver.TableExistsParams) (bool, error) { + ctx = informationSchemaParam(ctx) + + exists, err := dbsqlc.New().TableExists(ctx, e.dbtx, &dbsqlc.TableExistsParams{ + TableName: params.Table, + Schema: sql.NullString{String: params.Schema, Valid: params.Schema != ""}, + }) + if err != nil { + return false, interpretError(err) + } + return exists, nil +} + +func (e *Executor) TableTruncate(ctx context.Context, params *riverdriver.TableTruncateParams) error { + var maybeSchema string + if params.Schema != "" { + maybeSchema = mysqlIdentifier(params.Schema) + "." + } + + for _, table := range params.Table { + // MySQL's TRUNCATE TABLE is DDL and can't be used in transactions. + // Use DELETE FROM instead for transactional safety. + _, err := e.dbtx.ExecContext(ctx, "DELETE FROM "+maybeSchema+mysqlIdentifier(table)) + if err != nil { + return interpretError(err) + } + } + + return nil +} + +type ExecutorTx struct { + Executor + + tx *sql.Tx +} + +func (t *ExecutorTx) Begin(ctx context.Context) (riverdriver.ExecutorTx, error) { + executorSubTx := &ExecutorSubTx{ + beginOnce: &savepointutil.BeginOnlyOnce{}, + savepointNum: 0, + tx: t.tx, + } + executorSubTx.Executor = Executor{nil, templateReplaceWrapper{t.tx, &t.driver.replacer}, t.driver, executorSubTx} + return executorSubTx.Begin(ctx) +} + +func (t *ExecutorTx) Commit(ctx context.Context) error { + return t.tx.Commit() +} + +func (t *ExecutorTx) Rollback(ctx context.Context) error { + return t.tx.Rollback() +} + +type ExecutorSubTx struct { + Executor + + beginOnce *savepointutil.BeginOnlyOnce + savepointNum int + tx *sql.Tx +} + +const savepointPrefix = "river_savepoint_" + +func (t *ExecutorSubTx) Begin(ctx context.Context) (riverdriver.ExecutorTx, error) { + if err := t.beginOnce.Begin(); err != nil { + return nil, err + } + + nextSavepointNum := t.savepointNum + 1 + if err := t.Exec(ctx, fmt.Sprintf("SAVEPOINT %s%02d", savepointPrefix, nextSavepointNum)); err != nil { + return nil, err + } + + executorSubTx := &ExecutorSubTx{ + beginOnce: savepointutil.NewBeginOnlyOnce(t.beginOnce), + savepointNum: nextSavepointNum, + tx: t.tx, + } + executorSubTx.Executor = Executor{nil, templateReplaceWrapper{t.tx, &t.driver.replacer}, t.driver, executorSubTx} + + return executorSubTx, nil +} + +func (t *ExecutorSubTx) Commit(ctx context.Context) error { + defer t.beginOnce.Done() + + if t.beginOnce.IsDone() { + return errors.New("tx is closed") + } + + if err := t.Exec(ctx, fmt.Sprintf("RELEASE SAVEPOINT %s%02d", savepointPrefix, t.savepointNum)); err != nil { + // MySQL DDL statements (CREATE TABLE, ALTER TABLE, etc.) cause an + // implicit COMMIT which destroys savepoints. If the savepoint no + // longer exists, the work was already committed by the DDL. + if isSavepointDoesNotExist(err) { + return nil + } + return err + } + + return nil +} + +func (t *ExecutorSubTx) Rollback(ctx context.Context) error { + defer t.beginOnce.Done() + + if t.beginOnce.IsDone() { + return errors.New("tx is closed") + } + + if err := t.Exec(ctx, fmt.Sprintf("ROLLBACK TO SAVEPOINT %s%02d", savepointPrefix, t.savepointNum)); err != nil { + // MySQL DDL statements cause implicit COMMIT which destroys + // savepoints. If the savepoint no longer exists, there's nothing + // to roll back to. + if isSavepointDoesNotExist(err) { + return nil + } + return err + } + + return nil +} + +// isSavepointDoesNotExist checks if a MySQL error indicates that a savepoint +// does not exist (error 1305). This happens when DDL statements cause an +// implicit COMMIT that destroys active savepoints. +func isSavepointDoesNotExist(err error) bool { + var mysqlErr *mysql.MySQLError + return errors.As(err, &mysqlErr) && mysqlErr.Number == 1305 +} + +func isDuplicateEntry(err error) bool { + var mysqlErr *mysql.MySQLError + return errors.As(err, &mysqlErr) && mysqlErr.Number == 1062 +} + +func boolToInt64(b bool) int64 { + if b { + return 1 + } + return 0 +} + +func interpretError(err error) error { + if errors.Is(err, sql.ErrNoRows) { + return rivertype.ErrNotFound + } + return err +} + +// mysqlIdentifier quotes an identifier with backticks for MySQL. +// MySQL uses backticks instead of double quotes for identifier quoting. +func mysqlIdentifier(ident string) string { + return "`" + strings.ReplaceAll(ident, "`", "``") + "`" +} + +type templateReplaceWrapper struct { + dbtx dbsqlc.DBTX + replacer *sqlctemplate.Replacer +} + +func (w templateReplaceWrapper) ExecContext(ctx context.Context, rawSQL string, rawArgs ...any) (sql.Result, error) { + sqlStr, args := w.replacer.Run(ctx, argPlaceholder, rawSQL, rawArgs) + return w.dbtx.ExecContext(ctx, sqlStr, args...) +} + +func (w templateReplaceWrapper) PrepareContext(ctx context.Context, rawSQL string) (*sql.Stmt, error) { + sqlStr, _ := w.replacer.Run(ctx, argPlaceholder, rawSQL, nil) + return w.dbtx.PrepareContext(ctx, sqlStr) +} + +func (w templateReplaceWrapper) QueryContext(ctx context.Context, rawSQL string, rawArgs ...any) (*sql.Rows, error) { + sqlStr, args := w.replacer.Run(ctx, argPlaceholder, rawSQL, rawArgs) + return w.dbtx.QueryContext(ctx, sqlStr, args...) +} + +func (w templateReplaceWrapper) QueryRowContext(ctx context.Context, rawSQL string, rawArgs ...any) *sql.Row { + sqlStr, args := w.replacer.Run(ctx, argPlaceholder, rawSQL, rawArgs) + return w.dbtx.QueryRowContext(ctx, sqlStr, args...) +} + +func jobRowFromInternal(internal *dbsqlc.RiverJob) (*rivertype.JobRow, error) { + attemptedAt := ptrFromNullTime(internal.AttemptedAt) + if attemptedAt != nil { + t := attemptedAt.UTC() + attemptedAt = &t + } + + var attemptedBy []string + if internal.AttemptedBy != nil { + if err := json.Unmarshal(internal.AttemptedBy, &attemptedBy); err != nil { + return nil, fmt.Errorf("error unmarshaling `attempted_by`: %w", err) + } + } + + var errs []rivertype.AttemptError + if internal.Errors != nil { + if err := json.Unmarshal(internal.Errors, &errs); err != nil { + return nil, fmt.Errorf("error unmarshaling `errors`: %w", err) + } + } + + finalizedAt := ptrFromNullTime(internal.FinalizedAt) + if finalizedAt != nil { + t := finalizedAt.UTC() + finalizedAt = &t + } + + var tags []string + if err := json.Unmarshal(internal.Tags, &tags); err != nil { + return nil, fmt.Errorf("error unmarshaling `tags`: %w", err) + } + + var uniqueStatesByte byte + if internal.UniqueStates.Valid { + if internal.UniqueStates.Int16 < 0 || internal.UniqueStates.Int16 > 255 { + return nil, fmt.Errorf("value out of range for byte: %d", internal.UniqueStates.Int16) + } + uniqueStatesByte = byte(internal.UniqueStates.Int16) + } + + return &rivertype.JobRow{ + ID: internal.ID, + Attempt: max(int(internal.Attempt), 0), + AttemptedAt: attemptedAt, + AttemptedBy: attemptedBy, + CreatedAt: internal.CreatedAt.UTC(), + EncodedArgs: internal.Args, + Errors: errs, + FinalizedAt: finalizedAt, + Kind: internal.Kind, + MaxAttempts: max(int(internal.MaxAttempts), 0), + Metadata: internal.Metadata, + Priority: max(int(internal.Priority), 0), + Queue: internal.Queue, + ScheduledAt: internal.ScheduledAt.UTC(), + State: rivertype.JobState(internal.State), + Tags: tags, + UniqueKey: bytesFromNullString(internal.UniqueKey), + UniqueStates: uniquestates.UniqueBitmaskToStates(uniqueStatesByte), + }, nil +} + +func leaderFromInternal(internal *dbsqlc.RiverLeader) *riverdriver.Leader { + return &riverdriver.Leader{ + ElectedAt: internal.ElectedAt.UTC(), + ExpiresAt: internal.ExpiresAt.UTC(), + LeaderID: internal.LeaderID, + } +} + +func migrationFromInternal(internal *dbsqlc.RiverMigration) *riverdriver.Migration { + return &riverdriver.Migration{ + CreatedAt: internal.CreatedAt.UTC(), + Line: internal.Line, + Version: int(internal.Version), + } +} + +func schemaTemplateParam(ctx context.Context, schema string) context.Context { + if schema != "" { + schema = mysqlIdentifier(schema) + "." + } + + return sqlctemplate.WithReplacements(ctx, map[string]sqlctemplate.Replacement{ + "schema": {Value: schema}, + }, nil) +} + +func informationSchemaParam(ctx context.Context) context.Context { + return sqlctemplate.WithReplacements(ctx, map[string]sqlctemplate.Replacement{ + "information_schema": {Stable: true, Value: "INFORMATION_SCHEMA."}, + }, nil) +} + +func queueFromInternal(internal *dbsqlc.RiverQueue) *rivertype.Queue { + pausedAt := ptrFromNullTime(internal.PausedAt) + if pausedAt != nil { + t := pausedAt.UTC() + pausedAt = &t + } + return &rivertype.Queue{ + CreatedAt: internal.CreatedAt.UTC(), + Metadata: internal.Metadata, + Name: internal.Name, + PausedAt: pausedAt, + UpdatedAt: internal.UpdatedAt.UTC(), + } +} + +func nullTimeFromPtr(t *time.Time) sql.NullTime { + if t == nil { + return sql.NullTime{} + } + return sql.NullTime{Time: t.UTC(), Valid: true} +} + +func ptrFromNullTime(t sql.NullTime) *time.Time { + if !t.Valid { + return nil + } + return &t.Time +} + +func nullStringFromBytes(b []byte) sql.NullString { + if b == nil { + return sql.NullString{} + } + return sql.NullString{String: string(b), Valid: true} +} + +func bytesFromNullString(s sql.NullString) []byte { + if !s.Valid { + return nil + } + return []byte(s.String) +} + +func utcTimePtr(t *time.Time) *time.Time { + if t == nil { + return nil + } + return ptrutil.Ptr(t.UTC()) +} diff --git a/riverdriver/rivermysql/river_mysql_driver_test.go b/riverdriver/rivermysql/river_mysql_driver_test.go new file mode 100644 index 00000000..ab303ed4 --- /dev/null +++ b/riverdriver/rivermysql/river_mysql_driver_test.go @@ -0,0 +1,1088 @@ +package rivermysql + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + "fmt" + "sync" + "testing" + "time" + + "github.com/go-sql-driver/mysql" + "github.com/stretchr/testify/require" + + "github.com/riverqueue/river/riverdbtest" + "github.com/riverqueue/river/riverdriver" + "github.com/riverqueue/river/rivershared/riversharedtest" + "github.com/riverqueue/river/rivershared/testsignal" + "github.com/riverqueue/river/rivertype" +) + +// Verify interface compliance. +var _ riverdriver.Driver[*sql.Tx] = New(nil) + +func TestIsDuplicateEntry(t *testing.T) { + t.Parallel() + + require.True(t, isDuplicateEntry(&mysql.MySQLError{Number: 1062})) + require.False(t, isDuplicateEntry(&mysql.MySQLError{Number: 1305})) + require.False(t, isDuplicateEntry(errors.New("unrelated error containing 1062"))) +} + +func TestIsSavepointDoesNotExist(t *testing.T) { + t.Parallel() + + require.True(t, isSavepointDoesNotExist(&mysql.MySQLError{Number: 1305})) + require.False(t, isSavepointDoesNotExist(&mysql.MySQLError{Number: 1062})) + require.False(t, isSavepointDoesNotExist(errors.New("unrelated error containing 1305"))) +} + +func TestInterpretError(t *testing.T) { + t.Parallel() + + require.EqualError(t, interpretError(errors.New("an error")), "an error") + require.ErrorIs(t, interpretError(sql.ErrNoRows), rivertype.ErrNotFound) + require.NoError(t, interpretError(nil)) +} + +func TestSchemaTemplateParam(t *testing.T) { + t.Parallel() + + ctx := context.Background() + + t.Run("NoSchema", func(t *testing.T) { + t.Parallel() + ctx := schemaTemplateParam(ctx, "") + // Just verify it doesn't panic + _ = ctx + }) + + t.Run("WithSchema", func(t *testing.T) { + t.Parallel() + ctx := schemaTemplateParam(ctx, "custom_schema") + _ = ctx + }) +} + +func TestDriverProperties(t *testing.T) { + t.Parallel() + + driver := New(nil) + require.Equal(t, "?", driver.ArgPlaceholder()) + require.Equal(t, riverdriver.DatabaseNameMySQL, driver.DatabaseName()) + require.True(t, driver.SupportsListener()) + require.True(t, driver.SupportsListenNotify()) + require.Equal(t, time.Microsecond, driver.TimePrecision()) + require.False(t, driver.PoolIsSet()) +} + +func TestJobInsertFastManyUniqueConflictWithinBatch(t *testing.T) { + t.Parallel() + + riversharedtest.SkipIfMySQLNotEnabled(t) + + var ( + ctx = t.Context() + driver = New(riversharedtest.DBPoolMySQL(ctx, t)) + schema = riverdbtest.TestSchema(ctx, t, driver, nil) + exec = driver.GetExecutor() + ) + + newJobParams := func() *riverdriver.JobInsertFastParams { + return &riverdriver.JobInsertFastParams{ + EncodedArgs: []byte(`{"encoded": "args"}`), + Kind: "test_kind", + MaxAttempts: 25, + Metadata: []byte(`{"meta": "data"}`), + Priority: 1, + Queue: "default", + State: rivertype.JobStateAvailable, + Tags: []string{"tag"}, + UniqueKey: []byte("unique-key-fast-conflict-within-batch"), + UniqueStates: 0xff, + } + } + + results, err := exec.JobInsertFastMany(ctx, &riverdriver.JobInsertFastManyParams{ + Jobs: []*riverdriver.JobInsertFastParams{newJobParams(), newJobParams()}, + Schema: schema, + }) + require.NoError(t, err) + require.Len(t, results, 2) + require.False(t, results[0].UniqueSkippedAsDuplicate) + require.True(t, results[1].UniqueSkippedAsDuplicate) + require.Equal(t, results[0].Job.ID, results[1].Job.ID) +} + +func TestJobInsertFastManyUnrelatedUniqueConstraint(t *testing.T) { + t.Parallel() + + riversharedtest.SkipIfMySQLNotEnabled(t) + + var ( + ctx = t.Context() + driver = New(riversharedtest.DBPoolMySQL(ctx, t)) + schema = riverdbtest.TestSchema(ctx, t, driver, nil) + exec = driver.GetExecutor() + indexName = "river_job_test_kind_unique" + table = mysqlIdentifier(schema) + "." + mysqlIdentifier("river_job") + ) + + require.NoError(t, exec.Exec(ctx, "CREATE UNIQUE INDEX "+mysqlIdentifier(indexName)+" ON "+table+" (kind)")) + t.Cleanup(func() { + require.NoError(t, exec.Exec(context.Background(), "DROP INDEX "+mysqlIdentifier(indexName)+" ON "+table)) + }) + + newJobParams := func(uniqueStates byte) *riverdriver.JobInsertFastParams { + return &riverdriver.JobInsertFastParams{ + EncodedArgs: []byte(`{"encoded": "args"}`), + Kind: "test_kind", + MaxAttempts: 25, + Metadata: []byte(`{"meta": "data"}`), + Priority: 1, + Queue: "default", + State: rivertype.JobStateAvailable, + Tags: []string{"tag"}, + UniqueKey: []byte("same-unique-key"), + UniqueStates: uniqueStates, + } + } + + results, err := exec.JobInsertFastMany(ctx, &riverdriver.JobInsertFastManyParams{ + Jobs: []*riverdriver.JobInsertFastParams{newJobParams(0xff)}, + Schema: schema, + }) + require.NoError(t, err) + require.Len(t, results, 1) + + _, err = exec.JobInsertFastMany(ctx, &riverdriver.JobInsertFastManyParams{ + // The requested available job's River uniqueness constraint is inactive, + // so matching the existing key must not hide the unrelated kind conflict. + Jobs: []*riverdriver.JobInsertFastParams{newJobParams(1 << 2)}, + Schema: schema, + }) + require.EqualError(t, err, "job insert at index 0 conflicted on a non-River unique constraint") + + jobs, err := exec.JobGetByKindMany(ctx, &riverdriver.JobGetByKindManyParams{Kind: []string{"test_kind"}, Schema: schema}) + require.NoError(t, err) + require.Len(t, jobs, 1) + require.Equal(t, []byte("same-unique-key"), jobs[0].UniqueKey) +} + +func TestJobInsertAndGet(t *testing.T) { + t.Parallel() + + riversharedtest.SkipIfMySQLNotEnabled(t) + + var ( + ctx = t.Context() + driver = New(riversharedtest.DBPoolMySQL(ctx, t)) + schema = riverdbtest.TestSchema(ctx, t, driver, nil) + exec = driver.GetExecutor() + ) + + // Insert a job + job, err := exec.JobInsertFull(ctx, &riverdriver.JobInsertFullParams{ + EncodedArgs: []byte(`{"test": true}`), + Kind: "test_job", + MaxAttempts: 3, + Priority: 1, + Queue: "default", + Schema: schema, + State: rivertype.JobStateAvailable, + Tags: []string{"tag1"}, + }) + require.NoError(t, err) + require.NotNil(t, job) + require.Positive(t, job.ID) + require.Equal(t, "test_job", job.Kind) + require.Equal(t, rivertype.JobStateAvailable, job.State) + require.Equal(t, "default", job.Queue) + require.Equal(t, 1, job.Priority) + require.Equal(t, 3, job.MaxAttempts) + require.JSONEq(t, `{"test": true}`, string(job.EncodedArgs)) + require.Equal(t, []string{"tag1"}, job.Tags) + + // Get the job by ID + fetched, err := exec.JobGetByID(ctx, &riverdriver.JobGetByIDParams{ + ID: job.ID, + Schema: schema, + }) + require.NoError(t, err) + require.Equal(t, job.ID, fetched.ID) + require.Equal(t, job.Kind, fetched.Kind) +} + +func TestJobGetAvailable(t *testing.T) { + t.Parallel() + + riversharedtest.SkipIfMySQLNotEnabled(t) + + var ( + ctx = t.Context() + driver = New(riversharedtest.DBPoolMySQL(ctx, t)) + schema = riverdbtest.TestSchema(ctx, t, driver, nil) + exec = driver.GetExecutor() + ) + + // Insert some available jobs + for i := range 3 { + _, err := exec.JobInsertFull(ctx, &riverdriver.JobInsertFullParams{ + EncodedArgs: []byte(`{}`), + Kind: fmt.Sprintf("job_%d", i), + MaxAttempts: 3, + Priority: 1, + Queue: "default", + Schema: schema, + State: rivertype.JobStateAvailable, + Tags: []string{}, + }) + require.NoError(t, err) + } + + // Get available jobs (needs transaction for FOR UPDATE) + txExec, err := exec.Begin(ctx) + require.NoError(t, err) + defer txExec.Rollback(ctx) + + jobs, err := txExec.JobGetAvailable(ctx, &riverdriver.JobGetAvailableParams{ + ClientID: "test-client", + MaxAttemptedBy: 4, + MaxToLock: 2, + Queue: "default", + Schema: schema, + }) + require.NoError(t, err) + require.Len(t, jobs, 2) + + for _, job := range jobs { + require.Equal(t, rivertype.JobStateRunning, job.State) + } + + require.NoError(t, txExec.Commit(ctx)) +} + +func TestJobGetAvailableConcurrentClaimsAreUnique(t *testing.T) { + t.Parallel() + + riversharedtest.SkipIfMySQLNotEnabled(t) + + const ( + numJobs = 50 + numWorkers = 10 + ) + + var ( + ctx = t.Context() + driver = New(riversharedtest.DBPoolMySQL(ctx, t)) + schema = riverdbtest.TestSchema(ctx, t, driver, nil) + exec = driver.GetExecutor() + ) + + jobs := make([]*riverdriver.JobInsertFastParams, numJobs) + for i := range jobs { + jobs[i] = &riverdriver.JobInsertFastParams{ + EncodedArgs: []byte(`{}`), + Kind: "concurrent_claim", + MaxAttempts: 3, + Priority: 1, + Queue: "concurrent_claims", + State: rivertype.JobStateAvailable, + Tags: []string{}, + } + } + _, err := exec.JobInsertFastMany(ctx, &riverdriver.JobInsertFastManyParams{Jobs: jobs, Schema: schema}) + require.NoError(t, err) + + var ( + claimedIDs []int64 + errs []error + mu sync.Mutex + ready sync.WaitGroup + start = make(chan struct{}) + workers sync.WaitGroup + ) + ready.Add(numWorkers) + workers.Add(numWorkers) + for range numWorkers { + go func() { + defer workers.Done() + ready.Done() + <-start + + claimed, err := exec.JobGetAvailable(ctx, &riverdriver.JobGetAvailableParams{ + ClientID: "concurrent-client", + MaxAttemptedBy: 4, + MaxToLock: numJobs / numWorkers, + Queue: "concurrent_claims", + Schema: schema, + }) + + mu.Lock() + defer mu.Unlock() + if err != nil { + errs = append(errs, err) + return + } + for _, job := range claimed { + claimedIDs = append(claimedIDs, job.ID) + } + }() + } + + ready.Wait() + close(start) + workers.Wait() + + require.NoError(t, errors.Join(errs...)) + require.Len(t, claimedIDs, numJobs) + uniqueClaimedIDs := make(map[int64]struct{}, len(claimedIDs)) + for _, id := range claimedIDs { + uniqueClaimedIDs[id] = struct{}{} + } + require.Len(t, uniqueClaimedIDs, numJobs) +} + +func TestJobDeleteWaitsForConcurrentStateChange(t *testing.T) { + t.Parallel() + + riversharedtest.SkipIfMySQLNotEnabled(t) + + var ( + ctx = t.Context() + dbPool = riversharedtest.DBPoolMySQL(ctx, t) + driver = New(dbPool) + schema = riverdbtest.TestSchema(ctx, t, driver, nil) + exec = driver.GetExecutor() + ) + + job, err := exec.JobInsertFull(ctx, &riverdriver.JobInsertFullParams{ + EncodedArgs: []byte(`{}`), + Kind: "delete_concurrent_state_change", + MaxAttempts: 3, + Priority: 1, + Queue: "default", + Schema: schema, + State: rivertype.JobStateAvailable, + Tags: []string{}, + }) + require.NoError(t, err) + + tx, err := dbPool.BeginTx(ctx, nil) + require.NoError(t, err) + t.Cleanup(func() { _ = tx.Rollback() }) + + _, err = tx.ExecContext(ctx, "UPDATE "+mysqlIdentifier(schema)+".river_job SET state = 'running' WHERE id = ?", job.ID) //nolint:gosec + require.NoError(t, err) + + type deleteResult struct { + err error + job *rivertype.JobRow + } + deleteFinished := testsignal.TestSignal[deleteResult]{} + deleteFinished.Init(t) + deleteStarted := testsignal.TestSignal[struct{}]{} + deleteStarted.Init(t) + + go func() { + deleteStarted.Signal(struct{}{}) + job, err := exec.JobDelete(ctx, &riverdriver.JobDeleteParams{ID: job.ID, Schema: schema}) + deleteFinished.Signal(deleteResult{err: err, job: job}) + }() + + deleteStarted.WaitOrTimeout() + select { + case <-deleteFinished.WaitC(): + require.FailNow(t, "JobDelete returned while the concurrent update was uncommitted") + case <-time.After(250 * time.Millisecond): + } + + require.NoError(t, tx.Commit()) + + result := deleteFinished.WaitOrTimeout() + require.ErrorIs(t, result.err, rivertype.ErrJobRunning) + require.Nil(t, result.job) +} + +func TestJobCancel(t *testing.T) { + t.Parallel() + + riversharedtest.SkipIfMySQLNotEnabled(t) + + var ( + ctx = t.Context() + driver = New(riversharedtest.DBPoolMySQL(ctx, t)) + schema = riverdbtest.TestSchema(ctx, t, driver, nil) + exec = driver.GetExecutor() + ) + + // Insert a job + job, err := exec.JobInsertFull(ctx, &riverdriver.JobInsertFullParams{ + EncodedArgs: []byte(`{}`), + Kind: "test_cancel", + MaxAttempts: 3, + Priority: 1, + Queue: "default", + Schema: schema, + State: rivertype.JobStateAvailable, + Tags: []string{}, + }) + require.NoError(t, err) + + // Cancel the job + cancelled, err := exec.JobCancel(ctx, &riverdriver.JobCancelParams{ + ID: job.ID, + CancelAttemptedAt: time.Now().UTC(), + ControlTopic: "test_topic", + Schema: schema, + }) + require.NoError(t, err) + require.Equal(t, rivertype.JobStateCancelled, cancelled.State) + require.NotNil(t, cancelled.FinalizedAt) +} + +func TestJobDelete(t *testing.T) { + t.Parallel() + + riversharedtest.SkipIfMySQLNotEnabled(t) + + var ( + ctx = t.Context() + driver = New(riversharedtest.DBPoolMySQL(ctx, t)) + schema = riverdbtest.TestSchema(ctx, t, driver, nil) + exec = driver.GetExecutor() + ) + + // Insert a completed job + now := time.Now().UTC().Truncate(time.Microsecond) + job, err := exec.JobInsertFull(ctx, &riverdriver.JobInsertFullParams{ + EncodedArgs: []byte(`{}`), + FinalizedAt: &now, + Kind: "test_delete", + MaxAttempts: 3, + Priority: 1, + Queue: "default", + Schema: schema, + State: rivertype.JobStateCompleted, + Tags: []string{}, + }) + require.NoError(t, err) + + // Delete the job + deleted, err := exec.JobDelete(ctx, &riverdriver.JobDeleteParams{ + ID: job.ID, + Schema: schema, + }) + require.NoError(t, err) + require.Equal(t, job.ID, deleted.ID) + + // Verify it's gone + _, err = exec.JobGetByID(ctx, &riverdriver.JobGetByIDParams{ + ID: job.ID, + Schema: schema, + }) + require.ErrorIs(t, err, rivertype.ErrNotFound) +} + +func TestJobCountByState(t *testing.T) { + t.Parallel() + + riversharedtest.SkipIfMySQLNotEnabled(t) + + var ( + ctx = t.Context() + driver = New(riversharedtest.DBPoolMySQL(ctx, t)) + schema = riverdbtest.TestSchema(ctx, t, driver, nil) + exec = driver.GetExecutor() + ) + + // Insert jobs in different states + for range 3 { + _, err := exec.JobInsertFull(ctx, &riverdriver.JobInsertFullParams{ + EncodedArgs: []byte(`{}`), + Kind: "test_count", + MaxAttempts: 3, + Priority: 1, + Queue: "default", + Schema: schema, + State: rivertype.JobStateAvailable, + Tags: []string{}, + }) + require.NoError(t, err) + } + + count, err := exec.JobCountByState(ctx, &riverdriver.JobCountByStateParams{ + Schema: schema, + State: rivertype.JobStateAvailable, + }) + require.NoError(t, err) + require.Equal(t, 3, count) +} + +func TestLeaderElection(t *testing.T) { + t.Parallel() + + riversharedtest.SkipIfMySQLNotEnabled(t) + + var ( + ctx = t.Context() + driver = New(riversharedtest.DBPoolMySQL(ctx, t)) + schema = riverdbtest.TestSchema(ctx, t, driver, nil) + exec = driver.GetExecutor() + ) + + // Attempt to elect a leader + leader, err := exec.LeaderAttemptElect(ctx, &riverdriver.LeaderElectParams{ + LeaderID: "test-leader", + Schema: schema, + TTL: 30 * time.Second, + }) + require.NoError(t, err) + require.Equal(t, "test-leader", leader.LeaderID) + + // Get the elected leader + fetched, err := exec.LeaderGetElectedLeader(ctx, &riverdriver.LeaderGetElectedLeaderParams{ + Schema: schema, + }) + require.NoError(t, err) + require.Equal(t, "test-leader", fetched.LeaderID) + + // Re-elect should succeed + reelected, err := exec.LeaderAttemptReelect(ctx, &riverdriver.LeaderReelectParams{ + ElectedAt: leader.ElectedAt, + LeaderID: "test-leader", + Schema: schema, + TTL: 30 * time.Second, + }) + require.NoError(t, err) + require.Equal(t, "test-leader", reelected.LeaderID) + + // Resign + resigned, err := exec.LeaderResign(ctx, &riverdriver.LeaderResignParams{ + ElectedAt: leader.ElectedAt, + LeaderID: "test-leader", + LeadershipTopic: "leadership", + Schema: schema, + }) + require.NoError(t, err) + require.True(t, resigned) +} + +func TestListenerConnectDoesNotSkipLowerUncommittedID(t *testing.T) { + t.Parallel() + + riversharedtest.SkipIfMySQLNotEnabled(t) + + var ( + ctx = t.Context() + dbPool = riversharedtest.DBPoolMySQL(ctx, t) + driver = New(dbPool) + schema = riverdbtest.TestSchema(ctx, t, driver, nil) + listener = driver.GetListener(&riverdriver.GetListenenerParams{Schema: schema}).(*Listener) //nolint:forcetypeassert + table = mysqlIdentifier(schema) + "." + mysqlIdentifier("river_notification") + ) + insertNotificationSQL := "INSERT INTO " + table + " (payload, topic) VALUES (?, ?)" //nolint:gosec + + txLower, err := dbPool.BeginTx(ctx, nil) + require.NoError(t, err) + t.Cleanup(func() { _ = txLower.Rollback() }) + + _, err = txLower.ExecContext(ctx, insertNotificationSQL, "lower", "topic") + require.NoError(t, err) + + insertCtx, cancel := context.WithTimeout(ctx, riversharedtest.WaitTimeout()) + defer cancel() + + higherResult, err := dbPool.ExecContext(insertCtx, insertNotificationSQL, "higher", "topic") + require.NoError(t, err) + higherID, err := higherResult.LastInsertId() + require.NoError(t, err) + + connectFinished := testsignal.TestSignal[error]{} + connectFinished.Init(t) + connectStarted := testsignal.TestSignal[struct{}]{} + connectStarted.Init(t) + + go func() { + connectStarted.Signal(struct{}{}) + connectFinished.Signal(listener.Connect(ctx)) + }() + + connectStarted.WaitOrTimeout() + select { + case err := <-connectFinished.WaitC(): + require.NoError(t, err) + require.FailNow(t, "listener connected before the lower transaction committed") + case <-time.After(250 * time.Millisecond): + } + + require.NoError(t, txLower.Commit()) + require.NoError(t, connectFinished.WaitOrTimeout()) + require.Equal(t, higherID, listener.lastID) +} + +func TestListenerWaitForNotificationDoesNotSkipLowerUncommittedID(t *testing.T) { + t.Parallel() + + riversharedtest.SkipIfMySQLNotEnabled(t) + + type waitForNotificationResult struct { + err error + found bool + notification *riverdriver.Notification + } + + var ( + ctx = t.Context() + dbPool = riversharedtest.DBPoolMySQL(ctx, t) + driver = New(dbPool) + schema = riverdbtest.TestSchema(ctx, t, driver, nil) + listener = driver.GetListener(&riverdriver.GetListenenerParams{Schema: schema}).(*Listener) //nolint:forcetypeassert + table = mysqlIdentifier(schema) + "." + mysqlIdentifier("river_notification") + ) + insertNotificationSQL := "INSERT INTO " + table + " (payload, topic) VALUES (?, ?)" //nolint:gosec + + listener.pollInterval = time.Millisecond + require.NoError(t, listener.Connect(ctx)) + t.Cleanup(func() { require.NoError(t, listener.Close(context.Background())) }) + require.NoError(t, listener.Listen(ctx, "topic")) + + txLower, err := dbPool.BeginTx(ctx, nil) + require.NoError(t, err) + t.Cleanup(func() { _ = txLower.Rollback() }) + + _, err = txLower.ExecContext(ctx, insertNotificationSQL, "lower", "topic") + require.NoError(t, err) + + insertCtx, cancel := context.WithTimeout(ctx, riversharedtest.WaitTimeout()) + defer cancel() + + _, err = dbPool.ExecContext(insertCtx, insertNotificationSQL, "higher", "topic") + require.NoError(t, err) + + started := testsignal.TestSignal[struct{}]{} + started.Init(t) + waitForNotification := testsignal.TestSignal[waitForNotificationResult]{} + waitForNotification.Init(t) + + waitCtx, cancel := context.WithCancel(ctx) + defer cancel() + + go func() { + started.Signal(struct{}{}) + + notification, found, err := listener.waitForNotificationOnce(waitCtx) + waitForNotification.Signal(waitForNotificationResult{ + err: err, + found: found, + notification: notification, + }) + }() + + started.WaitOrTimeout() + + select { + case res := <-waitForNotification.WaitC(): + require.NoError(t, res.err) + require.FailNow(t, "listener returned before lower transaction committed", "notification: %+v", res.notification) + case <-time.After(250 * time.Millisecond): + } + + require.NoError(t, txLower.Commit()) + + res := waitForNotification.WaitOrTimeout() + require.NoError(t, res.err) + require.True(t, res.found) + require.Equal(t, &riverdriver.Notification{Payload: "lower", Topic: "topic"}, res.notification) + + waitHigherCtx, cancel := context.WithTimeout(ctx, riversharedtest.WaitTimeout()) + defer cancel() + + notification, err := listener.WaitForNotification(waitHigherCtx) + require.NoError(t, err) + require.Equal(t, &riverdriver.Notification{Payload: "higher", Topic: "topic"}, notification) +} + +func TestQueueOperations(t *testing.T) { + t.Parallel() + + riversharedtest.SkipIfMySQLNotEnabled(t) + + var ( + ctx = t.Context() + driver = New(riversharedtest.DBPoolMySQL(ctx, t)) + schema = riverdbtest.TestSchema(ctx, t, driver, nil) + exec = driver.GetExecutor() + ) + + now := time.Now().UTC().Truncate(time.Microsecond) + + // Create a queue + queue, err := exec.QueueCreateOrSetUpdatedAt(ctx, &riverdriver.QueueCreateOrSetUpdatedAtParams{ + Metadata: []byte(`{}`), + Name: "test_queue", + Now: &now, + Schema: schema, + }) + require.NoError(t, err) + require.Equal(t, "test_queue", queue.Name) + require.Nil(t, queue.PausedAt) + + // Get the queue + fetched, err := exec.QueueGet(ctx, &riverdriver.QueueGetParams{ + Name: "test_queue", + Schema: schema, + }) + require.NoError(t, err) + require.Equal(t, "test_queue", fetched.Name) + + // List queues + queues, err := exec.QueueList(ctx, &riverdriver.QueueListParams{ + Max: 100, + Schema: schema, + }) + require.NoError(t, err) + require.Len(t, queues, 1) + + // Pause queue + err = exec.QueuePause(ctx, &riverdriver.QueuePauseParams{ + Name: "test_queue", + Schema: schema, + }) + require.NoError(t, err) + + paused, err := exec.QueueGet(ctx, &riverdriver.QueueGetParams{ + Name: "test_queue", + Schema: schema, + }) + require.NoError(t, err) + require.NotNil(t, paused.PausedAt) + + // Resume queue + err = exec.QueueResume(ctx, &riverdriver.QueueResumeParams{ + Name: "test_queue", + Schema: schema, + }) + require.NoError(t, err) + + resumed, err := exec.QueueGet(ctx, &riverdriver.QueueGetParams{ + Name: "test_queue", + Schema: schema, + }) + require.NoError(t, err) + require.Nil(t, resumed.PausedAt) +} + +func TestServerTimestampsAreUTC(t *testing.T) { + t.Parallel() + + riversharedtest.SkipIfMySQLNotEnabled(t) + + var ( + ctx = t.Context() + driver = New(riversharedtest.DBPoolMySQL(ctx, t)) + schema = riverdbtest.TestSchema(ctx, t, driver, nil) + ) + + execTx, err := driver.GetExecutor().Begin(ctx) + require.NoError(t, err) + closed := false + t.Cleanup(func() { + if closed { + return + } + _ = execTx.Exec(context.Background(), "SET time_zone = '+00:00'") + _ = execTx.Rollback(context.Background()) + }) + + require.NoError(t, execTx.Exec(ctx, "SET time_zone = '-05:00'")) + + beforeInsert := time.Now().UTC().Add(-time.Second) + job, err := execTx.JobInsertFull(ctx, &riverdriver.JobInsertFullParams{ + EncodedArgs: []byte(`{}`), + Kind: "server_timestamp_utc", + MaxAttempts: 3, + Priority: 1, + Queue: "default", + Schema: schema, + State: rivertype.JobStateAvailable, + Tags: []string{}, + }) + require.NoError(t, err) + afterInsert := time.Now().UTC().Add(time.Second) + + require.WithinRange(t, job.CreatedAt, beforeInsert, afterInsert) + require.WithinRange(t, job.ScheduledAt, beforeInsert, afterInsert) + + require.NoError(t, execTx.Exec(ctx, "SET time_zone = '+00:00'")) + require.NoError(t, execTx.Rollback(ctx)) + closed = true +} + +func TestTimeParametersAreNormalizedToUTC(t *testing.T) { + t.Parallel() + + riversharedtest.SkipIfMySQLNotEnabled(t) + + var ( + ctx = t.Context() + driver = New(riversharedtest.DBPoolMySQL(ctx, t)) + schema = riverdbtest.TestSchema(ctx, t, driver, nil) + ) + + localTime := time.Date(2026, time.July, 24, 12, 34, 56, 123456000, time.FixedZone("UTC-7", -7*60*60)) + job, err := driver.GetExecutor().JobInsertFull(ctx, &riverdriver.JobInsertFullParams{ + CreatedAt: &localTime, + EncodedArgs: []byte(`{}`), + Kind: "time_parameters_utc", + MaxAttempts: 3, + Priority: 1, + Queue: "default", + ScheduledAt: &localTime, + Schema: schema, + State: rivertype.JobStateAvailable, + Tags: []string{}, + }) + require.NoError(t, err) + require.Equal(t, localTime.UTC(), job.CreatedAt) + require.Equal(t, localTime.UTC(), job.ScheduledAt) +} + +func TestTransactions(t *testing.T) { + t.Parallel() + + riversharedtest.SkipIfMySQLNotEnabled(t) + + var ( + ctx = t.Context() + driver = New(riversharedtest.DBPoolMySQL(ctx, t)) + schema = riverdbtest.TestSchema(ctx, t, driver, nil) + exec = driver.GetExecutor() + ) + + // Begin and commit + tx, err := exec.Begin(ctx) + require.NoError(t, err) + + _, err = tx.JobInsertFull(ctx, &riverdriver.JobInsertFullParams{ + EncodedArgs: []byte(`{}`), + Kind: "tx_test", + MaxAttempts: 3, + Priority: 1, + Queue: "default", + Schema: schema, + State: rivertype.JobStateAvailable, + Tags: []string{}, + }) + require.NoError(t, err) + + require.NoError(t, tx.Commit(ctx)) + + // Verify job exists + count, err := exec.JobCountByState(ctx, &riverdriver.JobCountByStateParams{ + Schema: schema, + State: rivertype.JobStateAvailable, + }) + require.NoError(t, err) + require.Equal(t, 1, count) + + // Begin and rollback + tx2, err := exec.Begin(ctx) + require.NoError(t, err) + + _, err = tx2.JobInsertFull(ctx, &riverdriver.JobInsertFullParams{ + EncodedArgs: []byte(`{}`), + Kind: "tx_test_rollback", + MaxAttempts: 3, + Priority: 1, + Queue: "default", + Schema: schema, + State: rivertype.JobStateAvailable, + Tags: []string{}, + }) + require.NoError(t, err) + + require.NoError(t, tx2.Rollback(ctx)) + + // Verify second job was rolled back + count, err = exec.JobCountByState(ctx, &riverdriver.JobCountByStateParams{ + Schema: schema, + State: rivertype.JobStateAvailable, + }) + require.NoError(t, err) + require.Equal(t, 1, count) // still 1 +} + +func TestJobInsertFastMany(t *testing.T) { + t.Parallel() + + riversharedtest.SkipIfMySQLNotEnabled(t) + + var ( + ctx = t.Context() + driver = New(riversharedtest.DBPoolMySQL(ctx, t)) + schema = riverdbtest.TestSchema(ctx, t, driver, nil) + exec = driver.GetExecutor() + ) + + results, err := exec.JobInsertFastMany(ctx, &riverdriver.JobInsertFastManyParams{ + Jobs: []*riverdriver.JobInsertFastParams{ + { + EncodedArgs: []byte(`{"i": 1}`), + Kind: "fast_job", + MaxAttempts: 3, + Priority: 1, + Queue: "default", + State: rivertype.JobStateAvailable, + Tags: []string{}, + }, + { + EncodedArgs: []byte(`{"i": 2}`), + Kind: "fast_job", + MaxAttempts: 3, + Priority: 1, + Queue: "default", + State: rivertype.JobStateAvailable, + Tags: []string{}, + }, + }, + Schema: schema, + }) + require.NoError(t, err) + require.Len(t, results, 2) + + for _, result := range results { + require.NotNil(t, result.Job) + require.False(t, result.UniqueSkippedAsDuplicate) + } +} + +func TestJobMetadata(t *testing.T) { + t.Parallel() + + riversharedtest.SkipIfMySQLNotEnabled(t) + + var ( + ctx = t.Context() + driver = New(riversharedtest.DBPoolMySQL(ctx, t)) + schema = riverdbtest.TestSchema(ctx, t, driver, nil) + exec = driver.GetExecutor() + ) + + // Insert a job with metadata + job, err := exec.JobInsertFull(ctx, &riverdriver.JobInsertFullParams{ + EncodedArgs: []byte(`{}`), + Kind: "metadata_test", + MaxAttempts: 3, + Metadata: []byte(`{"key": "value"}`), + Priority: 1, + Queue: "default", + Schema: schema, + State: rivertype.JobStateAvailable, + Tags: []string{}, + }) + require.NoError(t, err) + + var metadata map[string]any + require.NoError(t, json.Unmarshal(job.Metadata, &metadata)) + require.Equal(t, "value", metadata["key"]) + + // Update metadata + updated, err := exec.JobUpdate(ctx, &riverdriver.JobUpdateParams{ + ID: job.ID, + MetadataDoMerge: true, + Metadata: []byte(`{"new_key": "new_value"}`), + Schema: schema, + }) + require.NoError(t, err) + + require.NoError(t, json.Unmarshal(updated.Metadata, &metadata)) + require.Equal(t, "value", metadata["key"]) + require.Equal(t, "new_value", metadata["new_key"]) +} + +func TestJobSchedule(t *testing.T) { + t.Parallel() + + riversharedtest.SkipIfMySQLNotEnabled(t) + + var ( + ctx = t.Context() + driver = New(riversharedtest.DBPoolMySQL(ctx, t)) + schema = riverdbtest.TestSchema(ctx, t, driver, nil) + exec = driver.GetExecutor() + ) + + // Insert a scheduled job with a past scheduled_at + past := time.Now().UTC().Add(-1 * time.Hour).Truncate(time.Microsecond) + _, err := exec.JobInsertFull(ctx, &riverdriver.JobInsertFullParams{ + EncodedArgs: []byte(`{}`), + Kind: "scheduled_job", + MaxAttempts: 3, + Priority: 1, + Queue: "default", + ScheduledAt: &past, + Schema: schema, + State: rivertype.JobStateScheduled, + Tags: []string{}, + }) + require.NoError(t, err) + + // Schedule should find and transition the job + now := time.Now().UTC() + results, err := exec.JobSchedule(ctx, &riverdriver.JobScheduleParams{ + Max: 100, + Now: &now, + Schema: schema, + }) + require.NoError(t, err) + require.Len(t, results, 1) + require.Equal(t, rivertype.JobStateAvailable, results[0].Job.State) +} + +func TestNotifyMany(t *testing.T) { + t.Parallel() + + driver := New(nil) + require.NotNil(t, driver.GetListener(&riverdriver.GetListenenerParams{})) +} + +func TestNotifyManyInsertsNotifications(t *testing.T) { + t.Parallel() + + riversharedtest.SkipIfMySQLNotEnabled(t) + + var ( + ctx = t.Context() + driver = New(riversharedtest.DBPoolMySQL(ctx, t)) + exec = driver.GetExecutor() + schema = riverdbtest.TestSchema(ctx, t, driver, nil) + ) + + err := exec.NotifyMany(ctx, &riverdriver.NotifyManyParams{ + Payload: []string{"test1", "test2"}, + Schema: schema, + Topic: "test_topic", + }) + require.NoError(t, err) + + var count int + require.NoError(t, exec.QueryRow(ctx, "SELECT count(*) FROM "+mysqlIdentifier(schema)+".river_notification WHERE topic = ? AND payload IN (?, ?)", "test_topic", "test1", "test2").Scan(&count)) + require.Equal(t, 2, count) +} + +func TestPGAdvisoryXactLock(t *testing.T) { + t.Parallel() + + riversharedtest.SkipIfMySQLNotEnabled(t) + + var ( + ctx = t.Context() + exec = New(riversharedtest.DBPoolMySQL(ctx, t)).GetExecutor() + ) + + _, err := exec.PGAdvisoryXactLock(ctx, 12345) + require.ErrorIs(t, err, riverdriver.ErrNotImplemented) +} diff --git a/riverdriver/rivermysql/river_mysql_listener.go b/riverdriver/rivermysql/river_mysql_listener.go new file mode 100644 index 00000000..e1624dea --- /dev/null +++ b/riverdriver/rivermysql/river_mysql_listener.go @@ -0,0 +1,290 @@ +package rivermysql + +import ( + "context" + "database/sql" + "errors" + "sync" + "time" + + "github.com/riverqueue/river/riverdriver" + "github.com/riverqueue/river/riverdriver/rivermysql/internal/dbsqlc" + "github.com/riverqueue/river/rivershared/sqlctemplate" +) + +const notificationPollIntervalDefault = 50 * time.Millisecond + +// Listener receives MySQL notifications from the river_notification outbox +// table. MySQL doesn't have a native LISTEN/NOTIFY equivalent, so NotifyMany +// appends rows to river_notification and this listener polls for rows with IDs +// greater than its remembered lastID. Polls use a locking read because InnoDB +// may commit auto-increment IDs out of order across concurrent transactions; a +// normal consistent read could otherwise skip a lower uncommitted notification. +type Listener struct { + afterConnectExec string // should only ever be used in testing + dbPool *sql.DB + isConnected bool + lastID int64 + mu sync.Mutex + pollInterval time.Duration + replacer *sqlctemplate.Replacer + schema string + topics map[string]struct{} +} + +func (l *Listener) Close(context.Context) error { + l.mu.Lock() + defer l.mu.Unlock() + + l.isConnected = false + return nil +} + +func (l *Listener) Connect(ctx context.Context) error { + var ( + afterConnectExec string + dbPool *sql.DB + replacer *sqlctemplate.Replacer + schema string + ) + + l.mu.Lock() + if l.isConnected { + l.mu.Unlock() + return errors.New("connection already established") + } + afterConnectExec = l.afterConnectExec + dbPool = l.dbPool + replacer = l.replacer + schema = l.schema + l.mu.Unlock() + + if dbPool == nil { + return errors.New("database pool is nil") + } + if replacer == nil { + replacer = &sqlctemplate.Replacer{UnnumberedPlaceholders: true} + } + + if afterConnectExec != "" { + if _, err := dbPool.ExecContext(ctx, afterConnectExec); err != nil { + return err + } + } + + // Establish the initial high-water mark with a locking scan. InnoDB allocates + // auto-increment IDs before commit, so a higher ID can commit while a lower + // ID is still uncommitted. A normal MAX(id) read would advance past the lower + // row and permanently skip it when it commits. Scan the IDs in ascending + // order with FOR UPDATE so the read reaches and waits for every lower + // uncommitted row before establishing the high-water mark. In particular, + // this must not be reduced to an aggregate MAX or descending LIMIT query, + // neither of which has to encounter the lower row. + tx, err := dbPool.BeginTx(ctx, nil) + if err != nil { + return err + } + defer tx.Rollback() + + ids, err := dbsqlc.New().NotificationGetIDsForUpdate(schemaTemplateParam(ctx, schema), notificationDBTX(tx, replacer)) + if err != nil { + return err + } + var lastID int64 + if len(ids) > 0 { + lastID = ids[len(ids)-1] + } + if err := tx.Commit(); err != nil { + return err + } + + l.mu.Lock() + defer l.mu.Unlock() + + if l.isConnected { + return errors.New("connection already established") + } + + l.isConnected = true + l.lastID = lastID + + return nil +} + +func (l *Listener) Listen(_ context.Context, topic string) error { + l.mu.Lock() + defer l.mu.Unlock() + + if !l.isConnected { + return errors.New("listener is not connected") + } + + if l.topics == nil { + l.topics = make(map[string]struct{}) + } + + l.topics[topic] = struct{}{} + return nil +} + +func (l *Listener) Ping(ctx context.Context) error { + dbPool, err := l.stateDBPool() + if err != nil { + return err + } + return dbPool.PingContext(ctx) +} + +func (l *Listener) Schema() string { + l.mu.Lock() + defer l.mu.Unlock() + + return l.schema +} + +func (l *Listener) SetAfterConnectExec(sql string) { + l.mu.Lock() + defer l.mu.Unlock() + + l.afterConnectExec = sql +} + +func (l *Listener) Unlisten(_ context.Context, topic string) error { + l.mu.Lock() + defer l.mu.Unlock() + + if !l.isConnected { + return errors.New("listener is not connected") + } + + delete(l.topics, topic) + return nil +} + +func (l *Listener) WaitForNotification(ctx context.Context) (*riverdriver.Notification, error) { + for { + if err := ctx.Err(); err != nil { + return nil, err + } + + notification, found, err := l.waitForNotificationOnce(ctx) + if errors.Is(err, sql.ErrNoRows) { + if err := l.waitForNextPoll(ctx); err != nil { + return nil, err + } + continue + } + if err != nil { + if ctx.Err() != nil { + return nil, ctx.Err() + } + return nil, err + } + if found { + return notification, nil + } + } +} + +func (l *Listener) stateDBPool() (*sql.DB, error) { + l.mu.Lock() + defer l.mu.Unlock() + + if !l.isConnected { + return nil, errors.New("listener is not connected") + } + if l.dbPool == nil { + return nil, errors.New("database pool is nil") + } + + return l.dbPool, nil +} + +func (l *Listener) waitForNextPoll(ctx context.Context) error { + l.mu.Lock() + pollInterval := l.pollInterval + l.mu.Unlock() + + if pollInterval <= 0 { + pollInterval = notificationPollIntervalDefault + } + + timer := time.NewTimer(pollInterval) + defer timer.Stop() + + select { + case <-ctx.Done(): + return ctx.Err() + case <-timer.C: + return nil + } +} + +func (l *Listener) waitForNotificationOnce(ctx context.Context) (*riverdriver.Notification, bool, error) { + var ( + after int64 + dbPool *sql.DB + replacer *sqlctemplate.Replacer + schema string + ) + + l.mu.Lock() + if !l.isConnected { + l.mu.Unlock() + return nil, false, errors.New("listener is not connected") + } + after = l.lastID + dbPool = l.dbPool + replacer = l.replacer + schema = l.schema + l.mu.Unlock() + + if dbPool == nil { + return nil, false, errors.New("database pool is nil") + } + + tx, err := dbPool.BeginTx(ctx, nil) + if err != nil { + return nil, false, err + } + defer tx.Rollback() + + // This must be a locking read in ascending ID order. If a lower ID is still + // uncommitted while a higher ID has committed, the read waits for the lower + // row instead of returning the higher row and advancing lastID past it. + notification, err := dbsqlc.New().NotificationGetAfterForUpdate( + schemaTemplateParam(ctx, schema), + notificationDBTX(tx, replacer), + after, + ) + if err != nil { + return nil, false, err + } + + if err := tx.Commit(); err != nil { + return nil, false, err + } + + l.mu.Lock() + defer l.mu.Unlock() + + if notification.ID > l.lastID { + l.lastID = notification.ID + } + + if _, ok := l.topics[notification.Topic]; !ok { + return nil, false, nil + } + + return &riverdriver.Notification{ + Payload: notification.Payload, + Topic: notification.Topic, + }, true, nil +} + +func notificationDBTX(dbtx dbsqlc.DBTX, replacer *sqlctemplate.Replacer) templateReplaceWrapper { + if replacer == nil { + replacer = &sqlctemplate.Replacer{UnnumberedPlaceholders: true} + } + return templateReplaceWrapper{dbtx, replacer} +} diff --git a/riverdriver/riverpgxv5/go.sum b/riverdriver/riverpgxv5/go.sum index 3f4e22e5..04182411 100644 --- a/riverdriver/riverpgxv5/go.sum +++ b/riverdriver/riverpgxv5/go.sum @@ -15,14 +15,14 @@ github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= -github.com/riverqueue/river v0.40.0 h1:4dynKqqU1P22iPmwWDfDj/YZXnuUTZysXXF3wNHekNw= -github.com/riverqueue/river v0.40.0/go.mod h1:auvB4kHqM97tnshEQxzy2E7aFvJFhl010NB6N29DXXc= -github.com/riverqueue/river/riverdriver v0.40.0 h1:QjBHHh+kaxgUgK9tPumrfx7W14vhAHLIt+0dK40ikI8= -github.com/riverqueue/river/riverdriver v0.40.0/go.mod h1:7wEqxsqvtjGk3hBKGKK3IvdnULlPZltpAfkbMr9M2VY= -github.com/riverqueue/river/rivershared v0.40.0 h1:6ZX1Ok94Nkcx/WpHPhexpfdkRB1BqRWFEWrbDszg+JY= -github.com/riverqueue/river/rivershared v0.40.0/go.mod h1:Z77wB2/ctD+zSvlNDKjxC1SW9c1Y2jyE+PXR2CxBlzQ= -github.com/riverqueue/river/rivertype v0.40.0 h1:6sFLhs0OtkxDZ/iyZIUqaQJb/G1krnb2bh9lfXdCimg= -github.com/riverqueue/river/rivertype v0.40.0/go.mod h1:D1Ad+EaZiaXbQbJcJcfeicXJMBKno0n6UcfKI5Q7DIQ= +github.com/riverqueue/river v0.41.0 h1:E7Yfyhn74IgVaCBKsBpTHoC0LxWJvoG0jKt7Tb+lJZg= +github.com/riverqueue/river v0.41.0/go.mod h1:WXiAF1/2gfUPj3H+WqQG3Q5NW5w9JmRXRTe7U9zybNc= +github.com/riverqueue/river/riverdriver v0.41.0 h1:A5g80n6fCGu9Bp4hSq7gys7mvfOiuWkbQHuFx3iPp7w= +github.com/riverqueue/river/riverdriver v0.41.0/go.mod h1:OeGOhZYoH7Ac0CrVanNAAwD9C0WAXy4LzUEl65Ex0B4= +github.com/riverqueue/river/rivershared v0.41.0 h1:ax3uz5KiqfmcvRQ3kKB6iYs9Nt8J5L1RNfOKoVMWkqw= +github.com/riverqueue/river/rivershared v0.41.0/go.mod h1:FaZ7bxC2DORhyFDVyHKRiZzgQPf2b/IELwPBXnHZVLA= +github.com/riverqueue/river/rivertype v0.41.0 h1:dfscvt1asf1PpeeHTMFxQdV1ZfMoReiiMNFCPPfR0as= +github.com/riverqueue/river/rivertype v0.41.0/go.mod h1:D1Ad+EaZiaXbQbJcJcfeicXJMBKno0n6UcfKI5Q7DIQ= github.com/rogpeppe/go-internal v1.12.0 h1:exVL4IDcn6na9z1rAb56Vxr+CgyK3nn3O+epU5NdKM8= github.com/rogpeppe/go-internal v1.12.0/go.mod h1:E+RYuTGaKKdloAfM02xzb0FW3Paa99yedzYV+kq4uf4= github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= diff --git a/riverdriver/riverpgxv5/river_pgx_v5_driver.go b/riverdriver/riverpgxv5/river_pgx_v5_driver.go index e34fedcb..8016eff3 100644 --- a/riverdriver/riverpgxv5/river_pgx_v5_driver.go +++ b/riverdriver/riverpgxv5/river_pgx_v5_driver.go @@ -63,8 +63,9 @@ func New(dbPool *pgxpool.Pool) *Driver { const argPlaceholder = "$" -func (d *Driver) ArgPlaceholder() string { return argPlaceholder } -func (d *Driver) DatabaseName() string { return riverdriver.DatabaseNamePostgres } +func (d *Driver) ArgPlaceholder() string { return argPlaceholder } +func (d *Driver) DatabaseName() string { return riverdriver.DatabaseNamePostgres } +func (d *Driver) SafeIdentifier(ident string) string { return dbutil.SafeIdentifier(ident) } func (d *Driver) GetExecutor() riverdriver.Executor { return &Executor{templateReplaceWrapper{d.dbPool, &d.replacer}, d} diff --git a/riverdriver/riversqlite/go.sum b/riverdriver/riversqlite/go.sum index 6266a979..f61ee793 100644 --- a/riverdriver/riversqlite/go.sum +++ b/riverdriver/riversqlite/go.sum @@ -14,16 +14,16 @@ github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= -github.com/riverqueue/river v0.40.0 h1:4dynKqqU1P22iPmwWDfDj/YZXnuUTZysXXF3wNHekNw= -github.com/riverqueue/river v0.40.0/go.mod h1:auvB4kHqM97tnshEQxzy2E7aFvJFhl010NB6N29DXXc= -github.com/riverqueue/river/riverdriver v0.40.0 h1:QjBHHh+kaxgUgK9tPumrfx7W14vhAHLIt+0dK40ikI8= -github.com/riverqueue/river/riverdriver v0.40.0/go.mod h1:7wEqxsqvtjGk3hBKGKK3IvdnULlPZltpAfkbMr9M2VY= -github.com/riverqueue/river/riverdriver/riverpgxv5 v0.40.0 h1:WbzXgGukvOqYpsORX4rr03sPZfasjoMckLr4On+Je+4= -github.com/riverqueue/river/riverdriver/riverpgxv5 v0.40.0/go.mod h1:HtgFIcn/eJ9O3f371S1cSb0ttS7xJMMkkTD/K6gQKQg= -github.com/riverqueue/river/rivershared v0.40.0 h1:6ZX1Ok94Nkcx/WpHPhexpfdkRB1BqRWFEWrbDszg+JY= -github.com/riverqueue/river/rivershared v0.40.0/go.mod h1:Z77wB2/ctD+zSvlNDKjxC1SW9c1Y2jyE+PXR2CxBlzQ= -github.com/riverqueue/river/rivertype v0.40.0 h1:6sFLhs0OtkxDZ/iyZIUqaQJb/G1krnb2bh9lfXdCimg= -github.com/riverqueue/river/rivertype v0.40.0/go.mod h1:D1Ad+EaZiaXbQbJcJcfeicXJMBKno0n6UcfKI5Q7DIQ= +github.com/riverqueue/river v0.41.0 h1:E7Yfyhn74IgVaCBKsBpTHoC0LxWJvoG0jKt7Tb+lJZg= +github.com/riverqueue/river v0.41.0/go.mod h1:WXiAF1/2gfUPj3H+WqQG3Q5NW5w9JmRXRTe7U9zybNc= +github.com/riverqueue/river/riverdriver v0.41.0 h1:A5g80n6fCGu9Bp4hSq7gys7mvfOiuWkbQHuFx3iPp7w= +github.com/riverqueue/river/riverdriver v0.41.0/go.mod h1:OeGOhZYoH7Ac0CrVanNAAwD9C0WAXy4LzUEl65Ex0B4= +github.com/riverqueue/river/riverdriver/riverpgxv5 v0.41.0 h1:uHGluToMMWvCnpNkAbEWEmP4YAjdseuriEXGu1dwOp4= +github.com/riverqueue/river/riverdriver/riverpgxv5 v0.41.0/go.mod h1:0ueUi+3eW5fIsibS5VkD2HzQmwLKTAiM/uiiVMX/SME= +github.com/riverqueue/river/rivershared v0.41.0 h1:ax3uz5KiqfmcvRQ3kKB6iYs9Nt8J5L1RNfOKoVMWkqw= +github.com/riverqueue/river/rivershared v0.41.0/go.mod h1:FaZ7bxC2DORhyFDVyHKRiZzgQPf2b/IELwPBXnHZVLA= +github.com/riverqueue/river/rivertype v0.41.0 h1:dfscvt1asf1PpeeHTMFxQdV1ZfMoReiiMNFCPPfR0as= +github.com/riverqueue/river/rivertype v0.41.0/go.mod h1:D1Ad+EaZiaXbQbJcJcfeicXJMBKno0n6UcfKI5Q7DIQ= github.com/rogpeppe/go-internal v1.12.0 h1:exVL4IDcn6na9z1rAb56Vxr+CgyK3nn3O+epU5NdKM8= github.com/rogpeppe/go-internal v1.12.0/go.mod h1:E+RYuTGaKKdloAfM02xzb0FW3Paa99yedzYV+kq4uf4= github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= diff --git a/riverdriver/riversqlite/river_sqlite_driver.go b/riverdriver/riversqlite/river_sqlite_driver.go index 88ad262e..f15c68b7 100644 --- a/riverdriver/riversqlite/river_sqlite_driver.go +++ b/riverdriver/riversqlite/river_sqlite_driver.go @@ -76,8 +76,9 @@ func New(dbPool *sql.DB) *Driver { const argPlaceholder = "?" -func (d *Driver) ArgPlaceholder() string { return argPlaceholder } -func (d *Driver) DatabaseName() string { return riverdriver.DatabaseNameSQLite } +func (d *Driver) ArgPlaceholder() string { return argPlaceholder } +func (d *Driver) DatabaseName() string { return riverdriver.DatabaseNameSQLite } +func (d *Driver) SafeIdentifier(ident string) string { return dbutil.SafeIdentifier(ident) } func (d *Driver) GetExecutor() riverdriver.Executor { return &Executor{d.dbPool, templateReplaceWrapper{d.dbPool, &d.replacer}, d, nil} diff --git a/rivermigrate/river_migrate.go b/rivermigrate/river_migrate.go index 19d15fdf..0270c1a0 100644 --- a/rivermigrate/river_migrate.go +++ b/rivermigrate/river_migrate.go @@ -576,7 +576,7 @@ func (m *Migrator[TTx]) applyMigrations(ctx context.Context, exec riverdriver.Ex var schema string if m.schema != "" { - schema = dbutil.SafeIdentifier(m.schema) + "." + schema = m.driver.SafeIdentifier(m.schema) + "." } schemaReplacement := map[string]sqlctemplate.Replacement{ "schema": {Value: schema}, diff --git a/rivershared/go.sum b/rivershared/go.sum index ddf936f7..3e25c457 100644 --- a/rivershared/go.sum +++ b/rivershared/go.sum @@ -17,14 +17,14 @@ github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= -github.com/riverqueue/river v0.40.0 h1:4dynKqqU1P22iPmwWDfDj/YZXnuUTZysXXF3wNHekNw= -github.com/riverqueue/river v0.40.0/go.mod h1:auvB4kHqM97tnshEQxzy2E7aFvJFhl010NB6N29DXXc= -github.com/riverqueue/river/riverdriver v0.40.0 h1:QjBHHh+kaxgUgK9tPumrfx7W14vhAHLIt+0dK40ikI8= -github.com/riverqueue/river/riverdriver v0.40.0/go.mod h1:7wEqxsqvtjGk3hBKGKK3IvdnULlPZltpAfkbMr9M2VY= -github.com/riverqueue/river/riverdriver/riverpgxv5 v0.40.0 h1:WbzXgGukvOqYpsORX4rr03sPZfasjoMckLr4On+Je+4= -github.com/riverqueue/river/riverdriver/riverpgxv5 v0.40.0/go.mod h1:HtgFIcn/eJ9O3f371S1cSb0ttS7xJMMkkTD/K6gQKQg= -github.com/riverqueue/river/rivertype v0.40.0 h1:6sFLhs0OtkxDZ/iyZIUqaQJb/G1krnb2bh9lfXdCimg= -github.com/riverqueue/river/rivertype v0.40.0/go.mod h1:D1Ad+EaZiaXbQbJcJcfeicXJMBKno0n6UcfKI5Q7DIQ= +github.com/riverqueue/river v0.41.0 h1:E7Yfyhn74IgVaCBKsBpTHoC0LxWJvoG0jKt7Tb+lJZg= +github.com/riverqueue/river v0.41.0/go.mod h1:WXiAF1/2gfUPj3H+WqQG3Q5NW5w9JmRXRTe7U9zybNc= +github.com/riverqueue/river/riverdriver v0.41.0 h1:A5g80n6fCGu9Bp4hSq7gys7mvfOiuWkbQHuFx3iPp7w= +github.com/riverqueue/river/riverdriver v0.41.0/go.mod h1:OeGOhZYoH7Ac0CrVanNAAwD9C0WAXy4LzUEl65Ex0B4= +github.com/riverqueue/river/riverdriver/riverpgxv5 v0.41.0 h1:uHGluToMMWvCnpNkAbEWEmP4YAjdseuriEXGu1dwOp4= +github.com/riverqueue/river/riverdriver/riverpgxv5 v0.41.0/go.mod h1:0ueUi+3eW5fIsibS5VkD2HzQmwLKTAiM/uiiVMX/SME= +github.com/riverqueue/river/rivertype v0.41.0 h1:dfscvt1asf1PpeeHTMFxQdV1ZfMoReiiMNFCPPfR0as= +github.com/riverqueue/river/rivertype v0.41.0/go.mod h1:D1Ad+EaZiaXbQbJcJcfeicXJMBKno0n6UcfKI5Q7DIQ= github.com/rogpeppe/go-internal v1.12.0 h1:exVL4IDcn6na9z1rAb56Vxr+CgyK3nn3O+epU5NdKM8= github.com/rogpeppe/go-internal v1.12.0/go.mod h1:E+RYuTGaKKdloAfM02xzb0FW3Paa99yedzYV+kq4uf4= github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= diff --git a/rivershared/riversharedtest/riversharedtest.go b/rivershared/riversharedtest/riversharedtest.go index a5905ad1..78c8a227 100644 --- a/rivershared/riversharedtest/riversharedtest.go +++ b/rivershared/riversharedtest/riversharedtest.go @@ -112,6 +112,47 @@ var sqliteTestDir = sync.OnceValue(func() string { //nolint:gochecknoglobals return path.Join(rootDir, "sqlite") }) +// A pool and sync.Once to initialize a MySQL pool, invoked by DBPoolMySQL. +var ( + dbPoolMySQL *sql.DB //nolint:gochecknoglobals + dbPoolMySQLOnce sync.Once //nolint:gochecknoglobals + errDBPoolMySQL error //nolint:gochecknoglobals +) + +// SkipIfMySQLNotEnabled skips the current test if MySQL tests are not enabled. +// MySQL tests are opt-in because they require a running MySQL server. Set +// RIVER_MYSQL_TESTS_ENABLED=1 or RIVER_MYSQL_TESTS_ENABLED=true to enable. +func SkipIfMySQLNotEnabled(tb testing.TB) { + tb.Helper() + + val := os.Getenv("RIVER_MYSQL_TESTS_ENABLED") + if val != "1" && val != "true" { //nolint:goconst + tb.Skip("Skipping MySQL tests; set RIVER_MYSQL_TESTS_ENABLED=1 to enable") + } +} + +// DBPoolMySQL gets a lazily initialized database pool for MySQL testing. +// Uses TEST_MYSQL_URL or defaults to root@tcp(localhost:3306)/?. +func DBPoolMySQL(ctx context.Context, tb testing.TB) *sql.DB { + tb.Helper() + + dbPoolMySQLOnce.Do(func() { + dsn := cmp.Or( + os.Getenv("TEST_MYSQL_URL"), + "root@tcp(localhost:3306)/?parseTime=true&multiStatements=true&loc=UTC&time_zone=%27%2B00%3A00%27", + ) + + dbPoolMySQL, errDBPoolMySQL = sql.Open("mysql", dsn) + if errDBPoolMySQL == nil { + errDBPoolMySQL = dbPoolMySQL.PingContext(ctx) + } + }) + require.NoError(tb, errDBPoolMySQL) + require.NotNil(tb, dbPoolMySQL) + + return dbPoolMySQL +} + // DBPoolLibSQL gets a database pool appropriate for use with libSQL (a SQLite // fork) in testing. func DBPoolLibSQL(ctx context.Context, tb testing.TB, schema string) *sql.DB { diff --git a/rivershared/sqlctemplate/sqlc_template.go b/rivershared/sqlctemplate/sqlc_template.go index 1151b7d1..86fc1e11 100644 --- a/rivershared/sqlctemplate/sqlc_template.go +++ b/rivershared/sqlctemplate/sqlc_template.go @@ -89,6 +89,14 @@ type Replacement struct { // be initialized with a constructor. This lets it default to a usable instance // on drivers that may themselves not be initialized. type Replacer struct { + // UnnumberedPlaceholders, when true, causes Run to emit plain `?` + // placeholders instead of numbered `?1`, `?2`, etc. for named args + // injected via the template system. The args slice is reordered to + // match the positional order of placeholders in the SQL. This is + // needed for MySQL, whose database/sql driver does not support + // numbered `?N` syntax. + UnnumberedPlaceholders bool + cache map[replacerCacheKey]string cacheMu sync.RWMutex } @@ -142,6 +150,13 @@ func (r *Replacer) RunSafely(ctx context.Context, argPlaceholder, sql string, ar } cacheKey, cacheEligible := replacerCacheKeyFrom(sql, container) + + // In unnumbered mode, named args are interleaved with sqlc args during a + // left-to-right SQL walk. The cache can't reconstruct this ordering on hit. + if r.UnnumberedPlaceholders && len(container.NamedArgs) > 0 { + cacheEligible = false + } + if cacheEligible { r.cacheMu.RLock() var ( @@ -212,28 +227,36 @@ func (r *Replacer) RunSafely(ctx context.Context, argPlaceholder, sql string, ar } if len(container.NamedArgs) > 0 { - placeholderNum := len(args) - // For the benefit of the test suite's output being predictable, sort // named args before processing them. sortedNamedArgs := maputil.Keys(container.NamedArgs) slices.Sort(sortedNamedArgs) - for _, arg := range sortedNamedArgs { - placeholderNum++ - - var ( - symbol = "@" + arg - symbolIndex = strings.Index(updatedSQL, symbol) - val = container.NamedArgs[arg] - ) - if symbolIndex == -1 { - return "", nil, fmt.Errorf("sqltemplate expected to find named arg %q, but it wasn't present", symbol) + if r.UnnumberedPlaceholders { + var err error + updatedSQL, args, err = replaceUnnumberedArgs(updatedSQL, args, container.NamedArgs, sortedNamedArgs) + if err != nil { + return "", nil, err } + } else { + placeholderNum := len(args) + for _, arg := range sortedNamedArgs { + placeholderNum++ + + var ( + symbol = "@" + arg + symbolIndex = strings.Index(updatedSQL, symbol) + val = container.NamedArgs[arg] + ) + + if symbolIndex == -1 { + return "", nil, fmt.Errorf("sqltemplate expected to find named arg %q, but it wasn't present", symbol) + } - // ReplaceAll because an input parameter may appear multiple times. - updatedSQL = strings.ReplaceAll(updatedSQL, symbol, argPlaceholder+strconv.Itoa(placeholderNum)) - args = append(args, val) + // ReplaceAll because an input parameter may appear multiple times. + updatedSQL = strings.ReplaceAll(updatedSQL, symbol, argPlaceholder+strconv.Itoa(placeholderNum)) + args = append(args, val) + } } } @@ -249,6 +272,148 @@ func (r *Replacer) RunSafely(ctx context.Context, argPlaceholder, sql string, ar return updatedSQL, args, nil } +func replaceUnnumberedArgs(sql string, args []any, namedArgs map[string]any, sortedNamedArgs []string) (string, []any, error) { + // Prefer the longest name when one named argument is a prefix of another. + slices.SortFunc(sortedNamedArgs, func(left, right string) int { + switch { + case len(left) > len(right): + return -1 + case len(left) < len(right): + return 1 + default: + return strings.Compare(left, right) + } + }) + + var ( + inBlockComment bool + inLineComment bool + newArgs = make([]any, 0, len(args)+len(namedArgs)) + out strings.Builder + quote byte + sqlcArgIndex int + usedNamedArgs = make(map[string]bool, len(namedArgs)) + ) + out.Grow(len(sql)) + + for i := 0; i < len(sql); { + if inLineComment { + out.WriteByte(sql[i]) + if sql[i] == '\n' { + inLineComment = false + } + i++ + continue + } + + if inBlockComment { + if sql[i] == '*' && i+1 < len(sql) && sql[i+1] == '/' { + out.WriteString("*/") + i += 2 + inBlockComment = false + continue + } + out.WriteByte(sql[i]) + i++ + continue + } + + if quote != 0 { + out.WriteByte(sql[i]) + if sql[i] == '\\' && i+1 < len(sql) { + out.WriteByte(sql[i+1]) + i += 2 + continue + } + if sql[i] == quote { + if i+1 < len(sql) && sql[i+1] == quote { + out.WriteByte(sql[i+1]) + i += 2 + continue + } + quote = 0 + } + i++ + continue + } + + switch { + case sql[i] == '\'', sql[i] == '"', sql[i] == '`': + quote = sql[i] + out.WriteByte(sql[i]) + i++ + + case sql[i] == '#': + inLineComment = true + out.WriteByte(sql[i]) + i++ + + case sql[i] == '-' && i+2 < len(sql) && sql[i+1] == '-' && sql[i+2] <= ' ': + inLineComment = true + out.WriteString("--") + i += 2 + + case sql[i] == '/' && i+1 < len(sql) && sql[i+1] == '*': + inBlockComment = true + out.WriteString("/*") + i += 2 + + case sql[i] == '@': + matched := false + for _, name := range sortedNamedArgs { + symbol := "@" + name + if !strings.HasPrefix(sql[i:], symbol) { + continue + } + + end := i + len(symbol) + if end < len(sql) && isNamedArgChar(sql[end]) { + continue + } + + out.WriteByte('?') + newArgs = append(newArgs, namedArgs[name]) + usedNamedArgs[name] = true + i = end + matched = true + break + } + if !matched { + out.WriteByte(sql[i]) + i++ + } + + case sql[i] == '?': + if sqlcArgIndex >= len(args) { + return "", nil, errors.New("sqlctemplate found more unnumbered placeholders than sqlc arguments") + } + out.WriteByte('?') + newArgs = append(newArgs, args[sqlcArgIndex]) + sqlcArgIndex++ + i++ + + default: + out.WriteByte(sql[i]) + i++ + } + } + + if sqlcArgIndex != len(args) { + return "", nil, errors.New("sqlctemplate found fewer unnumbered placeholders than sqlc arguments") + } + for _, name := range sortedNamedArgs { + if !usedNamedArgs[name] { + return "", nil, fmt.Errorf("sqltemplate expected to find named arg %q, but it wasn't present", "@"+name) + } + } + + return out.String(), newArgs, nil +} + +func isNamedArgChar(c byte) bool { + return c >= '0' && c <= '9' || c >= 'A' && c <= 'Z' || c == '_' || c >= 'a' && c <= 'z' +} + // WithReplacements adds sqlctemplate templates to the given context (they go in // context because it's the only way to get them down into the innards of sqlc). // namedArgs can also be passed in to replace arguments found in diff --git a/rivershared/sqlctemplate/sqlc_template_test.go b/rivershared/sqlctemplate/sqlc_template_test.go index 119ebd54..5441b26c 100644 --- a/rivershared/sqlctemplate/sqlc_template_test.go +++ b/rivershared/sqlctemplate/sqlc_template_test.go @@ -259,7 +259,7 @@ func TestReplacer(t *testing.T) { AND state = @state; ` - // Initially cached value + // Initially cached value. { ctx := WithReplacements(ctx, map[string]Replacement{ "schema": {Stable: true, Value: "test_schema."}, @@ -420,6 +420,254 @@ func TestReplacer(t *testing.T) { `, updatedSQL) }) + t.Run("UnnumberedPlaceholders_NoNamedArgs", func(t *testing.T) { + t.Parallel() + + replacer := &Replacer{UnnumberedPlaceholders: true} + + ctx := WithReplacements(ctx, map[string]Replacement{ + "schema": {Value: "test_schema."}, + }, nil) + + updatedSQL, args, err := replacer.RunSafely(ctx, "?", ` + SELECT count(*) + FROM /* TEMPLATE: schema */river_job + WHERE state = ?; + `, []any{"available"}) + require.NoError(t, err) + require.Equal(t, []any{"available"}, args) + require.Equal(t, ` + SELECT count(*) + FROM test_schema.river_job + WHERE state = ?; + `, updatedSQL) + }) + + t.Run("UnnumberedPlaceholders_NamedArgsNoInitialArgs", func(t *testing.T) { + t.Parallel() + + replacer := &Replacer{UnnumberedPlaceholders: true} + + ctx := WithReplacements(ctx, map[string]Replacement{ + "where_clause": {Value: "kind = @kind"}, + }, map[string]any{ + "kind": "no_op", + }) + + updatedSQL, args, err := replacer.RunSafely(ctx, "?", ` + SELECT count(*) + FROM river_job + WHERE /* TEMPLATE_BEGIN: where_clause */ true /* TEMPLATE_END */; + `, nil) + require.NoError(t, err) + require.Equal(t, []any{"no_op"}, args) + require.Equal(t, ` + SELECT count(*) + FROM river_job + WHERE kind = ?; + `, updatedSQL) + }) + + t.Run("UnnumberedPlaceholders_NamedArgsWithInitialArgs", func(t *testing.T) { + t.Parallel() + + replacer := &Replacer{UnnumberedPlaceholders: true} + + ctx := WithReplacements(ctx, map[string]Replacement{ + "where_clause": {Value: "kind = @kind"}, + }, map[string]any{ + "kind": "no_op", + }) + + // The named arg @kind appears in the WHERE clause before the + // sqlc-generated ? for LIMIT. UnnumberedPlaceholders reorders args so + // that they match the positional ? order in the final SQL. + updatedSQL, args, err := replacer.RunSafely(ctx, "?", ` + SELECT count(*) + FROM river_job + WHERE /* TEMPLATE_BEGIN: where_clause */ true /* TEMPLATE_END */ + LIMIT ?; + `, []any{100}) + require.NoError(t, err) + require.Equal(t, []any{"no_op", 100}, args) + require.Equal(t, ` + SELECT count(*) + FROM river_job + WHERE kind = ? + LIMIT ?; + `, updatedSQL) + }) + + t.Run("UnnumberedPlaceholders_NamedArgRepeated", func(t *testing.T) { + t.Parallel() + + replacer := &Replacer{UnnumberedPlaceholders: true} + + ctx := WithReplacements(ctx, map[string]Replacement{ + "where_clause": {Value: "kind = @kind OR queue = @kind"}, + }, map[string]any{ + "kind": "no_op", + }) + + // When a named arg appears multiple times, it should produce a ? for + // each occurrence with the value duplicated in the args slice. + updatedSQL, args, err := replacer.RunSafely(ctx, "?", ` + SELECT * + FROM river_job + WHERE /* TEMPLATE_BEGIN: where_clause */ true /* TEMPLATE_END */ + LIMIT ?; + `, []any{100}) + require.NoError(t, err) + require.Equal(t, []any{"no_op", "no_op", 100}, args) + require.Equal(t, ` + SELECT * + FROM river_job + WHERE kind = ? OR queue = ? + LIMIT ?; + `, updatedSQL) + }) + + t.Run("UnnumberedPlaceholders_MultipleNamedArgs", func(t *testing.T) { + t.Parallel() + + replacer := &Replacer{UnnumberedPlaceholders: true} + + ctx := WithReplacements(ctx, map[string]Replacement{ + "schema": {Stable: true, Value: "test_schema."}, + "where_clause": {Value: "kind = @kind AND status = @status"}, + }, map[string]any{ + "kind": "no_op", + "status": "succeeded", + }) + + updatedSQL, args, err := replacer.RunSafely(ctx, "?", ` + SELECT count(*) + FROM /* TEMPLATE: schema */river_job + WHERE /* TEMPLATE_BEGIN: where_clause */ true /* TEMPLATE_END */ + LIMIT ?; + `, []any{100}) + require.NoError(t, err) + require.Equal(t, []any{"no_op", "succeeded", 100}, args) + require.Equal(t, ` + SELECT count(*) + FROM test_schema.river_job + WHERE kind = ? AND status = ? + LIMIT ?; + `, updatedSQL) + }) + + t.Run("UnnumberedPlaceholders_QuotedAndCommentedSymbols", func(t *testing.T) { + t.Parallel() + + replacer := &Replacer{UnnumberedPlaceholders: true} + + ctx := WithReplacements(ctx, map[string]Replacement{ + "where_clause": {Value: "kind = @kind"}, + }, map[string]any{ + "kind": "no_op", + }) + + updatedSQL, args, err := replacer.RunSafely(ctx, "?", ` + SELECT '?' AS question, '@kind' AS literal + FROM river_job + WHERE /* TEMPLATE_BEGIN: where_clause */ true /* TEMPLATE_END */ + /* ? @kind */ + LIMIT ?; + `, []any{100}) + require.NoError(t, err) + require.Equal(t, []any{"no_op", 100}, args) + require.Equal(t, ` + SELECT '?' AS question, '@kind' AS literal + FROM river_job + WHERE kind = ? + /* ? @kind */ + LIMIT ?; + `, updatedSQL) + }) + + t.Run("UnnumberedPlaceholders_NotCachedWithNamedArgs", func(t *testing.T) { + t.Parallel() + + replacer := &Replacer{UnnumberedPlaceholders: true} + + ctx := WithReplacements(ctx, map[string]Replacement{ + "schema": {Stable: true, Value: "test_schema."}, + "where_clause": {Stable: true, Value: "kind = @kind"}, + }, map[string]any{ + "kind": "no_op", + }) + + sql := ` + SELECT count(*) + FROM /* TEMPLATE: schema */river_job + WHERE /* TEMPLATE_BEGIN: where_clause */ true /* TEMPLATE_END */ + LIMIT ?; + ` + + // Unnumbered mode with named args skips caching because the + // cached SQL can't preserve the positional arg ordering. + updatedSQL, args, err := replacer.RunSafely(ctx, "?", sql, []any{100}) + require.NoError(t, err) + require.Equal(t, []any{"no_op", 100}, args) + require.Equal(t, ` + SELECT count(*) + FROM test_schema.river_job + WHERE kind = ? + LIMIT ?; + `, updatedSQL) + + require.Empty(t, replacer.cache) + + // Second call still produces correct results. + updatedSQL, args, err = replacer.RunSafely(ctx, "?", sql, []any{200}) + require.NoError(t, err) + require.Equal(t, []any{"no_op", 200}, args) + require.Equal(t, ` + SELECT count(*) + FROM test_schema.river_job + WHERE kind = ? + LIMIT ?; + `, updatedSQL) + }) + + t.Run("UnnumberedPlaceholders_CachedWithoutNamedArgs", func(t *testing.T) { + t.Parallel() + + replacer := &Replacer{UnnumberedPlaceholders: true} + + ctx := WithReplacements(ctx, map[string]Replacement{ + "schema": {Stable: true, Value: "test_schema."}, + }, nil) + + sql := ` + SELECT count(*) + FROM /* TEMPLATE: schema */river_job + LIMIT ?; + ` + + // Without named args, caching works normally in unnumbered mode. + updatedSQL, args, err := replacer.RunSafely(ctx, "?", sql, []any{100}) + require.NoError(t, err) + require.Equal(t, []any{100}, args) + require.Equal(t, ` + SELECT count(*) + FROM test_schema.river_job + LIMIT ?; + `, updatedSQL) + + require.Len(t, replacer.cache, 1) + + // Second call uses cache. + updatedSQL, args, err = replacer.RunSafely(ctx, "?", sql, []any{200}) + require.NoError(t, err) + require.Equal(t, []any{200}, args) + require.Equal(t, ` + SELECT count(*) + FROM test_schema.river_job + LIMIT ?; + `, updatedSQL) + }) + t.Run("Stress", func(t *testing.T) { t.Parallel()