fix(service): ParallelAgent 使用 goroutine 并发执行子 Agent
This commit is contained in:
@@ -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()
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// ============================================================
|
// ============================================================
|
||||||
|
|||||||
Reference in New Issue
Block a user