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>
This commit is contained in:
Roberto Jiménez Sánchez
2025-11-05 10:07:21 +00:00
committed by GitHub
co-authored by Copilot
parent 4ceb7eec52
commit 571e5c2e3c
4 changed files with 129 additions and 34 deletions
+74 -32
View File
@@ -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
}
@@ -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)
@@ -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
@@ -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)