diff --git a/backend/modules/evaluation/domain/service/expt_run_item_turn_impl.go b/backend/modules/evaluation/domain/service/expt_run_item_turn_impl.go index 0d50de3a1f..7dfca6f335 100644 --- a/backend/modules/evaluation/domain/service/expt_run_item_turn_impl.go +++ b/backend/modules/evaluation/domain/service/expt_run_item_turn_impl.go @@ -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) diff --git a/backend/modules/evaluation/domain/service/expt_run_item_turn_impl_test.go b/backend/modules/evaluation/domain/service/expt_run_item_turn_impl_test.go index 2eb2563ea9..f6547beae7 100644 --- a/backend/modules/evaluation/domain/service/expt_run_item_turn_impl_test.go +++ b/backend/modules/evaluation/domain/service/expt_run_item_turn_impl_test.go @@ -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() diff --git a/backend/pkg/lang/goroutine/pool.go b/backend/pkg/lang/goroutine/pool.go index e5e9709822..6e273f5d2c 100644 --- a/backend/pkg/lang/goroutine/pool.go +++ b/backend/pkg/lang/goroutine/pool.go @@ -5,9 +5,9 @@ package goroutine import ( "context" + "errors" "fmt" "sync" - "sync/atomic" "github.com/panjf2000/ants/v2" ) @@ -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] @@ -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 } @@ -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] } diff --git a/backend/pkg/lang/goroutine/pool_test.go b/backend/pkg/lang/goroutine/pool_test.go index 733e60ca39..98f0e7af77 100644 --- a/backend/pkg/lang/goroutine/pool_test.go +++ b/backend/pkg/lang/goroutine/pool_test.go @@ -6,6 +6,7 @@ package goroutine import ( "context" "errors" + "fmt" "math/rand" "sync" "testing" @@ -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) { @@ -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()