Prevent scheduler callbacks after cancellation

This commit is contained in:
2026-08-09 00:52:31 +00:00
parent 0d8017e23f
commit ee600975f0
2 changed files with 76 additions and 0 deletions

View File

@@ -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)
}

View File

@@ -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)