From 571e5c2e3c128b8d0f3f2d3f477eaae8d42782df Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Roberto=20Jim=C3=A9nez=20S=C3=A1nchez?= Date: Wed, 5 Nov 2025 11:07:21 +0100 Subject: [PATCH] Provisioning: Fix data race in job progress and leasing (#113157) * Fix data race in provisioning job execution * Fix TODO * Update pkg/registry/apis/provisioning/jobs/driver.go Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> * Update pkg/registry/apis/provisioning/jobs/driver.go Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> * Fix unlocking issue on panic --------- Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> --- pkg/registry/apis/provisioning/jobs/driver.go | 106 ++++++++++++------ .../jobs/job_progress_recorder_mock.go | 51 ++++++++- .../apis/provisioning/jobs/progress.go | 4 + pkg/registry/apis/provisioning/jobs/queue.go | 2 + 4 files changed, 129 insertions(+), 34 deletions(-) diff --git a/pkg/registry/apis/provisioning/jobs/driver.go b/pkg/registry/apis/provisioning/jobs/driver.go index bad1df13759..7902aad12fc 100644 --- a/pkg/registry/apis/provisioning/jobs/driver.go +++ b/pkg/registry/apis/provisioning/jobs/driver.go @@ -4,6 +4,7 @@ import ( "context" "errors" "strings" + "sync" "time" apierrors "k8s.io/apimachinery/pkg/api/errors" @@ -74,6 +75,11 @@ type jobDriver struct { // notifications channel for job create events notifications chan struct{} + + // Mutex to protect concurrent access to job processing + mu sync.Mutex + // currentJob is the job currently being processed + currentJob *provisioning.Job } func NewJobDriver( @@ -142,7 +148,7 @@ func (d *jobDriver) claimAndProcessOneJob(ctx context.Context) error { logger := logging.FromContext(ctx) // Claim a job to work on. - job, rollback, err := d.store.Claim(ctx) + claimedJob, rollback, err := d.store.Claim(ctx) if err != nil { return apifmt.Errorf("failed to claim job: %w", err) } @@ -150,14 +156,16 @@ func (d *jobDriver) claimAndProcessOneJob(ctx context.Context) error { // The rollback function does not care about cancellations. defer rollback() - logger = logger.With("job", job.GetName(), "namespace", job.GetNamespace()) + namespace := claimedJob.GetNamespace() + logger = logger.With("job", claimedJob.GetName(), "namespace", namespace) ctx = logging.Context(ctx, logger) logger.Debug("claimed a job") + d.currentJob = claimedJob // Now that we have a job, we need to augment our namespace to grant ourselves permission to work on it. // Incidentally, this also limits our permissions to only the namespace of the job. - ctx = request.WithNamespace(ctx, job.GetNamespace()) - ctx, _, err = identity.WithProvisioningIdentity(ctx, job.GetNamespace()) + ctx = request.WithNamespace(ctx, namespace) + ctx, _, err = identity.WithProvisioningIdentity(ctx, namespace) if err != nil { return apifmt.Errorf("failed to grant provisioning identity: %w", err) } @@ -169,37 +177,42 @@ func (d *jobDriver) claimAndProcessOneJob(ctx context.Context) error { leaseRenewalCtx, cancelLeaseRenewal := context.WithCancel(jobctx) leaseExpired := make(chan struct{}) - go d.leaseRenewalLoop(leaseRenewalCtx, job, logger, leaseExpired) + go d.leaseRenewalLoop(leaseRenewalCtx, logger, leaseExpired) defer cancelLeaseRenewal() - recorder := newJobProgressRecorder(d.onProgress(job)) + recorder := newJobProgressRecorder(d.onProgress()) + recorder.SetMessage(ctx, "start job") // Process the job with lease loss detection - start := time.Now() - job.Status.Started = start.UnixMilli() - err = d.processJobWithLeaseCheck(jobctx, job, recorder, leaseExpired) + err = d.processJobWithLeaseCheck(jobctx, recorder, leaseExpired) end := time.Now() - logger.Debug("job processed", "duration", end.Sub(start), "error", err) + logger.Debug("job processed", "duration", end.Sub(recorder.Started()), "error", err) // Capture job timeout if jobctx.Err() != nil && err == nil { err = jobctx.Err() } - job.Status = recorder.Complete(ctx, err) + // Complete the job + d.mu.Lock() + d.currentJob.Status = recorder.Complete(ctx, err) + defer func() { + d.currentJob = nil + d.mu.Unlock() + }() // Save the finished job - err = d.historicJobs.WriteJob(ctx, job.DeepCopy()) + err = d.historicJobs.WriteJob(ctx, d.currentJob.DeepCopy()) if err != nil { // We're not going to return this as it is not critical. Not ideal, but not critical. - logger.Warn("failed to create historic job", "historic_job", *job, "error", err) + logger.Warn("failed to create historic job", "historic_job", *d.currentJob, "error", err) } else { - logger.Debug("created historic job", "historic_job", *job) + logger.Debug("created historic job", "historic_job", *d.currentJob) } // Mark the job as completed. - if err := d.store.Complete(ctx, job); err != nil { - return apifmt.Errorf("failed to complete job '%s' in '%s': %w", job.GetName(), job.GetNamespace(), err) + if err := d.store.Complete(ctx, d.currentJob); err != nil { + return apifmt.Errorf("failed to complete job '%s' in '%s': %w", d.currentJob.GetName(), d.currentJob.GetNamespace(), err) } logger.Debug("job completed") @@ -208,7 +221,7 @@ func (d *jobDriver) claimAndProcessOneJob(ctx context.Context) error { // leaseRenewalLoop continuously renews the lease for a job until the context is cancelled. // If lease renewal fails persistently, it signals via the leaseExpired channel. -func (d *jobDriver) leaseRenewalLoop(ctx context.Context, job *provisioning.Job, logger logging.Logger, leaseExpired chan struct{}) { +func (d *jobDriver) leaseRenewalLoop(ctx context.Context, logger logging.Logger, leaseExpired chan struct{}) { ticker := time.NewTicker(d.leaseRenewalInterval) defer ticker.Stop() @@ -223,7 +236,15 @@ func (d *jobDriver) leaseRenewalLoop(ctx context.Context, job *provisioning.Job, logger.Debug("lease renewal loop stopping") return case <-ticker.C: - err := d.store.RenewLease(ctx, job) + d.mu.Lock() + if d.currentJob == nil { + d.mu.Unlock() + return + } + + err := d.store.RenewLease(ctx, d.currentJob) + d.mu.Unlock() + if err != nil { consecutiveFailures++ if apierrors.IsNotFound(err) || @@ -253,11 +274,11 @@ func (d *jobDriver) leaseRenewalLoop(ctx context.Context, job *provisioning.Job, } // processJobWithLeaseCheck processes a job but aborts if the lease expires. -func (d *jobDriver) processJobWithLeaseCheck(ctx context.Context, job *provisioning.Job, recorder JobProgressRecorder, leaseExpired <-chan struct{}) error { +func (d *jobDriver) processJobWithLeaseCheck(ctx context.Context, recorder JobProgressRecorder, leaseExpired <-chan struct{}) error { // Run the job processing in a goroutine so we can monitor lease expiry resultChan := make(chan error, 1) go func() { - resultChan <- d.processJob(ctx, job, recorder) + resultChan <- d.processJob(ctx, recorder) }() select { @@ -270,16 +291,28 @@ func (d *jobDriver) processJobWithLeaseCheck(ctx context.Context, job *provision } } -func (d *jobDriver) processJob(ctx context.Context, job *provisioning.Job, recorder JobProgressRecorder) error { +func (d *jobDriver) processJob(ctx context.Context, recorder JobProgressRecorder) error { logger := logging.FromContext(ctx) + d.mu.Lock() + if d.currentJob == nil { + d.mu.Unlock() + return nil + } + + // Here it's safe to copy as only job spec is used for processing + job := d.currentJob.DeepCopy() + repoName := d.currentJob.Spec.Repository + namespace := d.currentJob.Namespace + d.mu.Unlock() + for _, worker := range d.workers { if !worker.IsSupported(ctx, *job) { continue } - repo, err := d.repoGetter.GetRepository(ctx, job.Namespace, job.Spec.Repository) + repo, err := d.repoGetter.GetRepository(ctx, namespace, repoName) if err != nil { - return apifmt.Errorf("failed to get repository '%s': %w", job.Spec.Repository, err) + return apifmt.Errorf("failed to get repository '%s': %w", repoName, err) } r := repo.Config() @@ -298,42 +331,51 @@ func (d *jobDriver) processJob(ctx context.Context, job *provisioning.Job, recor return apifmt.Errorf("no workers were registered to handle the job") } -func (d *jobDriver) onProgress(job *provisioning.Job) ProgressFn { +func (d *jobDriver) onProgress() ProgressFn { return func(ctx context.Context, status provisioning.JobStatus) error { logging.FromContext(ctx).Debug("job progress", "status", status) const maxRetries = 3 for attempt := 0; attempt < maxRetries; attempt++ { - // Use the current job for the first attempt, fetch fresh for retries - currentJob := job + d.mu.Lock() + if d.currentJob == nil { + d.mu.Unlock() + return nil + } + + // Use the current job for the first attempt; on retry attempts, fetch fresh data from the store to resolve conflicts if attempt > 0 { // Fetch the latest version to resolve conflicts - latest, err := d.store.Get(ctx, job.GetNamespace(), job.GetName()) + latest, err := d.store.Get(ctx, d.currentJob.GetNamespace(), d.currentJob.GetName()) if err != nil { + d.mu.Unlock() if apierrors.IsNotFound(err) { // Job was completed/deleted, nothing to update return nil } return apifmt.Errorf("failed to fetch job for progress update: %w", err) } - currentJob = latest + + *d.currentJob = *latest } + job := d.currentJob // Update status on the current job - currentJob.Status = status - - updated, err := d.store.Update(ctx, currentJob) + job.Status = status + updated, err := d.store.Update(ctx, job) if err != nil { if apierrors.IsConflict(err) && attempt < maxRetries-1 { // Conflict detected, retry with fresh data logging.FromContext(ctx).Debug("progress update conflict, retrying", "attempt", attempt+1) continue } + d.mu.Unlock() return apifmt.Errorf("failed to update job progress: %w", err) } // Update succeeded, update our local copy - *job = *updated + *d.currentJob = *updated + d.mu.Unlock() return nil } diff --git a/pkg/registry/apis/provisioning/jobs/job_progress_recorder_mock.go b/pkg/registry/apis/provisioning/jobs/job_progress_recorder_mock.go index 3cefca9e5af..45d8572e94a 100644 --- a/pkg/registry/apis/provisioning/jobs/job_progress_recorder_mock.go +++ b/pkg/registry/apis/provisioning/jobs/job_progress_recorder_mock.go @@ -1,12 +1,14 @@ -// Code generated by mockery v2.52.4. DO NOT EDIT. +// Code generated by mockery v2.53.4. DO NOT EDIT. package jobs import ( context "context" + time "time" + + mock "github.com/stretchr/testify/mock" v0alpha1 "github.com/grafana/grafana/apps/provisioning/pkg/apis/provisioning/v0alpha1" - mock "github.com/stretchr/testify/mock" ) // MockJobProgressRecorder is an autogenerated mock type for the JobProgressRecorder type @@ -271,6 +273,51 @@ func (_c *MockJobProgressRecorder_SetTotal_Call) RunAndReturn(run func(context.C return _c } +// Started provides a mock function with no fields +func (_m *MockJobProgressRecorder) Started() time.Time { + ret := _m.Called() + + if len(ret) == 0 { + panic("no return value specified for Started") + } + + var r0 time.Time + if rf, ok := ret.Get(0).(func() time.Time); ok { + r0 = rf() + } else { + r0 = ret.Get(0).(time.Time) + } + + return r0 +} + +// MockJobProgressRecorder_Started_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'Started' +type MockJobProgressRecorder_Started_Call struct { + *mock.Call +} + +// Started is a helper method to define mock.On call +func (_e *MockJobProgressRecorder_Expecter) Started() *MockJobProgressRecorder_Started_Call { + return &MockJobProgressRecorder_Started_Call{Call: _e.mock.On("Started")} +} + +func (_c *MockJobProgressRecorder_Started_Call) Run(run func()) *MockJobProgressRecorder_Started_Call { + _c.Call.Run(func(args mock.Arguments) { + run() + }) + return _c +} + +func (_c *MockJobProgressRecorder_Started_Call) Return(_a0 time.Time) *MockJobProgressRecorder_Started_Call { + _c.Call.Return(_a0) + return _c +} + +func (_c *MockJobProgressRecorder_Started_Call) RunAndReturn(run func() time.Time) *MockJobProgressRecorder_Started_Call { + _c.Call.Return(run) + return _c +} + // StrictMaxErrors provides a mock function with given fields: maxErrors func (_m *MockJobProgressRecorder) StrictMaxErrors(maxErrors int) { _m.Called(maxErrors) diff --git a/pkg/registry/apis/provisioning/jobs/progress.go b/pkg/registry/apis/provisioning/jobs/progress.go index 95b99ea704f..97a293a5d8c 100644 --- a/pkg/registry/apis/provisioning/jobs/progress.go +++ b/pkg/registry/apis/provisioning/jobs/progress.go @@ -69,6 +69,10 @@ func newJobProgressRecorder(ProgressFn ProgressFn) JobProgressRecorder { } } +func (r *jobProgressRecorder) Started() time.Time { + return r.started +} + func (r *jobProgressRecorder) Record(ctx context.Context, result JobResourceResult) { var shouldLogError bool var logErr error diff --git a/pkg/registry/apis/provisioning/jobs/queue.go b/pkg/registry/apis/provisioning/jobs/queue.go index 14f6b8bd9c7..e1992395efd 100644 --- a/pkg/registry/apis/provisioning/jobs/queue.go +++ b/pkg/registry/apis/provisioning/jobs/queue.go @@ -2,6 +2,7 @@ package jobs import ( "context" + "time" provisioning "github.com/grafana/grafana/apps/provisioning/pkg/apis/provisioning/v0alpha1" "github.com/grafana/grafana/apps/provisioning/pkg/repository" @@ -18,6 +19,7 @@ type RepoGetter interface { // //go:generate mockery --name JobProgressRecorder --structname MockJobProgressRecorder --inpackage --filename job_progress_recorder_mock.go --with-expecter type JobProgressRecorder interface { + Started() time.Time Record(ctx context.Context, result JobResourceResult) ResetResults() SetFinalMessage(ctx context.Context, msg string)