From 7ef57fe6fd7808cdb62a25f3a905e707ed29d113 Mon Sep 17 00:00:00 2001 From: Blake Gentry Date: Sat, 1 Aug 2026 18:54:21 -0500 Subject: [PATCH 1/2] add tag filtering to job list Applications that need to find jobs by tag currently have to use a driver-specific `Where` predicate. That prevents shared consumers from supporting PostgreSQL and SQLite with the same list query. Add `JobListParams.Tags` with case-insensitive, match-any semantics and carry it through each driver. PostgreSQL compares unnested arrays, while SQLite compares the equivalent JSON array values. Keep the list and delete-many driver parameter layouts synchronized so the existing pointer conversion remains valid. Add shared driver coverage for combined filters, transactions, and tag matching across every supported driver. --- CHANGELOG.md | 4 +++ internal/dblist/db_list.go | 2 ++ job_list_params.go | 12 +++++++ job_list_params_test.go | 8 +++++ riverdriver/river_driver_interface.go | 2 ++ .../internal/dbsqlc/river_job.sql.go | 20 +++++++++-- .../river_database_sql_driver.go | 5 ++- .../riverdrivertest/driver_client_test.go | 34 ++++++++++++++++--- .../riverpgxv5/internal/dbsqlc/river_job.sql | 9 +++++ .../internal/dbsqlc/river_job.sql.go | 20 +++++++++-- riverdriver/riverpgxv5/river_pgx_v5_driver.go | 5 ++- .../riversqlite/internal/dbsqlc/river_job.sql | 9 +++++ .../internal/dbsqlc/river_job.sql.go | 20 +++++++++-- .../riversqlite/river_sqlite_driver.go | 10 +++++- 14 files changed, 144 insertions(+), 16 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index acbf907b..ae4a3413 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,10 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Added + +- Added `JobListParams.Tags` for filtering jobs by one or more case-insensitive tags. [PR #1339](https://github.com/riverqueue/river/pull/1339). + ## [0.42.0] - 2026-07-31 ### Added diff --git a/internal/dblist/db_list.go b/internal/dblist/db_list.go index 6116d8dc..14a2eaf0 100644 --- a/internal/dblist/db_list.go +++ b/internal/dblist/db_list.go @@ -33,6 +33,7 @@ type JobListParams struct { Queues []string Schema string States []rivertype.JobState + Tags []string Where []WherePredicate } @@ -186,6 +187,7 @@ func JobMakeDriverParams(ctx context.Context, params *JobListParams, sqlFragment NamedArgs: namedArgs, OrderByClause: orderByBuilder.String(), Schema: params.Schema, + Tags: params.Tags, WhereClause: whereBuilder.String(), }, nil } diff --git a/job_list_params.go b/job_list_params.go index c9f18016..6135be0b 100644 --- a/job_list_params.go +++ b/job_list_params.go @@ -177,6 +177,7 @@ type JobListParams struct { sortField JobListOrderByField sortOrder SortOrder states []rivertype.JobState + tags []string where []dblist.WherePredicate } @@ -214,6 +215,7 @@ func (p *JobListParams) copy() *JobListParams { sortOrder: p.sortOrder, schema: p.schema, states: append([]rivertype.JobState(nil), p.states...), + tags: append([]string(nil), p.tags...), where: append([]dblist.WherePredicate(nil), p.where...), } } @@ -294,6 +296,7 @@ func (p *JobListParams) toDBParams() (*dblist.JobListParams, error) { Queues: p.queues, Schema: p.schema, States: p.states, + Tags: p.tags, Where: p.where, }, nil } @@ -419,6 +422,15 @@ func (p *JobListParams) States(states ...rivertype.JobState) *JobListParams { return paramsCopy } +// Tags returns an updated filter set that will only return jobs containing at +// least one of the given tags. Tag matching is case-insensitive. +func (p *JobListParams) Tags(tags ...string) *JobListParams { + paramsCopy := p.copy() + paramsCopy.tags = make([]string, len(tags)) + copy(paramsCopy.tags, tags) + return paramsCopy +} + // NamedArgs are named arguments for use with JobListParams.Where. Keys should // look like "my_param", and map to parameters like "@my_param" in SQL queries. // "@" are present in the SQL, but not in the keys of this map. diff --git a/job_list_params_test.go b/job_list_params_test.go index 3665500c..00c50476 100644 --- a/job_list_params_test.go +++ b/job_list_params_test.go @@ -243,4 +243,12 @@ func Test_JobListParams_toDBParams(t *testing.T) { toDBParams() require.EqualError(t, err, "cannot order by finalized_at without finalized state filters") }) + + t.Run("Tags", func(t *testing.T) { + t.Parallel() + + dbParams, err := NewJobListParams().Tags("alpha", "beta").toDBParams() + require.NoError(t, err) + require.Equal(t, []string{"alpha", "beta"}, dbParams.Tags) + }) } diff --git a/riverdriver/river_driver_interface.go b/riverdriver/river_driver_interface.go index 64412aad..022a5074 100644 --- a/riverdriver/river_driver_interface.go +++ b/riverdriver/river_driver_interface.go @@ -409,6 +409,7 @@ type JobDeleteManyParams struct { NamedArgs map[string]any OrderByClause string Schema string + Tags []string WhereClause string } @@ -513,6 +514,7 @@ type JobListParams struct { NamedArgs map[string]any OrderByClause string Schema string + Tags []string WhereClause string } diff --git a/riverdriver/riverdatabasesql/internal/dbsqlc/river_job.sql.go b/riverdriver/riverdatabasesql/internal/dbsqlc/river_job.sql.go index c72d277a..3fa2ca55 100644 --- a/riverdriver/riverdatabasesql/internal/dbsqlc/river_job.sql.go +++ b/riverdriver/riverdatabasesql/internal/dbsqlc/river_job.sql.go @@ -1166,12 +1166,26 @@ 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 */ + AND ( + coalesce(cardinality($1::text[]), 0) = 0 + OR EXISTS ( + SELECT 1 + FROM unnest(tags) AS job_tag(value) + INNER JOIN unnest($1::text[]) AS filter_tag(value) + ON lower(job_tag.value) = lower(filter_tag.value) + ) + ) ORDER BY /* TEMPLATE_BEGIN: order_by_clause */ id /* TEMPLATE_END */ -LIMIT $1::int +LIMIT $2::int ` -func (q *Queries) JobList(ctx context.Context, db DBTX, max int32) ([]*RiverJob, error) { - rows, err := db.QueryContext(ctx, jobList, max) +type JobListParams struct { + Tags []string + Max int32 +} + +func (q *Queries) JobList(ctx context.Context, db DBTX, arg *JobListParams) ([]*RiverJob, error) { + rows, err := db.QueryContext(ctx, jobList, pq.Array(arg.Tags), arg.Max) if err != nil { return nil, err } diff --git a/riverdriver/riverdatabasesql/river_database_sql_driver.go b/riverdriver/riverdatabasesql/river_database_sql_driver.go index 8320b17a..22e01c6d 100644 --- a/riverdriver/riverdatabasesql/river_database_sql_driver.go +++ b/riverdriver/riverdatabasesql/river_database_sql_driver.go @@ -573,7 +573,10 @@ func (e *Executor) JobList(ctx context.Context, params *riverdriver.JobListParam "where_clause": {Value: params.WhereClause}, }, params.NamedArgs) - jobs, err := dbsqlc.New().JobList(schemaTemplateParam(ctx, params.Schema), e.dbtx, params.Max) + jobs, err := dbsqlc.New().JobList(schemaTemplateParam(ctx, params.Schema), e.dbtx, &dbsqlc.JobListParams{ + Max: params.Max, + Tags: params.Tags, + }) if err != nil { return nil, interpretError(err) } diff --git a/riverdriver/riverdrivertest/driver_client_test.go b/riverdriver/riverdrivertest/driver_client_test.go index 923c7a2b..2eda4a5d 100644 --- a/riverdriver/riverdrivertest/driver_client_test.go +++ b/riverdriver/riverdrivertest/driver_client_test.go @@ -516,7 +516,10 @@ func ExerciseClient[TTx any](ctx context.Context, t *testing.T, client, bundle := setup(t) - job := testfactory.Job(ctx, t, bundle.exec, &testfactory.JobOpts{Schema: bundle.schema}) + job := testfactory.Job(ctx, t, bundle.exec, &testfactory.JobOpts{ + Schema: bundle.schema, + Tags: []string{"all-args-tag"}, + }) listRes, err := client.JobList(ctx, river.NewJobListParams(). @@ -524,7 +527,8 @@ func ExerciseClient[TTx any](ctx context.Context, t *testing.T, Kinds(job.Kind). Priorities(int16(min(job.Priority, math.MaxInt16))). //nolint:gosec Queues(job.Queue). - States(job.State), + States(job.State). + Tags("ALL-ARGS-TAG"), ) require.NoError(t, err) require.Len(t, listRes.Jobs, 1) @@ -552,6 +556,24 @@ func ExerciseClient[TTx any](ctx context.Context, t *testing.T, require.Equal(t, job.ID, listRes.Jobs[0].ID) }) + t.Run("JobListTags", func(t *testing.T) { + t.Parallel() + + client, bundle := setup(t) + + var ( + job1 = testfactory.Job(ctx, t, bundle.exec, &testfactory.JobOpts{Schema: bundle.schema, Tags: []string{"alpha", "shared"}}) + job2 = testfactory.Job(ctx, t, bundle.exec, &testfactory.JobOpts{Schema: bundle.schema, Tags: []string{"beta"}}) + _ = testfactory.Job(ctx, t, bundle.exec, &testfactory.JobOpts{Schema: bundle.schema, Tags: []string{"gamma"}}) + ) + + listRes, err := client.JobList(ctx, river.NewJobListParams().Tags("ALPHA", "BETA")) + require.NoError(t, err) + require.Len(t, listRes.Jobs, 2) + require.Equal(t, job1.ID, listRes.Jobs[0].ID) + require.Equal(t, job2.ID, listRes.Jobs[1].ID) + }) + t.Run("JobListTx", func(t *testing.T) { t.Parallel() @@ -578,7 +600,10 @@ func ExerciseClient[TTx any](ctx context.Context, t *testing.T, tx, execTx := beginTx(ctx, t, bundle) - job := testfactory.Job(ctx, t, execTx, &testfactory.JobOpts{Schema: bundle.schema}) + job := testfactory.Job(ctx, t, execTx, &testfactory.JobOpts{ + Schema: bundle.schema, + Tags: []string{"all-args-tag"}, + }) listRes, err := client.JobListTx(ctx, tx, river.NewJobListParams(). @@ -586,7 +611,8 @@ func ExerciseClient[TTx any](ctx context.Context, t *testing.T, Kinds(job.Kind). Priorities(int16(min(job.Priority, math.MaxInt16))). //nolint:gosec Queues(job.Queue). - States(job.State), + States(job.State). + Tags("ALL-ARGS-TAG"), ) require.NoError(t, err) require.Len(t, listRes.Jobs, 1) diff --git a/riverdriver/riverpgxv5/internal/dbsqlc/river_job.sql b/riverdriver/riverpgxv5/internal/dbsqlc/river_job.sql index 509e479b..d9c0e80f 100644 --- a/riverdriver/riverpgxv5/internal/dbsqlc/river_job.sql +++ b/riverdriver/riverpgxv5/internal/dbsqlc/river_job.sql @@ -472,6 +472,15 @@ LIMIT @max; SELECT * FROM /* TEMPLATE: schema */river_job WHERE /* TEMPLATE_BEGIN: where_clause */ true /* TEMPLATE_END */ + AND ( + coalesce(cardinality(@tags::text[]), 0) = 0 + OR EXISTS ( + SELECT 1 + FROM unnest(tags) AS job_tag(value) + INNER JOIN unnest(@tags::text[]) AS filter_tag(value) + ON lower(job_tag.value) = lower(filter_tag.value) + ) + ) ORDER BY /* TEMPLATE_BEGIN: order_by_clause */ id /* TEMPLATE_END */ LIMIT @max::int; diff --git a/riverdriver/riverpgxv5/internal/dbsqlc/river_job.sql.go b/riverdriver/riverpgxv5/internal/dbsqlc/river_job.sql.go index a361baac..3d90dcab 100644 --- a/riverdriver/riverpgxv5/internal/dbsqlc/river_job.sql.go +++ b/riverdriver/riverpgxv5/internal/dbsqlc/river_job.sql.go @@ -1136,12 +1136,26 @@ 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 */ + AND ( + coalesce(cardinality($1::text[]), 0) = 0 + OR EXISTS ( + SELECT 1 + FROM unnest(tags) AS job_tag(value) + INNER JOIN unnest($1::text[]) AS filter_tag(value) + ON lower(job_tag.value) = lower(filter_tag.value) + ) + ) ORDER BY /* TEMPLATE_BEGIN: order_by_clause */ id /* TEMPLATE_END */ -LIMIT $1::int +LIMIT $2::int ` -func (q *Queries) JobList(ctx context.Context, db DBTX, max int32) ([]*RiverJob, error) { - rows, err := db.Query(ctx, jobList, max) +type JobListParams struct { + Tags []string + Max int32 +} + +func (q *Queries) JobList(ctx context.Context, db DBTX, arg *JobListParams) ([]*RiverJob, error) { + rows, err := db.Query(ctx, jobList, arg.Tags, arg.Max) if err != nil { return nil, err } diff --git a/riverdriver/riverpgxv5/river_pgx_v5_driver.go b/riverdriver/riverpgxv5/river_pgx_v5_driver.go index e34fedcb..7dd91bf9 100644 --- a/riverdriver/riverpgxv5/river_pgx_v5_driver.go +++ b/riverdriver/riverpgxv5/river_pgx_v5_driver.go @@ -567,7 +567,10 @@ func (e *Executor) JobList(ctx context.Context, params *riverdriver.JobListParam "where_clause": {Value: params.WhereClause}, }, params.NamedArgs) - jobs, err := dbsqlc.New().JobList(schemaTemplateParam(ctx, params.Schema), e.dbtx, params.Max) + jobs, err := dbsqlc.New().JobList(schemaTemplateParam(ctx, params.Schema), e.dbtx, &dbsqlc.JobListParams{ + Max: params.Max, + Tags: params.Tags, + }) if err != nil { return nil, interpretError(err) } diff --git a/riverdriver/riversqlite/internal/dbsqlc/river_job.sql b/riverdriver/riversqlite/internal/dbsqlc/river_job.sql index 94b433ca..7f2fd3e1 100644 --- a/riverdriver/riversqlite/internal/dbsqlc/river_job.sql +++ b/riverdriver/riversqlite/internal/dbsqlc/river_job.sql @@ -471,6 +471,15 @@ LIMIT @max; SELECT * FROM /* TEMPLATE: schema */river_job WHERE /* TEMPLATE_BEGIN: where_clause */ true /* TEMPLATE_END */ + AND ( + json_array_length(cast(@tags AS blob)) = 0 + OR EXISTS ( + SELECT 1 + FROM json_each(river_job.tags) AS job_tag + INNER JOIN json_each(cast(@tags AS blob)) AS filter_tag + ON lower(cast(job_tag.value AS text)) = lower(cast(filter_tag.value AS text)) + ) + ) ORDER BY /* TEMPLATE_BEGIN: order_by_clause */ id /* TEMPLATE_END */ LIMIT @max; diff --git a/riverdriver/riversqlite/internal/dbsqlc/river_job.sql.go b/riverdriver/riversqlite/internal/dbsqlc/river_job.sql.go index 5bd39b2c..124dd288 100644 --- a/riverdriver/riversqlite/internal/dbsqlc/river_job.sql.go +++ b/riverdriver/riversqlite/internal/dbsqlc/river_job.sql.go @@ -1194,12 +1194,26 @@ const jobList = `-- name: JobList :many SELECT id, json(args), attempt, attempted_at, json(attempted_by), created_at, json(errors), finalized_at, kind, max_attempts, json(metadata), priority, queue, state, scheduled_at, json(tags), unique_key, unique_states FROM /* TEMPLATE: schema */river_job WHERE /* TEMPLATE_BEGIN: where_clause */ true /* TEMPLATE_END */ + AND ( + json_array_length(cast(?1 AS blob)) = 0 + OR EXISTS ( + SELECT 1 + FROM json_each(river_job.tags) AS job_tag + INNER JOIN json_each(cast(?1 AS blob)) AS filter_tag + ON lower(cast(job_tag.value AS text)) = lower(cast(filter_tag.value AS text)) + ) + ) ORDER BY /* TEMPLATE_BEGIN: order_by_clause */ id /* TEMPLATE_END */ -LIMIT ?1 +LIMIT ?2 ` -func (q *Queries) JobList(ctx context.Context, db DBTX, max int64) ([]*RiverJob, error) { - rows, err := db.QueryContext(ctx, jobList, max) +type JobListParams struct { + Tags []byte + Max int64 +} + +func (q *Queries) JobList(ctx context.Context, db DBTX, arg *JobListParams) ([]*RiverJob, error) { + rows, err := db.QueryContext(ctx, jobList, arg.Tags, arg.Max) if err != nil { return nil, err } diff --git a/riverdriver/riversqlite/river_sqlite_driver.go b/riverdriver/riversqlite/river_sqlite_driver.go index 88ad262e..51c0cdff 100644 --- a/riverdriver/riversqlite/river_sqlite_driver.go +++ b/riverdriver/riversqlite/river_sqlite_driver.go @@ -676,7 +676,15 @@ func (e *Executor) JobList(ctx context.Context, params *riverdriver.JobListParam "where_clause": {Value: params.WhereClause}, }, params.NamedArgs) - jobs, err := dbsqlc.New().JobList(schemaTemplateParam(ctx, params.Schema), e.dbtx, int64(params.Max)) + tags, err := json.Marshal(sliceutil.FirstNonEmpty(params.Tags, []string{})) + if err != nil { + return nil, fmt.Errorf("error encoding tags: %w", err) + } + + jobs, err := dbsqlc.New().JobList(schemaTemplateParam(ctx, params.Schema), e.dbtx, &dbsqlc.JobListParams{ + Max: int64(params.Max), + Tags: tags, + }) if err != nil { return nil, interpretError(err) } From 478b870113b1e37342e3a0162b904d84a10f153b Mon Sep 17 00:00:00 2001 From: Blake Gentry Date: Sat, 1 Aug 2026 20:11:40 -0500 Subject: [PATCH 2/2] make tag list semantics explicit The tag list filter exposes match-any behavior through an ambiguous method name and folds case even though River preserves tag spelling. Its static SQL also adds work to every list query and prevents PostgreSQL from using native array operators. Replace it with exact TagsAny and TagsAll filters. Build their predicates only when requested through driver-specific fragments, using PostgreSQL overlap and containment operators and equivalent SQLite JSON predicates. Define delete parameters from list parameters so their pointer conversion stays structurally safe. Expand shared driver coverage for any, all, combined, case-sensitive, empty, pagination, transaction, and low-level fragment behavior. --- CHANGELOG.md | 2 +- client.go | 6 +- insert_opts.go | 6 +- internal/dblist/db_list.go | 56 +++++++++++++---- internal/dblist/db_list_test.go | 10 ++-- job_list_params.go | 36 ++++++++--- job_list_params_test.go | 38 +++++++++++- riverdriver/river_driver_interface.go | 26 +++++--- .../internal/dbsqlc/river_job.sql.go | 20 +------ .../river_database_sql_driver.go | 13 ++-- .../riverdrivertest/driver_client_test.go | 60 +++++++++++++++---- riverdriver/riverdrivertest/sql_fragments.go | 53 ++++++++++++++++ .../riverpgxv5/internal/dbsqlc/river_job.sql | 9 --- .../internal/dbsqlc/river_job.sql.go | 20 +------ riverdriver/riverpgxv5/river_pgx_v5_driver.go | 13 ++-- .../riversqlite/internal/dbsqlc/river_job.sql | 9 --- .../internal/dbsqlc/river_job.sql.go | 20 +------ .../riversqlite/river_sqlite_driver.go | 41 ++++++++++--- rivertype/river_type.go | 5 +- 19 files changed, 299 insertions(+), 144 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index ae4a3413..6df49a17 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -9,7 +9,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Added -- Added `JobListParams.Tags` for filtering jobs by one or more case-insensitive tags. [PR #1339](https://github.com/riverqueue/river/pull/1339). +- Added `JobListParams.TagsAll` and `JobListParams.TagsAny` for filtering jobs that match every or any exact tag, respectively. [PR #1339](https://github.com/riverqueue/river/pull/1339). ## [0.42.0] - 2026-07-31 diff --git a/client.go b/client.go index 6cee0af1..bf07bfc0 100644 --- a/client.go +++ b/client.go @@ -2398,7 +2398,7 @@ func (c *Client[TTx]) jobDeleteMany(ctx context.Context, exec riverdriver.Execut return nil, errors.New("delete with no filters not allowed to prevent accidental deletion of all jobs; either specify a predicate (e.g. JobDeleteManyParams.IDs, JobDeleteManyParams.Kinds, ...) or call JobDeleteManyParams.All") } - listParams, err := dblist.JobMakeDriverParams(ctx, params.toDBParams(), c.driver.SQLFragmentColumnIn) + listParams, err := dblist.JobMakeDriverParams(ctx, params.toDBParams(), c.driver) if err != nil { return nil, err } @@ -2451,7 +2451,7 @@ func (c *Client[TTx]) JobList(ctx context.Context, params *JobListParams) (*JobL return nil, err } - listParams, err := dblist.JobMakeDriverParams(ctx, dbParams, c.driver.SQLFragmentColumnIn) + listParams, err := dblist.JobMakeDriverParams(ctx, dbParams, c.driver) if err != nil { return nil, err } @@ -2492,7 +2492,7 @@ func (c *Client[TTx]) JobListTx(ctx context.Context, tx TTx, params *JobListPara return nil, err } - listParams, err := dblist.JobMakeDriverParams(ctx, dbParams, c.driver.SQLFragmentColumnIn) + listParams, err := dblist.JobMakeDriverParams(ctx, dbParams, c.driver) if err != nil { return nil, err } diff --git a/insert_opts.go b/insert_opts.go index bc1443d8..64fb770d 100644 --- a/insert_opts.go +++ b/insert_opts.go @@ -65,9 +65,9 @@ type InsertOpts struct { // JobArgsWithInsertOpts, however, it will work in both cases. ScheduledAt time.Time - // Tags are an arbitrary list of keywords to add to the job. They have no - // functional behavior and are meant entirely as a user-specified construct - // to help group and categorize jobs. + // Tags are an arbitrary list of keywords to add to the job. They don't + // affect job execution, but can be used with JobListParams.TagsAll and + // JobListParams.TagsAny to group and filter jobs. // // Tags should conform to the regex `\A[\w][\w\-]+[\w]\z` and be a maximum // of 255 characters long. No special characters are allowed. diff --git a/internal/dblist/db_list.go b/internal/dblist/db_list.go index 14a2eaf0..dba6780d 100644 --- a/internal/dblist/db_list.go +++ b/internal/dblist/db_list.go @@ -33,7 +33,8 @@ type JobListParams struct { Queues []string Schema string States []rivertype.JobState - Tags []string + TagsAll []string + TagsAny []string Where []WherePredicate } @@ -42,15 +43,21 @@ type WherePredicate struct { SQL string } +type sqlFragmentBuilder interface { + SQLFragmentColumnContainsAll(column, namedArg string, values []string) (string, any, error) + SQLFragmentColumnContainsAny(column, namedArg string, values []string) (string, any, error) + SQLFragmentColumnIn(column string, values any) (string, any, error) +} + // JobMakeDriverParams converts client-level parameters for job and delete to // driver-level parameters for use with an executor, which generally goes by // converting typed fields for IDs, kinds, queues, etc. to lower-level SQL. // // This was originally implemented for listing jobs, but since the logic is so // similar, it also performs the same function for JobDeleteMany. This works -// because `riverdriver.JobListParams` is identical to `JobDeleteMany` and -// therefore pointer-level converts to it. -func JobMakeDriverParams(ctx context.Context, params *JobListParams, sqlFragmentColumnIn func(column string, values any) (string, any, error)) (*riverdriver.JobListParams, error) { +// because `riverdriver.JobDeleteManyParams` has `JobListParams` as its +// underlying type and therefore pointer-level converts to it. +func JobMakeDriverParams(ctx context.Context, params *JobListParams, sqlFragmentBuilder sqlFragmentBuilder) (*riverdriver.JobListParams, error) { var ( namedArgs = make(map[string]any) whereBuilder strings.Builder @@ -76,7 +83,7 @@ func JobMakeDriverParams(ctx context.Context, params *JobListParams, sqlFragment writeAndAfterFirst() const column = "id" - sqlFragment, arg, err := sqlFragmentColumnIn(column, params.IDs) + sqlFragment, arg, err := sqlFragmentBuilder.SQLFragmentColumnIn(column, params.IDs) if err != nil { return nil, fmt.Errorf("error building SQL fragment for %q: %w", column, err) } @@ -88,7 +95,7 @@ func JobMakeDriverParams(ctx context.Context, params *JobListParams, sqlFragment writeAndAfterFirst() const column = "kind" - sqlFragment, arg, err := sqlFragmentColumnIn(column, params.Kinds) + sqlFragment, arg, err := sqlFragmentBuilder.SQLFragmentColumnIn(column, params.Kinds) if err != nil { return nil, fmt.Errorf("error building SQL fragment for %q: %w", column, err) } @@ -100,7 +107,7 @@ func JobMakeDriverParams(ctx context.Context, params *JobListParams, sqlFragment writeAndAfterFirst() const column = "priority" - sqlFragment, arg, err := sqlFragmentColumnIn(column, params.Priorities) + sqlFragment, arg, err := sqlFragmentBuilder.SQLFragmentColumnIn(column, params.Priorities) if err != nil { return nil, fmt.Errorf("error building SQL fragment for %q: %w", column, err) } @@ -112,7 +119,7 @@ func JobMakeDriverParams(ctx context.Context, params *JobListParams, sqlFragment writeAndAfterFirst() const column = "queue" - sqlFragment, arg, err := sqlFragmentColumnIn(column, params.Queues) + sqlFragment, arg, err := sqlFragmentBuilder.SQLFragmentColumnIn(column, params.Queues) if err != nil { return nil, fmt.Errorf("error building SQL fragment for %q: %w", column, err) } @@ -124,7 +131,7 @@ func JobMakeDriverParams(ctx context.Context, params *JobListParams, sqlFragment writeAndAfterFirst() const column = "state" - sqlFragment, arg, err := sqlFragmentColumnIn(column, + sqlFragment, arg, err := sqlFragmentBuilder.SQLFragmentColumnIn(column, sliceutil.Map(params.States, func(v rivertype.JobState) string { return string(v) })) if err != nil { return nil, fmt.Errorf("error building SQL fragment for %q: %w", column, err) @@ -133,6 +140,36 @@ func JobMakeDriverParams(ctx context.Context, params *JobListParams, sqlFragment namedArgs[column] = arg } + if len(params.TagsAll) > 0 { + writeAndAfterFirst() + + const ( + column = "tags" + namedArg = "tags_all" + ) + sqlFragment, arg, err := sqlFragmentBuilder.SQLFragmentColumnContainsAll(column, namedArg, params.TagsAll) + if err != nil { + return nil, fmt.Errorf("error building SQL fragment for %q: %w", namedArg, err) + } + whereBuilder.WriteString(sqlFragment) + namedArgs[namedArg] = arg + } + + if len(params.TagsAny) > 0 { + writeAndAfterFirst() + + const ( + column = "tags" + namedArg = "tags_any" + ) + sqlFragment, arg, err := sqlFragmentBuilder.SQLFragmentColumnContainsAny(column, namedArg, params.TagsAny) + if err != nil { + return nil, fmt.Errorf("error building SQL fragment for %q: %w", namedArg, err) + } + whereBuilder.WriteString(sqlFragment) + namedArgs[namedArg] = arg + } + for _, where := range params.Where { writeAndAfterFirst() @@ -187,7 +224,6 @@ func JobMakeDriverParams(ctx context.Context, params *JobListParams, sqlFragment NamedArgs: namedArgs, OrderByClause: orderByBuilder.String(), Schema: params.Schema, - Tags: params.Tags, WhereClause: whereBuilder.String(), }, nil } diff --git a/internal/dblist/db_list_test.go b/internal/dblist/db_list_test.go index 36767507..508acfd8 100644 --- a/internal/dblist/db_list_test.go +++ b/internal/dblist/db_list_test.go @@ -45,7 +45,7 @@ func TestJobListNoJobs(t *testing.T) { States: []rivertype.JobState{rivertype.JobStateCompleted}, LimitCount: 1, OrderBy: []JobListOrderBy{{Expr: "id", Order: SortOrderAsc}}, - }, bundle.driver.SQLFragmentColumnIn) + }, bundle.driver) require.NoError(t, err) _, err = bundle.exec.JobList(ctx, listParams) @@ -64,7 +64,7 @@ func TestJobListNoJobs(t *testing.T) { Where: []WherePredicate{ {NamedArgs: map[string]any{"foo": "bar"}, SQL: "queue = 'test' AND priority = 1 AND args->>'foo' = @foo"}, }, - }, bundle.driver.SQLFragmentColumnIn) + }, bundle.driver) require.NoError(t, err) _, err = bundle.exec.JobList(ctx, listParams) @@ -112,7 +112,7 @@ func TestJobListWithJobs(t *testing.T) { execTest := func(ctx context.Context, t *testing.T, bundle *testBundle, params *JobListParams, testFunc testListFunc) { t.Helper() - listParams, err := JobMakeDriverParams(ctx, params, bundle.driver.SQLFragmentColumnIn) + listParams, err := JobMakeDriverParams(ctx, params, bundle.driver) require.NoError(t, err) t.Logf("testing JobList in Executor") @@ -326,7 +326,7 @@ func TestJobListWithJobs(t *testing.T) { }, } - _, err := JobMakeDriverParams(ctx, params, bundle.driver.SQLFragmentColumnIn) + _, err := JobMakeDriverParams(ctx, params, bundle.driver) require.EqualError(t, err, `expected "1" to contain named arg symbol @not_present`) }) @@ -344,7 +344,7 @@ func TestJobListWithJobs(t *testing.T) { }, } - _, err := JobMakeDriverParams(ctx, params, bundle.driver.SQLFragmentColumnIn) + _, err := JobMakeDriverParams(ctx, params, bundle.driver) require.EqualError(t, err, "named argument @duplicate already registered") }) } diff --git a/job_list_params.go b/job_list_params.go index 6135be0b..4c275857 100644 --- a/job_list_params.go +++ b/job_list_params.go @@ -177,7 +177,8 @@ type JobListParams struct { sortField JobListOrderByField sortOrder SortOrder states []rivertype.JobState - tags []string + tagsAll []string + tagsAny []string where []dblist.WherePredicate } @@ -215,7 +216,8 @@ func (p *JobListParams) copy() *JobListParams { sortOrder: p.sortOrder, schema: p.schema, states: append([]rivertype.JobState(nil), p.states...), - tags: append([]string(nil), p.tags...), + tagsAll: append([]string(nil), p.tagsAll...), + tagsAny: append([]string(nil), p.tagsAny...), where: append([]dblist.WherePredicate(nil), p.where...), } } @@ -296,7 +298,8 @@ func (p *JobListParams) toDBParams() (*dblist.JobListParams, error) { Queues: p.queues, Schema: p.schema, States: p.states, - Tags: p.tags, + TagsAll: p.tagsAll, + TagsAny: p.tagsAny, Where: p.where, }, nil } @@ -422,12 +425,29 @@ func (p *JobListParams) States(states ...rivertype.JobState) *JobListParams { return paramsCopy } -// Tags returns an updated filter set that will only return jobs containing at -// least one of the given tags. Tag matching is case-insensitive. -func (p *JobListParams) Tags(tags ...string) *JobListParams { +// TagsAll returns an updated filter set that will only return jobs containing +// all of the given tags. Matching is exact and case-sensitive. TagsAll is +// combined with TagsAny and all other filters using AND. +// +// Calling TagsAll replaces any tags supplied to a previous TagsAll call. +// Calling it with no tags removes the filter. +func (p *JobListParams) TagsAll(tags ...string) *JobListParams { + paramsCopy := p.copy() + paramsCopy.tagsAll = make([]string, len(tags)) + copy(paramsCopy.tagsAll, tags) + return paramsCopy +} + +// TagsAny returns an updated filter set that will only return jobs containing +// at least one of the given tags. Matching is exact and case-sensitive. +// TagsAny is combined with TagsAll and all other filters using AND. +// +// Calling TagsAny replaces any tags supplied to a previous TagsAny call. +// Calling it with no tags removes the filter. +func (p *JobListParams) TagsAny(tags ...string) *JobListParams { paramsCopy := p.copy() - paramsCopy.tags = make([]string, len(tags)) - copy(paramsCopy.tags, tags) + paramsCopy.tagsAny = make([]string, len(tags)) + copy(paramsCopy.tagsAny, tags) return paramsCopy } diff --git a/job_list_params_test.go b/job_list_params_test.go index 00c50476..7bb5c9b0 100644 --- a/job_list_params_test.go +++ b/job_list_params_test.go @@ -244,11 +244,43 @@ func Test_JobListParams_toDBParams(t *testing.T) { require.EqualError(t, err, "cannot order by finalized_at without finalized state filters") }) - t.Run("Tags", func(t *testing.T) { + t.Run("TagsAll", func(t *testing.T) { t.Parallel() - dbParams, err := NewJobListParams().Tags("alpha", "beta").toDBParams() + tags := []string{"alpha", "beta"} + params := NewJobListParams().TagsAll(tags...) + tags[0] = "modified" + + dbParams, err := params.toDBParams() + require.NoError(t, err) + require.Equal(t, []string{"alpha", "beta"}, dbParams.TagsAll) + + dbParams, err = params.TagsAll("gamma").toDBParams() + require.NoError(t, err) + require.Equal(t, []string{"gamma"}, dbParams.TagsAll) + + dbParams, err = params.TagsAll().toDBParams() + require.NoError(t, err) + require.Empty(t, dbParams.TagsAll) + }) + + t.Run("TagsAny", func(t *testing.T) { + t.Parallel() + + tags := []string{"alpha", "beta"} + params := NewJobListParams().TagsAny(tags...) + tags[0] = "modified" + + dbParams, err := params.toDBParams() + require.NoError(t, err) + require.Equal(t, []string{"alpha", "beta"}, dbParams.TagsAny) + + dbParams, err = params.TagsAny("gamma").toDBParams() + require.NoError(t, err) + require.Equal(t, []string{"gamma"}, dbParams.TagsAny) + + dbParams, err = params.TagsAny().toDBParams() require.NoError(t, err) - require.Equal(t, []string{"alpha", "beta"}, dbParams.Tags) + require.Empty(t, dbParams.TagsAny) }) } diff --git a/riverdriver/river_driver_interface.go b/riverdriver/river_driver_interface.go index 022a5074..6ab9e812 100644 --- a/riverdriver/river_driver_interface.go +++ b/riverdriver/river_driver_interface.go @@ -128,6 +128,22 @@ type Driver[TTx any] interface { // API is not stable. DO NOT USE. PoolSet(dbPool any) error + // SQLFragmentColumnContainsAll generates an SQL fragment to be included as + // a predicate in a `WHERE` query for a collection column containing all of + // the given values. PostgreSQL uses array containment while SQLite compares + // values from a JSON array. + // + // API is not stable. DO NOT USE. + SQLFragmentColumnContainsAll(column, namedArg string, values []string) (string, any, error) + + // SQLFragmentColumnContainsAny generates an SQL fragment to be included as + // a predicate in a `WHERE` query for a collection column containing at least + // one of the given values. PostgreSQL uses array overlap while SQLite + // compares values from a JSON array. + // + // API is not stable. DO NOT USE. + SQLFragmentColumnContainsAny(column, namedArg string, values []string) (string, any, error) + // 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 @@ -404,14 +420,7 @@ type JobDeleteBeforeParams struct { Schema string } -type JobDeleteManyParams struct { - Max int32 - NamedArgs map[string]any - OrderByClause string - Schema string - Tags []string - WhereClause string -} +type JobDeleteManyParams JobListParams type JobGetAvailableParams struct { ClientID string @@ -514,7 +523,6 @@ type JobListParams struct { NamedArgs map[string]any OrderByClause string Schema string - Tags []string WhereClause string } diff --git a/riverdriver/riverdatabasesql/internal/dbsqlc/river_job.sql.go b/riverdriver/riverdatabasesql/internal/dbsqlc/river_job.sql.go index 3fa2ca55..c72d277a 100644 --- a/riverdriver/riverdatabasesql/internal/dbsqlc/river_job.sql.go +++ b/riverdriver/riverdatabasesql/internal/dbsqlc/river_job.sql.go @@ -1166,26 +1166,12 @@ 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 */ - AND ( - coalesce(cardinality($1::text[]), 0) = 0 - OR EXISTS ( - SELECT 1 - FROM unnest(tags) AS job_tag(value) - INNER JOIN unnest($1::text[]) AS filter_tag(value) - ON lower(job_tag.value) = lower(filter_tag.value) - ) - ) ORDER BY /* TEMPLATE_BEGIN: order_by_clause */ id /* TEMPLATE_END */ -LIMIT $2::int +LIMIT $1::int ` -type JobListParams struct { - Tags []string - Max int32 -} - -func (q *Queries) JobList(ctx context.Context, db DBTX, arg *JobListParams) ([]*RiverJob, error) { - rows, err := db.QueryContext(ctx, jobList, pq.Array(arg.Tags), arg.Max) +func (q *Queries) JobList(ctx context.Context, db DBTX, max int32) ([]*RiverJob, error) { + rows, err := db.QueryContext(ctx, jobList, max) if err != nil { return nil, err } diff --git a/riverdriver/riverdatabasesql/river_database_sql_driver.go b/riverdriver/riverdatabasesql/river_database_sql_driver.go index 22e01c6d..a37e9a88 100644 --- a/riverdriver/riverdatabasesql/river_database_sql_driver.go +++ b/riverdriver/riverdatabasesql/river_database_sql_driver.go @@ -82,6 +82,14 @@ func (d *Driver) GetMigrationTruncateTables(line string, version int) []string { func (d *Driver) PoolIsSet() bool { return d.dbPool != nil } func (d *Driver) PoolSet(dbPool any) error { return riverdriver.ErrNotImplemented } +func (d *Driver) SQLFragmentColumnContainsAll(column, namedArg string, values []string) (string, any, error) { + return fmt.Sprintf("%s @> @%s", column, namedArg), pq.Array(values), nil +} + +func (d *Driver) SQLFragmentColumnContainsAny(column, namedArg string, values []string) (string, any, error) { + return fmt.Sprintf("%s && @%s", column, namedArg), pq.Array(values), nil +} + func (d *Driver) SQLFragmentColumnIn(column string, values any) (string, any, error) { // Identical to the Pgx implementation except for use of `pg.Array`. return fmt.Sprintf("%s = any(@%s)", column, column), pq.Array(values), nil @@ -573,10 +581,7 @@ func (e *Executor) JobList(ctx context.Context, params *riverdriver.JobListParam "where_clause": {Value: params.WhereClause}, }, params.NamedArgs) - jobs, err := dbsqlc.New().JobList(schemaTemplateParam(ctx, params.Schema), e.dbtx, &dbsqlc.JobListParams{ - Max: params.Max, - Tags: params.Tags, - }) + jobs, err := dbsqlc.New().JobList(schemaTemplateParam(ctx, params.Schema), e.dbtx, params.Max) if err != nil { return nil, interpretError(err) } diff --git a/riverdriver/riverdrivertest/driver_client_test.go b/riverdriver/riverdrivertest/driver_client_test.go index 2eda4a5d..7081f042 100644 --- a/riverdriver/riverdrivertest/driver_client_test.go +++ b/riverdriver/riverdrivertest/driver_client_test.go @@ -518,7 +518,7 @@ func ExerciseClient[TTx any](ctx context.Context, t *testing.T, job := testfactory.Job(ctx, t, bundle.exec, &testfactory.JobOpts{ Schema: bundle.schema, - Tags: []string{"all-args-tag"}, + Tags: []string{"all-args-tag", "all-args-secondary"}, }) listRes, err := client.JobList(ctx, @@ -528,7 +528,8 @@ func ExerciseClient[TTx any](ctx context.Context, t *testing.T, Priorities(int16(min(job.Priority, math.MaxInt16))). //nolint:gosec Queues(job.Queue). States(job.State). - Tags("ALL-ARGS-TAG"), + TagsAll("all-args-tag"). + TagsAny("all-args-secondary"), ) require.NoError(t, err) require.Len(t, listRes.Jobs, 1) @@ -561,17 +562,49 @@ func ExerciseClient[TTx any](ctx context.Context, t *testing.T, client, bundle := setup(t) - var ( - job1 = testfactory.Job(ctx, t, bundle.exec, &testfactory.JobOpts{Schema: bundle.schema, Tags: []string{"alpha", "shared"}}) - job2 = testfactory.Job(ctx, t, bundle.exec, &testfactory.JobOpts{Schema: bundle.schema, Tags: []string{"beta"}}) - _ = testfactory.Job(ctx, t, bundle.exec, &testfactory.JobOpts{Schema: bundle.schema, Tags: []string{"gamma"}}) - ) + jobAlphaBeta := testfactory.Job(ctx, t, bundle.exec, &testfactory.JobOpts{Schema: bundle.schema, Tags: []string{"alpha", "beta", "shared"}}) + jobAlpha := testfactory.Job(ctx, t, bundle.exec, &testfactory.JobOpts{Schema: bundle.schema, Tags: []string{"alpha"}}) + jobBeta := testfactory.Job(ctx, t, bundle.exec, &testfactory.JobOpts{Schema: bundle.schema, Tags: []string{"beta"}}) + jobUpperAlpha := testfactory.Job(ctx, t, bundle.exec, &testfactory.JobOpts{Schema: bundle.schema, Tags: []string{"ALPHA"}}) + _ = testfactory.Job(ctx, t, bundle.exec, &testfactory.JobOpts{Schema: bundle.schema, Tags: []string{"gamma"}}) + _ = testfactory.Job(ctx, t, bundle.exec, &testfactory.JobOpts{Schema: bundle.schema}) - listRes, err := client.JobList(ctx, river.NewJobListParams().Tags("ALPHA", "BETA")) + listRes, err := client.JobList(ctx, river.NewJobListParams().TagsAny("alpha", "beta")) require.NoError(t, err) - require.Len(t, listRes.Jobs, 2) - require.Equal(t, job1.ID, listRes.Jobs[0].ID) - require.Equal(t, job2.ID, listRes.Jobs[1].ID) + require.Len(t, listRes.Jobs, 3) + require.Equal(t, jobAlphaBeta.ID, listRes.Jobs[0].ID) + require.Equal(t, jobAlpha.ID, listRes.Jobs[1].ID) + require.Equal(t, jobBeta.ID, listRes.Jobs[2].ID) + + listRes, err = client.JobList(ctx, river.NewJobListParams().TagsAll("alpha", "beta")) + require.NoError(t, err) + require.Len(t, listRes.Jobs, 1) + require.Equal(t, jobAlphaBeta.ID, listRes.Jobs[0].ID) + + listRes, err = client.JobList(ctx, river.NewJobListParams().TagsAll("shared").TagsAny("alpha", "gamma")) + require.NoError(t, err) + require.Len(t, listRes.Jobs, 1) + require.Equal(t, jobAlphaBeta.ID, listRes.Jobs[0].ID) + + listRes, err = client.JobList(ctx, river.NewJobListParams().TagsAny("ALPHA")) + require.NoError(t, err) + require.Len(t, listRes.Jobs, 1) + require.Equal(t, jobUpperAlpha.ID, listRes.Jobs[0].ID) + + params := river.NewJobListParams().TagsAny("alpha", "beta").First(1) + listRes, err = client.JobList(ctx, params) + require.NoError(t, err) + require.Len(t, listRes.Jobs, 1) + require.Equal(t, jobAlphaBeta.ID, listRes.Jobs[0].ID) + + listRes, err = client.JobList(ctx, params.After(listRes.LastCursor)) + require.NoError(t, err) + require.Len(t, listRes.Jobs, 1) + require.Equal(t, jobAlpha.ID, listRes.Jobs[0].ID) + + listRes, err = client.JobList(ctx, river.NewJobListParams().TagsAny("alpha").TagsAny()) + require.NoError(t, err) + require.Len(t, listRes.Jobs, 6) }) t.Run("JobListTx", func(t *testing.T) { @@ -602,7 +635,7 @@ func ExerciseClient[TTx any](ctx context.Context, t *testing.T, job := testfactory.Job(ctx, t, execTx, &testfactory.JobOpts{ Schema: bundle.schema, - Tags: []string{"all-args-tag"}, + Tags: []string{"all-args-tag", "all-args-secondary"}, }) listRes, err := client.JobListTx(ctx, tx, @@ -612,7 +645,8 @@ func ExerciseClient[TTx any](ctx context.Context, t *testing.T, Priorities(int16(min(job.Priority, math.MaxInt16))). //nolint:gosec Queues(job.Queue). States(job.State). - Tags("ALL-ARGS-TAG"), + TagsAll("all-args-tag"). + TagsAny("all-args-secondary"), ) require.NoError(t, err) require.Len(t, listRes.Jobs, 1) diff --git a/riverdriver/riverdrivertest/sql_fragments.go b/riverdriver/riverdrivertest/sql_fragments.go index ac1fe41e..72edc9ac 100644 --- a/riverdriver/riverdrivertest/sql_fragments.go +++ b/riverdriver/riverdrivertest/sql_fragments.go @@ -14,6 +14,59 @@ import ( func exerciseSQLFragments[TTx any](ctx context.Context, t *testing.T, executorWithTx func(ctx context.Context, t *testing.T) (riverdriver.Executor, riverdriver.Driver[TTx])) { t.Helper() + t.Run("SQLFragmentColumnContainsAll", func(t *testing.T) { + t.Parallel() + + exec, driver := executorWithTx(ctx, t) + + var ( + job1 = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{Tags: []string{"alpha", "beta"}}) + _ = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{Tags: []string{"alpha"}}) + _ = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{Tags: []string{"beta"}}) + ) + + const namedArg = "tags_all" + sqlFragment, arg, err := driver.SQLFragmentColumnContainsAll("tags", namedArg, []string{"alpha", "beta"}) + require.NoError(t, err) + + jobs, err := exec.JobList(ctx, &riverdriver.JobListParams{ + Max: 100, + NamedArgs: map[string]any{namedArg: arg}, + OrderByClause: "id", + WhereClause: sqlFragment, + }) + require.NoError(t, err) + require.Len(t, jobs, 1) + require.Equal(t, job1.ID, jobs[0].ID) + }) + + t.Run("SQLFragmentColumnContainsAny", func(t *testing.T) { + t.Parallel() + + exec, driver := executorWithTx(ctx, t) + + var ( + job1 = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{Tags: []string{"alpha"}}) + job2 = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{Tags: []string{"beta"}}) + _ = testfactory.Job(ctx, t, exec, &testfactory.JobOpts{Tags: []string{"gamma"}}) + ) + + const namedArg = "tags_any" + sqlFragment, arg, err := driver.SQLFragmentColumnContainsAny("tags", namedArg, []string{"alpha", "beta"}) + require.NoError(t, err) + + jobs, err := exec.JobList(ctx, &riverdriver.JobListParams{ + Max: 100, + NamedArgs: map[string]any{namedArg: arg}, + OrderByClause: "id", + WhereClause: sqlFragment, + }) + require.NoError(t, err) + require.Len(t, jobs, 2) + require.Equal(t, job1.ID, jobs[0].ID) + require.Equal(t, job2.ID, jobs[1].ID) + }) + t.Run("SQLFragmentColumnIn", func(t *testing.T) { t.Parallel() diff --git a/riverdriver/riverpgxv5/internal/dbsqlc/river_job.sql b/riverdriver/riverpgxv5/internal/dbsqlc/river_job.sql index d9c0e80f..509e479b 100644 --- a/riverdriver/riverpgxv5/internal/dbsqlc/river_job.sql +++ b/riverdriver/riverpgxv5/internal/dbsqlc/river_job.sql @@ -472,15 +472,6 @@ LIMIT @max; SELECT * FROM /* TEMPLATE: schema */river_job WHERE /* TEMPLATE_BEGIN: where_clause */ true /* TEMPLATE_END */ - AND ( - coalesce(cardinality(@tags::text[]), 0) = 0 - OR EXISTS ( - SELECT 1 - FROM unnest(tags) AS job_tag(value) - INNER JOIN unnest(@tags::text[]) AS filter_tag(value) - ON lower(job_tag.value) = lower(filter_tag.value) - ) - ) ORDER BY /* TEMPLATE_BEGIN: order_by_clause */ id /* TEMPLATE_END */ LIMIT @max::int; diff --git a/riverdriver/riverpgxv5/internal/dbsqlc/river_job.sql.go b/riverdriver/riverpgxv5/internal/dbsqlc/river_job.sql.go index 3d90dcab..a361baac 100644 --- a/riverdriver/riverpgxv5/internal/dbsqlc/river_job.sql.go +++ b/riverdriver/riverpgxv5/internal/dbsqlc/river_job.sql.go @@ -1136,26 +1136,12 @@ 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 */ - AND ( - coalesce(cardinality($1::text[]), 0) = 0 - OR EXISTS ( - SELECT 1 - FROM unnest(tags) AS job_tag(value) - INNER JOIN unnest($1::text[]) AS filter_tag(value) - ON lower(job_tag.value) = lower(filter_tag.value) - ) - ) ORDER BY /* TEMPLATE_BEGIN: order_by_clause */ id /* TEMPLATE_END */ -LIMIT $2::int +LIMIT $1::int ` -type JobListParams struct { - Tags []string - Max int32 -} - -func (q *Queries) JobList(ctx context.Context, db DBTX, arg *JobListParams) ([]*RiverJob, error) { - rows, err := db.Query(ctx, jobList, arg.Tags, arg.Max) +func (q *Queries) JobList(ctx context.Context, db DBTX, max int32) ([]*RiverJob, error) { + rows, err := db.Query(ctx, jobList, max) if err != nil { return nil, err } diff --git a/riverdriver/riverpgxv5/river_pgx_v5_driver.go b/riverdriver/riverpgxv5/river_pgx_v5_driver.go index 7dd91bf9..71dcc6d5 100644 --- a/riverdriver/riverpgxv5/river_pgx_v5_driver.go +++ b/riverdriver/riverpgxv5/river_pgx_v5_driver.go @@ -92,6 +92,14 @@ func (d *Driver) GetMigrationTruncateTables(line string, version int) []string { func (d *Driver) PoolIsSet() bool { return d.dbPool != nil } func (d *Driver) PoolSet(dbPool any) error { return riverdriver.ErrNotImplemented } +func (d *Driver) SQLFragmentColumnContainsAll(column, namedArg string, values []string) (string, any, error) { + return fmt.Sprintf("%s @> @%s", column, namedArg), values, nil +} + +func (d *Driver) SQLFragmentColumnContainsAny(column, namedArg string, values []string) (string, any, error) { + return fmt.Sprintf("%s && @%s", column, namedArg), values, nil +} + func (d *Driver) SQLFragmentColumnIn(column string, values any) (string, any, error) { return fmt.Sprintf("%s = any(@%s)", column, column), values, nil } @@ -567,10 +575,7 @@ func (e *Executor) JobList(ctx context.Context, params *riverdriver.JobListParam "where_clause": {Value: params.WhereClause}, }, params.NamedArgs) - jobs, err := dbsqlc.New().JobList(schemaTemplateParam(ctx, params.Schema), e.dbtx, &dbsqlc.JobListParams{ - Max: params.Max, - Tags: params.Tags, - }) + jobs, err := dbsqlc.New().JobList(schemaTemplateParam(ctx, params.Schema), e.dbtx, params.Max) if err != nil { return nil, interpretError(err) } diff --git a/riverdriver/riversqlite/internal/dbsqlc/river_job.sql b/riverdriver/riversqlite/internal/dbsqlc/river_job.sql index 7f2fd3e1..94b433ca 100644 --- a/riverdriver/riversqlite/internal/dbsqlc/river_job.sql +++ b/riverdriver/riversqlite/internal/dbsqlc/river_job.sql @@ -471,15 +471,6 @@ LIMIT @max; SELECT * FROM /* TEMPLATE: schema */river_job WHERE /* TEMPLATE_BEGIN: where_clause */ true /* TEMPLATE_END */ - AND ( - json_array_length(cast(@tags AS blob)) = 0 - OR EXISTS ( - SELECT 1 - FROM json_each(river_job.tags) AS job_tag - INNER JOIN json_each(cast(@tags AS blob)) AS filter_tag - ON lower(cast(job_tag.value AS text)) = lower(cast(filter_tag.value AS text)) - ) - ) ORDER BY /* TEMPLATE_BEGIN: order_by_clause */ id /* TEMPLATE_END */ LIMIT @max; diff --git a/riverdriver/riversqlite/internal/dbsqlc/river_job.sql.go b/riverdriver/riversqlite/internal/dbsqlc/river_job.sql.go index 124dd288..5bd39b2c 100644 --- a/riverdriver/riversqlite/internal/dbsqlc/river_job.sql.go +++ b/riverdriver/riversqlite/internal/dbsqlc/river_job.sql.go @@ -1194,26 +1194,12 @@ const jobList = `-- name: JobList :many SELECT id, json(args), attempt, attempted_at, json(attempted_by), created_at, json(errors), finalized_at, kind, max_attempts, json(metadata), priority, queue, state, scheduled_at, json(tags), unique_key, unique_states FROM /* TEMPLATE: schema */river_job WHERE /* TEMPLATE_BEGIN: where_clause */ true /* TEMPLATE_END */ - AND ( - json_array_length(cast(?1 AS blob)) = 0 - OR EXISTS ( - SELECT 1 - FROM json_each(river_job.tags) AS job_tag - INNER JOIN json_each(cast(?1 AS blob)) AS filter_tag - ON lower(cast(job_tag.value AS text)) = lower(cast(filter_tag.value AS text)) - ) - ) ORDER BY /* TEMPLATE_BEGIN: order_by_clause */ id /* TEMPLATE_END */ -LIMIT ?2 +LIMIT ?1 ` -type JobListParams struct { - Tags []byte - Max int64 -} - -func (q *Queries) JobList(ctx context.Context, db DBTX, arg *JobListParams) ([]*RiverJob, error) { - rows, err := db.QueryContext(ctx, jobList, arg.Tags, arg.Max) +func (q *Queries) JobList(ctx context.Context, db DBTX, max int64) ([]*RiverJob, error) { + rows, err := db.QueryContext(ctx, jobList, max) if err != nil { return nil, err } diff --git a/riverdriver/riversqlite/river_sqlite_driver.go b/riverdriver/riversqlite/river_sqlite_driver.go index 51c0cdff..df1f4acf 100644 --- a/riverdriver/riversqlite/river_sqlite_driver.go +++ b/riverdriver/riversqlite/river_sqlite_driver.go @@ -117,6 +117,37 @@ func (d *Driver) PoolSet(dbPool any) error { return nil } +func (d *Driver) SQLFragmentColumnContainsAll(column, namedArg string, values []string) (string, any, error) { + arg, err := json.Marshal(values) + if err != nil { + return "", nil, err + } + + return fmt.Sprintf(`NOT EXISTS ( + SELECT 1 + FROM json_each(cast(@%s AS blob)) AS filter_value + WHERE NOT EXISTS ( + SELECT 1 + FROM json_each(%s) AS column_value + WHERE column_value.value = filter_value.value + ) +)`, namedArg, column), arg, nil +} + +func (d *Driver) SQLFragmentColumnContainsAny(column, namedArg string, values []string) (string, any, error) { + arg, err := json.Marshal(values) + if err != nil { + return "", nil, err + } + + return fmt.Sprintf(`EXISTS ( + SELECT 1 + FROM json_each(%s) AS column_value + INNER JOIN json_each(cast(@%s AS blob)) AS filter_value + ON column_value.value = filter_value.value +)`, column, namedArg), arg, nil +} + func (d *Driver) SQLFragmentColumnIn(column string, values any) (string, any, error) { arg, err := json.Marshal(values) if err != nil { @@ -676,15 +707,7 @@ func (e *Executor) JobList(ctx context.Context, params *riverdriver.JobListParam "where_clause": {Value: params.WhereClause}, }, params.NamedArgs) - tags, err := json.Marshal(sliceutil.FirstNonEmpty(params.Tags, []string{})) - if err != nil { - return nil, fmt.Errorf("error encoding tags: %w", err) - } - - jobs, err := dbsqlc.New().JobList(schemaTemplateParam(ctx, params.Schema), e.dbtx, &dbsqlc.JobListParams{ - Max: int64(params.Max), - Tags: tags, - }) + jobs, err := dbsqlc.New().JobList(schemaTemplateParam(ctx, params.Schema), e.dbtx, int64(params.Max)) if err != nil { return nil, interpretError(err) } diff --git a/rivertype/river_type.go b/rivertype/river_type.go index 1abdf28e..fe74d707 100644 --- a/rivertype/river_type.go +++ b/rivertype/river_type.go @@ -124,9 +124,8 @@ type JobRow struct { // `available` when they're first inserted. State JobState - // Tags are an arbitrary list of keywords to add to the job. They have no - // functional behavior and are meant entirely as a user-specified construct - // to help group and categorize jobs. + // Tags are an arbitrary list of keywords attached to the job. They don't + // affect job execution, but clients can use them to group and filter jobs. Tags []string // UniqueKey is a unique key for the job within its kind that's used for