package llm import ( "context" "errors" "reflect" "runtime" "sync" "sync/atomic" "testing" "time" ) func TestSchedulerFIFOOrdering(t *testing.T) { s, err := NewScheduler(1) if err != nil { t.Fatalf("NewScheduler: %v", err) } firstRelease, err := s.Acquire(context.Background()) if err != nil { t.Fatalf("Acquire(first): %v", err) } gotOrder := make(chan int, 3) waitChans := []chan struct{}{ make(chan struct{}), make(chan struct{}), make(chan struct{}), } var wg sync.WaitGroup for i := 0; i < 3; i++ { id := i wg.Add(1) go func() { defer wg.Done() runErr := s.Run(context.Background(), func(context.Context) error { gotOrder <- id <-waitChans[id] return nil }) if runErr != nil { t.Errorf("Run[%d] error: %v", id, runErr) } }() waitForQueueDepth(t, s, i+1) } firstRelease() order := make([]int, 0, 3) for i := 0; i < 3; i++ { id := <-gotOrder order = append(order, id) close(waitChans[id]) } wg.Wait() if !reflect.DeepEqual(order, []int{0, 1, 2}) { t.Fatalf("expected FIFO order [0 1 2], got %v", order) } } func TestSchedulerEnforcesMaxConcurrency(t *testing.T) { s, err := NewScheduler(2) if err != nil { t.Fatalf("NewScheduler: %v", err) } var inFlight int32 var maxInFlight int32 release := make(chan struct{}) var wg sync.WaitGroup for i := 0; i < 12; i++ { wg.Add(1) go func() { defer wg.Done() err := s.Run(context.Background(), func(ctx context.Context) error { _ = ctx current := atomic.AddInt32(&inFlight, 1) for { seen := atomic.LoadInt32(&maxInFlight) if current <= seen || atomic.CompareAndSwapInt32(&maxInFlight, seen, current) { break } } <-release atomic.AddInt32(&inFlight, -1) return nil }) if err != nil { t.Errorf("Run error: %v", err) } }() } waitForMinInFlight(t, &maxInFlight, 2) close(release) wg.Wait() if got := atomic.LoadInt32(&maxInFlight); got > 2 { t.Fatalf("expected max in-flight <= 2, got %d", got) } } func TestSchedulerReleasesPermitOnSuccess(t *testing.T) { s, err := NewScheduler(1) if err != nil { t.Fatalf("NewScheduler: %v", err) } for i := 0; i < 3; i++ { if err := s.Run(context.Background(), func(ctx context.Context) error { _ = ctx return nil }); err != nil { t.Fatalf("Run[%d] error: %v", i, err) } } } func TestSchedulerReleasesPermitOnError(t *testing.T) { s, err := NewScheduler(1) if err != nil { t.Fatalf("NewScheduler: %v", err) } expectedErr := errors.New("boom") err = s.Run(context.Background(), func(ctx context.Context) error { _ = ctx return expectedErr }) if !errors.Is(err, expectedErr) { t.Fatalf("expected %v, got %v", expectedErr, err) } if err := s.Run(context.Background(), func(ctx context.Context) error { _ = ctx return nil }); err != nil { t.Fatalf("second Run error: %v", err) } } func TestSchedulerContextCancellationWhileQueued(t *testing.T) { s, err := NewScheduler(1) if err != nil { t.Fatalf("NewScheduler: %v", err) } release, err := s.Acquire(context.Background()) if err != nil { t.Fatalf("Acquire: %v", err) } defer release() ctx, cancel := context.WithTimeout(context.Background(), 25*time.Millisecond) defer cancel() _, err = s.Acquire(ctx) if err == nil { t.Fatalf("expected context cancellation error") } if !errors.Is(err, context.DeadlineExceeded) { t.Fatalf("expected deadline exceeded, got %v", err) } } func TestSchedulerNoPermitLeakAfterQueuedCancellation(t *testing.T) { s, err := NewScheduler(1) if err != nil { t.Fatalf("NewScheduler: %v", err) } firstRelease, err := s.Acquire(context.Background()) if err != nil { t.Fatalf("Acquire(first): %v", err) } ctx, cancel := context.WithCancel(context.Background()) cancel() if _, err := s.Acquire(ctx); !errors.Is(err, context.Canceled) { t.Fatalf("expected context canceled, got %v", err) } firstRelease() if err := s.Run(context.Background(), func(context.Context) error { return nil }); err != nil { t.Fatalf("expected scheduler to accept new work after cancellation, got %v", err) } } func waitForQueueDepth(t *testing.T, s *Scheduler, want int) { t.Helper() deadline := time.Now().Add(250 * time.Millisecond) for time.Now().Before(deadline) { s.mu.Lock() depth := len(s.queue) s.mu.Unlock() if depth >= want { return } runtime.Gosched() } s.mu.Lock() depth := len(s.queue) s.mu.Unlock() t.Fatalf("timed out waiting for queue depth >= %d (got %d)", want, depth) } func waitForMinInFlight(t *testing.T, maxInFlight *int32, want int32) { t.Helper() deadline := time.Now().Add(250 * time.Millisecond) for time.Now().Before(deadline) { if atomic.LoadInt32(maxInFlight) >= want { return } runtime.Gosched() } t.Fatalf("timed out waiting for max in-flight >= %d (got %d)", want, atomic.LoadInt32(maxInFlight)) }