diff --git a/CHANGELOG.md b/CHANGELOG.md index 05277e667..9649508b4 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] +### Fixed + +- Fixed `AddWorkerSafely` leaving a worker's primary kind and earlier aliases registered when a later kind alias conflicts with an existing registration. Failed registrations now leave the worker registry unchanged. [PR #1440](https://github.com/riverqueue/river/pull/1440). + ## [0.49.0] - 2026-10-05 ### Changed diff --git a/worker.go b/worker.go index f833b760d..0c1ede252 100644 --- a/worker.go +++ b/worker.go @@ -162,32 +162,26 @@ func NewWorkers() *Workers { } func (w Workers) add(jobArgs JobArgs, workUnitFactory workunit.WorkUnitFactory) error { - checkRegistered := func(kind string) error { - if _, ok := w.workersMap[kind]; ok { + kinds := []string{jobArgs.Kind()} + if args, ok := jobArgs.(JobArgsWithKindAliases); ok { + kinds = append(kinds, args.KindAliases()...) + } + + // Validate all kinds before changing the registry. + seen := make(map[string]bool, len(kinds)) + for _, kind := range kinds { + if _, ok := w.workersMap[kind]; ok || seen[kind] { return fmt.Errorf("worker for kind %q is already registered", kind) } - return nil + seen[kind] = true } - workerInfo := workerInfo{ + info := workerInfo{ jobArgs: jobArgs, workUnitFactory: workUnitFactory, } - - kind := jobArgs.Kind() - if err := checkRegistered(kind); err != nil { - return err - } - w.workersMap[kind] = workerInfo - - // Jobs can register an alternate kind to make renaming easier. - if jobArgsWithKindAliases, ok := jobArgs.(JobArgsWithKindAliases); ok { - for _, kind := range jobArgsWithKindAliases.KindAliases() { - if err := checkRegistered(kind); err != nil { - return err - } - w.workersMap[kind] = workerInfo - } + for _, kind := range kinds { + w.workersMap[kind] = info } return nil diff --git a/worker_test.go b/worker_test.go index 8654b21ba..5fa58121e 100644 --- a/worker_test.go +++ b/worker_test.go @@ -13,6 +13,19 @@ import ( "github.com/riverqueue/river/rivershared/util/testutil" ) +func TestAddWorkerSafelyAliasCollisionDoesNotRegisterKinds(t *testing.T) { + t.Parallel() + + workers := NewWorkers() + + require.NoError(t, AddWorkerSafely(workers, WorkFunc(func(context.Context, *Job[withKindAliasesArgs]) error { return nil }))) + before := len(workers.workersMap) + require.Error(t, AddWorkerSafely(workers, WorkFunc(func(context.Context, *Job[lateAliasCollisionArgs]) error { return nil }))) + require.Len(t, workers.workersMap, before) + require.NotContains(t, workers.workersMap, "candidate_primary") + require.NotContains(t, workers.workersMap, "candidate_early_alias") +} + func TestWork(t *testing.T) { t.Parallel() @@ -74,6 +87,13 @@ func (w *configurableWorker) Work(ctx context.Context, job *Job[configurableArgs return nil } +type lateAliasCollisionArgs struct{} + +func (lateAliasCollisionArgs) Kind() string { return "candidate_primary" } +func (lateAliasCollisionArgs) KindAliases() []string { + return []string{"candidate_early_alias", "with_kind_alternate_alternate"} +} + type withKindAliasesArgs struct{} func (a withKindAliasesArgs) Kind() string { return "with_kind_alternate" }