package llm import ( "context" "errors" "sync" "sync/atomic" "testing" "time" ) func TestSchedulerEnforcesMaxConcurrency(t *testing.T) { s, err := NewScheduler(2) if err != nil { t.Fatalf("NewScheduler: %v", err) } var inFlight int32 var maxInFlight int32 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 } } time.Sleep(20 * time.Millisecond) atomic.AddInt32(&inFlight, -1) return nil }) if err != nil { t.Errorf("Run error: %v", err) } }() } 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) } // Must be able to run again after an error, proving permit release. if err := s.Run(context.Background(), func(ctx context.Context) error { _ = ctx return nil }); err != nil { t.Fatalf("second Run error: %v", err) } } func TestSchedulerRespectsContextCancellation(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) } }