diff --git a/provisionerd/provisionerd_test.go b/provisionerd/provisionerd_test.go index 1999061dd4..c4b82472e0 100644 --- a/provisionerd/provisionerd_test.go +++ b/provisionerd/provisionerd_test.go @@ -130,8 +130,11 @@ func TestProvisionerd(t *testing.T) { // Ensures tars with "../../../etc/passwd" as the path // are not allowed to run, and will fail the job. t.Parallel() - var complete sync.Once - completeChan := make(chan struct{}) + var ( + completeChan = make(chan struct{}) + completeOnce sync.Once + ) + closer := createProvisionerd(t, func(ctx context.Context) (proto.DRPCProvisionerDaemonClient, error) { return createProvisionerDaemonClient(t, provisionerDaemonTestServer{ acquireJob: func(ctx context.Context, _ *proto.Empty) (*proto.AcquiredJob, error) { @@ -150,9 +153,7 @@ func TestProvisionerd(t *testing.T) { }, updateJob: noopUpdateJob, failJob: func(ctx context.Context, job *proto.FailedJob) (*proto.Empty, error) { - complete.Do(func() { - close(completeChan) - }) + completeOnce.Do(func() { close(completeChan) }) return &proto.Empty{}, nil }, }), nil @@ -165,8 +166,11 @@ func TestProvisionerd(t *testing.T) { t.Run("RunningPeriodicUpdate", func(t *testing.T) { t.Parallel() - var complete sync.Once - completeChan := make(chan struct{}) + var ( + completeChan = make(chan struct{}) + completeOnce sync.Once + ) + closer := createProvisionerd(t, func(ctx context.Context) (proto.DRPCProvisionerDaemonClient, error) { return createProvisionerDaemonClient(t, provisionerDaemonTestServer{ acquireJob: func(ctx context.Context, _ *proto.Empty) (*proto.AcquiredJob, error) { @@ -184,9 +188,7 @@ func TestProvisionerd(t *testing.T) { }, nil }, updateJob: func(ctx context.Context, update *proto.UpdateJobRequest) (*proto.UpdateJobResponse, error) { - complete.Do(func() { - close(completeChan) - }) + completeOnce.Do(func() { close(completeChan) }) return &proto.UpdateJobResponse{}, nil }, failJob: func(ctx context.Context, job *proto.FailedJob) (*proto.Empty, error) { @@ -212,19 +214,18 @@ func TestProvisionerd(t *testing.T) { didLog atomic.Bool didAcquireJob atomic.Bool didDryRun atomic.Bool + completeChan = make(chan struct{}) + completeOnce sync.Once ) - var complete sync.Once - completeChan := make(chan struct{}) + closer := createProvisionerd(t, func(ctx context.Context) (proto.DRPCProvisionerDaemonClient, error) { return createProvisionerDaemonClient(t, provisionerDaemonTestServer{ acquireJob: func(ctx context.Context, _ *proto.Empty) (*proto.AcquiredJob, error) { - if didAcquireJob.Load() { - complete.Do(func() { - close(completeChan) - }) + if !didAcquireJob.CAS(false, true) { + completeOnce.Do(func() { close(completeChan) }) return &proto.AcquiredJob{}, nil } - didAcquireJob.Store(true) + return &proto.AcquiredJob{ JobId: "test", Provisioner: "someprovisioner", @@ -315,19 +316,18 @@ func TestProvisionerd(t *testing.T) { didComplete atomic.Bool didLog atomic.Bool didAcquireJob atomic.Bool + completeChan = make(chan struct{}) + completeOnce sync.Once ) - var complete sync.Once - completeChan := make(chan struct{}) + closer := createProvisionerd(t, func(ctx context.Context) (proto.DRPCProvisionerDaemonClient, error) { return createProvisionerDaemonClient(t, provisionerDaemonTestServer{ acquireJob: func(ctx context.Context, _ *proto.Empty) (*proto.AcquiredJob, error) { - if didAcquireJob.Load() { - complete.Do(func() { - close(completeChan) - }) + if !didAcquireJob.CAS(false, true) { + completeOnce.Do(func() { close(completeChan) }) return &proto.AcquiredJob{}, nil } - didAcquireJob.Store(true) + return &proto.AcquiredJob{ JobId: "test", Provisioner: "someprovisioner", @@ -386,16 +386,18 @@ func TestProvisionerd(t *testing.T) { var ( didFail atomic.Bool didAcquireJob atomic.Bool + completeChan = make(chan struct{}) + completeOnce sync.Once ) - completeChan := make(chan struct{}) + closer := createProvisionerd(t, func(ctx context.Context) (proto.DRPCProvisionerDaemonClient, error) { return createProvisionerDaemonClient(t, provisionerDaemonTestServer{ acquireJob: func(ctx context.Context, _ *proto.Empty) (*proto.AcquiredJob, error) { - if didAcquireJob.Load() { - close(completeChan) + if !didAcquireJob.CAS(false, true) { + completeOnce.Do(func() { close(completeChan) }) return &proto.AcquiredJob{}, nil } - didAcquireJob.Store(true) + return &proto.AcquiredJob{ JobId: "test", Provisioner: "someprovisioner", @@ -585,10 +587,15 @@ func TestProvisionerd(t *testing.T) { t.Run("ReconnectAndFail", func(t *testing.T) { t.Parallel() - var second atomic.Bool - failChan := make(chan struct{}) - failedChan := make(chan struct{}) - completeChan := make(chan struct{}) + var ( + second atomic.Bool + failChan = make(chan struct{}) + failOnce sync.Once + failedChan = make(chan struct{}) + failedOnce sync.Once + completeChan = make(chan struct{}) + completeOnce sync.Once + ) server := createProvisionerd(t, func(ctx context.Context) (proto.DRPCProvisionerDaemonClient, error) { client := createProvisionerDaemonClient(t, provisionerDaemonTestServer{ acquireJob: func(ctx context.Context, _ *proto.Empty) (*proto.AcquiredJob, error) { @@ -613,10 +620,10 @@ func TestProvisionerd(t *testing.T) { }, failJob: func(ctx context.Context, job *proto.FailedJob) (*proto.Empty, error) { if second.Load() { - close(completeChan) + completeOnce.Do(func() { close(completeChan) }) return &proto.Empty{}, nil } - close(failChan) + failOnce.Do(func() { close(failChan) }) <-failedChan return &proto.Empty{}, nil }, @@ -626,7 +633,7 @@ func TestProvisionerd(t *testing.T) { <-failChan _ = client.DRPCConn().Close() second.Store(true) - close(failedChan) + failedOnce.Do(func() { close(failedChan) }) }() } return client, nil @@ -651,18 +658,20 @@ func TestProvisionerd(t *testing.T) { t.Run("ReconnectAndComplete", func(t *testing.T) { t.Parallel() - var completed sync.Once - var second atomic.Bool - failChan := make(chan struct{}) - failedChan := make(chan struct{}) - completeChan := make(chan struct{}) + var ( + second atomic.Bool + failChan = make(chan struct{}) + failOnce sync.Once + failedChan = make(chan struct{}) + failedOnce sync.Once + completeChan = make(chan struct{}) + completeOnce sync.Once + ) server := createProvisionerd(t, func(ctx context.Context) (proto.DRPCProvisionerDaemonClient, error) { client := createProvisionerDaemonClient(t, provisionerDaemonTestServer{ acquireJob: func(ctx context.Context, _ *proto.Empty) (*proto.AcquiredJob, error) { if second.Load() { - completed.Do(func() { - close(completeChan) - }) + completeOnce.Do(func() { close(completeChan) }) return &proto.AcquiredJob{}, nil } return &proto.AcquiredJob{ @@ -688,7 +697,7 @@ func TestProvisionerd(t *testing.T) { if second.Load() { return &proto.Empty{}, nil } - close(failChan) + failOnce.Do(func() { close(failChan) }) <-failedChan return &proto.Empty{}, nil }, @@ -698,7 +707,7 @@ func TestProvisionerd(t *testing.T) { <-failChan _ = client.DRPCConn().Close() second.Store(true) - close(failedChan) + failedOnce.Do(func() { close(failedChan) }) }() } return client, nil