grok2api / backend /internal /pkg /batch /executor_test.go
fromozuzhouzzz
Deploy grok2api v3.0.11 to HF Spaces
f1dd159
Raw
History Blame Contribute Delete
11.1 kB
package batch
import (
"context"
"errors"
"runtime"
"sync"
"sync/atomic"
"testing"
"time"
)
func TestSharedPoolBoundsConcurrentExecutors(t *testing.T) {
pool := NewPool(3)
var active atomic.Int64
var peak atomic.Int64
work := func(context.Context, int) error {
current := active.Add(1)
for {
value := peak.Load()
if current <= value || peak.CompareAndSwap(value, current) {
break
}
}
time.Sleep(5 * time.Millisecond)
active.Add(-1)
return nil
}
var wait sync.WaitGroup
for range 2 {
wait.Add(1)
go func() {
defer wait.Done()
_, summary, err := Run(context.Background(), make([]int, 10), Options{Workers: 10, Pool: pool}, work)
if err != nil || summary.Succeeded != 10 {
t.Errorf("summary = %#v, err = %v", summary, err)
}
}()
}
wait.Wait()
if peak.Load() != 3 || pool.Snapshot().Peak != 3 {
t.Fatalf("peak = %d, pool = %#v", peak.Load(), pool.Snapshot())
}
}
type leaseLimiterStub struct {
mu sync.Mutex
current map[string]int
}
func (s *leaseLimiterStub) Acquire(_ context.Context, key string, limit int) (func(), bool, error) {
s.mu.Lock()
if s.current == nil {
s.current = make(map[string]int)
}
if s.current[key] >= limit {
s.mu.Unlock()
return nil, false, nil
}
s.current[key]++
s.mu.Unlock()
return func() {
s.mu.Lock()
s.current[key]--
s.mu.Unlock()
}, true, nil
}
func TestDistributedLeaseBoundsSeparateProcessPools(t *testing.T) {
limiter := &leaseLimiterStub{}
pools := []*Pool{NewSharedPool(2, limiter, "bulk"), NewSharedPool(2, limiter, "bulk")}
var active atomic.Int64
var peak atomic.Int64
var wait sync.WaitGroup
for _, pool := range pools {
pool := pool
wait.Add(1)
go func() {
defer wait.Done()
_, _, err := Run(context.Background(), make([]int, 4), Options{Workers: 2, Pool: pool}, func(context.Context, int) error {
current := active.Add(1)
for {
value := peak.Load()
if current <= value || peak.CompareAndSwap(value, current) {
break
}
}
time.Sleep(5 * time.Millisecond)
active.Add(-1)
return nil
})
if err != nil {
t.Errorf("run: %v", err)
}
}()
}
wait.Wait()
if peak.Load() != 2 {
t.Fatalf("distributed peak = %d", peak.Load())
}
}
func TestSharedChildPoolBoundsCategoryAcrossProcesses(t *testing.T) {
limiter := &leaseLimiterStub{}
parents := []*Pool{NewSharedPool(3, limiter, "global"), NewSharedPool(3, limiter, "global")}
children := []*Pool{
NewSharedChildPool(1, limiter, "sync", parents[0]),
NewSharedChildPool(1, limiter, "sync", parents[1]),
}
var active atomic.Int64
var peak atomic.Int64
var wait sync.WaitGroup
for _, pool := range children {
wait.Add(1)
go func() {
defer wait.Done()
_, _, err := Run(context.Background(), make([]int, 2), Options{Workers: 2, Pool: pool}, func(context.Context, int) error {
current := active.Add(1)
for {
value := peak.Load()
if current <= value || peak.CompareAndSwap(value, current) {
break
}
}
time.Sleep(5 * time.Millisecond)
active.Add(-1)
return nil
})
if err != nil {
t.Errorf("run: %v", err)
}
}()
}
wait.Wait()
if peak.Load() != 1 {
t.Fatalf("category peak = %d", peak.Load())
}
}
func TestChildPoolsBoundConcurrentRequestsByCategoryAndGlobalLimit(t *testing.T) {
global := NewPool(3)
refresh := NewChildPool(2, global)
syncPool := NewChildPool(2, global)
var globalActive atomic.Int64
var globalPeak atomic.Int64
var refreshActive atomic.Int64
var refreshPeak atomic.Int64
var syncActive atomic.Int64
var syncPeak atomic.Int64
trackPeak := func(active, peak *atomic.Int64) {
current := active.Add(1)
for {
value := peak.Load()
if current <= value || peak.CompareAndSwap(value, current) {
return
}
}
}
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
start := make(chan struct{})
work := func(active, peak *atomic.Int64) func(context.Context, int) error {
return func(workCtx context.Context, _ int) error {
trackPeak(&globalActive, &globalPeak)
trackPeak(active, peak)
defer globalActive.Add(-1)
defer active.Add(-1)
select {
case <-start:
return nil
case <-workCtx.Done():
return workCtx.Err()
}
}
}
var wait sync.WaitGroup
for _, operation := range []struct {
pool *Pool
active *atomic.Int64
peak *atomic.Int64
}{
{pool: refresh, active: &refreshActive, peak: &refreshPeak},
{pool: refresh, active: &refreshActive, peak: &refreshPeak},
{pool: syncPool, active: &syncActive, peak: &syncPeak},
{pool: syncPool, active: &syncActive, peak: &syncPeak},
} {
operation := operation
wait.Add(1)
go func() {
defer wait.Done()
_, summary, err := Run(ctx, make([]int, 3), Options{Workers: 3, Pool: operation.pool}, work(operation.active, operation.peak))
if err != nil || summary.Succeeded != 3 {
t.Errorf("summary = %#v, err = %v", summary, err)
}
}()
}
deadline := time.After(time.Second)
for globalPeak.Load() < 3 {
select {
case <-deadline:
t.Fatal("并发任务未填满全局容量")
default:
runtime.Gosched()
}
}
close(start)
wait.Wait()
if globalPeak.Load() != 3 {
t.Fatalf("global peak = %d", globalPeak.Load())
}
if refreshPeak.Load() > 2 || syncPeak.Load() > 2 {
t.Fatalf("refresh peak = %d, sync peak = %d", refreshPeak.Load(), syncPeak.Load())
}
}
func TestMapIsolatesFailureAndPanic(t *testing.T) {
results, summary, err := Map(context.Background(), []int{1, 2, 3}, Options{Workers: 3}, func(_ context.Context, value int) (int, error) {
switch value {
case 2:
return 0, errors.New("failed")
case 3:
panic("broken")
default:
return value * 2, nil
}
})
if err != nil {
t.Fatal(err)
}
if results[0].Value != 2 || results[0].Err != nil || results[1].Err == nil {
t.Fatalf("results = %#v", results)
}
var panicErr *PanicError
if !errors.As(results[2].Err, &panicErr) || len(panicErr.Stack) == 0 {
t.Fatalf("panic result = %#v", results[2])
}
if summary.Succeeded != 1 || summary.Failed != 2 || summary.Panicked != 1 {
t.Fatalf("summary = %#v", summary)
}
}
func TestMapObservedReleasesPoolBeforeStartingDownstreamWork(t *testing.T) {
pool := NewPool(1)
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
var downstreamErr error
_, summary, err := MapObserved(ctx, []int{1}, Options{Workers: 1, Pool: pool}, func(context.Context, int) (int, error) {
return 42, nil
}, func(_ int, result Result[int]) {
if result.Err != nil {
downstreamErr = result.Err
return
}
downstreamErr = pool.Do(ctx, func(context.Context) error { return nil })
})
if err != nil || downstreamErr != nil || summary.Succeeded != 1 {
t.Fatalf("summary = %#v, err = %v, downstream = %v", summary, err, downstreamErr)
}
}
func TestForEachObservedReturnsSummaryWithoutCollectedResults(t *testing.T) {
pool := NewPool(2)
var observed atomic.Int64
summary, err := ForEachObserved(context.Background(), []int{1, 2, 3}, Options{Workers: 2, Pool: pool}, func(_ context.Context, value int) (int, error) {
if value == 2 {
return value, errors.New("rejected")
}
return value * 2, nil
}, func(_ int, result Result[int]) {
if !result.Completed {
t.Error("observer received incomplete result")
}
observed.Add(1)
})
if err != nil {
t.Fatal(err)
}
if observed.Load() != 3 || summary.Completed != 3 || summary.Succeeded != 2 || summary.Failed != 1 {
t.Fatalf("observed = %d, summary = %#v", observed.Load(), summary)
}
}
func TestMapStopsSubmittingAfterCancellation(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
_, summary, err := Run(ctx, make([]int, 100), Options{Workers: 2, QueueSize: 1}, func(context.Context, int) error {
cancel()
return nil
})
if !errors.Is(err, context.Canceled) || !summary.Canceled || summary.Submitted >= summary.Total {
t.Fatalf("summary = %#v, err = %v", summary, err)
}
}
func TestPoolHotResizeAndChildLimit(t *testing.T) {
global := NewPool(3)
child := NewChildPool(1, global)
started := make(chan struct{}, 3)
release := make(chan struct{})
done := make(chan error, 3)
for range 3 {
go func() {
done <- child.Do(context.Background(), func(context.Context) error {
started <- struct{}{}
<-release
return nil
})
}()
}
<-started
select {
case <-started:
t.Fatal("child exceeded its initial limit")
case <-time.After(20 * time.Millisecond):
}
child.UpdateLimit(3)
for range 2 {
select {
case <-started:
case <-time.After(time.Second):
t.Fatal("waiting work did not observe the increased limit")
}
}
child.UpdateLimit(1)
close(release)
for range 3 {
if err := <-done; err != nil {
t.Fatal(err)
}
}
if child.Snapshot().Limit != 1 || child.Snapshot().Peak != 3 || global.Snapshot().Peak != 3 {
t.Fatalf("child = %#v, global = %#v", child.Snapshot(), global.Snapshot())
}
}
func TestChildSnapshotSeparatesQueuedFromActiveWork(t *testing.T) {
global := NewPool(1)
child := NewChildPool(1, global)
blockerStarted := make(chan struct{})
releaseBlocker := make(chan struct{})
blockerDone := make(chan error, 1)
go func() {
blockerDone <- global.Do(context.Background(), func(context.Context) error {
close(blockerStarted)
<-releaseBlocker
return nil
})
}()
<-blockerStarted
childStarted := make(chan struct{})
childDone := make(chan error, 1)
go func() {
childDone <- child.Do(context.Background(), func(context.Context) error {
close(childStarted)
return nil
})
}()
deadline := time.After(time.Second)
for {
snapshot := child.Snapshot()
if snapshot.Queued == 1 {
if snapshot.Active != 0 {
t.Fatalf("waiting child snapshot = %#v", snapshot)
}
break
}
select {
case <-deadline:
t.Fatal("child task was not reported as queued")
default:
runtime.Gosched()
}
}
close(releaseBlocker)
if err := <-blockerDone; err != nil {
t.Fatal(err)
}
select {
case <-childStarted:
case <-time.After(time.Second):
t.Fatal("queued child task did not start")
}
if err := <-childDone; err != nil {
t.Fatal(err)
}
if snapshot := child.Snapshot(); snapshot.Active != 0 || snapshot.Queued != 0 || snapshot.Peak != 1 {
t.Fatalf("completed child snapshot = %#v", snapshot)
}
}
func TestPoolJitterHonorsCancellationBeforeWork(t *testing.T) {
pool := NewPool(2)
pool.UpdateJitter(time.Hour)
ctx, cancel := context.WithCancel(context.Background())
cancel()
called := false
err := pool.Do(ctx, func(context.Context) error {
called = true
return nil
})
if !errors.Is(err, context.Canceled) || called {
t.Fatalf("err = %v, called = %t", err, called)
}
}
func TestPoolJitterIsSkippedForSerialExecution(t *testing.T) {
pool := NewPool(1)
pool.UpdateJitter(time.Hour)
startedAt := time.Now()
if err := pool.Do(context.Background(), func(context.Context) error { return nil }); err != nil {
t.Fatal(err)
}
if elapsed := time.Since(startedAt); elapsed > 100*time.Millisecond {
t.Fatalf("serial execution was delayed by %s", elapsed)
}
}