From ee600975f03890e8c32b655495f29962c35360d5 Mon Sep 17 00:00:00 2001 From: Eric Rakestraw Date: Sun, 9 Aug 2026 00:52:31 +0000 Subject: [PATCH] Prevent scheduler callbacks after cancellation --- internal/framework/llm/scheduler.go | 3 + internal/framework/llm/scheduler_test.go | 73 ++++++++++++++++++++++++ 2 files changed, 76 insertions(+) diff --git a/internal/framework/llm/scheduler.go b/internal/framework/llm/scheduler.go index 9d30223..f9a1a40 100644 --- a/internal/framework/llm/scheduler.go +++ b/internal/framework/llm/scheduler.go @@ -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) } diff --git a/internal/framework/llm/scheduler_test.go b/internal/framework/llm/scheduler_test.go index d06ea00..cb1d53d 100644 --- a/internal/framework/llm/scheduler_test.go +++ b/internal/framework/llm/scheduler_test.go @@ -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)