From 1f8ba6144482ccb880a567b1110403ac0e949fea Mon Sep 17 00:00:00 2001 From: hhs <386998068@qq.com> Date: Wed, 10 Jun 2026 15:54:52 +0800 Subject: [PATCH] =?UTF-8?q?fix(service):=20ParallelAgent=20=E4=BD=BF?= =?UTF-8?q?=E7=94=A8=20goroutine=20=E5=B9=B6=E5=8F=91=E6=89=A7=E8=A1=8C?= =?UTF-8?q?=E5=AD=90=20Agent?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- backend/internal/service/agent.go | 66 +++++++++++++++++++++++++------ 1 file changed, 53 insertions(+), 13 deletions(-) diff --git a/backend/internal/service/agent.go b/backend/internal/service/agent.go index 28e4393..14b56ba 100644 --- a/backend/internal/service/agent.go +++ b/backend/internal/service/agent.go @@ -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 } // ============================================================