fix(service): ParallelAgent 使用 goroutine 并发执行子 Agent

This commit is contained in:
hhs
2026-06-10 15:54:52 +08:00
parent 8c515a8b3e
commit 1f8ba61444

View File

@@ -4,6 +4,7 @@ import (
"context" "context"
"fmt" "fmt"
"strings" "strings"
"sync"
"sync/atomic" "sync/atomic"
"ai-agent-scaffold-go/internal/model" "ai-agent-scaffold-go/internal/model"
@@ -283,28 +284,67 @@ func (a *ParallelAgent) Stream(ctx context.Context, content model.ChatContent, o
} }
func (a *ParallelAgent) runWithVars(ctx context.Context, content model.ChatContent, vars map[string]string) (string, error) { func (a *ParallelAgent) runWithVars(ctx context.Context, content model.ChatContent, vars map[string]string) (string, error) {
type result struct {
index int
name string
text string
err error
}
results := make([]result, len(a.subAgents))
var wg sync.WaitGroup
for i, sub := range a.subAgents {
wg.Add(1)
go func(idx int, s workflowSubAgent) {
defer wg.Done()
text, err := s.runWithVars(ctx, content, vars)
results[idx] = result{index: idx, name: s.Name(), text: text, err: err}
}(i, sub)
}
wg.Wait()
// 按原始顺序组装结果,遇到错误立即返回
parts := make([]string, 0, len(a.subAgents)) parts := make([]string, 0, len(a.subAgents))
for _, sub := range a.subAgents { for _, r := range results {
text, err := sub.runWithVars(ctx, content, vars) if r.err != nil {
if err != nil { return "", r.err
return "", err
} }
parts = append(parts, fmt.Sprintf("[%s] %s", sub.Name(), text)) parts = append(parts, fmt.Sprintf("[%s] %s", r.name, r.text))
} }
return strings.Join(parts, "\n"), nil return strings.Join(parts, "\n"), nil
} }
func (a *ParallelAgent) streamWithVars(ctx context.Context, content model.ChatContent, out chan<- string, vars map[string]string) error { func (a *ParallelAgent) streamWithVars(ctx context.Context, content model.ChatContent, out chan<- string, vars map[string]string) error {
text, err := a.runWithVars(ctx, content, vars) errCh := make(chan error, len(a.subAgents))
if err != nil { var wg sync.WaitGroup
for _, sub := range a.subAgents {
wg.Add(1)
go func(s workflowSubAgent) {
defer wg.Done()
// 流式输出前缀标识
prefix := fmt.Sprintf("[%s] ", s.Name())
select {
case out <- prefix:
case <-ctx.Done():
errCh <- ctx.Err()
return
}
if err := s.streamWithVars(ctx, content, out, vars); err != nil {
errCh <- err
}
}(sub)
}
wg.Wait()
close(errCh)
// 返回第一个错误(如有)
for err := range errCh {
return err return err
} }
select { return nil
case out <- text:
return nil
case <-ctx.Done():
return ctx.Err()
}
} }
// ============================================================ // ============================================================