Files
audita/internal/framework/llm/scheduler_test.go

224 lines
4.8 KiB
Go

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