diff --git a/internal/multiagent/eino_run_message_accumulator.go b/internal/multiagent/eino_run_message_accumulator.go new file mode 100644 index 00000000..2f4f7fbb --- /dev/null +++ b/internal/multiagent/eino_run_message_accumulator.go @@ -0,0 +1,69 @@ +package multiagent + +import ( + "strings" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +type einoRunMessageAccumulator struct { + baseCount int + msgs []adk.Message +} + +func newEinoRunMessageAccumulator(base []adk.Message) *einoRunMessageAccumulator { + msgs := append([]adk.Message(nil), base...) + return &einoRunMessageAccumulator{ + baseCount: len(msgs), + msgs: msgs, + } +} + +func (a *einoRunMessageAccumulator) Append(msg adk.Message) bool { + if a == nil || msg == nil { + return false + } + a.msgs = append(a.msgs, msg) + return true +} + +func (a *einoRunMessageAccumulator) AppendToolMessage(content, toolCallID string, opts ...schema.ToolMessageOption) bool { + if strings.TrimSpace(toolCallID) == "" { + return false + } + return a.Append(schema.ToolMessage(content, toolCallID, opts...)) +} + +func (a *einoRunMessageAccumulator) AppendAssistantText(content string) bool { + content = strings.TrimSpace(content) + if content == "" { + return false + } + return a.Append(schema.AssistantMessage(content, nil)) +} + +func (a *einoRunMessageAccumulator) AppendAssistantToolCalls(toolCalls []schema.ToolCall) bool { + if len(toolCalls) == 0 { + return false + } + return a.Append(schema.AssistantMessage("", toolCalls)) +} + +func (a *einoRunMessageAccumulator) Messages() []adk.Message { + if a == nil { + return nil + } + return a.msgs +} + +func (a *einoRunMessageAccumulator) BaseCount() int { + if a == nil { + return 0 + } + return a.baseCount +} + +func (a *einoRunMessageAccumulator) HasNewMessages() bool { + return a != nil && len(a.msgs) > a.baseCount +} diff --git a/internal/multiagent/eino_run_progress_tracker.go b/internal/multiagent/eino_run_progress_tracker.go new file mode 100644 index 00000000..ae612ced --- /dev/null +++ b/internal/multiagent/eino_run_progress_tracker.go @@ -0,0 +1,144 @@ +package multiagent + +import ( + "fmt" + "strings" + + "github.com/cloudwego/eino/schema" +) + +type einoRunProgressTracker struct { + orchMode string + orchestratorName string + conversationID string + progress func(eventType, message string, data interface{}) + + streamsMainAssistant func(agent string) bool + einoRoleTag func(agent string) string + + mainRound int + lastAgent string + toolEmitSeen map[string]struct{} + subAgentToolStep map[string]int + mainAgentToolStep map[string]int +} + +func newEinoRunProgressTracker( + orchMode, orchestratorName, conversationID string, + progress func(eventType, message string, data interface{}), + streamsMainAssistant func(agent string) bool, + einoRoleTag func(agent string) string, +) *einoRunProgressTracker { + if streamsMainAssistant == nil { + streamsMainAssistant = func(agent string) bool { + return agent == "" || agent == orchestratorName + } + } + if einoRoleTag == nil { + einoRoleTag = func(agent string) string { + if streamsMainAssistant(agent) { + return "orchestrator" + } + return "sub" + } + } + return &einoRunProgressTracker{ + orchMode: orchMode, + orchestratorName: orchestratorName, + conversationID: conversationID, + progress: progress, + streamsMainAssistant: streamsMainAssistant, + einoRoleTag: einoRoleTag, + toolEmitSeen: make(map[string]struct{}), + subAgentToolStep: make(map[string]int), + mainAgentToolStep: make(map[string]int), + } +} + +func (t *einoRunProgressTracker) ObserveAgent(agentName string) { + if t == nil || strings.TrimSpace(agentName) == "" || t.progress == nil { + return + } + iterEinoAgent := t.orchestratorName + if t.orchMode == "plan_execute" { + if a := strings.TrimSpace(agentName); a != "" { + iterEinoAgent = a + } + } + if t.streamsMainAssistant(agentName) { + mainIterKey := einoMainIterationKey(iterEinoAgent, t.orchestratorName) + if t.mainRound == 0 { + t.mainRound = 1 + t.mainAgentToolStep[mainIterKey] = 1 + t.emitMainIteration(iterEinoAgent, t.mainRound) + } else if t.lastAgent != "" { + needBump := false + if !t.streamsMainAssistant(t.lastAgent) { + needBump = true + } else if t.lastAgent != agentName { + needBump = true + } + if needBump { + t.mainRound++ + t.mainAgentToolStep[mainIterKey] = t.mainRound + t.emitMainIteration(iterEinoAgent, t.mainRound) + } + } + } + if t.lastAgent != agentName { + t.progress("progress", fmt.Sprintf("[Eino] %s", agentName), map[string]interface{}{ + "conversationId": t.conversationID, + "einoAgent": agentName, + "einoRole": t.einoRoleTag(agentName), + "orchestration": t.orchMode, + }) + } + t.lastAgent = agentName +} + +func (t *einoRunProgressTracker) MainIteration(agentName string) int { + if t == nil { + return 0 + } + key := einoMainIterationKey(agentName, t.orchestratorName) + if n := t.mainAgentToolStep[key]; n > 0 { + return n + } + return t.mainRound +} + +func (t *einoRunProgressTracker) EmitToolCalls(msg *schema.Message, agentName string, markPending func(toolCallPendingInfo)) { + if t == nil { + return + } + before := t.MainIteration(agentName) + tryEmitToolCallsOnce( + msg, + agentName, + t.orchestratorName, + t.conversationID, + t.orchMode, + t.progress, + t.toolEmitSeen, + t.subAgentToolStep, + t.mainAgentToolStep, + markPending, + ) + if t.streamsMainAssistant(agentName) { + if after := t.MainIteration(agentName); after > before { + t.mainRound = after + } + } +} + +func (t *einoRunProgressTracker) emitMainIteration(agentName string, iteration int) { + t.progress("iteration", "", map[string]interface{}{ + "iteration": iteration, + "einoScope": "main", + "einoRole": "orchestrator", + "einoAgent": agentName, + "orchestration": t.orchMode, + "conversationId": t.conversationID, + "source": "eino", + }) +} diff --git a/internal/multiagent/eino_run_progress_tracker_test.go b/internal/multiagent/eino_run_progress_tracker_test.go new file mode 100644 index 00000000..a3e9c635 --- /dev/null +++ b/internal/multiagent/eino_run_progress_tracker_test.go @@ -0,0 +1,185 @@ +package multiagent + +import ( + "testing" + + "github.com/cloudwego/eino/schema" +) + +func TestEinoRunProgressTrackerMainToolCallAdvancesResponseIteration(t *testing.T) { + var events []string + var iterations []int + progress := func(eventType, _ string, raw interface{}) { + events = append(events, eventType) + data, _ := raw.(map[string]interface{}) + if eventType == "iteration" { + if n, ok := data["iteration"].(int); ok { + iterations = append(iterations, n) + } + } + } + tracker := newEinoRunProgressTracker( + "eino_single", "main", "conv-1", progress, + func(agent string) bool { return agent == "" || agent == "main" }, + nil, + ) + + tracker.ObserveAgent("main") + if got := tracker.MainIteration("main"); got != 1 { + t.Fatalf("initial main iteration = %d, want 1", got) + } + tracker.EmitToolCalls(&schema.Message{ToolCalls: []schema.ToolCall{{ + ID: "call-1", + Type: "function", + Function: schema.FunctionCall{ + Name: "execute", + Arguments: `{"command":"pwd"}`, + }, + }}}, "main", nil) + if got := tracker.MainIteration("main"); got != 2 { + t.Fatalf("post-tool main iteration = %d, want 2", got) + } + if len(iterations) != 2 || iterations[0] != 1 || iterations[1] != 2 { + t.Fatalf("iteration events = %#v, want [1 2]; events=%#v", iterations, events) + } +} + +func TestEinoRunProgressTrackerMainAgentSwitchAdvancesIteration(t *testing.T) { + var iterations []int + progress := func(eventType, _ string, raw interface{}) { + if eventType != "iteration" { + return + } + data, _ := raw.(map[string]interface{}) + if n, ok := data["iteration"].(int); ok { + iterations = append(iterations, n) + } + } + tracker := newEinoRunProgressTracker( + "supervisor", "lead", "conv-1", progress, + func(agent string) bool { return agent == "" || agent == "lead" }, + nil, + ) + + tracker.ObserveAgent("lead") + tracker.ObserveAgent("sub") + tracker.ObserveAgent("lead") + + if got := tracker.MainIteration("lead"); got != 2 { + t.Fatalf("main iteration after sub->main = %d, want 2", got) + } + if len(iterations) != 2 || iterations[0] != 1 || iterations[1] != 2 { + t.Fatalf("iteration events = %#v, want [1 2]", iterations) + } +} + +func TestEinoRunProgressTrackerDedupesToolCalls(t *testing.T) { + var toolCalls int + progress := func(eventType, _ string, _ interface{}) { + if eventType == "tool_call" { + toolCalls++ + } + } + tracker := newEinoRunProgressTracker("deep", "lead", "conv-1", progress, nil, nil) + msg := &schema.Message{ToolCalls: []schema.ToolCall{{ + ID: "call-1", + Type: "function", + Function: schema.FunctionCall{ + Name: "search", + Arguments: `{"q":"x"}`, + }, + }}} + + tracker.EmitToolCalls(msg, "lead", nil) + tracker.EmitToolCalls(msg, "lead", nil) + + if toolCalls != 1 { + t.Fatalf("tool call events = %d, want 1", toolCalls) + } +} + +func TestEinoRunProgressTrackerHidesModelOutputRecoveryToolCalls(t *testing.T) { + var eventTypes []string + var marked []toolCallPendingInfo + progress := func(eventType, _ string, _ interface{}) { + eventTypes = append(eventTypes, eventType) + } + tracker := newEinoRunProgressTracker("deep", "lead", "conv-1", progress, nil, nil) + msg := &schema.Message{ToolCalls: []schema.ToolCall{{ + ID: "call-recovery", + Type: "function", + Function: schema.FunctionCall{ + Name: "task", + Arguments: `{"_cyberstrike_model_output_recovery":{"reason":"invalid_tool_arguments_json","repair_attempt":1}}`, + }, + }}} + + tracker.EmitToolCalls(msg, "lead", func(info toolCallPendingInfo) { + marked = append(marked, info) + }) + + if containsString(eventTypes, "tool_calls_detected") || containsString(eventTypes, "tool_call") { + t.Fatalf("event types = %#v, want no visible recovery tool call events", eventTypes) + } + if len(marked) != 0 { + t.Fatalf("marked pending = %#v, want none", marked) + } +} + +func TestEinoRunProgressTrackerHidesAnonymousToolCallFragments(t *testing.T) { + var eventTypes []string + var marked []toolCallPendingInfo + progress := func(eventType, _ string, _ interface{}) { + eventTypes = append(eventTypes, eventType) + } + tracker := newEinoRunProgressTracker("eino_single", "lead", "conv-1", progress, nil, nil) + idx := 0 + msg := &schema.Message{ToolCalls: []schema.ToolCall{{ + Type: "function", + Index: &idx, + Function: schema.FunctionCall{ + Arguments: `"`, + }, + }}} + + tracker.EmitToolCalls(msg, "lead", func(info toolCallPendingInfo) { + marked = append(marked, info) + }) + + if containsString(eventTypes, "tool_calls_detected") || containsString(eventTypes, "tool_call") { + t.Fatalf("event types = %#v, want no visible anonymous fragment tool call events", eventTypes) + } + if len(marked) != 0 { + t.Fatalf("marked pending = %#v, want none", marked) + } +} + +func TestEinoRunProgressTrackerKeepsNamedInvalidToolCallsVisible(t *testing.T) { + var toolCalls int + var marked []toolCallPendingInfo + progress := func(eventType, _ string, _ interface{}) { + if eventType == "tool_call" { + toolCalls++ + } + } + tracker := newEinoRunProgressTracker("eino_single", "lead", "conv-1", progress, nil, nil) + msg := &schema.Message{ToolCalls: []schema.ToolCall{{ + ID: "call-bad-args", + Type: "function", + Function: schema.FunctionCall{ + Name: "exec", + Arguments: `command`, + }, + }}} + + tracker.EmitToolCalls(msg, "lead", func(info toolCallPendingInfo) { + marked = append(marked, info) + }) + + if toolCalls != 1 { + t.Fatalf("tool call events = %d, want 1", toolCalls) + } + if len(marked) != 1 || marked[0].ToolName != "exec" { + t.Fatalf("marked pending = %#v, want one exec call", marked) + } +} diff --git a/internal/multiagent/eino_run_recovery_handler.go b/internal/multiagent/eino_run_recovery_handler.go new file mode 100644 index 00000000..367287b2 --- /dev/null +++ b/internal/multiagent/eino_run_recovery_handler.go @@ -0,0 +1,70 @@ +package multiagent + +import ( + "github.com/cloudwego/eino/adk" + "go.uber.org/zap" +) + +type einoRunRecoveryHandlerConfig struct { + ConversationID string + OrchMode string + Args *einoADKRunLoopArgs + BaseMsgs []adk.Message + Progress func(eventType, message string, data interface{}) + Logger *zap.Logger + RunError *einoRunErrorHandler + ContextOverflow *einoContextOverflowRetryHandler + Transient *einoTransientRunRetryHandler +} + +type einoRunRecoveryResult struct { + Handled bool + Restarted bool + RestartMsgs []adk.Message + Fatal error +} + +type einoRunRecoveryHandler struct { + cfg einoRunRecoveryHandlerConfig +} + +func newEinoRunRecoveryHandler(cfg einoRunRecoveryHandlerConfig) *einoRunRecoveryHandler { + if cfg.Args == nil { + cfg.Args = &einoADKRunLoopArgs{} + } + return &einoRunRecoveryHandler{cfg: cfg} +} + +func (h *einoRunRecoveryHandler) Handle(runErr error, accumulated []adk.Message, baseCount int) einoRunRecoveryResult { + if h == nil || runErr == nil { + return einoRunRecoveryResult{} + } + if willRetry, ok := isEinoNativeWillRetry(runErr); ok { + emitEinoNativeModelRetryProgress(h.cfg.ConversationID, h.cfg.OrchMode, willRetry, h.cfg.Progress, h.cfg.Logger, runErr) + return einoRunRecoveryResult{Handled: true} + } + if h.cfg.ContextOverflow != nil { + if overflowRetry := h.cfg.ContextOverflow.Prepare(runErr, accumulated, baseCount); overflowRetry.Handled { + return einoRunRecoveryResult{Handled: true, Restarted: true, RestartMsgs: overflowRetry.RestartMsgs} + } + } + if h.cfg.Transient != nil { + if runRetry := h.cfg.Transient.Prepare(runErr, accumulated, baseCount); runRetry.Handled { + if runRetry.Fatal != nil { + return einoRunRecoveryResult{Handled: true, Fatal: runRetry.Fatal} + } + if !runRetry.Restarted { + return einoRunRecoveryResult{Handled: true} + } + return einoRunRecoveryResult{Handled: true, Restarted: true, RestartMsgs: runRetry.RestartMsgs} + } + } + return einoRunRecoveryResult{Handled: true, Fatal: h.handleFatal(runErr)} +} + +func (h *einoRunRecoveryHandler) handleFatal(runErr error) error { + if h != nil && h.cfg.RunError != nil { + return h.cfg.RunError.Handle(runErr) + } + return runErr +} diff --git a/internal/multiagent/eino_run_recovery_handler_test.go b/internal/multiagent/eino_run_recovery_handler_test.go new file mode 100644 index 00000000..1d121503 --- /dev/null +++ b/internal/multiagent/eino_run_recovery_handler_test.go @@ -0,0 +1,49 @@ +package multiagent + +import ( + "context" + "errors" + "testing" + "time" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +func TestEinoRunRecoveryHandlerRoutesContextOverflowBeforeTransient(t *testing.T) { + baseMsgs := []adk.Message{schema.UserMessage("base")} + overflow := newEinoContextOverflowRetryHandler(einoContextOverflowRetryConfig{ + Context: context.Background(), + Args: &einoADKRunLoopArgs{}, + BaseMsgs: baseMsgs, + }) + transient := newEinoTransientRunRetryHandler(einoTransientRunRetryHandlerConfig{ + Args: &einoADKRunLoopArgs{}, + BaseMsgs: baseMsgs, + Policy: einoTransientRunRetryPolicy{maxAttempts: 1, maxBackoff: time.Nanosecond}, + }) + handler := newEinoRunRecoveryHandler(einoRunRecoveryHandlerConfig{ + ContextOverflow: overflow, + Transient: transient, + BaseMsgs: baseMsgs, + }) + + result := handler.Handle(errors.New("context length exceeded: upstream returned 503"), nil, 0) + if !result.Handled || !result.Restarted || result.Fatal != nil { + t.Fatalf("result = %+v, want context overflow restart", result) + } + second := handler.Handle(errors.New("upstream returned 503"), nil, 0) + if !second.Handled || !second.Restarted || second.Fatal != nil { + t.Fatalf("second result = %+v, want transient restart", second) + } +} + +func TestEinoRunRecoveryHandlerRoutesFatalFallback(t *testing.T) { + handler := newEinoRunRecoveryHandler(einoRunRecoveryHandlerConfig{ + RunError: newEinoRunErrorHandler(einoRunErrorHandlerConfig{}), + }) + result := handler.Handle(errors.New("invalid api key"), nil, 0) + if !result.Handled || result.Restarted || result.Fatal == nil { + t.Fatalf("result = %+v, want fatal fallback", result) + } +} diff --git a/internal/multiagent/eino_run_result_builder.go b/internal/multiagent/eino_run_result_builder.go new file mode 100644 index 00000000..958068e5 --- /dev/null +++ b/internal/multiagent/eino_run_result_builder.go @@ -0,0 +1,220 @@ +package multiagent + +import ( + "encoding/json" + "strings" + + "cyberstrike-ai/internal/agent" + "cyberstrike-ai/internal/einomcp" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +type einoRunResultBuilderConfig struct { + OrchMode string + EmptyHint string + RunMessages *einoRunMessageAccumulator + AssistantOutput *einoAssistantOutputAccumulator + SnapshotMCPIDs func() []string + ModelFacingTrace func() []adk.Message +} + +type einoRunResultBuilder struct { + cfg einoRunResultBuilderConfig +} + +func newEinoRunResultBuilder(cfg einoRunResultBuilderConfig) *einoRunResultBuilder { + return &einoRunResultBuilder{cfg: cfg} +} + +func (b *einoRunResultBuilder) BuildPartial(runErr error) (*RunResult, error) { + if b == nil || b.cfg.RunMessages == nil || !b.cfg.RunMessages.HasNewMessages() { + return nil, runErr + } + return b.build(true), runErr +} + +func (b *einoRunResultBuilder) BuildFinal() *RunResult { + if b == nil { + return &RunResult{} + } + return b.build(false) +} + +func (b *einoRunResultBuilder) build(partial bool) *RunResult { + var runMsgs []adk.Message + if b.cfg.RunMessages != nil { + runMsgs = b.cfg.RunMessages.Messages() + } + var lastAssistant string + var lastPlanExecuteExecutor string + if b.cfg.AssistantOutput != nil { + lastAssistant = b.cfg.AssistantOutput.LastAssistant() + lastPlanExecuteExecutor = b.cfg.AssistantOutput.LastPlanExecuteExecutor() + } + var modelFacing []adk.Message + if b.cfg.ModelFacingTrace != nil { + modelFacing = b.cfg.ModelFacingTrace() + } + var ids []string + if b.cfg.SnapshotMCPIDs != nil { + ids = b.cfg.SnapshotMCPIDs() + } + return buildEinoRunResultFromAccumulated( + b.cfg.OrchMode, + runMsgs, + modelFacing, + lastAssistant, + lastPlanExecuteExecutor, + b.cfg.EmptyHint, + ids, + partial, + ) +} + +func einoPartialRunLastOutputHint() string { + return "[执行未正常结束(用户停止、超时或异常)。续跑时请基于上文已产生的工具与结果继续,勿重复已完成步骤。]\n" + + "[Run ended abnormally; continue from the trace above without repeating completed steps.]" +} + +func buildEinoRunResultFromAccumulated( + orchMode string, + runAccumulatedMsgs []adk.Message, + persistMsgs []adk.Message, + lastAssistant string, + lastPlanExecuteExecutor string, + emptyHint string, + mcpIDs []string, + partial bool, +) *RunResult { + traceForJSON := persistMsgs + traceJSON := "" + if len(traceForJSON) > 0 { + traceForJSON = markModelFacingTraceForPersistence(traceForJSON) + if histJSON, err := json.Marshal(traceForJSON); err == nil { + traceJSON = string(histJSON) + } + } + cleaned := strings.TrimSpace(lastAssistant) + if orchMode == "plan_execute" { + if e := strings.TrimSpace(lastPlanExecuteExecutor); e != "" { + cleaned = e + } else { + cleaned = UnwrapPlanExecuteUserText(cleaned) + } + } + if cleaned == "" { + if fb := strings.TrimSpace(einoExtractFallbackAssistantFromMsgs(runAccumulatedMsgs)); fb != "" { + cleaned = fb + } + } + cleaned = dedupeRepeatedParagraphs(cleaned, 80) + cleaned = dedupeParagraphsByLineFingerprint(cleaned, 100) + const maxResponseRunes = 100000 + if rs := []rune(cleaned); len(rs) > maxResponseRunes { + cleaned = string(rs[:maxResponseRunes]) + "\n\n... (response truncated / 响应已截断)" + } + lastOut := cleaned + resp := cleaned + if partial && cleaned == "" { + lastOut = einoPartialRunLastOutputHint() + resp = emptyHint + } + out := &RunResult{ + Response: resp, + MCPExecutionIDs: mcpIDs, + LastAgentTraceInput: traceJSON, + LastAgentTraceOutput: lastOut, + } + if !partial && out.Response == "" { + out.Response = emptyHint + out.LastAgentTraceOutput = out.Response + } + return out +} + +func markModelFacingTraceForPersistence(msgs []adk.Message) []adk.Message { + out := cloneADKMessagesForTrace(msgs) + if len(out) == 0 || out[0] == nil { + return out + } + if out[0].Extra == nil { + out[0].Extra = make(map[string]any, 1) + } + out[0].Extra[agent.ModelFacingTraceVersionKey] = 1 + return out +} + +// einoExtractFallbackAssistantFromMsgs 在「主通道未产出助手正文」时,从 Eino ADK 轨迹中回填用户可见回复。 +// 典型场景:监督者仅调用 exit(final_result 落在 Tool 消息中),或工具结果已写入历史但 lastAssistant 未更新。 +// +// 优先级:最后一次 exit 工具输出 → 最后一条含 exit 的助手 tool_calls 参数中的 final_result。 +func einoExtractFallbackAssistantFromMsgs(msgs []adk.Message) string { + for i := len(msgs) - 1; i >= 0; i-- { + m := msgs[i] + if m == nil || m.Role != schema.Tool { + continue + } + if !strings.EqualFold(strings.TrimSpace(m.ToolName), adk.ToolInfoExit.Name) { + continue + } + content := strings.TrimSpace(m.Content) + if content == "" || strings.HasPrefix(content, einomcp.ToolErrorPrefix) { + continue + } + return content + } + for i := len(msgs) - 1; i >= 0; i-- { + m := msgs[i] + if m == nil || m.Role != schema.Assistant { + continue + } + if s := einoExtractExitFinalFromAssistantToolCalls(m); s != "" { + return s + } + } + return "" +} + +func einoExtractExitFinalFromAssistantToolCalls(msg *schema.Message) string { + if msg == nil || len(msg.ToolCalls) == 0 { + return "" + } + for i := len(msg.ToolCalls) - 1; i >= 0; i-- { + tc := msg.ToolCalls[i] + if !strings.EqualFold(strings.TrimSpace(tc.Function.Name), adk.ToolInfoExit.Name) { + continue + } + if s := einoParseExitFinalResultArguments(tc.Function.Arguments); s != "" { + return s + } + } + return "" +} + +func einoParseExitFinalResultArguments(arguments string) string { + arguments = strings.TrimSpace(arguments) + if arguments == "" { + return "" + } + var wrap struct { + FinalResult json.RawMessage `json:"final_result"` + } + if err := json.Unmarshal([]byte(arguments), &wrap); err != nil || len(wrap.FinalResult) == 0 { + return "" + } + var s string + if err := json.Unmarshal(wrap.FinalResult, &s); err == nil { + return strings.TrimSpace(s) + } + var anyVal interface{} + if err := json.Unmarshal(wrap.FinalResult, &anyVal); err != nil { + return "" + } + b, err := json.Marshal(anyVal) + if err != nil { + return "" + } + return strings.TrimSpace(string(b)) +} diff --git a/internal/multiagent/eino_run_runtime_session.go b/internal/multiagent/eino_run_runtime_session.go new file mode 100644 index 00000000..e582156d --- /dev/null +++ b/internal/multiagent/eino_run_runtime_session.go @@ -0,0 +1,372 @@ +package multiagent + +import ( + "context" + "sync/atomic" + + "github.com/cloudwego/eino/adk" + "go.uber.org/zap" +) + +type einoRunRuntimeSessionConfig struct { + Context context.Context + Args *einoADKRunLoopArgs + Drain *einoRunEventDrain + BaseMessages []adk.Message + EmptyHint string + SnapshotMCPIDs func() []string + EinoRoleTag func(agent string) string +} + +type einoRunRuntimeErrorResult struct { + Restarted bool + Result *RunResult + Err error +} + +type einoRunRuntimeSession struct { + ctx context.Context + args *einoADKRunLoopArgs + orchMode string + conversationID string + progress func(eventType, message string, data interface{}) + logger *zap.Logger + baseMsgs []adk.Message + msgs []adk.Message + drain *einoRunEventDrain + runMessages *einoRunMessageAccumulator + usage *einoRunUsageAccumulator + + iter *adk.AsyncIterator[*adk.AgentEvent] + startFreshIter einoAgentEventIteratorStarter + + unregisterAgentCancel func() + unregisterTurnLoopInterrupt func() + nativeCancelCause atomic.Value + + transientRetry *einoTransientRunRetryHandler + runRecoveryHandler *einoRunRecoveryHandler + resultBuilder *einoRunResultBuilder + streamErrorHandler *einoStreamErrorHandler + completionHandler *einoRunCompletionHandler + cancellationHandler *einoRunCancellationHandler +} + +func newEinoRunRuntimeSession(cfg einoRunRuntimeSessionConfig) *einoRunRuntimeSession { + if cfg.Context == nil { + cfg.Context = context.Background() + } + if cfg.Args == nil { + cfg.Args = &einoADKRunLoopArgs{} + } + if cfg.SnapshotMCPIDs == nil { + cfg.SnapshotMCPIDs = func() []string { return nil } + } + s := &einoRunRuntimeSession{ + ctx: cfg.Context, + args: cfg.Args, + orchMode: cfg.Args.OrchMode, + conversationID: cfg.Args.ConversationID, + progress: cfg.Args.Progress, + logger: cfg.Args.Logger, + baseMsgs: cfg.BaseMessages, + msgs: append([]adk.Message(nil), cfg.BaseMessages...), + drain: cfg.Drain, + } + if s.drain != nil { + s.runMessages = s.drain.RunMessages() + s.usage = s.drain.Usage() + } + if s.runMessages == nil { + s.runMessages = newEinoRunMessageAccumulator(s.msgs) + } + s.initIteratorRuntime() + s.initRecoveryRuntime() + s.initResultRuntime(cfg.EmptyHint, cfg.SnapshotMCPIDs, cfg.EinoRoleTag) + return s +} + +func (s *einoRunRuntimeSession) Iterator() *adk.AsyncIterator[*adk.AgentEvent] { + if s == nil { + return nil + } + return s.iter +} + +func (s *einoRunRuntimeSession) Close() { + if s == nil { + return + } + callAndClearUnregister(&s.unregisterAgentCancel) + callAndClearUnregister(&s.unregisterTurnLoopInterrupt) +} + +func (s *einoRunRuntimeSession) HandleIteratorContextError(err error) (*RunResult, error) { + if s == nil || s.cancellationHandler == nil { + return nil, err + } + return s.cancellationHandler.Handle(err) +} + +func (s *einoRunRuntimeSession) HandleIteratorEnd() (completed bool, result *RunResult, err error) { + if s == nil { + return true, nil, nil + } + if ctxErr := s.ctx.Err(); ctxErr != nil { + result, err = s.HandleIteratorContextError(ctxErr) + return false, result, err + } + if s.completionHandler != nil { + s.completionHandler.Complete() + } + return true, nil, nil +} + +func (s *einoRunRuntimeSession) HandleRunError(runErr error) einoRunRuntimeErrorResult { + if s == nil || runErr == nil { + return einoRunRuntimeErrorResult{} + } + restarted, fatal := s.maybeRestart(runErr) + if fatal != nil { + result, err := s.takePartial(fatal) + return einoRunRuntimeErrorResult{Result: result, Err: err} + } + return einoRunRuntimeErrorResult{Restarted: restarted} +} + +func (s *einoRunRuntimeSession) HandleStreamError(streamErr error, agentName string) einoRunRuntimeErrorResult { + if s == nil || s.streamErrorHandler == nil || streamErr == nil { + return einoRunRuntimeErrorResult{} + } + handled := s.streamErrorHandler.Handle(streamErr, agentName) + return einoRunRuntimeErrorResult{ + Restarted: handled.Restarted, + Result: handled.Result, + Err: handled.Err, + } +} + +func (s *einoRunRuntimeSession) ConfirmRecovery() { + if s != nil && s.transientRetry != nil { + s.transientRetry.ConfirmRecovery() + } +} + +func (s *einoRunRuntimeSession) BuildFinalResult() *RunResult { + if s == nil || s.resultBuilder == nil { + return &RunResult{} + } + s.emitUsageSummary("final") + return s.resultBuilder.BuildFinal() +} + +func (s *einoRunRuntimeSession) takePartial(err error) (*RunResult, error) { + if s == nil || s.resultBuilder == nil { + return nil, err + } + s.emitUsageSummary("partial") + return s.resultBuilder.BuildPartial(err) +} + +func (s *einoRunRuntimeSession) maybeRestart(runErr error) (restarted bool, fatal error) { + if s == nil || s.runRecoveryHandler == nil { + return false, runErr + } + recovery := s.runRecoveryHandler.Handle(runErr, s.runMessages.Messages(), s.runMessages.BaseCount()) + if recovery.Fatal != nil { + return false, recovery.Fatal + } + if !recovery.Restarted { + return false, nil + } + s.msgs = recovery.RestartMsgs + s.iter = s.startFreshIter(s.msgs) + return true, nil +} + +func (s *einoRunRuntimeSession) initIteratorRuntime() { + if s == nil || s.args == nil { + return + } + runnerCfg := adk.RunnerConfig{ + Agent: s.args.DA, + // 启用 ADK 流式事件:plan_execute 也需要输出 reasoning/response 流, + // 与 deep/supervisor/eino_single 的前端体验保持一致。 + EnableStreaming: true, + } + var cpStore *fileCheckPointStore + var checkPointID string + if checkpoint := newEinoCheckpointRuntime(s.args.CheckpointDir, s.conversationID, s.orchMode, s.logger); checkpoint != nil { + cpStore = checkpoint.Store + checkPointID = checkpoint.CheckPointID + runnerCfg.CheckPointStore = checkpoint.Store + } + runner := adk.NewRunner(s.ctx, runnerCfg) + runtimeCancelRegistrar := agentRuntimeCancelRegistrarFromContext(s.ctx) + turnLoopInterruptRegistrar := agentTurnLoopInterruptRegistrarFromContext(s.ctx) + runnerStarter := newEinoRunnerIteratorStarter(einoRunnerIteratorStarterConfig{ + Context: s.ctx, + ConversationID: s.conversationID, + OrchMode: s.orchMode, + Logger: s.logger, + Runner: runner, + CheckPointID: checkPointID, + NativeCancelCause: &s.nativeCancelCause, + UnregisterAgentCancel: &s.unregisterAgentCancel, + RuntimeCancelRegistrar: runtimeCancelRegistrar, + }) + turnLoopStarter := newEinoTurnLoopIteratorStarter(einoTurnLoopIteratorStarterConfig{ + Context: s.ctx, + Agent: s.args.DA, + ConversationID: s.conversationID, + OrchMode: s.orchMode, + Progress: s.progress, + Logger: s.logger, + Store: cpStore, + CheckPointID: checkPointID, + InterruptTimeout: s.args.TurnLoopInterruptTimeout, + NativeCancelCause: &s.nativeCancelCause, + UnregisterAgentCancel: &s.unregisterAgentCancel, + UnregisterTurnLoopInterrupt: &s.unregisterTurnLoopInterrupt, + RuntimeCancelRegistrar: runtimeCancelRegistrar, + TurnLoopInterruptRegistrar: turnLoopInterruptRegistrar, + }) + useTurnLoop := turnLoopInterruptRegistrar != nil + s.startFreshIter = runnerStarter.Start + if useTurnLoop { + s.startFreshIter = turnLoopStarter.Start + } + if !useTurnLoop && cpStore != nil && checkPointID != "" { + s.iter = newEinoCheckpointResumeHandler(einoCheckpointResumeHandlerConfig{ + Context: s.ctx, + ConversationID: s.conversationID, + OrchMode: s.orchMode, + Progress: s.progress, + Logger: s.logger, + Store: cpStore, + CheckPointID: checkPointID, + Resume: runnerStarter.Resume, + }).TryResume() + } + s.iter = newEinoInitialIteratorStartHandler(einoInitialIteratorStartHandlerConfig{ + ConversationID: s.conversationID, + OrchMode: s.orchMode, + Progress: s.progress, + UseTurnLoop: useTurnLoop, + StartRunner: runnerStarter.Start, + StartTurnLoop: turnLoopStarter.Start, + }).StartIfNeeded(s.iter, s.msgs) + + pending := s.pending() + s.completionHandler = newEinoRunCompletionHandler(einoRunCompletionHandlerConfig{ + ConversationID: s.conversationID, + OrchMode: s.orchMode, + Progress: s.progress, + Logger: s.logger, + Pending: pending, + Checkpoint: cpStore, + CheckpointID: checkPointID, + }) + s.cancellationHandler = newEinoRunCancellationHandler(einoRunCancellationHandlerConfig{ + Context: s.ctx, + ConversationID: s.conversationID, + Progress: s.progress, + Pending: pending, + TakePartial: s.takePartial, + }) +} + +func (s *einoRunRuntimeSession) initRecoveryRuntime() { + if s == nil || s.args == nil { + return + } + pending := s.pending() + contextOverflowRetry := newEinoContextOverflowRetryHandler(einoContextOverflowRetryConfig{ + Context: s.ctx, + ConversationID: s.conversationID, + OrchMode: s.orchMode, + Args: s.args, + BaseMsgs: s.baseMsgs, + Progress: s.progress, + Logger: s.logger, + }) + s.transientRetry = newEinoTransientRunRetryHandler(einoTransientRunRetryHandlerConfig{ + Context: s.ctx, + ConversationID: s.conversationID, + OrchMode: s.orchMode, + Args: s.args, + BaseMsgs: s.baseMsgs, + Progress: s.progress, + Logger: s.logger, + Pending: pending, + }) + runErrorHandler := newEinoRunErrorHandler(einoRunErrorHandlerConfig{ + ConversationID: s.conversationID, + OrchMode: s.orchMode, + Progress: s.progress, + Pending: pending, + NativeCancelFallback: s.nativeCancelCauseOrCanceled, + }) + s.runRecoveryHandler = newEinoRunRecoveryHandler(einoRunRecoveryHandlerConfig{ + ConversationID: s.conversationID, + OrchMode: s.orchMode, + Args: s.args, + BaseMsgs: s.baseMsgs, + Progress: s.progress, + Logger: s.logger, + RunError: runErrorHandler, + ContextOverflow: contextOverflowRetry, + Transient: s.transientRetry, + }) +} + +func (s *einoRunRuntimeSession) initResultRuntime(emptyHint string, snapshotMCPIDs func() []string, einoRoleTag func(agent string) string) { + if s == nil { + return + } + var assistantOutput *einoAssistantOutputAccumulator + if s.drain != nil { + assistantOutput = s.drain.AssistantOutput() + } + s.resultBuilder = newEinoRunResultBuilder(einoRunResultBuilderConfig{ + OrchMode: s.orchMode, + EmptyHint: emptyHint, + RunMessages: s.runMessages, + AssistantOutput: assistantOutput, + SnapshotMCPIDs: snapshotMCPIDs, + ModelFacingTrace: func() []adk.Message { return modelFacingTraceSnapshot(s.args) }, + }) + s.streamErrorHandler = newEinoStreamErrorHandler( + s.ctx, + s.conversationID, + s.progress, + einoRoleTag, + s.maybeRestart, + s.takePartial, + ) +} + +func (s *einoRunRuntimeSession) pending() *einoPendingToolCalls { + if s == nil || s.drain == nil { + return nil + } + return s.drain.PendingToolCalls() +} + +func (s *einoRunRuntimeSession) nativeCancelCauseOrCanceled() error { + if s != nil { + if v := s.nativeCancelCause.Load(); v != nil { + if err, ok := v.(error); ok && err != nil { + return err + } + } + } + return context.Canceled +} + +func (s *einoRunRuntimeSession) emitUsageSummary(reason string) bool { + if s == nil || s.usage == nil { + return false + } + return s.usage.EmitOnce(s.conversationID, s.orchMode, reason, s.progress, s.logger) +} diff --git a/internal/multiagent/eino_run_runtime_session_test.go b/internal/multiagent/eino_run_runtime_session_test.go new file mode 100644 index 00000000..9e14c573 --- /dev/null +++ b/internal/multiagent/eino_run_runtime_session_test.go @@ -0,0 +1,196 @@ +package multiagent + +import ( + "context" + "errors" + "testing" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +type fakeRuntimeSessionAgent struct { + runMessages []adk.Message + runOpts int +} + +func (a *fakeRuntimeSessionAgent) Name(context.Context) string { + return "lead" +} + +func (a *fakeRuntimeSessionAgent) Description(context.Context) string { + return "fake runtime session agent" +} + +func (a *fakeRuntimeSessionAgent) Run(_ context.Context, input *adk.AgentInput, opts ...adk.AgentRunOption) *adk.AsyncIterator[*adk.AgentEvent] { + if input != nil { + a.runMessages = input.Messages + } + a.runOpts = len(opts) + iter, gen := adk.NewAsyncIteratorPair[*adk.AgentEvent]() + gen.Close() + return iter +} + +func TestEinoRunRuntimeSessionStartsRunner(t *testing.T) { + agent := &fakeRuntimeSessionAgent{} + drain := newEinoRunEventDrain(einoRunEventDrainConfig{ + ConversationID: "conv-1", + OrchMode: "deep", + OrchestratorName: "lead", + BaseMessages: []adk.Message{schema.UserMessage("base")}, + }) + session := newEinoRunRuntimeSession(einoRunRuntimeSessionConfig{ + Context: context.Background(), + Args: &einoADKRunLoopArgs{ + ConversationID: "conv-1", + OrchMode: "deep", + OrchestratorName: "lead", + DA: agent, + }, + Drain: drain, + BaseMessages: []adk.Message{schema.UserMessage("base")}, + EmptyHint: "empty", + }) + defer session.Close() + + if session.Iterator() == nil { + t.Fatal("session should start an iterator") + } + if len(agent.runMessages) != 1 || agent.runMessages[0].Content != "base" { + t.Fatalf("run messages = %#v", agent.runMessages) + } + if agent.runOpts != 1 { + t.Fatalf("run opts = %d, want native cancel option", agent.runOpts) + } +} + +func TestEinoRunRuntimeSessionCompletionFlushesPending(t *testing.T) { + agent := &fakeRuntimeSessionAgent{} + var events []string + drain := newEinoRunEventDrain(einoRunEventDrainConfig{ + ConversationID: "conv-1", + OrchMode: "deep", + OrchestratorName: "lead", + Progress: func(eventType, _ string, _ interface{}) { + events = append(events, eventType) + }, + BaseMessages: []adk.Message{schema.UserMessage("base")}, + }) + session := newEinoRunRuntimeSession(einoRunRuntimeSessionConfig{ + Context: context.Background(), + Args: &einoADKRunLoopArgs{ + ConversationID: "conv-1", + OrchMode: "deep", + OrchestratorName: "lead", + Progress: func(eventType, _ string, _ interface{}) { + events = append(events, eventType) + }, + DA: agent, + }, + Drain: drain, + BaseMessages: []adk.Message{schema.UserMessage("base")}, + EmptyHint: "empty", + }) + defer session.Close() + + drain.PendingToolCalls().Mark(toolCallPendingInfo{ + ToolCallID: "call-1", + ToolName: "execute", + EinoAgent: "lead", + EinoRole: "orchestrator", + }) + completed, result, err := session.HandleIteratorEnd() + + if !completed || result != nil || err != nil { + t.Fatalf("completed=%v result=%#v err=%v", completed, result, err) + } + if !containsString(events, "tool_result") || !containsString(events, "eino_pending_orphaned") { + t.Fatalf("events = %#v, want orphan pending flush", events) + } +} + +func TestEinoRunRuntimeSessionCancellationReturnsPartialError(t *testing.T) { + agent := &fakeRuntimeSessionAgent{} + var events []string + drain := newEinoRunEventDrain(einoRunEventDrainConfig{ + ConversationID: "conv-1", + OrchMode: "deep", + OrchestratorName: "lead", + Progress: func(eventType, _ string, _ interface{}) { + events = append(events, eventType) + }, + BaseMessages: []adk.Message{schema.UserMessage("base")}, + }) + session := newEinoRunRuntimeSession(einoRunRuntimeSessionConfig{ + Context: context.Background(), + Args: &einoADKRunLoopArgs{ + ConversationID: "conv-1", + OrchMode: "deep", + OrchestratorName: "lead", + Progress: func(eventType, _ string, _ interface{}) { + events = append(events, eventType) + }, + DA: agent, + }, + Drain: drain, + BaseMessages: []adk.Message{schema.UserMessage("base")}, + EmptyHint: "empty", + }) + defer session.Close() + + stopErr := errors.New("stop") + result, err := session.HandleIteratorContextError(stopErr) + + if result != nil { + t.Fatalf("result = %#v, want nil without new messages", result) + } + if !errors.Is(err, stopErr) { + t.Fatalf("err = %v, want %v", err, stopErr) + } + if !containsString(events, "error") { + t.Fatalf("events = %#v, want cancellation error event", events) + } +} + +func TestEinoRunRuntimeSessionBuildFinalEmitsUsageSummary(t *testing.T) { + agent := &fakeRuntimeSessionAgent{} + var usageEvent map[string]interface{} + progress := func(eventType, _ string, data interface{}) { + if eventType != "eino_usage_summary" { + return + } + usageEvent, _ = data.(map[string]interface{}) + } + drain := newEinoRunEventDrain(einoRunEventDrainConfig{ + ConversationID: "conv-1", + OrchMode: "deep", + OrchestratorName: "lead", + Progress: progress, + BaseMessages: []adk.Message{schema.UserMessage("base")}, + }) + session := newEinoRunRuntimeSession(einoRunRuntimeSessionConfig{ + Context: context.Background(), + Args: &einoADKRunLoopArgs{ + ConversationID: "conv-1", + OrchMode: "deep", + OrchestratorName: "lead", + Progress: progress, + DA: agent, + }, + Drain: drain, + BaseMessages: []adk.Message{schema.UserMessage("base")}, + EmptyHint: "empty", + }) + defer session.Close() + + drain.Usage().AddUsage(&schema.TokenUsage{PromptTokens: 3, CompletionTokens: 4, TotalTokens: 7}) + _ = session.BuildFinalResult() + + if usageEvent == nil { + t.Fatal("usage summary event was not emitted") + } + if usageEvent["conversationId"] != "conv-1" || usageEvent["orchestration"] != "deep" || usageEvent["reason"] != "final" || usageEvent["totalTokens"] != 7 { + t.Fatalf("usage event = %#v", usageEvent) + } +} diff --git a/internal/multiagent/eino_run_trace.go b/internal/multiagent/eino_run_trace.go new file mode 100644 index 00000000..d94c742f --- /dev/null +++ b/internal/multiagent/eino_run_trace.go @@ -0,0 +1,41 @@ +package multiagent + +import ( + "fmt" + "strings" + + "github.com/google/uuid" +) + +func newEinoRunID() string { + return uuid.New().String() +} + +func withEinoRunIDProgress( + runID string, + progress func(eventType, message string, data interface{}), +) func(eventType, message string, data interface{}) { + runID = strings.TrimSpace(runID) + if progress == nil || runID == "" { + return progress + } + return func(eventType, message string, data interface{}) { + progress(eventType, message, addEinoRunIDToProgressData(runID, data)) + } +} + +func addEinoRunIDToProgressData(runID string, data interface{}) interface{} { + runID = strings.TrimSpace(runID) + if runID == "" { + return data + } + switch v := data.(type) { + case map[string]interface{}: + if existing, ok := v["runId"]; !ok || strings.TrimSpace(fmt.Sprint(existing)) == "" { + v["runId"] = runID + } + return v + default: + return data + } +} diff --git a/internal/multiagent/eino_run_trace_test.go b/internal/multiagent/eino_run_trace_test.go new file mode 100644 index 00000000..b1349e66 --- /dev/null +++ b/internal/multiagent/eino_run_trace_test.go @@ -0,0 +1,39 @@ +package multiagent + +import "testing" + +func TestWithEinoRunIDProgressAddsRunIDToMapData(t *testing.T) { + var gotType, gotMessage string + var gotData interface{} + progress := withEinoRunIDProgress("run-1", func(eventType, message string, data interface{}) { + gotType = eventType + gotMessage = message + gotData = data + }) + + progress("progress", "hello", map[string]interface{}{"source": "eino"}) + + if gotType != "progress" || gotMessage != "hello" { + t.Fatalf("event = (%q, %q)", gotType, gotMessage) + } + m, ok := gotData.(map[string]interface{}) + if !ok { + t.Fatalf("data type = %T", gotData) + } + if m["runId"] != "run-1" || m["source"] != "eino" { + t.Fatalf("data = %#v", m) + } +} + +func TestWithEinoRunIDProgressPreservesExistingRunID(t *testing.T) { + var got map[string]interface{} + progress := withEinoRunIDProgress("outer-run", func(_, _ string, data interface{}) { + got, _ = data.(map[string]interface{}) + }) + + progress("progress", "", map[string]interface{}{"runId": "inner-run"}) + + if got["runId"] != "inner-run" { + t.Fatalf("runId = %q, want inner-run", got["runId"]) + } +} diff --git a/internal/multiagent/eino_run_usage_accumulator.go b/internal/multiagent/eino_run_usage_accumulator.go new file mode 100644 index 00000000..8daa5b53 --- /dev/null +++ b/internal/multiagent/eino_run_usage_accumulator.go @@ -0,0 +1,150 @@ +package multiagent + +import ( + "sync" + + "github.com/cloudwego/eino/schema" + "go.uber.org/zap" +) + +type einoRunUsageSummary struct { + ModelCalls int + PromptTokens int + CompletionTokens int + TotalTokens int + CachedTokens int + ReasoningTokens int +} + +type einoRunUsageAccumulator struct { + mu sync.Mutex + summary einoRunUsageSummary + emitted bool +} + +func newEinoRunUsageAccumulator() *einoRunUsageAccumulator { + return &einoRunUsageAccumulator{} +} + +func (a *einoRunUsageAccumulator) AddMessage(msg *schema.Message) bool { + if msg == nil || msg.ResponseMeta == nil || msg.ResponseMeta.Usage == nil { + return false + } + return a.AddUsage(msg.ResponseMeta.Usage) +} + +func (a *einoRunUsageAccumulator) AddUsage(usage *schema.TokenUsage) bool { + if a == nil || usage == nil || tokenUsageEmpty(usage) { + return false + } + a.mu.Lock() + defer a.mu.Unlock() + a.summary.ModelCalls++ + a.summary.PromptTokens += usage.PromptTokens + a.summary.CompletionTokens += usage.CompletionTokens + a.summary.TotalTokens += usage.TotalTokens + a.summary.CachedTokens += usage.PromptTokenDetails.CachedTokens + a.summary.ReasoningTokens += usage.CompletionTokensDetails.ReasoningTokens + return true +} + +func (a *einoRunUsageAccumulator) Summary() einoRunUsageSummary { + if a == nil { + return einoRunUsageSummary{} + } + a.mu.Lock() + defer a.mu.Unlock() + return a.summary +} + +func (a *einoRunUsageAccumulator) EmitOnce( + conversationID string, + orchestration string, + reason string, + progress func(eventType, message string, data interface{}), + logger *zap.Logger, +) bool { + if a == nil { + return false + } + a.mu.Lock() + if a.emitted || a.summary.ModelCalls == 0 { + a.mu.Unlock() + return false + } + a.emitted = true + s := a.summary + a.mu.Unlock() + + data := map[string]interface{}{ + "conversationId": conversationID, + "source": "eino", + "orchestration": orchestration, + "reason": reason, + "modelCalls": s.ModelCalls, + "promptTokens": s.PromptTokens, + "completionTokens": s.CompletionTokens, + "totalTokens": s.TotalTokens, + "cachedTokens": s.CachedTokens, + "reasoningTokens": s.ReasoningTokens, + } + if progress != nil { + progress("eino_usage_summary", "Eino token usage summary", data) + } + if logger != nil { + logger.Info("eino token usage summary", + zap.String("conversationId", conversationID), + zap.String("orchestration", orchestration), + zap.String("reason", reason), + zap.Int("modelCalls", s.ModelCalls), + zap.Int("promptTokens", s.PromptTokens), + zap.Int("completionTokens", s.CompletionTokens), + zap.Int("totalTokens", s.TotalTokens), + zap.Int("cachedTokens", s.CachedTokens), + zap.Int("reasoningTokens", s.ReasoningTokens), + ) + } + return true +} + +func maxEinoTokenUsage(dst *schema.TokenUsage, src *schema.TokenUsage) *schema.TokenUsage { + if src == nil { + return dst + } + if dst == nil { + return cloneEinoTokenUsage(src) + } + if src.PromptTokens > dst.PromptTokens { + dst.PromptTokens = src.PromptTokens + } + if src.CompletionTokens > dst.CompletionTokens { + dst.CompletionTokens = src.CompletionTokens + } + if src.TotalTokens > dst.TotalTokens { + dst.TotalTokens = src.TotalTokens + } + if src.PromptTokenDetails.CachedTokens > dst.PromptTokenDetails.CachedTokens { + dst.PromptTokenDetails.CachedTokens = src.PromptTokenDetails.CachedTokens + } + if src.CompletionTokensDetails.ReasoningTokens > dst.CompletionTokensDetails.ReasoningTokens { + dst.CompletionTokensDetails.ReasoningTokens = src.CompletionTokensDetails.ReasoningTokens + } + return dst +} + +func cloneEinoTokenUsage(src *schema.TokenUsage) *schema.TokenUsage { + if src == nil { + return nil + } + out := *src + return &out +} + +func tokenUsageEmpty(u *schema.TokenUsage) bool { + return u == nil || + (u.PromptTokens == 0 && + u.CompletionTokens == 0 && + u.TotalTokens == 0 && + u.PromptTokenDetails.CachedTokens == 0 && + u.CompletionTokensDetails.ReasoningTokens == 0) +} diff --git a/internal/multiagent/eino_run_usage_accumulator_test.go b/internal/multiagent/eino_run_usage_accumulator_test.go new file mode 100644 index 00000000..57c23779 --- /dev/null +++ b/internal/multiagent/eino_run_usage_accumulator_test.go @@ -0,0 +1,81 @@ +package multiagent + +import ( + "testing" + + "github.com/cloudwego/eino/schema" +) + +func TestEinoRunUsageAccumulatorSumsModelCalls(t *testing.T) { + acc := newEinoRunUsageAccumulator() + acc.AddUsage(&schema.TokenUsage{ + PromptTokens: 10, + CompletionTokens: 4, + TotalTokens: 14, + PromptTokenDetails: schema.PromptTokenDetails{ + CachedTokens: 3, + }, + CompletionTokensDetails: schema.CompletionTokensDetails{ + ReasoningTokens: 2, + }, + }) + msg := schema.AssistantMessage("ok", nil) + msg.ResponseMeta = &schema.ResponseMeta{Usage: &schema.TokenUsage{ + PromptTokens: 7, + CompletionTokens: 5, + TotalTokens: 12, + CompletionTokensDetails: schema.CompletionTokensDetails{ + ReasoningTokens: 1, + }, + }} + acc.AddMessage(msg) + + got := acc.Summary() + if got.ModelCalls != 2 || got.PromptTokens != 17 || got.CompletionTokens != 9 || got.TotalTokens != 26 || got.CachedTokens != 3 || got.ReasoningTokens != 3 { + t.Fatalf("summary = %#v", got) + } +} + +func TestEinoRunUsageAccumulatorEmitOnce(t *testing.T) { + acc := newEinoRunUsageAccumulator() + acc.AddUsage(&schema.TokenUsage{PromptTokens: 1, CompletionTokens: 2, TotalTokens: 3}) + var events []map[string]interface{} + progress := func(eventType, _ string, data interface{}) { + if eventType != "eino_usage_summary" { + return + } + if m, ok := data.(map[string]interface{}); ok { + events = append(events, m) + } + } + + if !acc.EmitOnce("conv-1", "deep", "final", progress, nil) { + t.Fatal("first emit should return true") + } + if acc.EmitOnce("conv-1", "deep", "partial", progress, nil) { + t.Fatal("second emit should return false") + } + if len(events) != 1 { + t.Fatalf("events = %#v, want one usage summary", events) + } + if events[0]["conversationId"] != "conv-1" || events[0]["orchestration"] != "deep" || events[0]["reason"] != "final" || events[0]["totalTokens"] != 3 { + t.Fatalf("event = %#v", events[0]) + } +} + +func TestMaxEinoTokenUsageUsesLargestStreamChunkValues(t *testing.T) { + var got *schema.TokenUsage + got = maxEinoTokenUsage(got, &schema.TokenUsage{PromptTokens: 10, CompletionTokens: 2, TotalTokens: 12}) + got = maxEinoTokenUsage(got, &schema.TokenUsage{ + PromptTokens: 9, + CompletionTokens: 5, + TotalTokens: 14, + CompletionTokensDetails: schema.CompletionTokensDetails{ + ReasoningTokens: 3, + }, + }) + + if got.PromptTokens != 10 || got.CompletionTokens != 5 || got.TotalTokens != 14 || got.CompletionTokensDetails.ReasoningTokens != 3 { + t.Fatalf("usage = %#v", got) + } +} diff --git a/internal/multiagent/eino_runner_iterator_starter.go b/internal/multiagent/eino_runner_iterator_starter.go new file mode 100644 index 00000000..cb1d9356 --- /dev/null +++ b/internal/multiagent/eino_runner_iterator_starter.go @@ -0,0 +1,94 @@ +package multiagent + +import ( + "context" + "sync/atomic" + + "github.com/cloudwego/eino/adk" + "go.uber.org/zap" +) + +type einoRunnerControl interface { + Run(context.Context, []adk.Message, ...adk.AgentRunOption) *adk.AsyncIterator[*adk.AgentEvent] + Resume(context.Context, string, ...adk.AgentRunOption) (*adk.AsyncIterator[*adk.AgentEvent], error) +} + +type einoRunnerIteratorStarterConfig struct { + Context context.Context + ConversationID string + OrchMode string + Logger *zap.Logger + Runner einoRunnerControl + CheckPointID string + NativeCancelCause *atomic.Value + UnregisterAgentCancel *func() + RuntimeCancelRegistrar AgentRuntimeCancelRegistrar +} + +type einoRunnerIteratorStarter struct { + cfg einoRunnerIteratorStarterConfig +} + +func newEinoRunnerIteratorStarter(cfg einoRunnerIteratorStarterConfig) *einoRunnerIteratorStarter { + return &einoRunnerIteratorStarter{cfg: cfg} +} + +func (s *einoRunnerIteratorStarter) Start(runMsgs []adk.Message) *adk.AsyncIterator[*adk.AgentEvent] { + if s == nil || s.cfg.Runner == nil { + return nil + } + opts := s.newRunOptions() + if s.cfg.CheckPointID != "" { + opts = append(opts, adk.WithCheckPointID(s.cfg.CheckPointID)) + } + return s.cfg.Runner.Run(s.cfg.Context, runMsgs, opts...) +} + +func (s *einoRunnerIteratorStarter) Resume(checkPointID string) (*adk.AsyncIterator[*adk.AgentEvent], error) { + if s == nil || s.cfg.Runner == nil { + return nil, nil + } + return s.cfg.Runner.Resume(s.cfg.Context, checkPointID, s.newRunOptions()...) +} + +func (s *einoRunnerIteratorStarter) newRunOptions() []adk.AgentRunOption { + cancelOpt, cancelFn := adk.WithCancel() + callAndClearUnregister(s.cfg.UnregisterAgentCancel) + if s.cfg.RuntimeCancelRegistrar != nil && s.cfg.UnregisterAgentCancel != nil { + *s.cfg.UnregisterAgentCancel = s.cfg.RuntimeCancelRegistrar(func(cause error) bool { + s.storeNativeCancelCause(cause) + waitErr, submitted, handled := requestEinoNativeAgentCancel(cancelFn, cause) + s.logNativeCancelRequest(cause, waitErr, submitted, handled) + return handled + }) + } + return []adk.AgentRunOption{cancelOpt} +} + +func (s *einoRunnerIteratorStarter) storeNativeCancelCause(cause error) { + if s == nil || s.cfg.NativeCancelCause == nil || cause == nil { + return + } + s.cfg.NativeCancelCause.Store(cause) +} + +func (s *einoRunnerIteratorStarter) logNativeCancelRequest(cause error, waitErr error, submitted bool, handled bool) { + if s == nil || s.cfg.Logger == nil { + return + } + fields := []zap.Field{ + zap.String("conversation_id", s.cfg.ConversationID), + zap.String("orchestration", s.cfg.OrchMode), + zap.Bool("submitted", submitted), + zap.Bool("handled", handled), + } + if cause != nil { + fields = append(fields, zap.Error(cause)) + } + if waitErr != nil { + fields = append(fields, zap.NamedError("cancel_wait_error", waitErr)) + s.cfg.Logger.Debug("eino native cancel requested", fields...) + } else { + s.cfg.Logger.Info("eino native cancel requested", fields...) + } +} diff --git a/internal/multiagent/eino_single_runner.go b/internal/multiagent/eino_single_runner.go index c18b72ed..6a1b494b 100644 --- a/internal/multiagent/eino_single_runner.go +++ b/internal/multiagent/eino_single_runner.go @@ -3,30 +3,25 @@ package multiagent import ( "context" "fmt" - "net" - "net/http" - "strings" "sync" - "time" "cyberstrike-ai/internal/agent" "cyberstrike-ai/internal/config" "cyberstrike-ai/internal/database" "cyberstrike-ai/internal/einomcp" - "cyberstrike-ai/internal/openai" "cyberstrike-ai/internal/project" "cyberstrike-ai/internal/reasoning" - einoopenai "github.com/cloudwego/eino-ext/components/model/openai" "github.com/cloudwego/eino/adk" "github.com/cloudwego/eino/compose" + "github.com/cloudwego/eino/schema" "go.uber.org/zap" ) // einoSingleAgentName 与 ChatModelAgent.Name 一致,供流式事件映射主对话区。 const einoSingleAgentName = "cyberstrike-eino-single" -// RunEinoSingleChatModelAgent 使用 Eino adk.NewChatModelAgent + adk.NewRunner.Run(官方 Quick Start 的 Query 同属 Runner API;此处用历史 + 用户消息切片等价于多轮 Query)。 +// RunEinoSingleChatModelAgent 使用 Eino TypedChatModelAgent[*schema.AgenticMessage] + adk.NewRunner.Run(官方 Quick Start 的 Query 同属 Runner API;此处用历史 + 用户消息切片等价于多轮 Query)。 // 与 RunDeepAgent 共享 runEinoADKAgentLoop 的 SSE 映射与 MCP 桥。 func RunEinoSingleChatModelAgent( ctx context.Context, @@ -52,7 +47,7 @@ func RunEinoSingleChatModelAgent( } runtimeUserMessage := prepareLatestUserMessageForModel(userMessage, appCfg, &ma.EinoMiddleware, conversationID, logger) - einoLoc, einoSkillMW, einoFSTools, skillsRoot, einoErr := prepareEinoSkills(ctx, appCfg.SkillsDir, ma, logger) + einoLoc, einoSkillMW, einoFSTools, skillsRoot, einoErr := prepareEinoAgenticSkills(ctx, appCfg.SkillsDir, ma, logger) if einoErr != nil { return nil, einoErr } @@ -89,58 +84,43 @@ func RunEinoSingleChatModelAgent( return nil, err } - mainToolsForCfg, mainOrchestratorPre, singleToolSearchActive, err := prependEinoMiddlewares(ctx, &ma.EinoMiddleware, einoMWMain, mainTools, einoLoc, skillsRoot, conversationID, projectID, logger) + mainToolsForCfg, mainOrchestratorPre, singleToolSearchActive, err := prependEinoAgenticMiddlewares(ctx, &ma.EinoMiddleware, einoMWMain, mainTools, einoLoc, skillsRoot, conversationID, projectID, logger) if err != nil { return nil, fmt.Errorf("eino single eino 中间件: %w", err) } - httpClient := &http.Client{ - Timeout: 30 * time.Minute, - Transport: &http.Transport{ - DialContext: (&net.Dialer{ - Timeout: 300 * time.Second, - KeepAlive: 300 * time.Second, - }).DialContext, - MaxIdleConns: 100, - MaxIdleConnsPerHost: 10, - IdleConnTimeout: 90 * time.Second, - TLSHandshakeTimeout: 30 * time.Second, - ResponseHeaderTimeout: 60 * time.Minute, - }, - } - httpClient = openai.NewEinoHTTPClient(&appCfg.OpenAI, httpClient) - openai.AttachSummarizationDiagTransport(httpClient, logger) - - maxCompletionTokens := appCfg.OpenAI.MaxCompletionTokensEffective() - baseModelCfg := &einoopenai.ChatModelConfig{ - APIKey: appCfg.OpenAI.APIKey, - BaseURL: strings.TrimSuffix(appCfg.OpenAI.BaseURL, "/"), - Model: appCfg.OpenAI.Model, - HTTPClient: httpClient, - MaxCompletionTokens: &maxCompletionTokens, - } - reasoning.ApplyToEinoChatModelConfig(baseModelCfg, &appCfg.OpenAI, reasoningClient) - - baseMainModel, err := einoopenai.NewChatModel(ctx, baseModelCfg) + baseHTTPClient := newEinoBaseHTTPClient() + agenticModelFactory := newEinoOpenAIAgenticChatModelFactory(baseHTTPClient, reasoningClient, logger) + mainModel, err := agenticModelFactory(ctx, appCfg.OpenAI, einoModelModeNormal) if err != nil { - return nil, fmt.Errorf("eino single 模型: %w", err) + return nil, fmt.Errorf("eino single agentic 模型: %w", err) } - mainModel := newStreamToolCallIndexRepairModel(baseMainModel) - - mainSumMw, err := newEinoSummarizationMiddleware(ctx, mainModel, appCfg, &ma.EinoMiddleware, conversationID, db, projectID, logger) + modelRetryCfg := newEinoAgenticModelRetryConfig(&ma.EinoMiddleware, logger, "eino_single") + modelFailoverCfg, err := newEinoAgenticModelFailoverConfig(ctx, appCfg, &ma.EinoMiddleware, einoModelModeNormal, agenticModelFactory, logger, "eino_single", progress, "eino_single", conversationID) if err != nil { - return nil, fmt.Errorf("eino single summarization: %w", err) + return nil, err + } + logEinoAgenticModelGate( + logger, + "eino_single", + "eino_single", + evaluateEinoAgenticModelGate(agenticModelGateFactory(agenticModelFactory, appCfg.OpenAI, einoModelModeNormal), einoAgenticRuntimeSupportV0914()), + ) + + mainSumMw, err := newEinoAgenticSummarizationMiddleware(ctx, mainModel, appCfg, &ma.EinoMiddleware, conversationID, db, projectID, logger) + if err != nil { + return nil, fmt.Errorf("eino single agentic summarization: %w", err) } modelFacingTrace := newModelFacingTraceHolder() - handlers := make([]adk.ChatModelAgentMiddleware, 0, 8) + handlers := make([]adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage], 0, 8) if len(mainOrchestratorPre) > 0 { handlers = append(handlers, mainOrchestratorPre...) } if einoSkillMW != nil { if einoFSTools && einoLoc != nil { - fsMw, fsErr := subAgentFilesystemMiddleware(ctx, einoLoc, toolInvokeNotify, einoSingleAgentName, einoExecBegin, einoExecAppendPartial, einoExecRegisterCancel, einoExecUnregisterCancel, einoExecFinish, agentToolTimeoutMinutes(appCfg), agentToolWaitTimeoutSeconds(appCfg), agentShellNoOutputTimeoutSeconds(appCfg), nil) + fsMw, fsErr := subAgentAgenticFilesystemMiddleware(ctx, einoLoc, toolInvokeNotify, einoSingleAgentName, einoExecBegin, einoExecAppendPartial, einoExecRegisterCancel, einoExecUnregisterCancel, einoExecFinish, agentToolTimeoutMinutes(appCfg), agentToolWaitTimeoutSeconds(appCfg), agentShellNoOutputTimeoutSeconds(appCfg), nil) if fsErr != nil { return nil, fmt.Errorf("eino single filesystem 中间件: %w", fsErr) } @@ -148,16 +128,16 @@ func RunEinoSingleChatModelAgent( } handlers = append(handlers, einoSkillMW) } - handlers = appendEinoChatModelTailMiddlewares(handlers, einoChatModelTailConfig{ - logger: logger, - phase: "eino_single", - summarization: mainSumMw, - modelName: appCfg.OpenAI.Model, - maxTotalTokens: appCfg.OpenAI.MaxTotalTokens, - toolMaxBytes: toolMaxBytesFromMW(&ma.EinoMiddleware), - conversationID: conversationID, - trace: modelFacingTrace, - middlewareConfig: &ma.EinoMiddleware, + handlers = appendEinoAgenticChatModelTailMiddlewares(handlers, einoChatModelTailConfig{ + logger: logger, + phase: "eino_single", + agenticSummarization: mainSumMw, + modelName: appCfg.OpenAI.Model, + maxTotalTokens: appCfg.OpenAI.MaxTotalTokens, + toolMaxBytes: toolMaxBytesFromMW(&ma.EinoMiddleware), + conversationID: conversationID, + trace: modelFacingTrace, + middlewareConfig: &ma.EinoMiddleware, }) maxIter := agentMaxIterations(appCfg) @@ -189,24 +169,26 @@ func RunEinoSingleChatModelAgent( ) } - chatCfg := &adk.ChatModelAgentConfig{ - Name: einoSingleAgentName, - Description: "Eino ADK ChatModelAgent with MCP tools for authorized security testing.", - Instruction: ins, - GenModelInput: literalInstructionGenModelInput, - Model: mainModel, - ToolsConfig: mainToolsCfg, - MaxIterations: maxIter, - Handlers: handlers, + chatCfg := einoAgenticChatModelAgentConfig{ + Name: einoSingleAgentName, + Description: "Eino ADK ChatModelAgent with MCP tools for authorized security testing.", + Instruction: ins, + GenModelInput: literalAgenticInstructionGenModelInput, + Model: mainModel, + ToolsConfig: mainToolsCfg, + MaxIterations: maxIter, + Handlers: handlers, + ModelRetryConfig: modelRetryCfg, + ModelFailoverConfig: modelFailoverCfg, } outKey, _ := deepExtrasFromConfig(ma) if outKey != "" { chatCfg.OutputKey = outKey } - chatAgent, err := adk.NewChatModelAgent(ctx, chatCfg) + chatAgent, err := newEinoAgenticChatModelAgentAdapter(ctx, chatCfg) if err != nil { - return nil, fmt.Errorf("eino single NewChatModelAgent: %w", err) + return nil, fmt.Errorf("eino single Agentic ChatModelAgent: %w", err) } baseMsgs := historyToMessages(history, appCfg, &ma.EinoMiddleware) @@ -230,8 +212,8 @@ func RunEinoSingleChatModelAgent( StreamsMainAssistant: streamsMainAssistant, EinoRoleTag: einoRoleTag, CheckpointDir: ma.EinoMiddleware.CheckpointDir, - RunRetryMaxAttempts: ma.EinoMiddleware.RunRetryMaxAttempts, - RunRetryMaxBackoffSec: ma.EinoMiddleware.RunRetryMaxBackoffSec, + RunRetryMaxAttempts: RunRetryMaxAttemptsFromConfig(&ma.EinoMiddleware), + RunRetryMaxBackoffSec: int(einoRunRetryMaxBackoffFromConfig(&ma.EinoMiddleware).Seconds()), McpIDsMu: &mcpIDsMu, McpIDs: &mcpIDs, FilesystemMonitorAgent: ag, @@ -244,6 +226,7 @@ func RunEinoSingleChatModelAgent( MaxTotalTokens: appCfg.OpenAI.MaxTotalTokens, ToolMaxBytes: toolMaxBytesFromMW(&ma.EinoMiddleware), ModelName: appCfg.OpenAI.Model, + MiddlewareConfig: &ma.EinoMiddleware, EmptyResponseMessage: "(Eino ADK single-agent session completed but no assistant text was captured. Check process details or logs.) " + "(Eino ADK 单代理会话已完成,但未捕获到助手文本输出。请查看过程详情或日志。)", }, baseMsgs) diff --git a/internal/multiagent/eino_skills.go b/internal/multiagent/eino_skills.go index 724562ba..9cbaec95 100644 --- a/internal/multiagent/eino_skills.go +++ b/internal/multiagent/eino_skills.go @@ -15,19 +15,16 @@ import ( "github.com/cloudwego/eino/adk" "github.com/cloudwego/eino/adk/middlewares/filesystem" "github.com/cloudwego/eino/adk/middlewares/skill" + "github.com/cloudwego/eino/schema" "go.uber.org/zap" ) -// prepareEinoSkills builds Eino official skill backend + middleware, and a shared local disk backend. -// The local backend is also required by reduction, so reduction must not silently disappear merely -// because Skills are disabled or skills_dir is unavailable. -// skillsRoot is the absolute skills directory (empty when skills are not active). -func prepareEinoSkills( +func prepareEinoAgenticSkills( ctx context.Context, skillsDir string, ma *config.MultiAgentConfig, logger *zap.Logger, -) (loc *localbk.Local, skillMW adk.ChatModelAgentMiddleware, fsTools bool, skillsRoot string, err error) { +) (loc *localbk.Local, skillMW adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage], fsTools bool, skillsRoot string, err error) { if ma == nil { return nil, nil, false, "", nil } @@ -49,7 +46,7 @@ func prepareEinoSkills( root := strings.TrimSpace(skillsDir) if root == "" { if logger != nil { - logger.Warn("eino skills: skills_dir empty, skip") + logger.Warn("eino agentic skills: skills_dir empty, skip") } if !needLocalBackend { return nil, nil, false, "", nil @@ -63,7 +60,7 @@ func prepareEinoSkills( } if st, err := os.Stat(abs); err != nil || !st.IsDir() { if logger != nil { - logger.Warn("eino skills: directory missing, skip", zap.String("dir", abs), zap.Error(err)) + logger.Warn("eino agentic skills: directory missing, skip", zap.String("dir", abs), zap.Error(err)) } if !needLocalBackend { return nil, nil, false, "", nil @@ -82,26 +79,23 @@ func prepareEinoSkills( BaseDir: abs, }) if err != nil { - return nil, nil, false, "", fmt.Errorf("eino skill filesystem backend: %w", err) + return nil, nil, false, "", fmt.Errorf("eino agentic skill filesystem backend: %w", err) } - sc := &skill.Config{Backend: skillBE} + sc := &skill.TypedConfig[*schema.AgenticMessage]{Backend: skillBE} if name := strings.TrimSpace(ma.EinoSkills.SkillToolName); name != "" { sc.SkillToolName = &name } - skillMW, err = skill.NewMiddleware(ctx, sc) + skillMW, err = skill.NewTyped[*schema.AgenticMessage](ctx, sc) if err != nil { - return nil, nil, false, "", fmt.Errorf("eino skill middleware: %w", err) + return nil, nil, false, "", fmt.Errorf("eino agentic skill middleware: %w", err) } fsTools = ma.EinoSkills.EinoSkillFilesystemToolsEffective() return loc, skillMW, fsTools, abs, nil } -// subAgentFilesystemMiddleware returns filesystem middleware for a sub-agent when Deep itself -// does not set Backend (fsTools false on orchestrator) but we still want tools on subs — not used; -// when orchestrator has Backend, builtin FS is only on outer agent; subs need explicit FS for parity. -func subAgentFilesystemMiddleware( +func subAgentAgenticFilesystemMiddleware( ctx context.Context, loc *localbk.Local, invokeNotify *einomcp.ToolInvokeNotifyHolder, @@ -115,11 +109,11 @@ func subAgentFilesystemMiddleware( toolWaitTimeoutSeconds int, shellNoOutputTimeoutSec int, outputChunk func(toolName, toolCallID, chunk string), -) (adk.ChatModelAgentMiddleware, error) { +) (adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage], error) { if loc == nil { return nil, nil } - return filesystem.New(ctx, &filesystem.MiddlewareConfig{ + return filesystem.NewTyped[*schema.AgenticMessage](ctx, &filesystem.MiddlewareConfig{ Backend: loc, StreamingShell: &einoStreamingShellWrap{ inner: security.NewEinoStreamingShell(), diff --git a/internal/multiagent/eino_skills_reduction_test.go b/internal/multiagent/eino_skills_reduction_test.go index 7f9fe8b5..10345d0e 100644 --- a/internal/multiagent/eino_skills_reduction_test.go +++ b/internal/multiagent/eino_skills_reduction_test.go @@ -7,21 +7,21 @@ import ( "cyberstrike-ai/internal/config" ) -func TestPrepareEinoSkillsStillCreatesReductionBackendWhenSkillsDisabled(t *testing.T) { +func TestPrepareEinoAgenticSkillsStillCreatesReductionBackendWhenSkillsDisabled(t *testing.T) { ma := &config.MultiAgentConfig{ EinoSkills: config.MultiAgentEinoSkillsConfig{Disable: true}, EinoMiddleware: config.MultiAgentEinoMiddlewareConfig{ ReductionEnable: true, }, } - loc, skillMW, fsTools, skillsRoot, err := prepareEinoSkills(context.Background(), "", ma, nil) + loc, skillMW, fsTools, skillsRoot, err := prepareEinoAgenticSkills(context.Background(), "", ma, nil) if err != nil { t.Fatal(err) } if loc == nil { - t.Fatal("reduction backend must exist even when Skills are disabled") + t.Fatal("agentic reduction backend must exist even when Skills are disabled") } if skillMW != nil || fsTools || skillsRoot != "" { - t.Fatalf("Skills unexpectedly enabled: mw=%v fs=%v root=%q", skillMW, fsTools, skillsRoot) + t.Fatalf("Agentic Skills unexpectedly enabled: mw=%v fs=%v root=%q", skillMW, fsTools, skillsRoot) } } diff --git a/internal/multiagent/eino_stream_error_handler.go b/internal/multiagent/eino_stream_error_handler.go new file mode 100644 index 00000000..d3ca39a2 --- /dev/null +++ b/internal/multiagent/eino_stream_error_handler.go @@ -0,0 +1,81 @@ +package multiagent + +import "context" + +type einoStreamRetryFunc func(error) (restarted bool, fatal error) +type einoPartialResultFunc func(error) (*RunResult, error) + +type einoStreamErrorHandler struct { + ctx context.Context + conversationID string + progress func(eventType, message string, data interface{}) + einoRoleTag func(agent string) string + retry einoStreamRetryFunc + takePartial einoPartialResultFunc +} + +type einoStreamErrorHandleResult struct { + Handled bool + Restarted bool + Result *RunResult + Err error +} + +func newEinoStreamErrorHandler( + ctx context.Context, + conversationID string, + progress func(eventType, message string, data interface{}), + einoRoleTag func(agent string) string, + retry einoStreamRetryFunc, + takePartial einoPartialResultFunc, +) *einoStreamErrorHandler { + if einoRoleTag == nil { + einoRoleTag = func(string) string { return "" } + } + return &einoStreamErrorHandler{ + ctx: ctx, + conversationID: conversationID, + progress: progress, + einoRoleTag: einoRoleTag, + retry: retry, + takePartial: takePartial, + } +} + +func (h *einoStreamErrorHandler) Handle(streamErr error, agentName string) einoStreamErrorHandleResult { + if h == nil || streamErr == nil { + return einoStreamErrorHandleResult{} + } + if isInterruptContinue(h.ctx) { + result, err := h.partial(streamErr) + return einoStreamErrorHandleResult{Handled: true, Result: result, Err: err} + } + if h.progress != nil { + h.progress("eino_stream_error", streamErr.Error(), map[string]interface{}{ + "conversationId": h.conversationID, + "source": "eino", + "einoAgent": agentName, + "einoRole": h.einoRoleTag(agentName), + }) + } + restarted, retErr := h.retryStream(streamErr) + if retErr != nil { + result, err := h.partial(retErr) + return einoStreamErrorHandleResult{Handled: true, Result: result, Err: err} + } + return einoStreamErrorHandleResult{Handled: true, Restarted: restarted} +} + +func (h *einoStreamErrorHandler) retryStream(err error) (bool, error) { + if h == nil || h.retry == nil { + return false, nil + } + return h.retry(err) +} + +func (h *einoStreamErrorHandler) partial(err error) (*RunResult, error) { + if h == nil || h.takePartial == nil { + return nil, err + } + return h.takePartial(err) +} diff --git a/internal/multiagent/eino_stream_error_handler_test.go b/internal/multiagent/eino_stream_error_handler_test.go new file mode 100644 index 00000000..a39ea977 --- /dev/null +++ b/internal/multiagent/eino_stream_error_handler_test.go @@ -0,0 +1,111 @@ +package multiagent + +import ( + "context" + "errors" + "testing" +) + +func TestEinoStreamErrorHandlerEmitsProgressAndRestarts(t *testing.T) { + streamErr := errors.New("stream broken") + var progressEvents []map[string]interface{} + handler := newEinoStreamErrorHandler( + context.Background(), + "conv-1", + func(eventType, _ string, data interface{}) { + if eventType != "eino_stream_error" { + return + } + m, _ := data.(map[string]interface{}) + progressEvents = append(progressEvents, m) + }, + func(agent string) string { + if agent == "worker" { + return "sub" + } + return "orchestrator" + }, + func(err error) (bool, error) { + if !errors.Is(err, streamErr) { + t.Fatalf("retry err = %v", err) + } + return true, nil + }, + nil, + ) + + got := handler.Handle(streamErr, "worker") + if !got.Handled || !got.Restarted || got.Result != nil || got.Err != nil { + t.Fatalf("result = %+v", got) + } + if len(progressEvents) != 1 { + t.Fatalf("progress events = %#v", progressEvents) + } + if progressEvents[0]["conversationId"] != "conv-1" || progressEvents[0]["einoAgent"] != "worker" || progressEvents[0]["einoRole"] != "sub" { + t.Fatalf("progress data = %#v", progressEvents[0]) + } +} + +func TestEinoStreamErrorHandlerRetryFatalUsesPartial(t *testing.T) { + streamErr := errors.New("stream broken") + fatalErr := errors.New("retry exhausted") + wantResult := &RunResult{Response: "partial"} + handler := newEinoStreamErrorHandler( + context.Background(), + "conv-1", + nil, + nil, + func(error) (bool, error) { return false, fatalErr }, + func(err error) (*RunResult, error) { + if !errors.Is(err, fatalErr) { + t.Fatalf("partial err = %v", err) + } + return wantResult, err + }, + ) + + got := handler.Handle(streamErr, "lead") + if !got.Handled || got.Restarted || got.Result != wantResult || !errors.Is(got.Err, fatalErr) { + t.Fatalf("result = %+v", got) + } +} + +func TestEinoStreamErrorHandlerInterruptContinueUsesPartialWithoutProgress(t *testing.T) { + base := context.Background() + ctx, cancel := context.WithCancelCause(base) + cancel(ErrInterruptContinue) + streamErr := errors.New("context canceled while streaming") + var progressCalled bool + var retryCalled bool + handler := newEinoStreamErrorHandler( + ctx, + "conv-1", + func(string, string, interface{}) { progressCalled = true }, + nil, + func(error) (bool, error) { + retryCalled = true + return false, nil + }, + func(err error) (*RunResult, error) { + if !errors.Is(err, streamErr) { + t.Fatalf("partial err = %v", err) + } + return nil, err + }, + ) + + got := handler.Handle(streamErr, "lead") + if !got.Handled || got.Result != nil || !errors.Is(got.Err, streamErr) { + t.Fatalf("result = %+v", got) + } + if progressCalled || retryCalled { + t.Fatalf("progressCalled=%v retryCalled=%v, want both false", progressCalled, retryCalled) + } +} + +func TestEinoStreamErrorHandlerNilError(t *testing.T) { + got := newEinoStreamErrorHandler(context.Background(), "conv", nil, nil, nil, nil).Handle(nil, "lead") + if got.Handled || got.Restarted || got.Result != nil || got.Err != nil { + t.Fatalf("nil error result = %+v", got) + } +} diff --git a/internal/multiagent/eino_stream_tool_call_completion_handler.go b/internal/multiagent/eino_stream_tool_call_completion_handler.go new file mode 100644 index 00000000..0bd4e5fe --- /dev/null +++ b/internal/multiagent/eino_stream_tool_call_completion_handler.go @@ -0,0 +1,49 @@ +package multiagent + +import "github.com/cloudwego/eino/schema" + +type einoStreamToolCallCompletionHandlerConfig struct { + ConversationID string + OrchMode string + Progress func(eventType, message string, data interface{}) + RunProgress *einoRunProgressTracker + RunMessages *einoRunMessageAccumulator + MarkPending func(toolCallPendingInfo) +} + +type einoStreamToolCallCompletionHandler struct { + conversationID string + orchMode string + progress func(eventType, message string, data interface{}) + runProgress *einoRunProgressTracker + runMessages *einoRunMessageAccumulator + markPending func(toolCallPendingInfo) +} + +func newEinoStreamToolCallCompletionHandler(cfg einoStreamToolCallCompletionHandlerConfig) *einoStreamToolCallCompletionHandler { + return &einoStreamToolCallCompletionHandler{ + conversationID: cfg.ConversationID, + orchMode: cfg.OrchMode, + progress: cfg.Progress, + runProgress: cfg.RunProgress, + runMessages: cfg.RunMessages, + markPending: cfg.MarkPending, + } +} + +func (h *einoStreamToolCallCompletionHandler) Complete(fragments []schema.ToolCall, agentName string) *schema.Message { + if h == nil { + return nil + } + var lastToolChunk *schema.Message + if merged := mergeStreamingToolCallFragments(fragments); len(merged) > 0 { + lastToolChunk = mergeMessageToolCalls(&schema.Message{ToolCalls: merged}) + } + if h.runProgress != nil { + h.runProgress.EmitToolCalls(lastToolChunk, agentName, h.markPending) + } + if lastToolChunk != nil && len(lastToolChunk.ToolCalls) > 0 && h.runMessages != nil { + h.runMessages.AppendAssistantToolCalls(lastToolChunk.ToolCalls) + } + return lastToolChunk +} diff --git a/internal/multiagent/eino_stream_tool_call_completion_handler_test.go b/internal/multiagent/eino_stream_tool_call_completion_handler_test.go new file mode 100644 index 00000000..072acecf --- /dev/null +++ b/internal/multiagent/eino_stream_tool_call_completion_handler_test.go @@ -0,0 +1,152 @@ +package multiagent + +import ( + "strings" + "testing" + + "github.com/cloudwego/eino/schema" +) + +func TestEinoStreamToolCallCompletionHandlerMergesEmitsAndPersistsToolCalls(t *testing.T) { + idx := 0 + var eventTypes []string + var marked []toolCallPendingInfo + progress := func(eventType, _ string, _ interface{}) { + eventTypes = append(eventTypes, eventType) + } + runMessages := newEinoRunMessageAccumulator(nil) + runProgress := newEinoRunProgressTracker( + "deep", "lead", "conv-1", progress, + func(agent string) bool { return agent == "lead" }, + nil, + ) + handler := newEinoStreamToolCallCompletionHandler(einoStreamToolCallCompletionHandlerConfig{ + ConversationID: "conv-1", + OrchMode: "deep", + Progress: progress, + RunProgress: runProgress, + RunMessages: runMessages, + MarkPending: func(info toolCallPendingInfo) { + marked = append(marked, info) + }, + }) + + chunk := handler.Complete([]schema.ToolCall{ + { + ID: "call-1", + Type: "function", + Index: &idx, + Function: schema.FunctionCall{ + Name: "execute", + Arguments: `{"command":`, + }, + }, + { + Index: &idx, + Function: schema.FunctionCall{ + Arguments: `"pwd"}`, + }, + }, + }, "lead") + + if chunk == nil || len(chunk.ToolCalls) != 1 { + t.Fatalf("merged chunk = %#v, want one tool call", chunk) + } + if got := chunk.ToolCalls[0].Function.Arguments; got != `{"command":"pwd"}` { + t.Fatalf("arguments = %q", got) + } + msgs := runMessages.Messages() + if len(msgs) != 1 || len(msgs[0].ToolCalls) != 1 { + t.Fatalf("run messages = %#v, want persisted assistant tool call", msgs) + } + if len(marked) != 1 || marked[0].ToolCallID != "call-1" || marked[0].ToolName != "execute" { + t.Fatalf("marked pending = %#v", marked) + } + if !containsString(eventTypes, "tool_call") { + t.Fatalf("event types = %#v, want tool_call", eventTypes) + } +} + +func TestEinoStreamToolCallCompletionHandlerPreservesStreamingToolArgumentsForToolLayerRecovery(t *testing.T) { + idx := 0 + var eventTypes []string + var marked []toolCallPendingInfo + progress := func(eventType, _ string, _ interface{}) { + eventTypes = append(eventTypes, eventType) + } + runMessages := newEinoRunMessageAccumulator(nil) + runProgress := newEinoRunProgressTracker( + "deep", "lead", "conv-1", progress, + func(agent string) bool { return agent == "lead" }, + nil, + ) + handler := newEinoStreamToolCallCompletionHandler(einoStreamToolCallCompletionHandlerConfig{ + ConversationID: "conv-1", + OrchMode: "deep", + Progress: progress, + RunProgress: runProgress, + RunMessages: runMessages, + MarkPending: func(info toolCallPendingInfo) { + marked = append(marked, info) + }, + }) + + chunk := handler.Complete([]schema.ToolCall{ + { + ID: "call-stream-unsafe", + Type: "function", + Index: &idx, + Function: schema.FunctionCall{ + Name: "execute", + Arguments: `{"command":"`, + }, + }, + { + Index: &idx, + Function: schema.FunctionCall{ + Arguments: strings.Repeat("x", 256) + `"}`, + }, + }, + }, "lead") + + if chunk == nil || len(chunk.ToolCalls) != 1 { + t.Fatalf("chunk = %#v, want one tool call", chunk) + } + args := chunk.ToolCalls[0].Function.Arguments + if !strings.Contains(args, strings.Repeat("x", 32)) { + t.Fatalf("streaming arguments were unexpectedly rewritten: %q", args) + } + msgs := runMessages.Messages() + if len(msgs) != 1 || len(msgs[0].ToolCalls) != 1 { + t.Fatalf("run messages = %#v, want assistant tool call", msgs) + } + if got := msgs[0].ToolCalls[0].Function.Arguments; got != args { + t.Fatalf("persisted tool call arguments = %q, want %q", got, args) + } + if len(marked) != 1 || marked[0].ToolCallID != "call-stream-unsafe" || marked[0].ToolName != "execute" { + t.Fatalf("marked pending = %#v", marked) + } + if containsString(eventTypes, "model_output_rejected") || !containsString(eventTypes, "tool_call") { + t.Fatalf("event types = %#v, want real tool_call without model-output recovery", eventTypes) + } +} + +func TestEinoStreamToolCallCompletionHandlerIgnoresEmptyFragments(t *testing.T) { + runMessages := newEinoRunMessageAccumulator(nil) + called := false + handler := newEinoStreamToolCallCompletionHandler(einoStreamToolCallCompletionHandlerConfig{ + RunMessages: runMessages, + Progress: func(string, string, interface{}) { + called = true + }, + }) + if chunk := handler.Complete(nil, "lead"); chunk != nil { + t.Fatalf("chunk = %#v, want nil", chunk) + } + if len(runMessages.Messages()) != 0 { + t.Fatalf("run messages = %#v, want empty", runMessages.Messages()) + } + if called { + t.Fatal("progress should not be called for empty fragments") + } +} diff --git a/internal/multiagent/eino_sub_agent_reply_emitter.go b/internal/multiagent/eino_sub_agent_reply_emitter.go new file mode 100644 index 00000000..8ab3ff3a --- /dev/null +++ b/internal/multiagent/eino_sub_agent_reply_emitter.go @@ -0,0 +1,100 @@ +package multiagent + +import ( + "strings" + + "cyberstrike-ai/internal/openai" +) + +type einoSubAgentReplyEmitter struct { + progress func(eventType, message string, data interface{}) + conversationID string + agentName string + nextStreamID func() string + + streamID string + buf string +} + +func newEinoSubAgentReplyEmitter( + conversationID, agentName string, + progress func(eventType, message string, data interface{}), + nextStreamID func() string, +) *einoSubAgentReplyEmitter { + return &einoSubAgentReplyEmitter{ + progress: progress, + conversationID: conversationID, + agentName: agentName, + nextStreamID: nextStreamID, + } +} + +func (e *einoSubAgentReplyEmitter) EmitDelta(content string) bool { + if e == nil || content == "" { + return false + } + var delta string + e.buf, delta = normalizeStreamingDelta(e.buf, content) + if delta == "" || e.progress == nil { + return false + } + if e.streamID == "" { + if e.nextStreamID != nil { + e.streamID = e.nextStreamID() + } + if e.streamID == "" { + e.streamID = "eino-sub-reply" + } + e.progress("eino_agent_reply_stream_start", "", map[string]interface{}{ + "streamId": e.streamID, + "einoAgent": e.agentName, + "einoRole": "sub", + "conversationId": e.conversationID, + "source": "eino", + }) + } + e.progress("eino_agent_reply_stream_delta", delta, openai.WithSSEAccumulated(map[string]interface{}{ + "streamId": e.streamID, + "conversationId": e.conversationID, + }, e.buf)) + return true +} + +func (e *einoSubAgentReplyEmitter) Finish() string { + if e == nil { + return "" + } + body := strings.TrimSpace(e.buf) + if body == "" || e.progress == nil { + return body + } + if e.streamID != "" { + e.progress("eino_agent_reply_stream_end", body, map[string]interface{}{ + "streamId": e.streamID, + "einoAgent": e.agentName, + "einoRole": "sub", + "conversationId": e.conversationID, + "source": "eino", + }) + } else { + e.EmitComplete(body) + } + return body +} + +func (e *einoSubAgentReplyEmitter) EmitComplete(body string) bool { + if e == nil || e.progress == nil { + return false + } + body = strings.TrimSpace(body) + if body == "" { + return false + } + e.progress("eino_agent_reply", body, map[string]interface{}{ + "conversationId": e.conversationID, + "einoAgent": e.agentName, + "einoRole": "sub", + "source": "eino", + }) + return true +} diff --git a/internal/multiagent/eino_sub_agent_reply_emitter_test.go b/internal/multiagent/eino_sub_agent_reply_emitter_test.go new file mode 100644 index 00000000..1ba3f000 --- /dev/null +++ b/internal/multiagent/eino_sub_agent_reply_emitter_test.go @@ -0,0 +1,84 @@ +package multiagent + +import ( + "testing" + + "cyberstrike-ai/internal/openai" +) + +func TestEinoSubAgentReplyEmitterStreamingLifecycle(t *testing.T) { + type progressEvent struct { + eventType string + message string + data map[string]interface{} + } + var events []progressEvent + progress := func(eventType, message string, data interface{}) { + m, _ := data.(map[string]interface{}) + events = append(events, progressEvent{eventType: eventType, message: message, data: m}) + } + emitter := newEinoSubAgentReplyEmitter("conv-1", "worker", progress, func() string { return "stream-1" }) + + if !emitter.EmitDelta("he") { + t.Fatal("first delta should emit") + } + if !emitter.EmitDelta("hello") { + t.Fatal("cumulative chunk should emit tail") + } + if got := emitter.Finish(); got != "hello" { + t.Fatalf("finish body = %q, want hello", got) + } + + if len(events) != 4 { + t.Fatalf("events = %#v, want start + 2 deltas + end", events) + } + if events[0].eventType != "eino_agent_reply_stream_start" { + t.Fatalf("event[0] = %s", events[0].eventType) + } + if events[1].eventType != "eino_agent_reply_stream_delta" || events[1].message != "he" { + t.Fatalf("event[1] = %#v", events[1]) + } + if events[2].eventType != "eino_agent_reply_stream_delta" || events[2].message != "llo" { + t.Fatalf("event[2] = %#v", events[2]) + } + if got := events[2].data[openai.SSEAccumulatedKey]; got != "hello" { + t.Fatalf("accumulated = %#v, want hello", got) + } + if events[3].eventType != "eino_agent_reply_stream_end" || events[3].message != "hello" { + t.Fatalf("event[3] = %#v", events[3]) + } + if got := events[0].data["einoAgent"]; got != "worker" { + t.Fatalf("einoAgent = %#v", got) + } +} + +func TestEinoSubAgentReplyEmitterComplete(t *testing.T) { + var eventType, message string + var data map[string]interface{} + progress := func(et, msg string, raw interface{}) { + eventType = et + message = msg + data, _ = raw.(map[string]interface{}) + } + + ok := newEinoSubAgentReplyEmitter("conv-1", "worker", progress, nil).EmitComplete(" done ") + if !ok { + t.Fatal("complete reply should emit") + } + if eventType != "eino_agent_reply" || message != "done" { + t.Fatalf("event = %s %q", eventType, message) + } + if data["conversationId"] != "conv-1" || data["einoAgent"] != "worker" || data["einoRole"] != "sub" { + t.Fatalf("bad event data: %#v", data) + } +} + +func TestEinoSubAgentReplyEmitterNoProgressStillBuffers(t *testing.T) { + emitter := newEinoSubAgentReplyEmitter("conv", "worker", nil, nil) + if emitter.EmitDelta("hello") { + t.Fatal("nil progress should not emit") + } + if got := emitter.Finish(); got != "hello" { + t.Fatalf("finish body = %q, want hello", got) + } +} diff --git a/internal/multiagent/eino_summarize.go b/internal/multiagent/eino_summarize.go index 308a3577..114a7626 100644 --- a/internal/multiagent/eino_summarize.go +++ b/internal/multiagent/eino_summarize.go @@ -124,11 +124,6 @@ func newEinoSummarizationMiddleware( trigger = 4096 } } - preserveMax := trigger / 3 - if preserveMax < 2048 { - preserveMax = 2048 - } - modelName := strings.TrimSpace(appCfg.OpenAI.Model) if modelName == "" { modelName = "gpt-4o" @@ -238,10 +233,6 @@ func newEinoSummarizationMiddleware( UserInstruction: einoSummarizeUserInstruction, EmitInternalEvents: emitInternalEvents, TranscriptFilePath: transcriptPath, - PreserveUserMessages: &summarization.PreserveUserMessages{ - Enabled: true, - MaxTokens: preserveMax, - }, Retry: &summarization.RetryConfig{ MaxRetries: &retryMax, ShouldRetry: func(_ context.Context, _ adk.Message, err error) bool { @@ -265,9 +256,17 @@ func newEinoSummarizationMiddleware( }, }, Finalize: func(ctx context.Context, originalMessages []adk.Message, summary adk.Message) ([]adk.Message, error) { + compactionMessages := stripOriginalUserIntentLedgerFromMessages(originalMessages) + defaultFinalized, derr := summarization.DefaultFinalize(ctx, compactionMessages, summary) + if derr != nil { + return nil, derr + } + if len(defaultFinalized) == 0 { + return nil, fmt.Errorf("summarization default finalize returned no messages") + } + summary = appendTranscriptPathToSummarizationMessage(defaultFinalized[len(defaultFinalized)-1], transcriptPath) summary = stripAnalysisFromSummarizationMessage(summary) userLedger := buildOriginalUserIntentLedgerMessage(originalMessages, userLedgerMaxRunes, userLedgerEntryMaxRunes) - compactionMessages := stripOriginalUserIntentLedgerFromMessages(originalMessages) out, ferr := summarizeFinalizeWithRecentAssistantToolTrail(ctx, compactionMessages, summary, tokenCounter, recentTrailMax) if ferr != nil { return nil, ferr diff --git a/internal/multiagent/eino_summarize_output.go b/internal/multiagent/eino_summarize_output.go index 079b5097..cc2c0697 100644 --- a/internal/multiagent/eino_summarize_output.go +++ b/internal/multiagent/eino_summarize_output.go @@ -19,6 +19,8 @@ var ( const ( userIntentLedgerStartMarker = "" userIntentLedgerEndMarker = "" + + summarizationTranscriptPathInstructionZh = "如果你需要压缩之前的具体细节(如精确的代码片段、错误消息或你生成的内容),完整的对话记录位于:%s" ) // stripAnalysisFromSummarizationMessage removes the block from a post-processed @@ -62,6 +64,45 @@ func stripAnalysisFromSummarizationText(text string) string { return stripped } +func appendTranscriptPathToSummarizationMessage(msg adk.Message, transcriptPath string) adk.Message { + transcriptPath = strings.TrimSpace(transcriptPath) + if msg == nil || transcriptPath == "" { + return msg + } + section := fmt.Sprintf(summarizationTranscriptPathInstructionZh, transcriptPath) + cloned := *msg + if cloned.Content != "" && !strings.Contains(cloned.Content, transcriptPath) { + cloned.Content = appendSummarizationSection(cloned.Content, section) + } + if len(cloned.UserInputMultiContent) > 0 { + parts := make([]schema.MessageInputPart, len(cloned.UserInputMultiContent)) + copy(parts, cloned.UserInputMultiContent) + for i := range parts { + if parts[i].Type != schema.ChatMessagePartTypeText { + continue + } + if parts[i].Text != "" && !strings.Contains(parts[i].Text, transcriptPath) { + parts[i].Text = appendSummarizationSection(parts[i].Text, section) + } + break + } + cloned.UserInputMultiContent = parts + } + return &cloned +} + +func appendSummarizationSection(text, section string) string { + text = strings.TrimSpace(text) + section = strings.TrimSpace(section) + if text == "" { + return section + } + if section == "" { + return text + } + return text + "\n\n" + section +} + // extractSummarizationSummaryBody returns the inner text of the last block when present. // Used by tests and optional strict compaction paths. func extractSummarizationSummaryBody(text string) (string, bool) { diff --git a/internal/multiagent/eino_tool_result_event_handler.go b/internal/multiagent/eino_tool_result_event_handler.go new file mode 100644 index 00000000..f07f90a4 --- /dev/null +++ b/internal/multiagent/eino_tool_result_event_handler.go @@ -0,0 +1,86 @@ +package multiagent + +import ( + "context" + "strings" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" + "go.uber.org/zap" +) + +type einoToolResultEventHandlerConfig struct { + Context context.Context + Logger *zap.Logger + RunMessages *einoRunMessageAccumulator + Emitter *einoToolResultProgressEmitter + ConfirmRecovery func() +} + +type einoToolResultEventHandler struct { + ctx context.Context + logger *zap.Logger + runMessages *einoRunMessageAccumulator + emitter *einoToolResultProgressEmitter + confirmRecovery func() +} + +func newEinoToolResultEventHandler(cfg einoToolResultEventHandlerConfig) *einoToolResultEventHandler { + if cfg.Context == nil { + cfg.Context = context.Background() + } + return &einoToolResultEventHandler{ + ctx: cfg.Context, + logger: cfg.Logger, + runMessages: cfg.RunMessages, + emitter: cfg.Emitter, + confirmRecovery: cfg.ConfirmRecovery, + } +} + +func (h *einoToolResultEventHandler) HandleStreaming(mv *adk.MessageVariant, agentName string) bool { + if h == nil || mv == nil || !mv.IsStreaming || mv.MessageStream == nil || mv.Role != schema.Tool { + return false + } + toolName := strings.TrimSpace(mv.ToolName) + content, streamToolCallID, streamToolName, recvErr := recvSchemaMessageStream(h.ctx, mv.MessageStream) + if toolName == "" { + toolName = streamToolName + } + isErr := einoToolResultIsError(toolName, content) + content = einoToolResultBody(content) + if streamToolCallID != "" && h.runMessages != nil { + h.runMessages.AppendToolMessage(content, streamToolCallID, schema.WithToolName(toolName)) + } + if h.emitter != nil { + h.emitter.Emit(h.ctx, toolName, content, streamToolCallID, isErr, agentName) + } + if recvErr != nil && h.logger != nil { + h.logger.Warn("eino tool result stream recv error", + zap.Error(recvErr), + zap.String("agent", agentName), + zap.String("tool", toolName)) + } + if recvErr == nil && h.confirmRecovery != nil { + h.confirmRecovery() + } + return true +} + +func (h *einoToolResultEventHandler) HandleMaterialized(mv *adk.MessageVariant, msg adk.Message, agentName string) bool { + if h == nil || mv == nil || msg == nil || (mv.Role != schema.Tool && msg.Role != schema.Tool) { + return false + } + toolName := msg.ToolName + if toolName == "" { + toolName = mv.ToolName + } + content := msg.Content + isErr := einoToolResultIsError(toolName, content) + content = einoToolResultBody(content) + toolCallID := strings.TrimSpace(msg.ToolCallID) + if h.emitter != nil { + h.emitter.Emit(h.ctx, toolName, content, toolCallID, isErr, agentName) + } + return true +} diff --git a/internal/multiagent/eino_tool_result_event_handler_test.go b/internal/multiagent/eino_tool_result_event_handler_test.go new file mode 100644 index 00000000..8532980c --- /dev/null +++ b/internal/multiagent/eino_tool_result_event_handler_test.go @@ -0,0 +1,95 @@ +package multiagent + +import ( + "testing" + + "cyberstrike-ai/internal/einomcp" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +func TestEinoToolResultEventHandlerHandlesStreamingToolResult(t *testing.T) { + var events []map[string]interface{} + runMessages := newEinoRunMessageAccumulator(nil) + recovered := false + emitter := newEinoToolResultProgressEmitter(einoToolResultProgressEmitterConfig{ + ConversationID: "conv-1", + Progress: func(eventType, _ string, data interface{}) { + if eventType != "tool_result" { + return + } + m, _ := data.(map[string]interface{}) + events = append(events, m) + }, + }) + handler := newEinoToolResultEventHandler(einoToolResultEventHandlerConfig{ + RunMessages: runMessages, + Emitter: emitter, + ConfirmRecovery: func() { + recovered = true + }, + }) + stream := schema.StreamReaderFromArray([]*schema.Message{ + {Role: schema.Tool, Content: "hello ", ToolCallID: "call-1"}, + {Role: schema.Tool, Content: "world", ToolCallID: "call-1"}, + }) + mv := &adk.MessageVariant{ + IsStreaming: true, + Role: schema.Tool, + ToolName: "execute", + MessageStream: stream, + } + + if !handler.HandleStreaming(mv, "worker") { + t.Fatal("streaming tool result was not handled") + } + if !recovered { + t.Fatal("expected retry recovery confirmation") + } + msgs := runMessages.Messages() + if len(msgs) != 1 || msgs[0].Role != schema.Tool || msgs[0].Content != "hello world" || msgs[0].ToolCallID != "call-1" { + t.Fatalf("run messages = %#v", msgs) + } + if len(events) != 1 { + t.Fatalf("events = %#v, want one tool_result", events) + } + if events[0]["toolName"] != "execute" || events[0]["toolCallId"] != "call-1" || events[0]["result"] != "hello world" { + t.Fatalf("event data = %#v", events[0]) + } +} + +func TestEinoToolResultEventHandlerHandlesMaterializedToolResult(t *testing.T) { + var event map[string]interface{} + emitter := newEinoToolResultProgressEmitter(einoToolResultProgressEmitterConfig{ + ConversationID: "conv-1", + Progress: func(eventType, _ string, data interface{}) { + if eventType == "tool_result" { + event, _ = data.(map[string]interface{}) + } + }, + }) + handler := newEinoToolResultEventHandler(einoToolResultEventHandlerConfig{Emitter: emitter}) + msg := schema.ToolMessage(einomcp.ToolErrorPrefix+"bad command", "call-2", schema.WithToolName("execute")) + mv := &adk.MessageVariant{Role: schema.Tool} + + if !handler.HandleMaterialized(mv, msg, "worker") { + t.Fatal("materialized tool result was not handled") + } + if event["toolName"] != "execute" || event["toolCallId"] != "call-2" { + t.Fatalf("event identity = %#v", event) + } + if event["result"] != "bad command" || event["isError"] != true || event["success"] != false { + t.Fatalf("event result flags = %#v", event) + } +} + +func TestEinoToolResultEventHandlerIgnoresNonToolOutput(t *testing.T) { + handler := newEinoToolResultEventHandler(einoToolResultEventHandlerConfig{}) + if handler.HandleStreaming(&adk.MessageVariant{IsStreaming: true, Role: schema.Assistant}, "worker") { + t.Fatal("assistant stream should not be handled as tool result") + } + if handler.HandleMaterialized(&adk.MessageVariant{Role: schema.Assistant}, schema.AssistantMessage("hi", nil), "worker") { + t.Fatal("assistant message should not be handled as tool result") + } +} diff --git a/internal/multiagent/eino_tool_result_progress_emitter.go b/internal/multiagent/eino_tool_result_progress_emitter.go new file mode 100644 index 00000000..b425d471 --- /dev/null +++ b/internal/multiagent/eino_tool_result_progress_emitter.go @@ -0,0 +1,154 @@ +package multiagent + +import ( + "context" + "fmt" + "strings" + "sync" + + "cyberstrike-ai/internal/agent" + "cyberstrike-ai/internal/einomcp" + + "github.com/cloudwego/eino/adk" +) + +type einoToolResultProgressEmitter struct { + conversationID string + orchestratorName string + progress func(eventType, message string, data interface{}) + einoRoleTag func(agent string) string + + pending *einoPendingToolCalls + executeStdoutDup *einoExecuteStdoutSuppressor + runMessages *einoRunMessageAccumulator + + filesystemMonitorAgent *agent.Agent + filesystemMonitorRecord einomcp.ExecutionRecorder + mcpExecutionBinder *MCPExecutionBinder + + sent sync.Map +} + +type einoToolResultProgressEmitterConfig struct { + ConversationID string + OrchestratorName string + Progress func(eventType, message string, data interface{}) + EinoRoleTag func(agent string) string + + Pending *einoPendingToolCalls + ExecuteStdoutDup *einoExecuteStdoutSuppressor + RunMessages *einoRunMessageAccumulator + + FilesystemMonitorAgent *agent.Agent + FilesystemMonitorRecord einomcp.ExecutionRecorder + MCPExecutionBinder *MCPExecutionBinder +} + +func newEinoToolResultProgressEmitter(cfg einoToolResultProgressEmitterConfig) *einoToolResultProgressEmitter { + if cfg.EinoRoleTag == nil { + cfg.EinoRoleTag = func(string) string { return "" } + } + return &einoToolResultProgressEmitter{ + conversationID: cfg.ConversationID, + orchestratorName: cfg.OrchestratorName, + progress: cfg.Progress, + einoRoleTag: cfg.EinoRoleTag, + pending: cfg.Pending, + executeStdoutDup: cfg.ExecuteStdoutDup, + runMessages: cfg.RunMessages, + filesystemMonitorAgent: cfg.FilesystemMonitorAgent, + filesystemMonitorRecord: cfg.FilesystemMonitorRecord, + mcpExecutionBinder: cfg.MCPExecutionBinder, + } +} + +func (e *einoToolResultProgressEmitter) Emit(ctx context.Context, toolName, content, toolCallID string, isErr bool, agentName string) bool { + if e == nil { + return false + } + if strings.HasPrefix(strings.TrimSpace(content), modelOutputRejectedResultPrefix) { + return false + } + toolName = strings.TrimSpace(toolName) + if toolName == "" { + toolName = "unknown" + } + preview := content + if len(preview) > 200 { + preview = preview[:200] + "..." + } + backgroundRunning := isErr && isMCPBackgroundWaitResult(content) + displayIsErr := isErr && !backgroundRunning + data := map[string]interface{}{ + "toolName": toolName, + "success": !displayIsErr, + "isError": displayIsErr, + "result": content, + "resultPreview": preview, + "agentFacing": true, + "conversationId": e.conversationID, + "einoAgent": agentName, + "einoRole": e.einoRoleTag(agentName), + "source": "eino", + } + if backgroundRunning { + data["status"] = "background_running" + data["modelFacingIsError"] = isErr + if execID := mcpExecutionIDFromWaitResult(content); execID != "" { + data["executionId"] = execID + } + } + tid := strings.TrimSpace(toolCallID) + if tid == "" { + tid = e.inferToolCallID(agentName) + } + if tid != "" { + if e.pending != nil { + e.pending.RemoveByID(tid) + } + if _, loaded := e.sent.LoadOrStore(tid, struct{}{}); loaded { + return false + } + data["toolCallId"] = tid + toolCallID = tid + } + if e.executeStdoutDup != nil { + e.executeStdoutDup.Record(toolName, content, displayIsErr) + } + recordEinoADKFilesystemToolMonitor(ctx, e.filesystemMonitorAgent, e.filesystemMonitorRecord, e.mcpExecutionBinder, toolName, toolCallID, e.messages(), content, displayIsErr) + if e.filesystemMonitorAgent != nil && e.mcpExecutionBinder != nil { + if execID := e.mcpExecutionBinder.ExecutionID(toolCallID); execID != "" { + e.filesystemMonitorAgent.UpdateMCPExecutionDisplayResult(execID, content) + } + } + if e.progress != nil { + e.progress("tool_result", fmt.Sprintf("工具结果 (%s)", toolName), data) + } + return true +} + +func (e *einoToolResultProgressEmitter) inferToolCallID(agentName string) string { + if e.pending == nil { + return "" + } + if inferred, ok := e.pending.PopNextForAgent(agentName); ok { + return inferred.ToolCallID + } + if inferred, ok := e.pending.PopNextForAgent(e.orchestratorName); ok { + return inferred.ToolCallID + } + if inferred, ok := e.pending.PopNextForAgent(""); ok { + return inferred.ToolCallID + } + if inferred, ok := e.pending.PopAny(); ok { + return inferred.ToolCallID + } + return "" +} + +func (e *einoToolResultProgressEmitter) messages() []adk.Message { + if e == nil || e.runMessages == nil { + return nil + } + return e.runMessages.Messages() +} diff --git a/internal/multiagent/eino_tool_result_progress_emitter_test.go b/internal/multiagent/eino_tool_result_progress_emitter_test.go new file mode 100644 index 00000000..71470d1b --- /dev/null +++ b/internal/multiagent/eino_tool_result_progress_emitter_test.go @@ -0,0 +1,137 @@ +package multiagent + +import "testing" + +func TestEinoToolResultProgressEmitterInfersPendingAndDedupes(t *testing.T) { + var events []map[string]interface{} + progress := func(eventType, _ string, data interface{}) { + if eventType != "tool_result" { + return + } + m, _ := data.(map[string]interface{}) + events = append(events, m) + } + pending := newEinoPendingToolCalls("conv-1", nil) + pending.Mark(toolCallPendingInfo{ + ToolCallID: "call-1", + ToolName: "execute", + EinoAgent: "worker", + EinoRole: "sub", + }) + stdoutDup := newEinoExecuteStdoutSuppressor() + emitter := newEinoToolResultProgressEmitter(einoToolResultProgressEmitterConfig{ + ConversationID: "conv-1", + OrchestratorName: "lead", + Progress: progress, + EinoRoleTag: func(agent string) string { + if agent == "worker" { + return "sub" + } + return "orchestrator" + }, + Pending: pending, + ExecuteStdoutDup: stdoutDup, + }) + + if !emitter.Emit(nil, "execute", "hello", "", false, "worker") { + t.Fatal("first tool result should emit") + } + if !emitter.Emit(nil, "execute", "duplicate without id", "", false, "worker") { + t.Fatal("id-less result should still emit after pending queue is empty") + } + if emitter.Emit(nil, "execute", "duplicate", "call-1", false, "worker") { + t.Fatal("duplicate toolCallId should not emit") + } + + if len(events) != 2 { + t.Fatalf("events = %#v, want two emitted results", events) + } + if events[0]["toolCallId"] != "call-1" || events[0]["einoRole"] != "sub" { + t.Fatalf("first event data = %#v", events[0]) + } + if _, ok := events[1]["toolCallId"]; ok { + t.Fatalf("second event should not invent toolCallId: %#v", events[1]) + } + if got := stdoutDup.Peek(); got != "duplicate without id" { + t.Fatalf("execute stdout suppressor = %q, want last emitted execute stdout", got) + } + if pending.Count() != 0 { + t.Fatalf("pending count = %d, want 0", pending.Count()) + } +} + +func TestEinoToolResultProgressEmitterBackgroundWaitDisplaysRunning(t *testing.T) { + var data map[string]interface{} + progress := func(eventType, _ string, raw interface{}) { + if eventType == "tool_result" { + data, _ = raw.(map[string]interface{}) + } + } + body := `工具已提交到后台执行,但本次等待已到达上限。 + +execution_id: 3eaaa391-050b-4be1-a870-48a855923cb7 +tool: exec +status: running +wait_timeout: 10s` + emitter := newEinoToolResultProgressEmitter(einoToolResultProgressEmitterConfig{ + ConversationID: "conv-1", + Progress: progress, + }) + + if !emitter.Emit(nil, "exec", body, "call-1", true, "lead") { + t.Fatal("background wait result should emit") + } + if data["success"] != true || data["isError"] != false || data["status"] != "background_running" { + t.Fatalf("background display flags = %#v", data) + } + if data["modelFacingIsError"] != true { + t.Fatalf("modelFacingIsError = %#v", data["modelFacingIsError"]) + } + if data["executionId"] != "3eaaa391-050b-4be1-a870-48a855923cb7" { + t.Fatalf("executionId = %#v", data["executionId"]) + } +} + +func TestEinoToolResultProgressEmitterHidesModelOutputRejectedResult(t *testing.T) { + called := false + emitter := newEinoToolResultProgressEmitter(einoToolResultProgressEmitterConfig{ + ConversationID: "conv-1", + Progress: func(eventType, _ string, _ interface{}) { + if eventType == "tool_result" { + called = true + } + }, + }) + + if emitter.Emit(nil, "task", modelOutputRejectedResultPrefix+" Tool call was not executed.", "call-1", true, "lead") { + t.Fatal("model output rejected result should not emit") + } + if called { + t.Fatal("progress should not receive model output rejected tool_result") + } +} + +func TestEinoToolResultProgressEmitterTruncatesPreview(t *testing.T) { + var data map[string]interface{} + progress := func(eventType, _ string, raw interface{}) { + if eventType == "tool_result" { + data, _ = raw.(map[string]interface{}) + } + } + long := "" + for i := 0; i < 205; i++ { + long += "x" + } + emitter := newEinoToolResultProgressEmitter(einoToolResultProgressEmitterConfig{ + ConversationID: "conv-1", + Progress: progress, + }) + emitter.Emit(nil, "", long, "", false, "") + + if data["toolName"] != "unknown" { + t.Fatalf("tool name = %#v", data["toolName"]) + } + if got, _ := data["resultPreview"].(string); len(got) != 203 || got[200:] != "..." { + t.Fatalf("preview = %q len=%d", got, len(got)) + } +} diff --git a/internal/multiagent/eino_transient_retry.go b/internal/multiagent/eino_transient_retry.go index 5046c52f..64a74bac 100644 --- a/internal/multiagent/eino_transient_retry.go +++ b/internal/multiagent/eino_transient_retry.go @@ -34,9 +34,16 @@ func isEinoTransientRunError(err error) bool { if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { return false } + if errors.Is(err, adk.ErrExceedMaxRetries) { + return false + } + if _, ok := isEinoNativeWillRetry(err); ok { + return false + } if isEinoIterationLimitError(err) { return false } + err = unwrapEinoRetryExhausted(err) var apiErr *einoopenai.APIError if errors.As(err, &apiErr) && apiErr.HTTPStatusCode > 0 { return isRetryableHTTPStatus(apiErr.HTTPStatusCode) @@ -198,13 +205,9 @@ func einoTransientRunRetryPolicyFromArgs(args *einoADKRunLoopArgs) einoTransient } func einoTransientRunRetryPolicyFromMW(mw *config.MultiAgentEinoMiddlewareConfig) einoTransientRunRetryPolicy { - maxBackoff := defaultEinoRunRetryMaxBackoff - if mw != nil && mw.RunRetryMaxBackoffSec > 0 { - maxBackoff = time.Duration(mw.RunRetryMaxBackoffSec) * time.Second - } return einoTransientRunRetryPolicy{ maxAttempts: RunRetryMaxAttemptsFromConfig(mw), - maxBackoff: maxBackoff, + maxBackoff: einoRunRetryMaxBackoffFromConfig(mw), } } @@ -257,10 +260,15 @@ func einoRunRetryMaxAttempts(args *einoADKRunLoopArgs) int { return defaultEinoRunRetryMaxAttempts } -// RunRetryMaxAttemptsFromConfig 与 eino_middleware.run_retry_max_attempts 一致。 +// RunRetryMaxAttemptsFromConfig returns the native model retry count, with legacy run_retry_max_attempts as a fallback. func RunRetryMaxAttemptsFromConfig(mw *config.MultiAgentEinoMiddlewareConfig) int { - if mw != nil && mw.RunRetryMaxAttempts > 0 { - return mw.RunRetryMaxAttempts + if mw != nil { + if mw.ModelRetryMaxRetries > 0 { + return mw.ModelRetryMaxRetries + } + if mw.RunRetryMaxAttempts > 0 { + return mw.RunRetryMaxAttempts + } } return defaultEinoRunRetryMaxAttempts } @@ -272,6 +280,18 @@ func einoRunRetryMaxBackoff(args *einoADKRunLoopArgs) time.Duration { return defaultEinoRunRetryMaxBackoff } +func einoRunRetryMaxBackoffFromConfig(mw *config.MultiAgentEinoMiddlewareConfig) time.Duration { + if mw != nil { + if mw.ModelRetryMaxBackoffSec > 0 { + return time.Duration(mw.ModelRetryMaxBackoffSec) * time.Second + } + if mw.RunRetryMaxBackoffSec > 0 { + return time.Duration(mw.RunRetryMaxBackoffSec) * time.Second + } + } + return defaultEinoRunRetryMaxBackoff +} + // einoRunRestartContextSource 描述无 checkpoint Resume 时 Run 使用的消息来源(日志/SSE)。 type einoRunRestartContextSource string diff --git a/internal/multiagent/eino_transient_run_retry_handler.go b/internal/multiagent/eino_transient_run_retry_handler.go new file mode 100644 index 00000000..fbe7f307 --- /dev/null +++ b/internal/multiagent/eino_transient_run_retry_handler.go @@ -0,0 +1,144 @@ +package multiagent + +import ( + "context" + "fmt" + "time" + + "github.com/cloudwego/eino/adk" + "go.uber.org/zap" +) + +type einoTransientRunRetryHandlerConfig struct { + Context context.Context + ConversationID string + OrchMode string + Args *einoADKRunLoopArgs + BaseMsgs []adk.Message + Progress func(eventType, message string, data interface{}) + Logger *zap.Logger + Pending *einoPendingToolCalls + Policy einoTransientRunRetryPolicy +} + +type einoTransientRunRetryResult struct { + Handled bool + Restarted bool + RestartMsgs []adk.Message + ContextSrc einoRunRestartContextSource + Fatal error +} + +type einoTransientRunRetryHandler struct { + cfg einoTransientRunRetryHandlerConfig + retrier *einoTransientRunRetrier +} + +func newEinoTransientRunRetryHandler(cfg einoTransientRunRetryHandlerConfig) *einoTransientRunRetryHandler { + if cfg.Context == nil { + cfg.Context = context.Background() + } + if cfg.Args == nil { + cfg.Args = &einoADKRunLoopArgs{} + } + if cfg.Policy.maxAttempts <= 0 { + cfg.Policy = einoTransientRunRetryPolicyFromArgs(cfg.Args) + } + return &einoTransientRunRetryHandler{ + cfg: cfg, + retrier: newEinoTransientRunRetrier(cfg.Policy), + } +} + +func (h *einoTransientRunRetryHandler) Prepare( + runErr error, + accumulated []adk.Message, + baseCount int, +) einoTransientRunRetryResult { + if h == nil || !isEinoTransientRunError(runErr) { + return einoTransientRunRetryResult{} + } + restarted, restartMsgs, ctxSource, backoff, retErr := h.retrier.tryRetry( + h.cfg.Context, runErr, h.cfg.Args, h.cfg.BaseMsgs, accumulated, baseCount, + ) + if retErr != nil { + if h.cfg.Pending != nil { + h.cfg.Pending.FlushAsFailed(runErr) + } + if h.cfg.Logger != nil { + h.cfg.Logger.Warn("eino transient retry exhausted", + zap.Error(retErr), + zap.String("orchestration", h.cfg.OrchMode), + zap.Int("maxAttempts", h.retrier.maxAttempts())) + } + return einoTransientRunRetryResult{Handled: true, Fatal: retErr} + } + if !restarted { + return einoTransientRunRetryResult{Handled: true} + } + attemptNo := h.retrier.attempt() + maxAttempts := h.retrier.maxAttempts() + if h.cfg.Logger != nil { + h.cfg.Logger.Warn("eino transient error, retrying after backoff", + zap.Error(runErr), + zap.String("orchestration", h.cfg.OrchMode), + zap.Int("attempt", attemptNo), + zap.Int("maxAttempts", maxAttempts), + zap.Duration("backoff", backoff)) + } + emitEinoRunRetryProgress( + h.cfg.Progress, + h.cfg.ConversationID, + h.cfg.OrchMode, + runErr, + attemptNo, + maxAttempts, + backoff, + ctxSource, + ) + return einoTransientRunRetryResult{ + Handled: true, + Restarted: true, + RestartMsgs: restartMsgs, + ContextSrc: ctxSource, + } +} + +func (h *einoTransientRunRetryHandler) ConfirmRecovery() { + if h != nil && h.retrier != nil && h.retrier.attempt() > 0 { + h.retrier.reset() + } +} + +func emitEinoRunRetryProgress( + progress func(eventType, message string, data interface{}), + conversationID, orchMode string, + runErr error, + attemptNo, maxAttempts int, + backoff time.Duration, + ctxSource einoRunRestartContextSource, +) int { + if progress == nil || runErr == nil { + return 0 + } + errorKind, errorSummary := einoTransientRunErrorUserDetail(runErr) + data := map[string]interface{}{ + "conversationId": conversationID, + "source": "eino", + "orchestration": orchMode, + "error": runErr.Error(), + "errorKind": errorKind, + "errorSummary": errorSummary, + "attempt": attemptNo, + "maxAttempts": maxAttempts, + "backoffSec": int(backoff.Seconds()), + } + progress("eino_run_retry", fmt.Sprintf("遇到临时错误,%d 秒后第 %d/%d 次重试。原因:%s", int(backoff.Seconds()), attemptNo, maxAttempts, errorSummary), data) + restartedData := make(map[string]interface{}, len(data)+1) + for k, v := range data { + restartedData[k] = v + } + restartedData["contextSource"] = string(ctxSource) + progress("eino_run_retry", "已恢复上下文,正在重试…", restartedData) + return 2 +} diff --git a/internal/multiagent/eino_transient_run_retry_handler_test.go b/internal/multiagent/eino_transient_run_retry_handler_test.go new file mode 100644 index 00000000..188f38fe --- /dev/null +++ b/internal/multiagent/eino_transient_run_retry_handler_test.go @@ -0,0 +1,150 @@ +package multiagent + +import ( + "errors" + "strings" + "testing" + "time" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" + "go.uber.org/zap" + "go.uber.org/zap/zaptest/observer" +) + +func TestEinoTransientRunRetryHandlerPreparesRetry(t *testing.T) { + baseMsgs := []adk.Message{schema.UserMessage("base")} + accumulated := []adk.Message{ + schema.UserMessage("base"), + schema.AssistantMessage("partial", nil), + } + runErr := errors.New("HTTP 503 Service Unavailable") + var events []capturedTransientRetryEvent + core, logs := observer.New(zap.WarnLevel) + handler := newEinoTransientRunRetryHandler(einoTransientRunRetryHandlerConfig{ + ConversationID: "conv-1", + OrchMode: "deep_agent", + Args: &einoADKRunLoopArgs{}, + BaseMsgs: baseMsgs, + Progress: func(eventType, message string, data interface{}) { + m, ok := data.(map[string]interface{}) + if !ok { + t.Fatalf("progress data type = %T, want map[string]interface{}", data) + } + events = append(events, capturedTransientRetryEvent{eventType: eventType, message: message, data: m}) + }, + Logger: zap.New(core), + Policy: einoTransientRunRetryPolicy{maxAttempts: 2, maxBackoff: time.Nanosecond}, + }) + + result := handler.Prepare(runErr, accumulated, len(baseMsgs)) + if !result.Handled || !result.Restarted { + t.Fatalf("result = %+v, want handled restarted", result) + } + if result.Fatal != nil { + t.Fatalf("fatal = %v, want nil", result.Fatal) + } + if result.ContextSrc != einoRestartContextAccumulated { + t.Fatalf("context source = %q, want %q", result.ContextSrc, einoRestartContextAccumulated) + } + if len(result.RestartMsgs) != len(accumulated) { + t.Fatalf("restart messages = %d, want %d", len(result.RestartMsgs), len(accumulated)) + } + if len(events) != 2 { + t.Fatalf("events = %d, want 2", len(events)) + } + if events[0].eventType != "eino_run_retry" || events[1].eventType != "eino_run_retry" { + t.Fatalf("event types = %q/%q", events[0].eventType, events[1].eventType) + } + if !strings.Contains(events[0].message, "第 1/2 次重试") { + t.Fatalf("first message = %q", events[0].message) + } + if events[1].message != "已恢复上下文,正在重试…" { + t.Fatalf("second message = %q", events[1].message) + } + assertTransientRetryMapValue(t, events[0].data, "conversationId", "conv-1") + assertTransientRetryMapValue(t, events[0].data, "source", "eino") + assertTransientRetryMapValue(t, events[0].data, "orchestration", "deep_agent") + assertTransientRetryMapValue(t, events[0].data, "error", runErr.Error()) + assertTransientRetryMapValue(t, events[0].data, "errorKind", "upstream_server") + assertTransientRetryMapValue(t, events[0].data, "attempt", 1) + assertTransientRetryMapValue(t, events[0].data, "maxAttempts", 2) + assertTransientRetryMapValue(t, events[0].data, "backoffSec", 0) + assertTransientRetryMapValue(t, events[1].data, "contextSource", string(einoRestartContextAccumulated)) + if logs.FilterMessage("eino transient error, retrying after backoff").Len() != 1 { + t.Fatalf("expected one retry log, got %d", logs.Len()) + } +} + +func TestEinoTransientRunRetryHandlerExhaustsAndFlushesPending(t *testing.T) { + runErr := errors.New("HTTP 503 Service Unavailable") + var progressEvents []string + pending := newEinoPendingToolCalls("conv-1", func(eventType, _ string, _ interface{}) { + progressEvents = append(progressEvents, eventType) + }) + pending.Mark(toolCallPendingInfo{ToolCallID: "call-1", ToolName: "execute", EinoAgent: "agent"}) + core, logs := observer.New(zap.WarnLevel) + handler := newEinoTransientRunRetryHandler(einoTransientRunRetryHandlerConfig{ + OrchMode: "deep_agent", + Args: &einoADKRunLoopArgs{}, + BaseMsgs: []adk.Message{schema.UserMessage("base")}, + Logger: zap.New(core), + Pending: pending, + Policy: einoTransientRunRetryPolicy{maxAttempts: 1, maxBackoff: time.Nanosecond}, + }) + + first := handler.Prepare(runErr, nil, 0) + if !first.Restarted { + t.Fatalf("first result = %+v, want restarted", first) + } + second := handler.Prepare(runErr, nil, 0) + if !second.Handled || second.Fatal == nil { + t.Fatalf("second result = %+v, want fatal exhaustion", second) + } + if pending.Count() != 0 { + t.Fatalf("pending count = %d, want 0", pending.Count()) + } + if len(progressEvents) != 1 || progressEvents[0] != "tool_result" { + t.Fatalf("pending flush events = %#v, want one tool_result", progressEvents) + } + if logs.FilterMessage("eino transient retry exhausted").Len() != 1 { + t.Fatalf("expected one exhausted log, got %d", logs.Len()) + } +} + +func TestEinoTransientRunRetryHandlerConfirmRecoveryResetsAttempts(t *testing.T) { + runErr := errors.New("HTTP 503 Service Unavailable") + handler := newEinoTransientRunRetryHandler(einoTransientRunRetryHandlerConfig{ + Args: &einoADKRunLoopArgs{}, + BaseMsgs: []adk.Message{schema.UserMessage("base")}, + Policy: einoTransientRunRetryPolicy{maxAttempts: 1, maxBackoff: time.Nanosecond}, + }) + if result := handler.Prepare(runErr, nil, 0); !result.Restarted { + t.Fatalf("first result = %+v, want restarted", result) + } + handler.ConfirmRecovery() + if result := handler.Prepare(runErr, nil, 0); !result.Restarted || result.Fatal != nil { + t.Fatalf("after reset result = %+v, want restarted without fatal", result) + } +} + +func TestEinoTransientRunRetryHandlerIgnoresOtherErrors(t *testing.T) { + handler := newEinoTransientRunRetryHandler(einoTransientRunRetryHandlerConfig{}) + result := handler.Prepare(errors.New("invalid api key"), nil, 0) + if result.Handled { + t.Fatalf("result = %+v, want unhandled", result) + } +} + +type capturedTransientRetryEvent struct { + eventType string + message string + data map[string]interface{} +} + +func assertTransientRetryMapValue(t *testing.T, data map[string]interface{}, key string, want interface{}) { + t.Helper() + if got := data[key]; got != want { + t.Fatalf("%s = %v, want %v", key, got, want) + } +} diff --git a/internal/multiagent/eino_turn_loop_bridge_test.go b/internal/multiagent/eino_turn_loop_bridge_test.go new file mode 100644 index 00000000..31b06d86 --- /dev/null +++ b/internal/multiagent/eino_turn_loop_bridge_test.go @@ -0,0 +1,117 @@ +package multiagent + +import ( + "context" + "sync" + "testing" + "time" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +func TestRunEinoADKAgentLoopUsesTurnLoopInterruptPush(t *testing.T) { + baseCtx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + + pushCh := make(chan func(string) bool, 1) + ctx := WithAgentTurnLoopInterruptRegistrar(baseCtx, func(push func(string) bool) func() { + pushCh <- push + return func() {} + }) + + mockModel := newTurnLoopBlockingModel() + agent, err := adk.NewChatModelAgent(ctx, &adk.ChatModelAgentConfig{ + Name: "turn-loop-agent", + Model: mockModel, + }) + if err != nil { + t.Fatalf("NewChatModelAgent: %v", err) + } + + var mu sync.Mutex + var eventTypes []string + var rawInterruptReason string + var rawInterruptRunID string + progress := func(eventType, _ string, data interface{}) { + mu.Lock() + defer mu.Unlock() + eventTypes = append(eventTypes, eventType) + if eventType == "user_interrupt_continue" { + if m, ok := data.(map[string]interface{}); ok { + rawInterruptReason, _ = m["rawReason"].(string) + rawInterruptRunID, _ = m["runId"].(string) + } + } + } + + done := make(chan struct{}) + var result *RunResult + var runErr error + go func() { + defer close(done) + result, runErr = runEinoADKAgentLoop(ctx, &einoADKRunLoopArgs{ + OrchMode: "eino_single", + OrchestratorName: "turn-loop-agent", + ConversationID: "conv-turn-loop", + Progress: progress, + DA: agent, + EmptyResponseMessage: "empty", + TurnLoopInterruptTimeout: 20 * time.Millisecond, + }, []*schema.Message{schema.UserMessage("initial task")}) + }() + + select { + case <-mockModel.started: + case <-ctx.Done(): + t.Fatal("first model call did not start") + } + var push func(string) bool + select { + case push = <-pushCh: + case <-ctx.Done(): + t.Fatal("turn loop interrupt hook was not registered") + } + if !push("focus ssh") { + t.Fatal("turn loop interrupt push was rejected") + } + + select { + case <-done: + case <-ctx.Done(): + t.Fatal("run loop did not finish") + } + if runErr != nil { + t.Fatalf("runErr = %v", runErr) + } + if result == nil || result.Response != "done" { + t.Fatalf("result = %#v, err=%v", result, runErr) + } + if rawInterruptReason != "focus ssh" { + t.Fatalf("raw interrupt reason = %q, want focus ssh", rawInterruptReason) + } + if rawInterruptRunID == "" { + t.Fatal("interrupt progress should include runId") + } + if !containsString(eventTypes, "user_interrupt_continue") { + t.Fatalf("events = %#v, want user_interrupt_continue", eventTypes) + } + + inputs := mockModel.snapshotInputs() + if len(inputs) < 2 { + t.Fatalf("model calls = %d, want at least 2", len(inputs)) + } + last := inputs[len(inputs)-1] + if len(last) == 0 || last[len(last)-1].Role != schema.User || last[len(last)-1].Content == "initial task" { + t.Fatalf("last model input = %#v, want interrupt supplement turn", last) + } +} + +func containsString(items []string, target string) bool { + for _, item := range items { + if item == target { + return true + } + } + return false +} diff --git a/internal/multiagent/eino_turn_loop_event_bridge.go b/internal/multiagent/eino_turn_loop_event_bridge.go new file mode 100644 index 00000000..3a5d9e90 --- /dev/null +++ b/internal/multiagent/eino_turn_loop_event_bridge.go @@ -0,0 +1,103 @@ +package multiagent + +import ( + "context" + "errors" + "strings" + "sync/atomic" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +type einoTurnLoopEventBridge struct { + conversationID string + orchestration string + progress func(eventType, message string, data interface{}) + gen *adk.AsyncGenerator[*adk.AgentEvent] + forwardedErr atomic.Bool +} + +func newEinoTurnLoopEventBridge( + conversationID string, + orchestration string, + progress func(eventType, message string, data interface{}), + gen *adk.AsyncGenerator[*adk.AgentEvent], +) *einoTurnLoopEventBridge { + return &einoTurnLoopEventBridge{ + conversationID: conversationID, + orchestration: orchestration, + progress: progress, + gen: gen, + } +} + +func (b *einoTurnLoopEventBridge) OnAgentEvents( + _ context.Context, + tc *adk.TurnContext[EinoTurnLoopItem, *schema.Message], + events *adk.AsyncIterator[*adk.AgentEvent], +) error { + for { + ev, ok := events.Next() + if !ok { + return nil + } + if ev == nil { + continue + } + if ev.Err != nil && isEinoTurnLoopPreemptCancel(tc, ev.Err) { + b.emitPreempted() + continue + } + if b.gen != nil { + b.gen.Send(ev) + } + if ev.Err != nil { + b.forwardedErr.Store(true) + return ev.Err + } + } +} + +func (b *einoTurnLoopEventBridge) ForwardedError() bool { + if b == nil { + return false + } + return b.forwardedErr.Load() +} + +func (b *einoTurnLoopEventBridge) emitPreempted() { + if b == nil || b.progress == nil { + return + } + b.progress("progress", "Eino TurnLoop 已在安全点切换到用户补充后的下一轮。", map[string]interface{}{ + "conversationId": b.conversationID, + "source": "eino", + "orchestration": b.orchestration, + "kind": "turn_loop_preempted", + }) +} + +func isEinoTurnLoopPreemptCancel(tc *adk.TurnContext[EinoTurnLoopItem, *schema.Message], err error) bool { + if tc == nil || err == nil { + return false + } + var cancelErr *adk.CancelError + if !errors.As(err, &cancelErr) { + return false + } + select { + case <-tc.Preempted: + return true + default: + return false + } +} + +func einoTurnLoopInterruptTimelineSummary(note string) string { + note = strings.TrimSpace(note) + if note == "" { + return "用户选择「中断并继续」,未填写说明;已推入 Eino TurnLoop 并等待安全点续跑。" + } + return "用户中断说明(Eino TurnLoop 原生续跑):\n\n" + note +} diff --git a/internal/multiagent/eino_turn_loop_event_bridge_test.go b/internal/multiagent/eino_turn_loop_event_bridge_test.go new file mode 100644 index 00000000..78f70f15 --- /dev/null +++ b/internal/multiagent/eino_turn_loop_event_bridge_test.go @@ -0,0 +1,93 @@ +package multiagent + +import ( + "context" + "errors" + "testing" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +func TestEinoTurnLoopEventBridgeSwallowsPreemptCancel(t *testing.T) { + iter, gen := adk.NewAsyncIteratorPair[*adk.AgentEvent]() + outIter, outGen := adk.NewAsyncIteratorPair[*adk.AgentEvent]() + + preempted := make(chan struct{}) + close(preempted) + var eventTypes []string + bridge := newEinoTurnLoopEventBridge("conv", "eino_single", func(eventType, _ string, _ interface{}) { + eventTypes = append(eventTypes, eventType) + }, outGen) + + gen.Send(&adk.AgentEvent{Err: &adk.CancelError{Info: &adk.AgentCancelInfo{}}}) + gen.Close() + + err := bridge.OnAgentEvents(context.Background(), &adk.TurnContext[EinoTurnLoopItem, *schema.Message]{ + Preempted: preempted, + }, iter) + if err != nil { + t.Fatalf("preempt cancel should be swallowed, got %v", err) + } + if bridge.ForwardedError() { + t.Fatal("preempt cancel should not be marked as forwarded") + } + if !containsString(eventTypes, "progress") { + t.Fatalf("events = %#v, want progress", eventTypes) + } + outGen.Close() + if ev, ok := outIter.Next(); ok || ev != nil { + t.Fatalf("preempt cancel should not be forwarded, got ok=%v ev=%#v", ok, ev) + } +} + +func TestEinoTurnLoopEventBridgeForwardsRegularError(t *testing.T) { + iter, gen := adk.NewAsyncIteratorPair[*adk.AgentEvent]() + outIter, outGen := adk.NewAsyncIteratorPair[*adk.AgentEvent]() + + want := errors.New("model failed") + bridge := newEinoTurnLoopEventBridge("conv", "eino_single", nil, outGen) + gen.Send(&adk.AgentEvent{Err: want}) + gen.Close() + + err := bridge.OnAgentEvents(context.Background(), &adk.TurnContext[EinoTurnLoopItem, *schema.Message]{ + Preempted: make(chan struct{}), + }, iter) + if !errors.Is(err, want) { + t.Fatalf("err = %v, want %v", err, want) + } + if !bridge.ForwardedError() { + t.Fatal("regular error should be marked as forwarded") + } + outGen.Close() + ev, ok := outIter.Next() + if !ok || ev == nil || !errors.Is(ev.Err, want) { + t.Fatalf("forwarded event = %#v ok=%v", ev, ok) + } +} + +func TestEinoTurnLoopEventBridgeForwardsNormalEvents(t *testing.T) { + iter, gen := adk.NewAsyncIteratorPair[*adk.AgentEvent]() + outIter, outGen := adk.NewAsyncIteratorPair[*adk.AgentEvent]() + + bridge := newEinoTurnLoopEventBridge("conv", "eino_single", nil, outGen) + gen.Send(&adk.AgentEvent{ + AgentName: "agent", + Output: &adk.AgentOutput{MessageOutput: &adk.MessageVariant{ + Message: schema.AssistantMessage("ok", nil), + Role: schema.Assistant, + }}, + }) + gen.Close() + + if err := bridge.OnAgentEvents(context.Background(), &adk.TurnContext[EinoTurnLoopItem, *schema.Message]{ + Preempted: make(chan struct{}), + }, iter); err != nil { + t.Fatalf("OnAgentEvents: %v", err) + } + outGen.Close() + ev, ok := outIter.Next() + if !ok || ev == nil || ev.AgentName != "agent" { + t.Fatalf("forwarded event = %#v ok=%v", ev, ok) + } +} diff --git a/internal/multiagent/eino_turn_loop_iterator_starter.go b/internal/multiagent/eino_turn_loop_iterator_starter.go new file mode 100644 index 00000000..6326e3cd --- /dev/null +++ b/internal/multiagent/eino_turn_loop_iterator_starter.go @@ -0,0 +1,162 @@ +package multiagent + +import ( + "context" + "errors" + "strings" + "sync/atomic" + "time" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" + "go.uber.org/zap" +) + +type einoTurnLoopRuntimeControl interface { + Run(context.Context) + PushInterruptContinue(string) bool + StopImmediate(string) + StopWhenIdle() + Wait() *adk.TurnLoopExitState[EinoTurnLoopItem, *schema.Message] +} + +type einoTurnLoopRuntimeFactory func(EinoTurnLoopRuntimeConfig) einoTurnLoopRuntimeControl + +type einoTurnLoopIteratorStarterConfig struct { + Context context.Context + Agent adk.Agent + ConversationID string + OrchMode string + Progress func(eventType, message string, data interface{}) + Logger *zap.Logger + Store adk.CheckPointStore + CheckPointID string + InterruptTimeout time.Duration + NativeCancelCause *atomic.Value + UnregisterAgentCancel *func() + UnregisterTurnLoopInterrupt *func() + RuntimeCancelRegistrar AgentRuntimeCancelRegistrar + TurnLoopInterruptRegistrar AgentTurnLoopInterruptRegistrar + RuntimeFactory einoTurnLoopRuntimeFactory +} + +type einoTurnLoopIteratorStarter struct { + cfg einoTurnLoopIteratorStarterConfig +} + +func newEinoTurnLoopIteratorStarter(cfg einoTurnLoopIteratorStarterConfig) *einoTurnLoopIteratorStarter { + if cfg.RuntimeFactory == nil { + cfg.RuntimeFactory = func(runtimeCfg EinoTurnLoopRuntimeConfig) einoTurnLoopRuntimeControl { + return NewEinoTurnLoopRuntime(runtimeCfg) + } + } + return &einoTurnLoopIteratorStarter{cfg: cfg} +} + +func (s *einoTurnLoopIteratorStarter) Start(runMsgs []adk.Message) *adk.AsyncIterator[*adk.AgentEvent] { + if s == nil { + return nil + } + callAndClearUnregister(s.cfg.UnregisterTurnLoopInterrupt) + callAndClearUnregister(s.cfg.UnregisterAgentCancel) + + iter, gen := adk.NewAsyncIteratorPair[*adk.AgentEvent]() + eventsBridge := newEinoTurnLoopEventBridge(s.cfg.ConversationID, s.cfg.OrchMode, s.cfg.Progress, gen) + runtime := s.cfg.RuntimeFactory(EinoTurnLoopRuntimeConfig{ + Agent: s.cfg.Agent, + InitialMessages: runMsgs, + Store: s.cfg.Store, + CheckpointID: s.turnLoopCheckpointID(), + EnableStreaming: true, + InterruptTimeout: s.cfg.InterruptTimeout, + OnAgentEvents: eventsBridge.OnAgentEvents, + }) + s.bindTurnLoopInterrupt(runtime) + s.bindRuntimeCancel(runtime) + runtime.Run(s.cfg.Context) + runtime.StopWhenIdle() + go func() { + defer gen.Close() + state := runtime.Wait() + if state == nil || state.ExitReason == nil || eventsBridge.ForwardedError() { + return + } + gen.Send(&adk.AgentEvent{Err: state.ExitReason}) + }() + return iter +} + +func (s *einoTurnLoopIteratorStarter) turnLoopCheckpointID() string { + if s == nil || s.cfg.CheckPointID == "" { + return "" + } + return buildEinoTurnLoopCheckpointID(s.cfg.OrchMode) +} + +func (s *einoTurnLoopIteratorStarter) bindTurnLoopInterrupt(runtime einoTurnLoopRuntimeControl) { + if s == nil || runtime == nil || s.cfg.TurnLoopInterruptRegistrar == nil || s.cfg.UnregisterTurnLoopInterrupt == nil { + return + } + *s.cfg.UnregisterTurnLoopInterrupt = s.cfg.TurnLoopInterruptRegistrar(func(note string) bool { + ok := runtime.PushInterruptContinue(note) + if ok { + s.emitInterruptContinueProgress(note) + } + return ok + }) +} + +func (s *einoTurnLoopIteratorStarter) bindRuntimeCancel(runtime einoTurnLoopRuntimeControl) { + if s == nil || runtime == nil || s.cfg.RuntimeCancelRegistrar == nil || s.cfg.UnregisterAgentCancel == nil { + return + } + *s.cfg.UnregisterAgentCancel = s.cfg.RuntimeCancelRegistrar(func(cause error) bool { + s.storeNativeCancelCause(cause) + if errors.Is(cause, ErrInterruptContinue) { + return runtime.PushInterruptContinue("") + } + runtime.StopImmediate("task_cancelled") + if s.cfg.Logger != nil { + s.cfg.Logger.Info("eino turn loop stop requested", + zap.String("conversation_id", s.cfg.ConversationID), + zap.String("orchestration", s.cfg.OrchMode), + zap.Error(cause)) + } + return true + }) +} + +func (s *einoTurnLoopIteratorStarter) storeNativeCancelCause(cause error) { + if s == nil || s.cfg.NativeCancelCause == nil || cause == nil { + return + } + s.cfg.NativeCancelCause.Store(cause) +} + +func (s *einoTurnLoopIteratorStarter) emitInterruptContinueProgress(note string) { + if s == nil || s.cfg.Progress == nil { + return + } + trimmed := strings.TrimSpace(note) + s.cfg.Progress("user_interrupt_continue", einoTurnLoopInterruptTimelineSummary(note), map[string]interface{}{ + "conversationId": s.cfg.ConversationID, + "rawReason": trimmed, + "emptyReason": trimmed == "", + "kind": "turn_loop_preempt", + "source": "eino", + "orchestration": s.cfg.OrchMode, + }) + s.cfg.Progress("progress", "已将用户补充推入 Eino TurnLoop,正在等待安全点切换…", map[string]interface{}{ + "conversationId": s.cfg.ConversationID, + "source": "eino", + "orchestration": s.cfg.OrchMode, + }) +} + +func callAndClearUnregister(target *func()) { + if target == nil || *target == nil { + return + } + (*target)() + *target = nil +} diff --git a/internal/multiagent/eino_turn_loop_iterator_starter_test.go b/internal/multiagent/eino_turn_loop_iterator_starter_test.go new file mode 100644 index 00000000..db63e6b7 --- /dev/null +++ b/internal/multiagent/eino_turn_loop_iterator_starter_test.go @@ -0,0 +1,188 @@ +package multiagent + +import ( + "context" + "errors" + "sync" + "sync/atomic" + "testing" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +type fakeTurnLoopRuntimeControl struct { + mu sync.Mutex + runCalled bool + stopIdle bool + stopped string + pushedNotes []string + pushOK bool +} + +func (f *fakeTurnLoopRuntimeControl) Run(context.Context) { + f.mu.Lock() + defer f.mu.Unlock() + f.runCalled = true +} + +func (f *fakeTurnLoopRuntimeControl) PushInterruptContinue(note string) bool { + f.mu.Lock() + defer f.mu.Unlock() + f.pushedNotes = append(f.pushedNotes, note) + return f.pushOK +} + +func (f *fakeTurnLoopRuntimeControl) StopImmediate(cause string) { + f.mu.Lock() + defer f.mu.Unlock() + f.stopped = cause +} + +func (f *fakeTurnLoopRuntimeControl) StopWhenIdle() { + f.mu.Lock() + defer f.mu.Unlock() + f.stopIdle = true +} + +func (f *fakeTurnLoopRuntimeControl) Wait() *adk.TurnLoopExitState[EinoTurnLoopItem, *schema.Message] { + return nil +} + +func (f *fakeTurnLoopRuntimeControl) snapshot() (runCalled bool, stopIdle bool, stopped string, pushed []string) { + f.mu.Lock() + defer f.mu.Unlock() + return f.runCalled, f.stopIdle, f.stopped, append([]string(nil), f.pushedNotes...) +} + +func TestEinoTurnLoopIteratorStarterBindsRegistrarsAndProgress(t *testing.T) { + fakeRuntime := &fakeTurnLoopRuntimeControl{pushOK: true} + oldAgentCleared := false + oldTurnCleared := false + unregisterAgent := func() { oldAgentCleared = true } + unregisterTurn := func() { oldTurnCleared = true } + var interruptPush func(string) bool + var cancelPush func(error) bool + var createdCfg EinoTurnLoopRuntimeConfig + var events []struct { + eventType string + message string + data map[string]interface{} + } + + iter := newEinoTurnLoopIteratorStarter(einoTurnLoopIteratorStarterConfig{ + Context: context.Background(), + ConversationID: "conv", + OrchMode: "deep", + CheckPointID: "runner-checkpoint", + UnregisterAgentCancel: &unregisterAgent, + UnregisterTurnLoopInterrupt: &unregisterTurn, + RuntimeCancelRegistrar: func(push func(error) bool) func() { + cancelPush = push + return func() {} + }, + TurnLoopInterruptRegistrar: func(push func(string) bool) func() { + interruptPush = push + return func() {} + }, + RuntimeFactory: func(cfg EinoTurnLoopRuntimeConfig) einoTurnLoopRuntimeControl { + createdCfg = cfg + return fakeRuntime + }, + Progress: func(eventType, message string, data interface{}) { + item := struct { + eventType string + message string + data map[string]interface{} + }{eventType: eventType, message: message} + if m, ok := data.(map[string]interface{}); ok { + item.data = m + } + events = append(events, item) + }, + }).Start([]adk.Message{}) + + if iter == nil { + t.Fatal("iterator should be created") + } + if !oldAgentCleared || !oldTurnCleared { + t.Fatalf("oldAgentCleared=%v oldTurnCleared=%v, want both true", oldAgentCleared, oldTurnCleared) + } + if interruptPush == nil { + t.Fatal("turn loop interrupt registrar was not bound") + } + if cancelPush == nil { + t.Fatal("runtime cancel registrar was not bound") + } + if createdCfg.CheckpointID != buildEinoTurnLoopCheckpointID("deep") { + t.Fatalf("checkpoint id = %q, want turn loop checkpoint id", createdCfg.CheckpointID) + } + if !interruptPush(" focus ssh ") { + t.Fatal("interrupt push should return runtime result") + } + + runCalled, stopIdle, _, pushed := fakeRuntime.snapshot() + if !runCalled || !stopIdle { + t.Fatalf("runCalled=%v stopIdle=%v, want both true", runCalled, stopIdle) + } + if len(pushed) != 1 || pushed[0] != " focus ssh " { + t.Fatalf("pushed notes = %#v", pushed) + } + if len(events) != 2 { + t.Fatalf("events = %#v, want user interrupt and progress", events) + } + if events[0].eventType != "user_interrupt_continue" || events[0].data["rawReason"] != "focus ssh" { + t.Fatalf("first event = %#v", events[0]) + } + if events[1].eventType != "progress" { + t.Fatalf("second event = %#v", events[1]) + } +} + +func TestEinoTurnLoopIteratorStarterRuntimeCancel(t *testing.T) { + fakeRuntime := &fakeTurnLoopRuntimeControl{pushOK: true} + var nativeCancelCause atomic.Value + var cancelPush func(error) bool + var unregisterAgent func() + + newEinoTurnLoopIteratorStarter(einoTurnLoopIteratorStarterConfig{ + Context: context.Background(), + ConversationID: "conv", + OrchMode: "eino_single", + NativeCancelCause: &nativeCancelCause, + UnregisterAgentCancel: &unregisterAgent, + RuntimeCancelRegistrar: func(push func(error) bool) func() { + cancelPush = push + return func() {} + }, + RuntimeFactory: func(EinoTurnLoopRuntimeConfig) einoTurnLoopRuntimeControl { + return fakeRuntime + }, + }).Start(nil) + if cancelPush == nil { + t.Fatal("runtime cancel registrar was not bound") + } + + if !cancelPush(ErrInterruptContinue) { + t.Fatal("interrupt continue cancel should be handled by TurnLoop push") + } + _, _, stopped, pushed := fakeRuntime.snapshot() + if stopped != "" { + t.Fatalf("stopped = %q, want no immediate stop for interrupt continue", stopped) + } + if len(pushed) != 1 || pushed[0] != "" { + t.Fatalf("pushed notes = %#v, want empty interrupt continue note", pushed) + } + + stopErr := errors.New("stop now") + if !cancelPush(stopErr) { + t.Fatal("regular cancel should be handled") + } + _, _, stopped, _ = fakeRuntime.snapshot() + if stopped != "task_cancelled" { + t.Fatalf("stopped = %q, want task_cancelled", stopped) + } + if got, _ := nativeCancelCause.Load().(error); !errors.Is(got, stopErr) { + t.Fatalf("native cancel cause = %v, want %v", got, stopErr) + } +} diff --git a/internal/multiagent/eino_turn_loop_runtime.go b/internal/multiagent/eino_turn_loop_runtime.go new file mode 100644 index 00000000..4ee38743 --- /dev/null +++ b/internal/multiagent/eino_turn_loop_runtime.go @@ -0,0 +1,178 @@ +package multiagent + +import ( + "context" + "strings" + "time" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/schema" +) + +const ( + einoTurnLoopInterruptPreemptTimeout = 3 * time.Second + einoTurnLoopIdleStop = 250 * time.Millisecond +) + +// EinoTurnLoopItem is the conversation-level input unit consumed by an Eino +// TurnLoop. The item is gob-friendly so it can be checkpointed by TurnLoop when +// a CheckPointStore is configured. +type EinoTurnLoopItem struct { + Messages []*schema.Message + Kind string + Note string +} + +// EinoTurnLoopRuntime wraps Eino's native TurnLoop with the semantics this +// project needs: persistent per-conversation runtime, user-supplement preempt, +// and graceful idle shutdown. +type EinoTurnLoopRuntime struct { + loop *adk.TurnLoop[EinoTurnLoopItem, *schema.Message] + interruptTimeout time.Duration +} + +type EinoTurnLoopRuntimeConfig struct { + Agent adk.Agent + InitialMessages []*schema.Message + Store adk.CheckPointStore + CheckpointID string + EnableStreaming bool + PrepareAgent func(context.Context, *adk.TurnLoop[EinoTurnLoopItem, *schema.Message], []EinoTurnLoopItem) (adk.Agent, error) + OnAgentEvents func(context.Context, *adk.TurnContext[EinoTurnLoopItem, *schema.Message], *adk.AsyncIterator[*adk.AgentEvent]) error + InterruptTimeout time.Duration +} + +func NewEinoTurnLoopRuntime(cfg EinoTurnLoopRuntimeConfig) *EinoTurnLoopRuntime { + timeout := cfg.InterruptTimeout + if timeout <= 0 { + timeout = einoTurnLoopInterruptPreemptTimeout + } + enableStreaming := cfg.EnableStreaming + prepareAgent := cfg.PrepareAgent + if prepareAgent == nil { + prepareAgent = func(context.Context, *adk.TurnLoop[EinoTurnLoopItem, *schema.Message], []EinoTurnLoopItem) (adk.Agent, error) { + return cfg.Agent, nil + } + } + loop := adk.NewTurnLoop[EinoTurnLoopItem, *schema.Message](adk.TurnLoopConfig[EinoTurnLoopItem, *schema.Message]{ + Store: cfg.Store, + CheckpointID: cfg.CheckpointID, + GenInput: func(ctx context.Context, _ *adk.TurnLoop[EinoTurnLoopItem, *schema.Message], items []EinoTurnLoopItem) (*adk.GenInputResult[EinoTurnLoopItem, *schema.Message], error) { + msgs := mergeEinoTurnLoopMessages(items) + return &adk.GenInputResult[EinoTurnLoopItem, *schema.Message]{ + RunCtx: ctx, + Input: &adk.AgentInput{ + Messages: msgs, + EnableStreaming: enableStreaming, + }, + Consumed: items, + }, nil + }, + GenResume: func(ctx context.Context, _ *adk.TurnLoop[EinoTurnLoopItem, *schema.Message], interruptedItems, unhandledItems, newItems []EinoTurnLoopItem) (*adk.GenResumeResult[EinoTurnLoopItem, *schema.Message], error) { + consumed := make([]EinoTurnLoopItem, 0, len(interruptedItems)+len(newItems)) + consumed = append(consumed, interruptedItems...) + consumed = append(consumed, newItems...) + remaining := append([]EinoTurnLoopItem(nil), unhandledItems...) + return &adk.GenResumeResult[EinoTurnLoopItem, *schema.Message]{ + RunCtx: ctx, + Consumed: consumed, + Remaining: remaining, + }, nil + }, + PrepareAgent: prepareAgent, + OnAgentEvents: cfg.OnAgentEvents, + }) + if len(cfg.InitialMessages) > 0 { + loop.Push(EinoTurnLoopItem{Kind: "initial", Messages: cloneSchemaMessages(cfg.InitialMessages)}) + } + return &EinoTurnLoopRuntime{loop: loop, interruptTimeout: timeout} +} + +func (r *EinoTurnLoopRuntime) Run(ctx context.Context) { + if r == nil || r.loop == nil { + return + } + r.loop.Run(ctx) +} + +func (r *EinoTurnLoopRuntime) PushInterruptContinue(note string) bool { + if r == nil || r.loop == nil { + return false + } + item := EinoTurnLoopItem{ + Kind: "interrupt_continue", + Note: strings.TrimSpace(note), + Messages: []*schema.Message{schema.UserMessage(formatInterruptContinuePrompt(note))}, + } + ok, ack := r.loop.Push(item, adk.WithPreemptTimeout[EinoTurnLoopItem, *schema.Message](adk.AnySafePoint, r.interruptTimeout)) + if ack != nil { + go func() { <-ack }() + } + return ok +} + +func (r *EinoTurnLoopRuntime) StopImmediate(cause string) { + if r == nil || r.loop == nil { + return + } + r.loop.Stop(adk.WithImmediate(), adk.WithStopCause(cause)) +} + +func (r *EinoTurnLoopRuntime) StopWhenIdle() { + if r == nil || r.loop == nil { + return + } + r.loop.Stop(adk.UntilIdleFor(einoTurnLoopIdleStop)) +} + +func (r *EinoTurnLoopRuntime) Wait() *adk.TurnLoopExitState[EinoTurnLoopItem, *schema.Message] { + if r == nil || r.loop == nil { + return nil + } + return r.loop.Wait() +} + +func mergeEinoTurnLoopMessages(items []EinoTurnLoopItem) []*schema.Message { + var msgs []*schema.Message + for _, item := range items { + msgs = append(msgs, cloneSchemaMessages(item.Messages)...) + } + return msgs +} + +func formatInterruptContinuePrompt(note string) string { + note = strings.TrimSpace(note) + if note == "" { + return "用户请求中断当前推理并继续。请基于已经完成的步骤继续,不要重复已完成工具调用。" + } + return "用户请求中断当前推理并补充上下文后继续:\n" + note + + "\n\n请基于已经完成的步骤继续,不要重复已完成工具调用。" +} + +func cloneSchemaMessages(in []*schema.Message) []*schema.Message { + if len(in) == 0 { + return nil + } + out := make([]*schema.Message, 0, len(in)) + for _, msg := range in { + if msg == nil { + continue + } + cp := *msg + if len(msg.ToolCalls) > 0 { + cp.ToolCalls = append([]schema.ToolCall(nil), msg.ToolCalls...) + } + if len(msg.MultiContent) > 0 { + cp.MultiContent = append([]schema.ChatMessagePart(nil), msg.MultiContent...) + } + if len(msg.UserInputMultiContent) > 0 { + cp.UserInputMultiContent = append([]schema.MessageInputPart(nil), msg.UserInputMultiContent...) + } + if len(msg.AssistantGenMultiContent) > 0 { + cp.AssistantGenMultiContent = append([]schema.MessageOutputPart(nil), msg.AssistantGenMultiContent...) + } + cp.Extra = cloneAnyMap(msg.Extra) + out = append(out, &cp) + } + return out +} diff --git a/internal/multiagent/eino_turn_loop_runtime_test.go b/internal/multiagent/eino_turn_loop_runtime_test.go new file mode 100644 index 00000000..62d5f1d2 --- /dev/null +++ b/internal/multiagent/eino_turn_loop_runtime_test.go @@ -0,0 +1,145 @@ +package multiagent + +import ( + "context" + "strings" + "sync" + "testing" + "time" + + "github.com/cloudwego/eino/adk" + "github.com/cloudwego/eino/components/model" + "github.com/cloudwego/eino/schema" +) + +type turnLoopBlockingModel struct { + started chan struct{} + release chan struct{} + + mu sync.Mutex + inputs [][]*schema.Message +} + +func newTurnLoopBlockingModel() *turnLoopBlockingModel { + return &turnLoopBlockingModel{ + started: make(chan struct{}, 8), + release: make(chan struct{}), + } +} + +func (m *turnLoopBlockingModel) Generate(ctx context.Context, input []*schema.Message, _ ...model.Option) (*schema.Message, error) { + m.mu.Lock() + m.inputs = append(m.inputs, cloneSchemaMessages(input)) + callNo := len(m.inputs) + m.mu.Unlock() + + select { + case m.started <- struct{}{}: + default: + } + if callNo == 1 { + select { + case <-ctx.Done(): + return nil, ctx.Err() + case <-m.release: + } + } + return schema.AssistantMessage("done", nil), nil +} + +func (m *turnLoopBlockingModel) Stream(ctx context.Context, input []*schema.Message, opts ...model.Option) (*schema.StreamReader[*schema.Message], error) { + msg, err := m.Generate(ctx, input, opts...) + if err != nil { + return nil, err + } + return schema.StreamReaderFromArray([]*schema.Message{msg}), nil +} + +func (m *turnLoopBlockingModel) snapshotInputs() [][]*schema.Message { + m.mu.Lock() + defer m.mu.Unlock() + out := make([][]*schema.Message, len(m.inputs)) + for i := range m.inputs { + out[i] = cloneSchemaMessages(m.inputs[i]) + } + return out +} + +func TestEinoTurnLoopRuntimePushInterruptStartsNextTurn(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + + mockModel := newTurnLoopBlockingModel() + agent, err := adk.NewChatModelAgent(ctx, &adk.ChatModelAgentConfig{ + Name: "turn-loop-agent", + Model: mockModel, + }) + if err != nil { + t.Fatalf("NewChatModelAgent: %v", err) + } + + runtime := NewEinoTurnLoopRuntime(EinoTurnLoopRuntimeConfig{ + Agent: agent, + InitialMessages: []*schema.Message{schema.UserMessage("initial task")}, + InterruptTimeout: 20 * time.Millisecond, + }) + runtime.Run(ctx) + + select { + case <-mockModel.started: + case <-ctx.Done(): + t.Fatal("first model call did not start") + } + if !runtime.PushInterruptContinue("focus on ssh") { + t.Fatal("interrupt continue push was rejected") + } + select { + case <-mockModel.started: + case <-ctx.Done(): + t.Fatal("second model call did not start after interrupt push") + } + + runtime.StopWhenIdle() + state := runtime.Wait() + if state == nil { + t.Fatal("expected turn loop exit state") + } + if state.ExitReason != nil { + t.Fatalf("exit reason = %v", state.ExitReason) + } + + inputs := mockModel.snapshotInputs() + if len(inputs) < 2 { + t.Fatalf("model calls = %d, want at least 2", len(inputs)) + } + if got := inputs[0][0].Content; got != "initial task" { + t.Fatalf("first input = %q, want initial task", got) + } + lastInput := inputs[len(inputs)-1] + if len(lastInput) == 0 || !strings.Contains(lastInput[len(lastInput)-1].Content, "focus on ssh") { + t.Fatalf("last input = %#v, want interrupt note", lastInput) + } +} + +func TestMergeEinoTurnLoopMessagesClonesInput(t *testing.T) { + original := schema.UserMessage("hello") + msgs := mergeEinoTurnLoopMessages([]EinoTurnLoopItem{{Messages: []*schema.Message{original}}}) + if len(msgs) != 1 || msgs[0].Content != "hello" { + t.Fatalf("merged = %#v", msgs) + } + msgs[0].Content = "changed" + if original.Content != "hello" { + t.Fatalf("original message was mutated: %#v", original) + } +} + +func TestFormatInterruptContinuePrompt(t *testing.T) { + got := formatInterruptContinuePrompt("focus ports") + if !strings.Contains(got, "focus ports") || !strings.Contains(got, "不要重复") { + t.Fatalf("prompt = %q", got) + } + empty := formatInterruptContinuePrompt(" ") + if !strings.Contains(empty, "不要重复") { + t.Fatalf("empty prompt = %q", empty) + } +} diff --git a/internal/multiagent/gen_model_input.go b/internal/multiagent/gen_model_input.go index 2c7a9fe9..bdf69619 100644 --- a/internal/multiagent/gen_model_input.go +++ b/internal/multiagent/gen_model_input.go @@ -21,3 +21,12 @@ func literalInstructionGenModelInput(ctx context.Context, instruction string, in msgs = append(msgs, input.Messages...) return msgs, nil } + +func literalAgenticInstructionGenModelInput(ctx context.Context, instruction string, input *adk.TypedAgentInput[*schema.AgenticMessage]) ([]*schema.AgenticMessage, error) { + msgs := make([]*schema.AgenticMessage, 0, len(input.Messages)+1) + if instruction != "" { + msgs = append(msgs, schema.SystemAgenticMessage(instruction)) + } + msgs = append(msgs, input.Messages...) + return msgs, nil +} diff --git a/internal/multiagent/model_output_recovery_compat.go b/internal/multiagent/model_output_recovery_compat.go new file mode 100644 index 00000000..2393b922 --- /dev/null +++ b/internal/multiagent/model_output_recovery_compat.go @@ -0,0 +1,78 @@ +package multiagent + +import ( + "context" + "encoding/json" + "fmt" + "strings" + + "github.com/cloudwego/eino/compose" + "github.com/cloudwego/eino/schema" +) + +const ( + modelOutputRecoveryKey = "_cyberstrike_model_output_recovery" + modelOutputRejectedResultPrefix = "[Model Output Rejected]" +) + +type modelOutputRecoveryMarker struct { + Reason string `json:"reason"` + RepairAttempt int `json:"repair_attempt"` +} + +// modelOutputExecutionGuardMiddleware is a compatibility shim for old persisted +// recovery-marker tool calls. New runs should let the tool layer return normal +// soft errors to the model instead of pre-rewriting model output. +func modelOutputExecutionGuardMiddleware() compose.ToolMiddleware { + messageFor := func(input *compose.ToolInput) (string, bool) { + if input == nil { + return "", false + } + var envelope map[string]json.RawMessage + if json.Unmarshal([]byte(input.Arguments), &envelope) != nil { + return "", false + } + raw, ok := envelope[modelOutputRecoveryKey] + if !ok { + return "", false + } + var marker modelOutputRecoveryMarker + _ = json.Unmarshal(raw, &marker) + return fmt.Sprintf("%s Tool call '%s' was not executed because it is a legacy model-output recovery marker (%s). Repair attempt %d.", + modelOutputRejectedResultPrefix, input.Name, marker.Reason, marker.RepairAttempt), true + } + return compose.ToolMiddleware{ + Invokable: func(next compose.InvokableToolEndpoint) compose.InvokableToolEndpoint { + return func(ctx context.Context, input *compose.ToolInput) (*compose.ToolOutput, error) { + if msg, reject := messageFor(input); reject { + return &compose.ToolOutput{Result: msg}, nil + } + return next(ctx, input) + } + }, + Streamable: func(next compose.StreamableToolEndpoint) compose.StreamableToolEndpoint { + return func(ctx context.Context, input *compose.ToolInput) (*compose.StreamToolOutput, error) { + if msg, reject := messageFor(input); reject { + return &compose.StreamToolOutput{Result: schema.StreamReaderFromArray([]string{msg})}, nil + } + return next(ctx, input) + } + }, + } +} + +func modelOutputRecoveryFromToolCall(tc schema.ToolCall) (modelOutputRecoveryMarker, bool) { + var envelope map[string]json.RawMessage + if json.Unmarshal([]byte(tc.Function.Arguments), &envelope) != nil { + return modelOutputRecoveryMarker{}, false + } + raw, ok := envelope[modelOutputRecoveryKey] + if !ok { + return modelOutputRecoveryMarker{}, false + } + var marker modelOutputRecoveryMarker + if json.Unmarshal(raw, &marker) != nil { + return modelOutputRecoveryMarker{}, false + } + return marker, strings.TrimSpace(marker.Reason) != "" || marker.RepairAttempt > 0 +} diff --git a/internal/multiagent/model_output_recovery_compat_test.go b/internal/multiagent/model_output_recovery_compat_test.go new file mode 100644 index 00000000..f9456d6a --- /dev/null +++ b/internal/multiagent/model_output_recovery_compat_test.go @@ -0,0 +1,69 @@ +package multiagent + +import ( + "context" + "io" + "strings" + "testing" + + "github.com/cloudwego/eino/compose" +) + +func TestModelOutputExecutionGuardMiddlewareBlocksLegacyRecoveryMarker(t *testing.T) { + called := false + markerJSON := `{"` + modelOutputRecoveryKey + `":{"reason":"invalid_tool_arguments_json","repair_attempt":1}}` + wrapped := modelOutputExecutionGuardMiddleware().Invokable(func(context.Context, *compose.ToolInput) (*compose.ToolOutput, error) { + called = true + return &compose.ToolOutput{Result: "executed"}, nil + }) + + out, err := wrapped(context.Background(), &compose.ToolInput{Name: "task", Arguments: markerJSON}) + if err != nil { + t.Fatalf("guard returned error: %v", err) + } + if called { + t.Fatal("legacy recovery marker should not reach the real tool endpoint") + } + if out == nil || !strings.HasPrefix(out.Result, modelOutputRejectedResultPrefix) { + t.Fatalf("output = %#v, want legacy rejected result", out) + } +} + +func TestModelOutputExecutionGuardMiddlewarePassesNormalToolCall(t *testing.T) { + called := false + wrapped := modelOutputExecutionGuardMiddleware().Invokable(func(context.Context, *compose.ToolInput) (*compose.ToolOutput, error) { + called = true + return &compose.ToolOutput{Result: "executed"}, nil + }) + + out, err := wrapped(context.Background(), &compose.ToolInput{Name: "exec", Arguments: `{"command":"pwd"}`}) + if err != nil { + t.Fatalf("guard returned error: %v", err) + } + if !called || out == nil || out.Result != "executed" { + t.Fatalf("called=%v output=%#v, want normal execution", called, out) + } +} + +func TestModelOutputExecutionGuardMiddlewareBlocksLegacyRecoveryMarkerStream(t *testing.T) { + markerJSON := `{"` + modelOutputRecoveryKey + `":{"reason":"shell_command_too_large","repair_attempt":1}}` + wrapped := modelOutputExecutionGuardMiddleware().Streamable(func(context.Context, *compose.ToolInput) (*compose.StreamToolOutput, error) { + t.Fatal("legacy recovery marker should not reach the stream endpoint") + return nil, nil + }) + + out, err := wrapped(context.Background(), &compose.ToolInput{Name: "execute", Arguments: markerJSON}) + if err != nil { + t.Fatalf("guard returned error: %v", err) + } + if out == nil || out.Result == nil { + t.Fatal("expected stream output") + } + got, recvErr := out.Result.Recv() + if recvErr != nil && recvErr != io.EOF { + t.Fatalf("recv: %v", recvErr) + } + if !strings.HasPrefix(got, modelOutputRejectedResultPrefix) { + t.Fatalf("stream output = %q, want legacy rejected result", got) + } +} diff --git a/internal/multiagent/no_nested_task.go b/internal/multiagent/no_nested_task.go index d6cb63aa..2deeaaef 100644 --- a/internal/multiagent/no_nested_task.go +++ b/internal/multiagent/no_nested_task.go @@ -6,6 +6,7 @@ import ( "github.com/cloudwego/eino/adk" "github.com/cloudwego/eino/components/tool" + "github.com/cloudwego/eino/schema" ) // noNestedTaskMiddleware 禁止在已经处于 task(sub-agent) 执行链中再次调用 task, @@ -23,10 +24,36 @@ func newNoNestedTaskMiddleware() adk.ChatModelAgentMiddleware { return &noNestedTaskMiddleware{} } +type noNestedAgenticTaskMiddleware struct { + *adk.TypedBaseChatModelAgentMiddleware[*schema.AgenticMessage] +} + +func newNoNestedAgenticTaskMiddleware() adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage] { + return &noNestedAgenticTaskMiddleware{ + TypedBaseChatModelAgentMiddleware: &adk.TypedBaseChatModelAgentMiddleware[*schema.AgenticMessage]{}, + } +} + func (m *noNestedTaskMiddleware) WrapInvokableToolCall( ctx context.Context, endpoint adk.InvokableToolCallEndpoint, tCtx *adk.ToolContext, +) (adk.InvokableToolCallEndpoint, error) { + return wrapNoNestedTaskCall(ctx, endpoint, tCtx) +} + +func (m *noNestedAgenticTaskMiddleware) WrapInvokableToolCall( + ctx context.Context, + endpoint adk.InvokableToolCallEndpoint, + tCtx *adk.ToolContext, +) (adk.InvokableToolCallEndpoint, error) { + return wrapNoNestedTaskCall(ctx, endpoint, tCtx) +} + +func wrapNoNestedTaskCall( + ctx context.Context, + endpoint adk.InvokableToolCallEndpoint, + tCtx *adk.ToolContext, ) (adk.InvokableToolCallEndpoint, error) { if tCtx == nil || strings.TrimSpace(tCtx.Name) == "" { return endpoint, nil diff --git a/internal/multiagent/plan_execute_executor.go b/internal/multiagent/plan_execute_executor.go index 5bb133ae..947850e2 100644 --- a/internal/multiagent/plan_execute_executor.go +++ b/internal/multiagent/plan_execute_executor.go @@ -6,29 +6,29 @@ import ( "github.com/cloudwego/eino/adk" "github.com/cloudwego/eino/adk/prebuilt/planexecute" + "github.com/cloudwego/eino/components/model" + "github.com/cloudwego/eino/schema" ) -// newPlanExecuteExecutor builds the Plan-Execute Executor as an Eino ChatModelAgent. -// -// Eino's planexecute.Config accepts any adk.Agent as Executor; this implementation -// keeps the official Executor contract (Plan/UserInput/ExecutedSteps session keys -// and ExecutedStepSessionKey output) while using ChatModelAgentConfig.Handlers so -// the executor can run the same ADK middleware stack as Deep/Supervisor. As of -// Eino v0.9.12/v0.10.0-alpha.10, planexecute.NewExecutor still does not expose a -// Handlers field, so this custom Executor is the best-practice extension point -// that preserves middleware without forking the whole planexecute loop. -func newPlanExecuteExecutor(ctx context.Context, cfg *planexecute.ExecutorConfig, handlers []adk.ChatModelAgentMiddleware) (adk.Agent, error) { +func newPlanExecuteAgenticExecutor( + ctx context.Context, + cfg *planexecute.ExecutorConfig, + agenticModel model.AgenticModel, + handlers []adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage], + modelRetryCfg *adk.TypedModelRetryConfig[*schema.AgenticMessage], + modelFailoverCfg *adk.ModelFailoverConfig[*schema.AgenticMessage], +) (adk.Agent, error) { if cfg == nil { return nil, fmt.Errorf("plan_execute: ExecutorConfig 为空") } - if cfg.Model == nil { - return nil, fmt.Errorf("plan_execute: Executor Model 为空") + if agenticModel == nil { + return nil, fmt.Errorf("plan_execute: Executor AgenticModel 为空") } genInputFn := cfg.GenInputFn if genInputFn == nil { genInputFn = planExecuteDefaultGenExecutorInput } - genInput := func(ctx context.Context, instruction string, _ *adk.AgentInput) ([]adk.Message, error) { + genInput := func(ctx context.Context, instruction string, _ *adk.TypedAgentInput[*schema.AgenticMessage]) ([]*schema.AgenticMessage, error) { plan, ok := adk.GetSessionValue(ctx, planexecute.PlanSessionKey) if !ok { return nil, fmt.Errorf("plan_execute executor: session value %q missing (possible session corruption)", planexecute.PlanSessionKey) @@ -61,22 +61,29 @@ func newPlanExecuteExecutor(ctx context.Context, cfg *planexecute.ExecutorConfig Plan: plan_, ExecutedSteps: executedSteps_, } - return genInputFn(ctx, in) + msgs, err := genInputFn(ctx, in) + if err != nil { + return nil, err + } + if instruction != "" { + msgs = normalizeSingleLeadingSystemMessage(msgs, instruction) + } + return EinoMessagesToAgentic(msgs), nil } - agentCfg := &adk.ChatModelAgentConfig{ - Name: "executor", - Description: "an executor agent", - Model: cfg.Model, - ToolsConfig: cfg.ToolsConfig, - GenModelInput: genInput, - MaxIterations: cfg.MaxIterations, - OutputKey: planexecute.ExecutedStepSessionKey, + agentCfg := einoAgenticChatModelAgentConfig{ + Name: "executor", + Description: "an executor agent", + Model: agenticModel, + ToolsConfig: cfg.ToolsConfig, + GenModelInput: genInput, + MaxIterations: cfg.MaxIterations, + OutputKey: planexecute.ExecutedStepSessionKey, + Handlers: handlers, + ModelRetryConfig: modelRetryCfg, + ModelFailoverConfig: modelFailoverCfg, } - if len(handlers) > 0 { - agentCfg.Handlers = handlers - } - return adk.NewChatModelAgent(ctx, agentCfg) + return newEinoAgenticChatModelAgentAdapter(ctx, agentCfg) } // planExecuteDefaultGenExecutorInput 对齐 Eino planexecute.defaultGenExecutorInputFn(包外不可引用默认实现)。 diff --git a/internal/multiagent/plan_execute_executor_handlers_test.go b/internal/multiagent/plan_execute_executor_handlers_test.go index 3252831e..a9c6854e 100644 --- a/internal/multiagent/plan_execute_executor_handlers_test.go +++ b/internal/multiagent/plan_execute_executor_handlers_test.go @@ -9,37 +9,41 @@ import ( "github.com/cloudwego/eino/adk" "github.com/cloudwego/eino/components/tool" + "github.com/cloudwego/eino/schema" ) -type stubChatModelAgentMiddleware struct { - adk.BaseChatModelAgentMiddleware +type stubAgenticChatModelAgentMiddleware struct { + *adk.TypedBaseChatModelAgentMiddleware[*schema.AgenticMessage] tag string } -func stubMW(tag string) adk.ChatModelAgentMiddleware { - return &stubChatModelAgentMiddleware{tag: tag} +func stubAgenticMW(tag string) adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage] { + return &stubAgenticChatModelAgentMiddleware{ + TypedBaseChatModelAgentMiddleware: &adk.TypedBaseChatModelAgentMiddleware[*schema.AgenticMessage]{}, + tag: tag, + } } -func TestBuildPlanExecuteExecutorHandlers_IncludesExecPreMiddlewares(t *testing.T) { +func TestBuildPlanExecuteAgenticExecutorHandlers_IncludesExecPreMiddlewares(t *testing.T) { t.Parallel() - pre := []adk.ChatModelAgentMiddleware{ - stubMW("patch"), - stubMW("reduction"), + pre := []adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage]{ + stubAgenticMW("patch"), + stubAgenticMW("reduction"), } - got, err := buildPlanExecuteExecutorHandlers(context.Background(), &PlanExecuteRootArgs{ - ExecPreMiddlewares: pre, - FilesystemMiddleware: stubMW("filesystem"), - SkillMiddleware: stubMW("skill"), + got, err := buildPlanExecuteAgenticExecutorHandlers(context.Background(), &PlanExecuteRootArgs{ + AgenticExecPreMiddlewares: pre, + AgenticFilesystemMiddleware: stubAgenticMW("filesystem"), + AgenticSkillMiddleware: stubAgenticMW("skill"), }) if err != nil { - t.Fatalf("buildPlanExecuteExecutorHandlers: %v", err) + t.Fatalf("buildPlanExecuteAgenticExecutorHandlers: %v", err) } if len(got) != 4 { t.Fatalf("expected 4 pre-tail handlers (2 pre + fs + skill), got %d", len(got)) } for i, want := range []string{"patch", "reduction", "filesystem", "skill"} { - st, ok := got[i].(*stubChatModelAgentMiddleware) + st, ok := got[i].(*stubAgenticChatModelAgentMiddleware) if !ok || st.tag != want { t.Fatalf("handler[%d]: got %#v want tag %q", i, got[i], want) } @@ -54,9 +58,9 @@ func stubTools(n int) []tool.BaseTool { return out } -func TestBuildPlanExecuteExecutorHandlers_NilArgs(t *testing.T) { +func TestBuildPlanExecuteAgenticExecutorHandlers_NilArgs(t *testing.T) { t.Parallel() - if _, err := buildPlanExecuteExecutorHandlers(context.Background(), nil); err == nil { + if _, err := buildPlanExecuteAgenticExecutorHandlers(context.Background(), nil); err == nil { t.Fatal("expected error for nil args") } } diff --git a/internal/multiagent/runner.go b/internal/multiagent/runner.go index fb919f82..1052e3a8 100644 --- a/internal/multiagent/runner.go +++ b/internal/multiagent/runner.go @@ -5,13 +5,10 @@ import ( "context" "encoding/json" "fmt" - "net" - "net/http" "sort" "strings" "sync" "sync/atomic" - "time" "unicode/utf8" "cyberstrike-ai/internal/agent" @@ -19,16 +16,16 @@ import ( "cyberstrike-ai/internal/config" "cyberstrike-ai/internal/database" "cyberstrike-ai/internal/einomcp" - "cyberstrike-ai/internal/openai" "cyberstrike-ai/internal/project" "cyberstrike-ai/internal/reasoning" "cyberstrike-ai/internal/security" - einoopenai "github.com/cloudwego/eino-ext/components/model/openai" "github.com/cloudwego/eino/adk" "github.com/cloudwego/eino/adk/filesystem" "github.com/cloudwego/eino/adk/prebuilt/deep" "github.com/cloudwego/eino/adk/prebuilt/supervisor" + "github.com/cloudwego/eino/components/model" + "github.com/cloudwego/eino/components/tool" "github.com/cloudwego/eino/compose" "github.com/cloudwego/eino/schema" "go.uber.org/zap" @@ -122,7 +119,7 @@ func RunDeepAgent( }) } - einoLoc, einoSkillMW, einoFSTools, skillsRoot, einoErr := prepareEinoSkills(ctx, appCfg.SkillsDir, ma, logger) + agenticLoc, agenticSkillMW, agenticFSTools, agenticSkillsRoot, einoErr := prepareEinoAgenticSkills(ctx, appCfg.SkillsDir, ma, logger) if einoErr != nil { return nil, einoErr } @@ -156,40 +153,28 @@ func RunDeepAgent( toolInvokeNotify := einomcp.NewToolInvokeNotifyHolder() mainDefs := ag.ToolsForRole(roleTools) - httpClient := &http.Client{ - Timeout: 30 * time.Minute, - Transport: &http.Transport{ - DialContext: (&net.Dialer{ - Timeout: 300 * time.Second, - KeepAlive: 300 * time.Second, - }).DialContext, - MaxIdleConns: 100, - MaxIdleConnsPerHost: 10, - IdleConnTimeout: 90 * time.Second, - TLSHandshakeTimeout: 30 * time.Second, - ResponseHeaderTimeout: 60 * time.Minute, - }, + baseHTTPClient := newEinoBaseHTTPClient() + modelFactory := newEinoOpenAIChatModelFactory(baseHTTPClient, reasoningClient, logger) + agenticModelFactory := newEinoOpenAIAgenticChatModelFactory(baseHTTPClient, reasoningClient, logger) + agenticModelRetryCfg := newEinoAgenticModelRetryConfig(&ma.EinoMiddleware, logger, "multiagent") + agenticModelFailoverCfg, err := newEinoAgenticModelFailoverConfig(ctx, appCfg, &ma.EinoMiddleware, einoModelModeNormal, agenticModelFactory, logger, "multiagent", progress, orchMode, conversationID) + if err != nil { + return nil, err } - - // 若配置为 Claude provider,注入自动桥接 transport,对 Eino 透明走 Anthropic Messages API - httpClient = openai.NewEinoHTTPClient(&appCfg.OpenAI, httpClient) - openai.AttachSummarizationDiagTransport(httpClient, logger) - - maxCompletionTokens := appCfg.OpenAI.MaxCompletionTokensEffective() - baseModelCfg := &einoopenai.ChatModelConfig{ - APIKey: appCfg.OpenAI.APIKey, - BaseURL: strings.TrimSuffix(appCfg.OpenAI.BaseURL, "/"), - Model: appCfg.OpenAI.Model, - HTTPClient: httpClient, - MaxCompletionTokens: &maxCompletionTokens, - } - reasoning.ApplyToEinoChatModelConfig(baseModelCfg, &appCfg.OpenAI, reasoningClient) + logEinoAgenticModelGate( + logger, + "multiagent", + orchMode, + evaluateEinoAgenticModelGate(agenticModelGateFactory(agenticModelFactory, appCfg.OpenAI, einoModelModeNormal), einoAgenticRuntimeSupportV0914()), + ) deepMaxIter := agentMaxIterations(appCfg) - var subAgents []adk.Agent + var subAgents []adk.TypedAgent[*schema.AgenticMessage] + var supervisorSubAgents []adk.Agent if orchMode != "plan_execute" { - subAgents = make([]adk.Agent, 0, len(effectiveSubs)) + subAgents = make([]adk.TypedAgent[*schema.AgenticMessage], 0, len(effectiveSubs)) + supervisorSubAgents = make([]adk.Agent, 0, len(effectiveSubs)) for _, sub := range effectiveSubs { id := strings.TrimSpace(sub.ID) if id == "" { @@ -218,11 +203,10 @@ func RunDeepAgent( } } - baseSubModel, err := einoopenai.NewChatModel(ctx, baseModelCfg) + subModel, err := agenticModelFactory(ctx, appCfg.OpenAI, einoModelModeNormal) if err != nil { - return nil, fmt.Errorf("子代理 %q ChatModel: %w", id, err) + return nil, fmt.Errorf("子代理 %q AgenticModel: %w", id, err) } - subModel := newStreamToolCallIndexRepairModel(baseSubModel) subDefs := ag.ToolsForRole(roleTools) subTools, err := einomcp.ToolsFromDefinitions(ag, holder, subDefs, recorder, nil, toolInvokeNotify, id) @@ -230,41 +214,41 @@ func RunDeepAgent( return nil, fmt.Errorf("子代理 %q 工具: %w", id, err) } - subToolsForCfg, subPre, subToolSearchActive, err := prependEinoMiddlewares(ctx, &ma.EinoMiddleware, einoMWSub, subTools, einoLoc, skillsRoot, conversationID, projectID, logger) + subToolsForCfg, subPre, subToolSearchActive, err := prependEinoAgenticMiddlewares(ctx, &ma.EinoMiddleware, einoMWSub, subTools, agenticLoc, agenticSkillsRoot, conversationID, projectID, logger) if err != nil { return nil, fmt.Errorf("子代理 %q eino 中间件: %w", id, err) } subMax := resolveMaxIterations(appCfg, sub.MaxIterations) - subSumMw, err := newEinoSummarizationMiddleware(ctx, subModel, appCfg, &ma.EinoMiddleware, conversationID, db, projectID, logger) + subSumMw, err := newEinoAgenticSummarizationMiddleware(ctx, subModel, appCfg, &ma.EinoMiddleware, conversationID, db, projectID, logger) if err != nil { - return nil, fmt.Errorf("子代理 %q summarization 中间件: %w", id, err) + return nil, fmt.Errorf("子代理 %q agentic summarization 中间件: %w", id, err) } - var subHandlers []adk.ChatModelAgentMiddleware + var subHandlers []adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage] if len(subPre) > 0 { subHandlers = append(subHandlers, subPre...) } - if einoSkillMW != nil { - if einoFSTools && einoLoc != nil { - subFs, fsErr := subAgentFilesystemMiddleware(ctx, einoLoc, toolInvokeNotify, id, einoExecBegin, einoExecAppendPartial, einoExecRegisterCancel, einoExecUnregisterCancel, einoExecFinish, agentToolTimeoutMinutes(appCfg), agentToolWaitTimeoutSeconds(appCfg), agentShellNoOutputTimeoutSeconds(appCfg), nil) + if agenticSkillMW != nil { + if agenticFSTools && agenticLoc != nil { + subFs, fsErr := subAgentAgenticFilesystemMiddleware(ctx, agenticLoc, toolInvokeNotify, id, einoExecBegin, einoExecAppendPartial, einoExecRegisterCancel, einoExecUnregisterCancel, einoExecFinish, agentToolTimeoutMinutes(appCfg), agentToolWaitTimeoutSeconds(appCfg), agentShellNoOutputTimeoutSeconds(appCfg), nil) if fsErr != nil { return nil, fmt.Errorf("子代理 %q filesystem 中间件: %w", id, fsErr) } subHandlers = append(subHandlers, subFs) } - subHandlers = append(subHandlers, einoSkillMW) + subHandlers = append(subHandlers, agenticSkillMW) } - subHandlers = appendEinoChatModelTailMiddlewares(subHandlers, einoChatModelTailConfig{ - logger: logger, - phase: "sub_agent:" + id, - summarization: subSumMw, - modelName: appCfg.OpenAI.Model, - maxTotalTokens: appCfg.OpenAI.MaxTotalTokens, - toolMaxBytes: toolMaxBytesFromMW(&ma.EinoMiddleware), - conversationID: conversationID, - middlewareConfig: &ma.EinoMiddleware, + subHandlers = appendEinoAgenticChatModelTailMiddlewares(subHandlers, einoChatModelTailConfig{ + logger: logger, + phase: "sub_agent:" + id, + agenticSummarization: subSumMw, + modelName: appCfg.OpenAI.Model, + maxTotalTokens: appCfg.OpenAI.MaxTotalTokens, + toolMaxBytes: toolMaxBytesFromMW(&ma.EinoMiddleware), + conversationID: conversationID, + middlewareConfig: &ma.EinoMiddleware, }) subInstrFinal := project.AppendVisionImageAnalysisIfReady(instr, appCfg.Vision.Ready()) @@ -280,11 +264,11 @@ func RunDeepAgent( zap.Bool("tool_search_middleware", subToolSearchActive), ) } - sa, err := adk.NewChatModelAgent(ctx, &adk.ChatModelAgentConfig{ + sa, err := newEinoAgenticChatModelAgent(ctx, einoAgenticChatModelAgentConfig{ Name: id, Description: desc, Instruction: subInstrFinal, - GenModelInput: literalInstructionGenModelInput, + GenModelInput: literalAgenticInstructionGenModelInput, Model: subModel, ToolsConfig: adk.ToolsConfig{ ToolsNodeConfig: compose.ToolsNodeConfig{ @@ -299,27 +283,21 @@ func RunDeepAgent( }, EmitInternalEvents: true, }, - MaxIterations: subMax, - Handlers: subHandlers, + MaxIterations: subMax, + Handlers: subHandlers, + ModelRetryConfig: agenticModelRetryCfg, + ModelFailoverConfig: agenticModelFailoverCfg, }) if err != nil { return nil, fmt.Errorf("子代理 %q: %w", id, err) } subAgents = append(subAgents, sa) + if adapted := newEinoAgenticMessageAgentAdapter(sa); adapted != nil { + supervisorSubAgents = append(supervisorSubAgents, adapted) + } } } - baseMainModel, err := einoopenai.NewChatModel(ctx, baseModelCfg) - if err != nil { - return nil, fmt.Errorf("多代理主模型: %w", err) - } - mainModel := newStreamToolCallIndexRepairModel(baseMainModel) - - mainSumMw, err := newEinoSummarizationMiddleware(ctx, mainModel, appCfg, &ma.EinoMiddleware, conversationID, db, projectID, logger) - if err != nil { - return nil, fmt.Errorf("多代理主 summarization 中间件: %w", err) - } - modelFacingTrace := newModelFacingTraceHolder() // 与 deep.Config.Name / supervisor 主代理 Name 一致。 @@ -346,7 +324,10 @@ func RunDeepAgent( if err != nil { return nil, err } - mainToolsForCfg, mainOrchestratorPre, mainToolSearchActive, err := prependEinoMiddlewares(ctx, &ma.EinoMiddleware, einoMWMain, mainTools, einoLoc, skillsRoot, conversationID, projectID, logger) + var mainToolsForCfg []tool.BaseTool + var mainToolSearchActive bool + var mainAgenticOrchestratorPre []adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage] + mainToolsForCfg, mainAgenticOrchestratorPre, mainToolSearchActive, err = prependEinoAgenticMiddlewares(ctx, &ma.EinoMiddleware, einoMWMain, mainTools, agenticLoc, agenticSkillsRoot, conversationID, projectID, logger) if err != nil { return nil, err } @@ -388,8 +369,8 @@ func RunDeepAgent( var deepBackend filesystem.Backend var deepShell filesystem.StreamingShell - if einoLoc != nil && einoFSTools { - deepBackend = einoLoc + if agenticLoc != nil && agenticFSTools { + deepBackend = agenticLoc deepShell = &einoStreamingShellWrap{ inner: security.NewEinoStreamingShell(), invokeNotify: toolInvokeNotify, @@ -406,8 +387,21 @@ func RunDeepAgent( } } + var mainModel model.AgenticModel + var mainSumMw adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage] + if orchMode != "plan_execute" { + mainModel, err = agenticModelFactory(ctx, appCfg.OpenAI, einoModelModeNormal) + if err != nil { + return nil, fmt.Errorf("多代理主 AgenticModel: %w", err) + } + mainSumMw, err = newEinoAgenticSummarizationMiddleware(ctx, mainModel, appCfg, &ma.EinoMiddleware, conversationID, db, projectID, logger) + if err != nil { + return nil, fmt.Errorf("多代理主 agentic summarization 中间件: %w", err) + } + } + // noNestedTaskMiddleware 必须在最外层(最先拦截),防止 skill 或其他中间件内部触发 task 调用绕过检测。 - deepHandlers := []adk.ChatModelAgentMiddleware{newNoNestedTaskMiddleware()} + deepHandlers := []adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage]{newNoNestedAgenticTaskMiddleware()} var taskBlackboardSupplement string if appCfg.Project.Enabled && db != nil { if pid := strings.TrimSpace(projectID); pid != "" { @@ -416,44 +410,44 @@ func RunDeepAgent( } } } - if mw := newTaskContextEnrichMiddleware(runtimeUserMessage, history, ma.SubAgentUserContextMaxRunesEffective(), taskBlackboardSupplement); mw != nil { + if mw := newAgenticTaskContextEnrichMiddleware(runtimeUserMessage, history, ma.SubAgentUserContextMaxRunesEffective(), taskBlackboardSupplement); mw != nil { deepHandlers = append(deepHandlers, mw) } - if len(mainOrchestratorPre) > 0 { - deepHandlers = append(deepHandlers, mainOrchestratorPre...) + if len(mainAgenticOrchestratorPre) > 0 { + deepHandlers = append(deepHandlers, mainAgenticOrchestratorPre...) } - if einoSkillMW != nil { - deepHandlers = append(deepHandlers, einoSkillMW) + if agenticSkillMW != nil { + deepHandlers = append(deepHandlers, agenticSkillMW) } - deepHandlers = appendEinoChatModelTailMiddlewares(deepHandlers, einoChatModelTailConfig{ - logger: logger, - phase: "deep_orchestrator", - summarization: mainSumMw, - modelName: appCfg.OpenAI.Model, - maxTotalTokens: appCfg.OpenAI.MaxTotalTokens, - toolMaxBytes: toolMaxBytesFromMW(&ma.EinoMiddleware), - conversationID: conversationID, - trace: modelFacingTrace, - middlewareConfig: &ma.EinoMiddleware, + deepHandlers = appendEinoAgenticChatModelTailMiddlewares(deepHandlers, einoChatModelTailConfig{ + logger: logger, + phase: "deep_orchestrator", + agenticSummarization: mainSumMw, + modelName: appCfg.OpenAI.Model, + maxTotalTokens: appCfg.OpenAI.MaxTotalTokens, + toolMaxBytes: toolMaxBytesFromMW(&ma.EinoMiddleware), + conversationID: conversationID, + trace: modelFacingTrace, + middlewareConfig: &ma.EinoMiddleware, }) - supHandlers := []adk.ChatModelAgentMiddleware{} - if len(mainOrchestratorPre) > 0 { - supHandlers = append(supHandlers, mainOrchestratorPre...) + supHandlers := []adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage]{} + if len(mainAgenticOrchestratorPre) > 0 { + supHandlers = append(supHandlers, mainAgenticOrchestratorPre...) } - if einoSkillMW != nil { - supHandlers = append(supHandlers, einoSkillMW) + if agenticSkillMW != nil { + supHandlers = append(supHandlers, agenticSkillMW) } - supHandlers = appendEinoChatModelTailMiddlewares(supHandlers, einoChatModelTailConfig{ - logger: logger, - phase: "supervisor_orchestrator", - summarization: mainSumMw, - modelName: appCfg.OpenAI.Model, - maxTotalTokens: appCfg.OpenAI.MaxTotalTokens, - toolMaxBytes: toolMaxBytesFromMW(&ma.EinoMiddleware), - conversationID: conversationID, - trace: modelFacingTrace, - middlewareConfig: &ma.EinoMiddleware, + supHandlers = appendEinoAgenticChatModelTailMiddlewares(supHandlers, einoChatModelTailConfig{ + logger: logger, + phase: "supervisor_orchestrator", + agenticSummarization: mainSumMw, + modelName: appCfg.OpenAI.Model, + maxTotalTokens: appCfg.OpenAI.MaxTotalTokens, + toolMaxBytes: toolMaxBytesFromMW(&ma.EinoMiddleware), + conversationID: conversationID, + trace: modelFacingTrace, + middlewareConfig: &ma.EinoMiddleware, }) mainToolsCfg := adk.ToolsConfig{ @@ -470,45 +464,42 @@ func RunDeepAgent( EmitInternalEvents: true, } - deepOutKey, taskGen := deepExtrasFromConfig(ma) + deepAgenticOutKey, agenticTaskGen := deepAgenticExtrasFromConfig(ma) var da adk.Agent switch orchMode { case "plan_execute": - plannerModelCfg := &einoopenai.ChatModelConfig{ - APIKey: appCfg.OpenAI.APIKey, - BaseURL: strings.TrimSuffix(appCfg.OpenAI.BaseURL, "/"), - Model: appCfg.OpenAI.Model, - HTTPClient: httpClient, - MaxCompletionTokens: &maxCompletionTokens, - } - reasoning.ApplyPlanExecutePlannerModelConfig(plannerModelCfg, &appCfg.OpenAI) - basePEMainModel, perr := einoopenai.NewChatModel(ctx, plannerModelCfg) + peMainModel, perr := modelFactory(ctx, appCfg.OpenAI, einoModelModePlanner) if perr != nil { return nil, fmt.Errorf("plan_execute 规划模型: %w", perr) } - peMainModel := newStreamToolCallIndexRepairModel(basePEMainModel) if logger != nil { logger.Info("plan_execute: planner/replanner 使用无 reasoning 的独立 ChatModel(ToolChoiceForced 兼容)", zap.String("model", appCfg.OpenAI.Model), ) } - baseExecModel, perr := einoopenai.NewChatModel(ctx, baseModelCfg) + execModel, perr := modelFactory(ctx, appCfg.OpenAI, einoModelModeNormal) if perr != nil { return nil, fmt.Errorf("plan_execute 执行器模型: %w", perr) } - execModel := newStreamToolCallIndexRepairModel(baseExecModel) - // 构建 filesystem 中间件(与 Deep sub-agent 一致) - var peFsMw adk.ChatModelAgentMiddleware - if einoSkillMW != nil && einoFSTools && einoLoc != nil { - peFsMw, err = subAgentFilesystemMiddleware(ctx, einoLoc, toolInvokeNotify, "executor", einoExecBegin, einoExecAppendPartial, einoExecRegisterCancel, einoExecUnregisterCancel, einoExecFinish, agentToolTimeoutMinutes(appCfg), agentToolWaitTimeoutSeconds(appCfg), agentShellNoOutputTimeoutSeconds(appCfg), nil) + agenticExecModel, perr := agenticModelFactory(ctx, appCfg.OpenAI, einoModelModeNormal) + if perr != nil { + return nil, fmt.Errorf("plan_execute 执行器 AgenticModel: %w", perr) + } + planRewriteSumMw, perr := newEinoSummarizationMiddleware(ctx, execModel, appCfg, &ma.EinoMiddleware, conversationID, db, projectID, logger) + if perr != nil { + return nil, fmt.Errorf("plan_execute planner/replanner summarization: %w", perr) + } + var peFsMw adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage] + if agenticSkillMW != nil && agenticFSTools && agenticLoc != nil { + peFsMw, err = subAgentAgenticFilesystemMiddleware(ctx, agenticLoc, toolInvokeNotify, "executor", einoExecBegin, einoExecAppendPartial, einoExecRegisterCancel, einoExecUnregisterCancel, einoExecFinish, agentToolTimeoutMinutes(appCfg), agentToolWaitTimeoutSeconds(appCfg), agentShellNoOutputTimeoutSeconds(appCfg), nil) if err != nil { - return nil, fmt.Errorf("plan_execute filesystem 中间件: %w", err) + return nil, fmt.Errorf("plan_execute agentic filesystem 中间件: %w", err) } } peRoot, perr := NewPlanExecuteRoot(ctx, &PlanExecuteRootArgs{ MainToolCallingModel: peMainModel, - ExecModel: execModel, + AgenticExecModel: agenticExecModel, OrchInstruction: orchInstruction, ToolsCfg: mainToolsCfg, ExecMaxIter: deepMaxIter, @@ -520,15 +511,15 @@ func RunDeepAgent( ProjectID: projectID, Logger: logger, ModelName: appCfg.OpenAI.Model, - // 与 Deep/Supervisor 主代理同源:patch / reduction / toolsearch / plantask(见 buildPlanExecuteExecutorHandlers)。 - ExecPreMiddlewares: mainOrchestratorPre, - SkillMiddleware: einoSkillMW, - FilesystemMiddleware: peFsMw, - ModelFacingTrace: modelFacingTrace, + // 与 Deep/Supervisor 主代理同源:typed patch / reduction / toolsearch / plantask(见 buildPlanExecuteAgenticExecutorHandlers)。 + AgenticExecPreMiddlewares: mainAgenticOrchestratorPre, + AgenticSkillMiddleware: agenticSkillMW, + AgenticFilesystemMiddleware: peFsMw, + ModelFacingTrace: modelFacingTrace, PlannerReplannerRewriteHandlers: appendEinoChatModelTailMiddlewares(nil, einoChatModelTailConfig{ logger: logger, phase: "plan_execute_planner_replanner", - summarization: mainSumMw, + summarization: planRewriteSumMw, modelName: appCfg.OpenAI.Model, maxTotalTokens: appCfg.OpenAI.MaxTotalTokens, toolMaxBytes: toolMaxBytesFromMW(&ma.EinoMiddleware), @@ -536,40 +527,44 @@ func RunDeepAgent( skipTrace: true, middlewareConfig: &ma.EinoMiddleware, }), + AgenticModelRetryConfig: agenticModelRetryCfg, + AgenticModelFailoverConfig: agenticModelFailoverCfg, }) if perr != nil { return nil, perr } da = peRoot case "supervisor": - supCfg := &adk.ChatModelAgentConfig{ - Name: orchestratorName, - Description: orchDescription, - Instruction: supInstr, - GenModelInput: literalInstructionGenModelInput, - Model: mainModel, - ToolsConfig: mainToolsCfg, - MaxIterations: deepMaxIter, - Handlers: supHandlers, - Exit: &adk.ExitTool{}, + supCfg := einoAgenticChatModelAgentConfig{ + Name: orchestratorName, + Description: orchDescription, + Instruction: supInstr, + GenModelInput: literalAgenticInstructionGenModelInput, + Model: mainModel, + ToolsConfig: mainToolsCfg, + MaxIterations: deepMaxIter, + Handlers: supHandlers, + Exit: &adk.ExitTool{}, + ModelRetryConfig: agenticModelRetryCfg, + ModelFailoverConfig: agenticModelFailoverCfg, } - if deepOutKey != "" { - supCfg.OutputKey = deepOutKey + if deepAgenticOutKey != "" { + supCfg.OutputKey = deepAgenticOutKey } - superChat, serr := adk.NewChatModelAgent(ctx, supCfg) + superChat, serr := newEinoAgenticChatModelAgentAdapter(ctx, supCfg) if serr != nil { - return nil, fmt.Errorf("supervisor 主代理: %w", serr) + return nil, fmt.Errorf("supervisor agentic 主代理: %w", serr) } supRoot, serr := supervisor.New(ctx, &supervisor.Config{ Supervisor: superChat, - SubAgents: subAgents, + SubAgents: supervisorSubAgents, }) if serr != nil { return nil, fmt.Errorf("supervisor.New: %w", serr) } da = supRoot default: - dcfg := &deep.Config{ + dcfg := &deep.TypedConfig[*schema.AgenticMessage]{ Name: orchestratorName, Description: orchDescription, ChatModel: mainModel, @@ -582,18 +577,20 @@ func RunDeepAgent( StreamingShell: deepShell, Handlers: deepHandlers, ToolsConfig: mainToolsCfg, + ModelRetryConfig: agenticModelRetryCfg, + ModelFailoverConfig: agenticModelFailoverCfg, } - if deepOutKey != "" { - dcfg.OutputKey = deepOutKey + if deepAgenticOutKey != "" { + dcfg.OutputKey = deepAgenticOutKey } - if taskGen != nil { - dcfg.TaskToolDescriptionGenerator = taskGen + if agenticTaskGen != nil { + dcfg.TaskToolDescriptionGenerator = agenticTaskGen } - dDeep, derr := deep.New(ctx, dcfg) + dDeep, derr := deep.NewTyped[*schema.AgenticMessage](ctx, dcfg) if derr != nil { - return nil, fmt.Errorf("deep.New: %w", derr) + return nil, fmt.Errorf("deep.NewTyped[AgenticMessage]: %w", derr) } - da = dDeep + da = newEinoAgenticMessageAgentAdapter(dDeep) } baseMsgs := historyToMessages(history, appCfg, &ma.EinoMiddleware) @@ -625,8 +622,8 @@ func RunDeepAgent( StreamsMainAssistant: streamsMainAssistant, EinoRoleTag: einoRoleTag, CheckpointDir: ma.EinoMiddleware.CheckpointDir, - RunRetryMaxAttempts: ma.EinoMiddleware.RunRetryMaxAttempts, - RunRetryMaxBackoffSec: ma.EinoMiddleware.RunRetryMaxBackoffSec, + RunRetryMaxAttempts: RunRetryMaxAttemptsFromConfig(&ma.EinoMiddleware), + RunRetryMaxBackoffSec: int(einoRunRetryMaxBackoffFromConfig(&ma.EinoMiddleware).Seconds()), McpIDsMu: &mcpIDsMu, McpIDs: &mcpIDs, FilesystemMonitorAgent: ag, @@ -639,6 +636,7 @@ func RunDeepAgent( MaxTotalTokens: appCfg.OpenAI.MaxTotalTokens, ToolMaxBytes: toolMaxBytesFromMW(&ma.EinoMiddleware), ModelName: appCfg.OpenAI.Model, + MiddlewareConfig: &ma.EinoMiddleware, EmptyResponseMessage: "(Eino multi-agent orchestration completed but no assistant text was captured. Check process details or logs.) " + "(Eino 多代理编排已完成,但未捕获到助手文本输出。请查看过程详情或日志。)", }, baseMsgs) @@ -819,7 +817,9 @@ func toolCallStableID(tc schema.ToolCall) string { return "" } -// toolCallDisplayName 避免前端「未知工具」:DeepAgent 内置 task 等可能延迟写入 function.name。 +// toolCallDisplayName returns the visible tool name once the model stream has +// produced a concrete function name. Anonymous stream fragments are filtered +// before progress emission instead of being guessed as task calls. func toolCallDisplayName(tc schema.ToolCall) string { if n := strings.TrimSpace(tc.Function.Name); n != "" { return n @@ -827,7 +827,7 @@ func toolCallDisplayName(tc schema.ToolCall) string { if n := strings.TrimSpace(tc.Type); n != "" && !strings.EqualFold(n, "function") { return n } - return "task" + return "" } // toolCallsSignatureFlush 用于去重键;无 id/index 时用占位 pos,避免流末帧缺 id 时整条工具事件丢失。 @@ -835,13 +835,24 @@ func toolCallsSignatureFlush(msg *schema.Message) string { if msg == nil || len(msg.ToolCalls) == 0 { return "" } - parts := make([]string, 0, len(msg.ToolCalls)) - for i, tc := range msg.ToolCalls { + visible := filterVisibleToolCallsForProgress(msg.ToolCalls) + if len(visible) == 0 { + return "" + } + parts := make([]string, 0, len(visible)) + for i, tc := range visible { id := toolCallStableID(tc) if id == "" { id = fmt.Sprintf("pos:%d", i) } - parts = append(parts, id+"|"+toolCallDisplayName(tc)) + name := toolCallDisplayName(tc) + if name == "" { + continue + } + parts = append(parts, id+"|"+name) + } + if len(parts) == 0 { + return "" } sort.Strings(parts) return strings.Join(parts, ";") @@ -853,8 +864,9 @@ func toolCallsRichSignature(msg *schema.Message) string { if base == "" { return "" } - parts := make([]string, 0, len(msg.ToolCalls)) - for _, tc := range msg.ToolCalls { + visible := filterVisibleToolCallsForProgress(msg.ToolCalls) + parts := make([]string, 0, len(visible)) + for _, tc := range visible { id := toolCallStableID(tc) arg := tc.Function.Arguments if len(arg) > 240 { @@ -909,6 +921,10 @@ func emitToolCallsFromMessage( if msg == nil || len(msg.ToolCalls) == 0 || progress == nil { return } + visibleToolCalls := filterVisibleToolCallsForProgress(msg.ToolCalls) + if len(visibleToolCalls) == 0 { + return + } if subAgentToolStep == nil { subAgentToolStep = make(map[string]int) } @@ -945,14 +961,14 @@ func emitToolCallsFromMessage( if isSubToolRound { role = "sub" } - progress("tool_calls_detected", fmt.Sprintf("检测到 %d 个工具调用", len(msg.ToolCalls)), map[string]interface{}{ - "count": len(msg.ToolCalls), + progress("tool_calls_detected", fmt.Sprintf("检测到 %d 个工具调用", len(visibleToolCalls)), map[string]interface{}{ + "count": len(visibleToolCalls), "conversationId": conversationID, "source": "eino", "einoAgent": agentName, "einoRole": role, }) - for idx, tc := range msg.ToolCalls { + for idx, tc := range visibleToolCalls { argStr := strings.TrimSpace(tc.Function.Arguments) if argStr == "" && len(tc.Extra) > 0 { if b, mErr := json.Marshal(tc.Extra); mErr == nil { @@ -973,8 +989,7 @@ func emitToolCallsFromMessage( // with an earlier batch in the same agent run. toolCallID = fmt.Sprintf("eino-stream-%d-%d", fallbackToolCallSequence.Add(1), *tc.Index) } - // Record pending tool calls for later tool_result correlation / recovery flushing. - // We intentionally record even for unknown tools to avoid "running" badge getting stuck. + // Record visible pending tool calls for later tool_result correlation / recovery flushing. if markPending != nil && toolCallID != "" { markPending(toolCallPendingInfo{ ToolCallID: toolCallID, @@ -989,7 +1004,7 @@ func emitToolCallsFromMessage( "argumentsObj": argsObj, "toolCallId": toolCallID, "index": idx + 1, - "total": len(msg.ToolCalls), + "total": len(visibleToolCalls), "conversationId": conversationID, "source": "eino", "einoAgent": agentName, @@ -998,6 +1013,23 @@ func emitToolCallsFromMessage( } } +func filterVisibleToolCallsForProgress(calls []schema.ToolCall) []schema.ToolCall { + if len(calls) == 0 { + return nil + } + out := make([]schema.ToolCall, 0, len(calls)) + for _, tc := range calls { + if _, ok := modelOutputRecoveryFromToolCall(tc); ok { + continue + } + if toolCallDisplayName(tc) == "" { + continue + } + out = append(out, tc) + } + return out +} + // dedupeRepeatedParagraphs 去掉完全相同的连续/重复段落,缓解多代理各自复述同一列表。 func dedupeRepeatedParagraphs(s string, minLen int) string { if s == "" || minLen <= 0 { diff --git a/internal/multiagent/sub_agent_context.go b/internal/multiagent/sub_agent_context.go index 22802d7e..d7fba52b 100644 --- a/internal/multiagent/sub_agent_context.go +++ b/internal/multiagent/sub_agent_context.go @@ -10,6 +10,7 @@ import ( "github.com/cloudwego/eino/adk" "github.com/cloudwego/eino/components/tool" + "github.com/cloudwego/eino/schema" ) const userContextSupplementHeader = "\n\n## 用户历史输入(原文,子代理必读)\n" @@ -47,10 +48,54 @@ func newTaskContextEnrichMiddleware(userMessage string, history []agent.ChatMess return &taskContextEnrichMiddleware{supplement: supplement} } +type agenticTaskContextEnrichMiddleware struct { + *adk.TypedBaseChatModelAgentMiddleware[*schema.AgenticMessage] + supplement string +} + +func newAgenticTaskContextEnrichMiddleware(userMessage string, history []agent.ChatMessage, maxRunes int, projectBlackboard string) adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage] { + supplement := buildUserContextSupplement(userMessage, history, maxRunes) + if bb := strings.TrimSpace(projectBlackboard); bb != "" { + if supplement != "" { + supplement += "\n\n" + bb + } else { + supplement = "\n\n" + bb + } + } + if supplement == "" { + return nil + } + return &agenticTaskContextEnrichMiddleware{ + TypedBaseChatModelAgentMiddleware: &adk.TypedBaseChatModelAgentMiddleware[*schema.AgenticMessage]{}, + supplement: supplement, + } +} + func (m *taskContextEnrichMiddleware) WrapInvokableToolCall( ctx context.Context, endpoint adk.InvokableToolCallEndpoint, tCtx *adk.ToolContext, +) (adk.InvokableToolCallEndpoint, error) { + return wrapTaskContextEnrichCall(m, ctx, endpoint, tCtx) +} + +func (m *agenticTaskContextEnrichMiddleware) WrapInvokableToolCall( + ctx context.Context, + endpoint adk.InvokableToolCallEndpoint, + tCtx *adk.ToolContext, +) (adk.InvokableToolCallEndpoint, error) { + return wrapTaskContextEnrichCall(m, ctx, endpoint, tCtx) +} + +type taskContextEnricher interface { + enrichTaskDescription(argsJSON string) string +} + +func wrapTaskContextEnrichCall( + m taskContextEnricher, + ctx context.Context, + endpoint adk.InvokableToolCallEndpoint, + tCtx *adk.ToolContext, ) (adk.InvokableToolCallEndpoint, error) { if tCtx == nil || !strings.EqualFold(strings.TrimSpace(tCtx.Name), "task") { return endpoint, nil @@ -65,6 +110,14 @@ func (m *taskContextEnrichMiddleware) WrapInvokableToolCall( // to the "description" field, and re-serializes. Falls back to the original // JSON if parsing fails or no description field exists. func (m *taskContextEnrichMiddleware) enrichTaskDescription(argsJSON string) string { + return enrichTaskDescriptionWithSupplement(argsJSON, m.supplement) +} + +func (m *agenticTaskContextEnrichMiddleware) enrichTaskDescription(argsJSON string) string { + return enrichTaskDescriptionWithSupplement(argsJSON, m.supplement) +} + +func enrichTaskDescriptionWithSupplement(argsJSON, supplement string) string { var raw map[string]interface{} if err := json.Unmarshal([]byte(argsJSON), &raw); err != nil { return argsJSON @@ -73,7 +126,7 @@ func (m *taskContextEnrichMiddleware) enrichTaskDescription(argsJSON string) str if !ok { return argsJSON } - raw["description"] = desc + m.supplement + raw["description"] = desc + supplement enriched, err := json.Marshal(raw) if err != nil { return argsJSON