Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -491,7 +491,7 @@ func (e *DefaultExptTurnEvaluationImpl) callEvaluators(ctx context.Context, exec
}
}

err = pool.Exec(ctx)
err = pool.ExecAll(ctx)
records := make(map[int64]*entity.EvaluatorRecord, len(expt.Evaluators))
recordMap.Range(func(key, value interface{}) bool {
record, _ := value.(*entity.EvaluatorRecord)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -3213,6 +3213,126 @@ func TestDefaultExptTurnEvaluationImpl_callEvaluators_EdgeCases(t *testing.T) {
}
}

// TestDefaultExptTurnEvaluationImpl_callEvaluators_ExecAll covers the change in callEvaluators
// from pool.Exec to pool.ExecAll: when multiple sync evaluators run concurrently, even if one
// of them fails, the rest should still run to completion (instead of failing fast and skipping
// the others), aggregating the errors on return while still collecting the successful results.
func TestDefaultExptTurnEvaluationImpl_callEvaluators_ExecAll(t *testing.T) {
t.Parallel()

mockContent := &entity.Content{Text: gptr.Of("value1")}
mockTargetResult := &entity.EvalTargetRecord{
EvalTargetOutputData: &entity.EvalTargetOutputData{
OutputFields: map[string]*entity.Content{"field1": mockContent},
},
}

newEvaluatorConf := func(verID int64) *entity.EvaluatorConf {
return &entity.EvaluatorConf{
EvaluatorVersionID: verID,
IngressConf: &entity.EvaluatorIngressConf{
EvalSetAdapter: &entity.FieldAdapter{
FieldConfs: []*entity.FieldConf{{FieldName: "field1", FromField: "field1"}},
},
TargetAdapter: &entity.FieldAdapter{
FieldConfs: []*entity.FieldConf{{FieldName: "field1", FromField: "field1"}},
},
},
}
}

// Two sync evaluators: version 1 fails, version 2 succeeds.
newEtec := func() *entity.ExptTurnEvalCtx {
return &entity.ExptTurnEvalCtx{
ExptItemEvalCtx: &entity.ExptItemEvalCtx{
EvalSetItem: &entity.EvaluationSetItem{ItemID: 1},
Event: &entity.ExptItemEvalEvent{ExptID: 1, SpaceID: 2},
Expt: &entity.Experiment{
SpaceID: 2,
Evaluators: []*entity.Evaluator{
{ID: 1, EvaluatorType: entity.EvaluatorTypePrompt, PromptEvaluatorVersion: &entity.PromptEvaluatorVersion{ID: 1}},
{ID: 2, EvaluatorType: entity.EvaluatorTypePrompt, PromptEvaluatorVersion: &entity.PromptEvaluatorVersion{ID: 2}},
},
EvalConf: &entity.EvaluationConfiguration{
ItemConcurNum: gptr.Of(1),
ConnectorConf: entity.Connector{
EvaluatorsConf: &entity.EvaluatorsConf{
// concurrency 2 so both evaluators can be submitted concurrently
EvaluatorConcurNum: gptr.Of(2),
EvaluatorConf: []*entity.EvaluatorConf{newEvaluatorConf(1), newEvaluatorConf(2)},
},
},
},
},
},
ExptTurnRunResult: &entity.ExptTurnRunResult{},
Turn: &entity.Turn{FieldDataList: []*entity.FieldData{{Name: "field1", Content: mockContent}}},
}
}

t.Run("one evaluator fails, the other still runs and result is collected", func(t *testing.T) {
ctrl := gomock.NewController(t)
defer ctrl.Finish()

mockMetric := metricsmocks.NewMockExptMetric(ctrl)
mockEvaluatorService := svcmocks.NewMockEvaluatorService(ctrl)
service := &DefaultExptTurnEvaluationImpl{metric: mockMetric, evaluatorService: mockEvaluatorService}

mockMetric.EXPECT().EmitTurnExecEvaluatorResult(gomock.Any(), gomock.Any()).AnyTimes()
mockEvaluatorService.EXPECT().ShouldInterceptEvaluator(gomock.Any(), gomock.Any()).Return(nil, false, nil).AnyTimes()

// version 1 fails, version 2 succeeds.
mockEvaluatorService.EXPECT().RunEvaluator(gomock.Any(), gomock.Any()).
DoAndReturn(func(_ context.Context, req *entity.RunEvaluatorRequest) (*entity.EvaluatorRecord, error) {
if req.EvaluatorVersionID == 1 {
return nil, errors.New("run evaluator 1 failed")
}
return &entity.EvaluatorRecord{ID: 2, Status: entity.EvaluatorRunStatusSuccess}, nil
}).AnyTimes()

records, err := service.callEvaluators(context.Background(), []int64{1, 2}, newEtec(), mockTargetResult, []*entity.Message{})

// The failure of the first evaluator should be returned.
assert.Error(t, err)
// The second evaluator still runs successfully; its result should be collected into recordMap.
assert.NotNil(t, records[2])
assert.Equal(t, int64(2), records[2].ID)
})

t.Run("multiple evaluators fail and all errors are aggregated", func(t *testing.T) {
// Deterministic difference of ExecAll: when multiple sync evaluators fail at the same
// time, pool.ExecAll uses errors.Join to aggregate all errors and passes them through
// as-is, whereas pool.Exec returns only one of them. Here we assert that the returned
// error contains both evaluators' error messages, pinning down the ExecAll semantics.
ctrl := gomock.NewController(t)
defer ctrl.Finish()

mockMetric := metricsmocks.NewMockExptMetric(ctrl)
mockEvaluatorService := svcmocks.NewMockEvaluatorService(ctrl)
service := &DefaultExptTurnEvaluationImpl{metric: mockMetric, evaluatorService: mockEvaluatorService}

err1Msg := "run evaluator 1 failed"
err2Msg := "run evaluator 2 failed"

mockMetric.EXPECT().EmitTurnExecEvaluatorResult(gomock.Any(), gomock.Any()).AnyTimes()
mockEvaluatorService.EXPECT().ShouldInterceptEvaluator(gomock.Any(), gomock.Any()).Return(nil, false, nil).AnyTimes()
mockEvaluatorService.EXPECT().RunEvaluator(gomock.Any(), gomock.Any()).
DoAndReturn(func(_ context.Context, req *entity.RunEvaluatorRequest) (*entity.EvaluatorRecord, error) {
if req.EvaluatorVersionID == 1 {
return nil, errors.New(err1Msg)
}
return nil, errors.New(err2Msg)
}).AnyTimes()

_, err := service.callEvaluators(context.Background(), []int64{1, 2}, newEtec(), mockTargetResult, []*entity.Message{})

assert.Error(t, err)
// Both error messages should appear after errors.Join aggregation; Exec would return only one.
assert.Contains(t, err.Error(), err1Msg)
assert.Contains(t, err.Error(), err2Msg)
})
}

func Test_deepCopyEvaluatorInputData(t *testing.T) {
t.Parallel()

Expand Down
38 changes: 27 additions & 11 deletions backend/pkg/lang/goroutine/pool.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,9 +5,9 @@ package goroutine

import (
"context"
"errors"
"fmt"
"sync"
"sync/atomic"

"github.com/panjf2000/ants/v2"
)
Expand Down Expand Up @@ -55,13 +55,26 @@ func (p *pool) exec(ctx context.Context, ignoreErr bool) error {
defer p.p.Release()

var (
gerr atomic.Value
mu sync.Mutex
errs []error
wg sync.WaitGroup
)

appendErr := func(err error) {
mu.Lock()
errs = append(errs, err)
mu.Unlock()
}

hasErr := func() bool {
mu.Lock()
defer mu.Unlock()
return len(errs) > 0
}

for idx := range p.tasks {
if !ignoreErr && gerr.Load() != nil {
return gerr.Load().(error)
if !ignoreErr && hasErr() {
break
}

t := p.tasks[idx]
Expand All @@ -73,15 +86,15 @@ func (p *pool) exec(ctx context.Context, ignoreErr bool) error {

select {
case <-ctx.Done():
gerr.Store(ctx.Err())
appendErr(ctx.Err())
return

default:
if !ignoreErr && gerr.Load() != nil {
if !ignoreErr && hasErr() {
return
}
if err := t(); err != nil {
gerr.Store(err)
appendErr(err)
}
return
}
Expand All @@ -91,9 +104,12 @@ func (p *pool) exec(ctx context.Context, ignoreErr bool) error {
}

wg.Wait()
if gerr.Load() != nil {
return gerr.Load().(error)
}

return nil
if len(errs) == 0 {
return nil
}
if ignoreErr {
return errors.Join(errs...)
}
return errs[0]
}
105 changes: 105 additions & 0 deletions backend/pkg/lang/goroutine/pool_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ package goroutine
import (
"context"
"errors"
"fmt"
"math/rand"
"sync"
"testing"
Expand Down Expand Up @@ -103,6 +104,23 @@ func TestPool_Exec(t *testing.T) {
assert.Error(t, err)
assert.Equal(t, context.Canceled, err)
})

t.Run("fail fast returns a single error not aggregated", func(t *testing.T) {
// Exec fails fast and should return a single error, not errors.Join of multiple errors.
ctx := context.Background()
sentinel := errors.New("sentinel")

pool, err := NewPool(1)
assert.NoError(t, err)

pool.Add(func() error { return sentinel })
pool.Add(func() error { return errors.New("should not run or be aggregated") })

err = pool.Exec(ctx)
assert.Error(t, err)
// The returned value is the error itself, not wrapped by Join (Join would make Error() multi-line).
assert.Equal(t, sentinel, err)
})
}

func TestPool_ExecAll(t *testing.T) {
Expand Down Expand Up @@ -137,8 +155,95 @@ func TestPool_ExecAll(t *testing.T) {
gslice.ToMap(results, func(v int) (int, bool) { return v, true }),
)
})

t.Run("aggregate multiple errors", func(t *testing.T) {
ctx := context.Background()
err1 := errors.New("err1")
err2 := errors.New("err2")
err3 := errors.New("err3")

pool, err := NewPool(3)
assert.NoError(t, err)

pool.Add(func() error { return err1 })
pool.Add(func() error { return nil })
pool.Add(func() error { return err2 })
pool.Add(func() error { return err3 })

err = pool.ExecAll(ctx)
assert.Error(t, err)
// errors.Join aggregates; each error should still match via errors.Is
assert.True(t, errors.Is(err, err1))
assert.True(t, errors.Is(err, err2))
assert.True(t, errors.Is(err, err3))
})

t.Run("aggregated error supports errors.As", func(t *testing.T) {
ctx := context.Background()

pool, err := NewPool(2)
assert.NoError(t, err)

pool.Add(func() error { return errors.New("plain error") })
pool.Add(func() error { return &customError{msg: "custom"} })

err = pool.ExecAll(ctx)
assert.Error(t, err)
var ce *customError
assert.True(t, errors.As(err, &ce))
assert.Equal(t, "custom", ce.msg)
})

t.Run("no error returns nil", func(t *testing.T) {
ctx := context.Background()

pool, err := NewPool(3)
assert.NoError(t, err)

for i := 0; i < 5; i++ {
pool.Add(func() error { return nil })
}

err = pool.ExecAll(ctx)
assert.NoError(t, err)
})

t.Run("concurrent errors of different concrete types do not panic", func(t *testing.T) {
// Regression: an earlier implementation stored errors in an atomic.Value, where
// concurrently storing errors of different concrete types would panic. Here we
// concurrently return errors of multiple concrete types to verify it no longer panics.
ctx := context.Background()

pool, err := NewPool(8)
assert.NoError(t, err)

for i := 0; i < 50; i++ {
idx := i
pool.Add(func() error {
switch idx % 3 {
case 0:
return errors.New("plain")
case 1:
return &customError{msg: "custom"}
default:
return fmt.Errorf("wrapped %d: %w", idx, errors.New("inner"))
}
})
}

assert.NotPanics(t, func() {
err = pool.ExecAll(ctx)
})
assert.Error(t, err)
})
}

type customError struct {
msg string
}

func (e *customError) Error() string { return e.msg }

func Test_pool_execute(t *testing.T) {
t.Run("execute tasks with pool size equal to task count", func(t *testing.T) {
ctx := context.Background()
Expand Down
Loading