Prevent scheduler callbacks after cancellation
This commit is contained in:
@@ -83,6 +83,9 @@ func (s *Scheduler) Run(ctx context.Context, fn func(context.Context) error) err
|
||||
return err
|
||||
}
|
||||
defer release()
|
||||
if err := ctx.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
return fn(ctx)
|
||||
}
|
||||
|
||||
|
||||
@@ -133,6 +133,79 @@ func TestSchedulerRunReleasesPermitAfterError(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
type cancelAfterAdmissionContext struct {
|
||||
done chan struct{}
|
||||
checks atomic.Int32
|
||||
once sync.Once
|
||||
}
|
||||
|
||||
func newCancelAfterAdmissionContext() *cancelAfterAdmissionContext {
|
||||
return &cancelAfterAdmissionContext{done: make(chan struct{})}
|
||||
}
|
||||
|
||||
func (*cancelAfterAdmissionContext) Deadline() (time.Time, bool) { return time.Time{}, false }
|
||||
|
||||
func (c *cancelAfterAdmissionContext) Done() <-chan struct{} { return c.done }
|
||||
|
||||
func (c *cancelAfterAdmissionContext) Err() error {
|
||||
if c.checks.Add(1) == 1 {
|
||||
return nil
|
||||
}
|
||||
c.once.Do(func() { close(c.done) })
|
||||
return context.Canceled
|
||||
}
|
||||
|
||||
func (*cancelAfterAdmissionContext) Value(any) any { return nil }
|
||||
|
||||
func TestSchedulerDoesNotDispatchCanceledAdmissionAndReleasesNextWaiter(t *testing.T) {
|
||||
s, err := NewScheduler(1)
|
||||
if err != nil {
|
||||
t.Fatalf("NewScheduler: %v", err)
|
||||
}
|
||||
hold, err := s.Acquire(context.Background())
|
||||
if err != nil {
|
||||
t.Fatalf("Acquire: %v", err)
|
||||
}
|
||||
defer hold()
|
||||
|
||||
ctx := newCancelAfterAdmissionContext()
|
||||
var canceledCalls atomic.Int32
|
||||
canceledErr := make(chan error, 1)
|
||||
go func() {
|
||||
canceledErr <- s.Run(ctx, func(context.Context) error {
|
||||
canceledCalls.Add(1)
|
||||
return nil
|
||||
})
|
||||
}()
|
||||
waitForQueueDepth(t, s, 1)
|
||||
|
||||
nextStarted := make(chan struct{})
|
||||
nextErr := make(chan error, 1)
|
||||
go func() {
|
||||
nextErr <- s.Run(context.Background(), func(context.Context) error {
|
||||
close(nextStarted)
|
||||
return nil
|
||||
})
|
||||
}()
|
||||
waitForQueueDepth(t, s, 2)
|
||||
|
||||
hold()
|
||||
if err := <-canceledErr; !errors.Is(err, context.Canceled) {
|
||||
t.Fatalf("Run() error = %v, want context.Canceled", err)
|
||||
}
|
||||
if got := canceledCalls.Load(); got != 0 {
|
||||
t.Fatalf("canceled callback calls = %d, want none", got)
|
||||
}
|
||||
select {
|
||||
case <-nextStarted:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatalf("timed out waiting for next FIFO waiter")
|
||||
}
|
||||
if err := <-nextErr; err != nil {
|
||||
t.Fatalf("next Run() error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func waitForAtomicAtLeast(t *testing.T, value *int32, want int32) {
|
||||
t.Helper()
|
||||
deadline := time.Now().Add(time.Second)
|
||||
|
||||
Reference in New Issue
Block a user