fix(service): ParallelAgent 使用 goroutine 并发执行子 Agent
This commit is contained in:
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
|
||||
"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) {
|
||||
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))
|
||||
for _, sub := range a.subAgents {
|
||||
text, err := sub.runWithVars(ctx, content, vars)
|
||||
if err != nil {
|
||||
return "", err
|
||||
for _, r := range results {
|
||||
if r.err != nil {
|
||||
return "", r.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
|
||||
}
|
||||
|
||||
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)
|
||||
if err != nil {
|
||||
errCh := make(chan error, len(a.subAgents))
|
||||
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
|
||||
}
|
||||
select {
|
||||
case out <- text:
|
||||
return nil
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ============================================================
|
||||
|
||||
Reference in New Issue
Block a user