package llm import ( "context" "errors" "runtime" "sync" "sync/atomic" "testing" "time" ) func TestNewSchedulerValidation(t *testing.T) { if _, err := NewScheduler(0); err == nil { t.Fatalf("expected validation error") } } func TestSchedulerMaxConcurrency(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() runErr := s.Run(context.Background(), func(context.Context) error { 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 runErr != nil { t.Errorf("Run error: %v", runErr) } }() } waitForAtomicAtLeast(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 TestSchedulerCancellationWhileQueued(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.WithCancel(context.Background()) errCh := make(chan error, 1) go func() { _, acquireErr := s.Acquire(ctx) errCh <- acquireErr }() waitForQueueDepth(t, s, 1) cancel() select { case acquireErr := <-errCh: if !errors.Is(acquireErr, context.Canceled) { t.Fatalf("expected context canceled, got %v", acquireErr) } case <-time.After(time.Second): t.Fatalf("timed out waiting for queued acquire to cancel") } release() if err := s.Run(context.Background(), func(context.Context) error { return nil }); err != nil { t.Fatalf("expected scheduler to accept work after cancellation, got %v", err) } } func TestSchedulerReleaseFunctionIsIdempotent(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) } release() release() if err := s.Run(context.Background(), func(context.Context) error { return nil }); err != nil { t.Fatalf("expected permit to be released once, got %v", err) } } func TestSchedulerRunReleasesPermitAfterError(t *testing.T) { s, err := NewScheduler(1) if err != nil { t.Fatalf("NewScheduler: %v", err) } expected := errors.New("failed") err = s.Run(context.Background(), func(context.Context) error { return expected }) if !errors.Is(err, expected) { t.Fatalf("expected %v, got %v", expected, err) } if err := s.Run(context.Background(), func(context.Context) error { return nil }); err != nil { t.Fatalf("expected permit to be released after error, got %v", err) } } func waitForAtomicAtLeast(t *testing.T, value *int32, want int32) { t.Helper() deadline := time.Now().Add(time.Second) for time.Now().Before(deadline) { if atomic.LoadInt32(value) >= want { return } runtime.Gosched() } t.Fatalf("timed out waiting for value >= %d; got %d", want, atomic.LoadInt32(value)) } func waitForQueueDepth(t *testing.T, s *Scheduler, want int) { t.Helper() deadline := time.Now().Add(time.Second) 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) }