Add files via upload

This commit is contained in:
公明
2026-08-15 01:39:58 +08:00
committed by GitHub
parent 4fe6defa28
commit 70b01206e4
47 changed files with 5598 additions and 1191 deletions
File diff suppressed because it is too large Load Diff
@@ -15,7 +15,7 @@ func TestRecvSchemaMessageStream_EOF(t *testing.T) {
_ = sw.Send(schema.ToolMessage("hello", "tc-1"), nil)
sw.Close()
content, tid, err := recvSchemaMessageStream(context.Background(), sr)
content, tid, toolName, err := recvSchemaMessageStream(context.Background(), sr)
if err != nil {
t.Fatalf("unexpected err: %v", err)
}
@@ -25,6 +25,23 @@ func TestRecvSchemaMessageStream_EOF(t *testing.T) {
if tid != "tc-1" {
t.Fatalf("toolCallID=%q want tc-1", tid)
}
if toolName != "" {
t.Fatalf("toolName=%q want empty", toolName)
}
}
func TestRecvSchemaMessageStream_CapturesToolName(t *testing.T) {
sr, sw := schema.Pipe[*schema.Message](4)
_ = sw.Send(schema.ToolMessage("hello", "tc-1", schema.WithToolName("execute")), nil)
sw.Close()
content, tid, toolName, err := recvSchemaMessageStream(context.Background(), sr)
if err != nil {
t.Fatalf("unexpected err: %v", err)
}
if content != "hello" || tid != "tc-1" || toolName != "execute" {
t.Fatalf("content=%q tid=%q toolName=%q", content, tid, toolName)
}
}
func TestRecvSchemaMessageStream_ContextCancel(t *testing.T) {
@@ -37,7 +54,7 @@ func TestRecvSchemaMessageStream_ContextCancel(t *testing.T) {
cancel()
}()
content, _, err := recvSchemaMessageStream(ctx, sr)
content, _, _, err := recvSchemaMessageStream(ctx, sr)
if !errors.Is(err, context.Canceled) {
t.Fatalf("want context.Canceled, got %v content=%q", err, content)
}
@@ -49,16 +66,16 @@ func TestRecvSchemaMessageStream_RecvError(t *testing.T) {
_ = sw.Send(nil, want)
sw.Close()
_, _, err := recvSchemaMessageStream(context.Background(), sr)
_, _, _, err := recvSchemaMessageStream(context.Background(), sr)
if !errors.Is(err, want) {
t.Fatalf("want %v, got %v", want, err)
}
}
func TestRecvSchemaMessageStream_NilStream(t *testing.T) {
content, tid, err := recvSchemaMessageStream(context.Background(), nil)
if err != nil || content != "" || tid != "" {
t.Fatalf("nil stream: content=%q tid=%q err=%v", content, tid, err)
content, tid, toolName, err := recvSchemaMessageStream(context.Background(), nil)
if err != nil || content != "" || tid != "" || toolName != "" {
t.Fatalf("nil stream: content=%q tid=%q toolName=%q err=%v", content, tid, toolName, err)
}
}
@@ -67,8 +84,39 @@ func TestRecvSchemaMessageStream_EOFViaEmptyRead(t *testing.T) {
_ = sw.Send(nil, io.EOF)
sw.Close()
_, _, err := recvSchemaMessageStream(context.Background(), sr)
_, _, _, err := recvSchemaMessageStream(context.Background(), sr)
if err != nil {
t.Fatalf("EOF should not surface as error, got %v", err)
}
}
func TestRecvEinoSchemaMessageStreamWithContext_SkipsNilChunks(t *testing.T) {
sr, sw := schema.Pipe[*schema.Message](4)
_ = sw.Send(nil, nil)
_ = sw.Send(schema.AssistantMessage("hello", nil), nil)
sw.Close()
var got []string
err := recvEinoSchemaMessageStreamWithContext(context.Background(), sr, 1, func(chunk *schema.Message) {
got = append(got, chunk.Content)
})
if err != nil {
t.Fatalf("unexpected err: %v", err)
}
if len(got) != 1 || got[0] != "hello" {
t.Fatalf("chunks = %#v, want [hello]", got)
}
}
func TestRecvEinoSchemaMessageStreamWithContext_NilStream(t *testing.T) {
called := false
err := recvEinoSchemaMessageStreamWithContext(context.Background(), nil, 0, func(*schema.Message) {
called = true
})
if err != nil {
t.Fatalf("nil stream should not error, got %v", err)
}
if called {
t.Fatal("nil stream should not call handler")
}
}
@@ -0,0 +1,81 @@
package multiagent
import (
"context"
"fmt"
"github.com/cloudwego/eino/adk"
"github.com/cloudwego/eino/schema"
)
type einoAgenticMessageAgentAdapter struct {
inner adk.TypedAgent[*schema.AgenticMessage]
}
func newEinoAgenticMessageAgentAdapter(inner adk.TypedAgent[*schema.AgenticMessage]) adk.Agent {
if inner == nil {
return nil
}
return &einoAgenticMessageAgentAdapter{inner: inner}
}
func (a *einoAgenticMessageAgentAdapter) Name(ctx context.Context) string {
if a == nil || a.inner == nil {
return ""
}
return a.inner.Name(ctx)
}
func (a *einoAgenticMessageAgentAdapter) Description(ctx context.Context) string {
if a == nil || a.inner == nil {
return ""
}
return a.inner.Description(ctx)
}
func (a *einoAgenticMessageAgentAdapter) Run(ctx context.Context, input *adk.AgentInput, opts ...adk.AgentRunOption) *adk.AsyncIterator[*adk.AgentEvent] {
return a.runTyped(ctx, input, nil, opts...)
}
func (a *einoAgenticMessageAgentAdapter) Resume(ctx context.Context, info *adk.ResumeInfo, opts ...adk.AgentRunOption) *adk.AsyncIterator[*adk.AgentEvent] {
return a.runTyped(ctx, nil, info, opts...)
}
func (a *einoAgenticMessageAgentAdapter) runTyped(ctx context.Context, input *adk.AgentInput, resumeInfo *adk.ResumeInfo, opts ...adk.AgentRunOption) *adk.AsyncIterator[*adk.AgentEvent] {
iter, gen := adk.NewAsyncIteratorPair[*adk.AgentEvent]()
go func() {
defer gen.Close()
if a == nil || a.inner == nil {
gen.Send(&adk.AgentEvent{Err: fmt.Errorf("agentic adapter: inner agent is nil")})
return
}
var agenticIter *adk.AsyncIterator[*adk.TypedAgentEvent[*schema.AgenticMessage]]
if resumeInfo != nil {
resumable, ok := a.inner.(adk.TypedResumableAgent[*schema.AgenticMessage])
if !ok {
gen.Send(&adk.AgentEvent{Err: fmt.Errorf("agentic adapter: inner agent does not support resume")})
return
}
agenticIter = resumable.Resume(ctx, resumeInfo, opts...)
} else {
agenticInput := &adk.TypedAgentInput[*schema.AgenticMessage]{}
if input != nil {
agenticInput.EnableStreaming = input.EnableStreaming
agenticInput.Messages = EinoMessagesToAgentic(input.Messages)
}
agenticIter = a.inner.Run(ctx, agenticInput, opts...)
}
for {
ev, ok := agenticIter.Next()
if !ok {
return
}
for _, adapted := range adaptAgenticEventToEinoEvents(ev) {
if adapted != nil {
gen.Send(adapted)
}
}
}
}()
return iter
}
@@ -0,0 +1,145 @@
package multiagent
import (
"context"
"testing"
"github.com/cloudwego/eino/adk"
"github.com/cloudwego/eino/schema"
)
type fakeAgenticMessageAgent struct {
name string
description string
captured *adk.TypedAgentInput[*schema.AgenticMessage]
resumeInfo *adk.ResumeInfo
events []*adk.TypedAgentEvent[*schema.AgenticMessage]
}
func (f *fakeAgenticMessageAgent) Name(context.Context) string {
return f.name
}
func (f *fakeAgenticMessageAgent) Description(context.Context) string {
return f.description
}
func (f *fakeAgenticMessageAgent) Run(_ context.Context, input *adk.TypedAgentInput[*schema.AgenticMessage], _ ...adk.AgentRunOption) *adk.AsyncIterator[*adk.TypedAgentEvent[*schema.AgenticMessage]] {
f.captured = input
iter, gen := adk.NewAsyncIteratorPair[*adk.TypedAgentEvent[*schema.AgenticMessage]]()
go func() {
defer gen.Close()
for _, ev := range f.events {
gen.Send(ev)
}
}()
return iter
}
func (f *fakeAgenticMessageAgent) Resume(_ context.Context, info *adk.ResumeInfo, _ ...adk.AgentRunOption) *adk.AsyncIterator[*adk.TypedAgentEvent[*schema.AgenticMessage]] {
f.resumeInfo = info
iter, gen := adk.NewAsyncIteratorPair[*adk.TypedAgentEvent[*schema.AgenticMessage]]()
go func() {
defer gen.Close()
for _, ev := range f.events {
gen.Send(ev)
}
}()
return iter
}
func TestEinoAgenticMessageAgentAdapterConvertsInputAndEvents(t *testing.T) {
inner := &fakeAgenticMessageAgent{
name: "agentic",
description: "typed agent",
events: []*adk.TypedAgentEvent[*schema.AgenticMessage]{
{
AgentName: "agentic",
Output: &adk.TypedAgentOutput[*schema.AgenticMessage]{
MessageOutput: &adk.TypedMessageVariant[*schema.AgenticMessage]{
Message: &schema.AgenticMessage{
Role: schema.AgenticRoleTypeAssistant,
ContentBlocks: []*schema.ContentBlock{
schema.NewContentBlock(&schema.AssistantGenText{Text: "hello"}),
},
},
},
},
},
},
}
agent := newEinoAgenticMessageAgentAdapter(inner)
if agent.Name(context.Background()) != "agentic" || agent.Description(context.Background()) != "typed agent" {
t.Fatalf("adapter metadata name=%q desc=%q", agent.Name(context.Background()), agent.Description(context.Background()))
}
iter := agent.Run(context.Background(), &adk.AgentInput{
EnableStreaming: true,
Messages: []*schema.Message{
schema.UserMessage("hi"),
},
})
ev, ok := iter.Next()
if !ok {
t.Fatal("expected adapted event")
}
if inner.captured == nil || !inner.captured.EnableStreaming || len(inner.captured.Messages) != 1 {
t.Fatalf("captured input = %#v", inner.captured)
}
if inner.captured.Messages[0].Role != schema.AgenticRoleTypeUser || inner.captured.Messages[0].ContentBlocks[0].UserInputText.Text != "hi" {
t.Fatalf("captured message = %#v", inner.captured.Messages[0])
}
if ev.AgentName != "agentic" || ev.Output == nil || ev.Output.MessageOutput == nil {
t.Fatalf("event = %#v", ev)
}
if ev.Output.MessageOutput.Role != schema.Assistant || ev.Output.MessageOutput.Message.Content != "hello" {
t.Fatalf("message output = %#v", ev.Output.MessageOutput)
}
if _, ok := iter.Next(); ok {
t.Fatal("expected iterator to close")
}
}
func TestEinoAgenticMessageAgentAdapterNilInnerReturnsNil(t *testing.T) {
if got := newEinoAgenticMessageAgentAdapter(nil); got != nil {
t.Fatalf("adapter = %#v, want nil", got)
}
}
func TestEinoAgenticMessageAgentAdapterResumeConvertsEvents(t *testing.T) {
inner := &fakeAgenticMessageAgent{
name: "agentic",
events: []*adk.TypedAgentEvent[*schema.AgenticMessage]{
{
AgentName: "agentic",
Output: &adk.TypedAgentOutput[*schema.AgenticMessage]{
MessageOutput: &adk.TypedMessageVariant[*schema.AgenticMessage]{
Message: &schema.AgenticMessage{
Role: schema.AgenticRoleTypeAssistant,
ContentBlocks: []*schema.ContentBlock{
schema.NewContentBlock(&schema.AssistantGenText{Text: "resumed"}),
},
},
},
},
},
},
}
agent, ok := newEinoAgenticMessageAgentAdapter(inner).(adk.ResumableAgent)
if !ok {
t.Fatal("adapter must implement adk.ResumableAgent")
}
info := &adk.ResumeInfo{WasInterrupted: true}
iter := agent.Resume(context.Background(), info)
ev, ok := iter.Next()
if !ok {
t.Fatal("expected adapted resume event")
}
if inner.resumeInfo != info {
t.Fatalf("resume info = %#v, want original pointer", inner.resumeInfo)
}
if ev.Output == nil || ev.Output.MessageOutput == nil || ev.Output.MessageOutput.Message.Content != "resumed" {
t.Fatalf("resume event = %#v", ev)
}
}
@@ -0,0 +1,64 @@
package multiagent
import (
"context"
"fmt"
"github.com/cloudwego/eino/adk"
"github.com/cloudwego/eino/components/model"
"github.com/cloudwego/eino/components/tool"
"github.com/cloudwego/eino/schema"
)
type einoAgenticChatModelAgentConfig struct {
Name string
Description string
Instruction string
Model model.AgenticModel
ToolsConfig adk.ToolsConfig
MaxIterations int
Exit tool.BaseTool
GenModelInput adk.TypedGenModelInput[*schema.AgenticMessage]
Handlers []adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage]
ModelRetryConfig *adk.TypedModelRetryConfig[*schema.AgenticMessage]
ModelFailoverConfig *adk.ModelFailoverConfig[*schema.AgenticMessage]
OutputKey string
}
func newEinoAgenticChatModelAgent(ctx context.Context, cfg einoAgenticChatModelAgentConfig) (adk.TypedResumableAgent[*schema.AgenticMessage], error) {
if cfg.Model == nil {
return nil, fmt.Errorf("eino agentic ChatModelAgent: model is required")
}
typedCfg := &adk.TypedChatModelAgentConfig[*schema.AgenticMessage]{
Name: cfg.Name,
Description: cfg.Description,
Instruction: cfg.Instruction,
Model: cfg.Model,
ToolsConfig: cfg.ToolsConfig,
MaxIterations: cfg.MaxIterations,
Exit: cfg.Exit,
GenModelInput: cfg.GenModelInput,
Handlers: cfg.Handlers,
ModelRetryConfig: cfg.ModelRetryConfig,
ModelFailoverConfig: cfg.ModelFailoverConfig,
OutputKey: cfg.OutputKey,
}
typedAgent, err := adk.NewTypedChatModelAgent(ctx, typedCfg)
if err != nil {
return nil, fmt.Errorf("eino agentic NewTypedChatModelAgent: %w", err)
}
return typedAgent, nil
}
func newEinoAgenticChatModelAgentAdapter(ctx context.Context, cfg einoAgenticChatModelAgentConfig) (adk.Agent, error) {
typedAgent, err := newEinoAgenticChatModelAgent(ctx, cfg)
if err != nil {
return nil, err
}
agent := newEinoAgenticMessageAgentAdapter(typedAgent)
if agent == nil {
return nil, fmt.Errorf("eino agentic ChatModelAgent: adapter is nil")
}
return agent, nil
}
@@ -0,0 +1,163 @@
package multiagent
import (
"context"
"strings"
"sync"
"testing"
"github.com/cloudwego/eino/adk"
"github.com/cloudwego/eino/components/model"
"github.com/cloudwego/eino/schema"
)
type capturingAgenticChatModel struct {
mu sync.Mutex
inputs [][]*schema.AgenticMessage
output *schema.AgenticMessage
}
func (m *capturingAgenticChatModel) Generate(_ context.Context, input []*schema.AgenticMessage, _ ...model.Option) (*schema.AgenticMessage, error) {
m.mu.Lock()
m.inputs = append(m.inputs, input)
m.mu.Unlock()
if m.output != nil {
return m.output, nil
}
return &schema.AgenticMessage{
Role: schema.AgenticRoleTypeAssistant,
ContentBlocks: []*schema.ContentBlock{schema.NewContentBlock(&schema.AssistantGenText{Text: "agentic answer"})},
}, nil
}
func (m *capturingAgenticChatModel) Stream(_ context.Context, input []*schema.AgenticMessage, _ ...model.Option) (*schema.StreamReader[*schema.AgenticMessage], error) {
msg, err := m.Generate(context.Background(), input)
if err != nil {
return nil, err
}
return schema.StreamReaderFromArray([]*schema.AgenticMessage{msg}), nil
}
func (m *capturingAgenticChatModel) snapshotInputs() [][]*schema.AgenticMessage {
m.mu.Lock()
defer m.mu.Unlock()
out := make([][]*schema.AgenticMessage, len(m.inputs))
copy(out, m.inputs)
return out
}
func TestNewEinoAgenticChatModelAgentAdapterRunsThroughClassicAgentBoundary(t *testing.T) {
t.Parallel()
ctx := context.Background()
trace := newModelFacingTraceHolder()
fakeModel := &capturingAgenticChatModel{}
agent, err := newEinoAgenticChatModelAgentAdapter(ctx, einoAgenticChatModelAgentConfig{
Name: "agentic",
Description: "agentic adapter test",
Instruction: "system instruction",
Model: fakeModel,
Handlers: appendEinoAgenticChatModelTailMiddlewares(nil, einoChatModelTailConfig{
phase: "agentic",
trace: trace,
}),
})
if err != nil {
t.Fatalf("newEinoAgenticChatModelAgentAdapter: %v", err)
}
iter := agent.Run(ctx, &adk.AgentInput{
Messages: []*schema.Message{schema.UserMessage("classic input")},
})
var last *adk.AgentEvent
for {
ev, ok := iter.Next()
if !ok {
break
}
if ev.Err != nil {
t.Fatalf("agent event error: %v", ev.Err)
}
last = ev
}
if last == nil || last.Output == nil || last.Output.MessageOutput == nil {
t.Fatalf("last event = %#v, want message output", last)
}
if got := last.Output.MessageOutput.Message.Content; got != "agentic answer" {
t.Fatalf("classic output content = %q, want agentic answer", got)
}
inputs := fakeModel.snapshotInputs()
if len(inputs) != 1 {
t.Fatalf("model calls = %d, want 1", len(inputs))
}
if len(inputs[0]) != 2 {
t.Fatalf("model input messages = %d, want instruction + user", len(inputs[0]))
}
if inputs[0][0].Role != schema.AgenticRoleTypeSystem || agenticMessageText(inputs[0][0]) != "system instruction" {
t.Fatalf("first agentic input = %#v", inputs[0][0])
}
if inputs[0][1].Role != schema.AgenticRoleTypeUser || agenticMessageText(inputs[0][1]) != "classic input" {
t.Fatalf("second agentic input = %#v", inputs[0][1])
}
snapshot := trace.Snapshot()
if len(snapshot) != 2 || snapshot[0].Role != schema.System || snapshot[1].Role != schema.User {
t.Fatalf("trace snapshot = %#v, want classic system + user trace", snapshot)
}
}
func TestNewEinoAgenticChatModelAgentAdapterPreservesTypedToolCallsForToolLayerRecovery(t *testing.T) {
t.Parallel()
ctx := context.Background()
fakeModel := &capturingAgenticChatModel{
output: &schema.AgenticMessage{
Role: schema.AgenticRoleTypeAssistant,
ContentBlocks: []*schema.ContentBlock{schema.NewContentBlock(&schema.FunctionToolCall{
CallID: "call-1",
Name: "exec",
Arguments: `{"command":"` + strings.Repeat("x", 20000) + `"}`,
})},
},
}
agent, err := newEinoAgenticChatModelAgentAdapter(ctx, einoAgenticChatModelAgentConfig{
Name: "agentic",
Description: "agentic adapter test",
Model: fakeModel,
Handlers: appendEinoAgenticChatModelTailMiddlewares(nil, einoChatModelTailConfig{
phase: "agentic",
}),
})
if err != nil {
t.Fatalf("newEinoAgenticChatModelAgentAdapter: %v", err)
}
iter := agent.Run(ctx, &adk.AgentInput{Messages: []*schema.Message{schema.UserMessage("run")}})
var last *adk.AgentEvent
for {
ev, ok := iter.Next()
if !ok {
break
}
if ev.Err != nil {
t.Fatalf("agent event error: %v", ev.Err)
}
last = ev
}
if last == nil || last.Output == nil || last.Output.MessageOutput == nil {
t.Fatalf("last event = %#v, want message output", last)
}
msg := last.Output.MessageOutput.Message
if len(msg.ToolCalls) != 1 {
t.Fatalf("tool calls = %#v, want one tool call", msg.ToolCalls)
}
args := msg.ToolCalls[0].Function.Arguments
if !strings.Contains(args, strings.Repeat("x", 32)) || strings.Contains(args, modelOutputRecoveryKey) {
t.Fatalf("agentic tool args were unexpectedly rewritten: %q", args)
}
}
func TestNewEinoAgenticChatModelAgentAdapterRequiresModel(t *testing.T) {
t.Parallel()
if _, err := newEinoAgenticChatModelAgentAdapter(context.Background(), einoAgenticChatModelAgentConfig{}); err == nil {
t.Fatal("expected missing model error")
}
}
@@ -0,0 +1,209 @@
package multiagent
import (
"context"
"strings"
"github.com/cloudwego/eino/adk"
"github.com/cloudwego/eino/schema"
"go.uber.org/zap"
)
// appendEinoAgenticChatModelTailMiddlewares appends protocol-neutral handlers for
// TypedChatModelAgent[*schema.AgenticMessage]. Classic ReAct history repair
// handlers stay on the schema.Message path because AgenticMessage has native
// content blocks for function calls/results.
func appendEinoAgenticChatModelTailMiddlewares(
handlers []adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage],
cfg einoChatModelTailConfig,
) []adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage] {
handlers = append(handlers, newAgenticSystemMessageNormalizerMiddleware(cfg.logger, cfg.phase))
handlers = append(handlers, newAgenticContinuationUserDedupMiddleware(cfg.logger, cfg.phase))
if cfg.agenticSummarization != nil {
handlers = append(handlers, cfg.agenticSummarization)
}
if !cfg.skipTrace && cfg.trace != nil {
if capMw := newAgenticModelFacingTraceMiddleware(cfg.trace); capMw != nil {
handlers = append(handlers, capMw)
}
}
return handlers
}
type agenticSystemMessageNormalizerMiddleware struct {
*adk.TypedBaseChatModelAgentMiddleware[*schema.AgenticMessage]
logger *zap.Logger
phase string
}
func newAgenticSystemMessageNormalizerMiddleware(logger *zap.Logger, phase string) adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage] {
return &agenticSystemMessageNormalizerMiddleware{
TypedBaseChatModelAgentMiddleware: &adk.TypedBaseChatModelAgentMiddleware[*schema.AgenticMessage]{},
logger: logger,
phase: phase,
}
}
func (m *agenticSystemMessageNormalizerMiddleware) BeforeModelRewriteState(
ctx context.Context,
state *adk.TypedChatModelAgentState[*schema.AgenticMessage],
mc *adk.TypedModelContext[*schema.AgenticMessage],
) (context.Context, *adk.TypedChatModelAgentState[*schema.AgenticMessage], error) {
_ = mc
if m == nil || state == nil || len(state.Messages) == 0 {
return ctx, state, nil
}
before := countAgenticSystemMessages(state.Messages)
if before <= 1 {
return ctx, state, nil
}
normalized := normalizeSingleLeadingAgenticSystemMessage(state.Messages)
if len(normalized) == len(state.Messages) && countAgenticSystemMessages(normalized) >= before {
return ctx, state, nil
}
if m.logger != nil {
m.logger.Info("eino agentic system messages merged",
zap.String("phase", m.phase),
zap.Int("system_before", before),
zap.Int("system_after", countAgenticSystemMessages(normalized)),
zap.Int("messages_before", len(state.Messages)),
zap.Int("messages_after", len(normalized)),
)
}
out := *state
out.Messages = normalized
return ctx, &out, nil
}
type agenticContinuationUserDedupMiddleware struct {
*adk.TypedBaseChatModelAgentMiddleware[*schema.AgenticMessage]
logger *zap.Logger
phase string
}
func newAgenticContinuationUserDedupMiddleware(logger *zap.Logger, phase string) adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage] {
return &agenticContinuationUserDedupMiddleware{
TypedBaseChatModelAgentMiddleware: &adk.TypedBaseChatModelAgentMiddleware[*schema.AgenticMessage]{},
logger: logger,
phase: phase,
}
}
func (m *agenticContinuationUserDedupMiddleware) BeforeModelRewriteState(
ctx context.Context,
state *adk.TypedChatModelAgentState[*schema.AgenticMessage],
mc *adk.TypedModelContext[*schema.AgenticMessage],
) (context.Context, *adk.TypedChatModelAgentState[*schema.AgenticMessage], error) {
_ = mc
if m == nil || state == nil || len(state.Messages) == 0 {
return ctx, state, nil
}
deduped, dropped := dedupAgenticContinuationUserMessages(state.Messages)
if dropped == 0 {
return ctx, state, nil
}
if m.logger != nil {
m.logger.Info("eino agentic continuation user messages deduplicated",
zap.String("phase", m.phase),
zap.Int("dropped", dropped),
zap.Int("messages_before", len(state.Messages)),
zap.Int("messages_after", len(deduped)),
)
}
out := *state
out.Messages = deduped
return ctx, &out, nil
}
func countAgenticSystemMessages(msgs []*schema.AgenticMessage) int {
n := 0
for _, msg := range msgs {
if msg != nil && msg.Role == schema.AgenticRoleTypeSystem {
n++
}
}
return n
}
func normalizeSingleLeadingAgenticSystemMessage(msgs []*schema.AgenticMessage) []*schema.AgenticMessage {
var systemParts []string
out := make([]*schema.AgenticMessage, 0, len(msgs))
for _, msg := range msgs {
if msg == nil {
continue
}
if msg.Role == schema.AgenticRoleTypeSystem {
if text := strings.TrimSpace(agenticMessageText(msg)); text != "" {
systemParts = append(systemParts, text)
}
continue
}
out = append(out, msg)
}
if len(systemParts) == 0 {
return out
}
merged := schema.SystemAgenticMessage(strings.Join(systemParts, "\n\n"))
return append([]*schema.AgenticMessage{merged}, out...)
}
func dedupAgenticContinuationUserMessages(msgs []*schema.AgenticMessage) ([]*schema.AgenticMessage, int) {
lastIdx := -1
contCount := 0
for i, msg := range msgs {
if !isAgenticContinuationUserMessage(msg) {
continue
}
contCount++
lastIdx = i
}
if contCount <= 1 {
return msgs, 0
}
out := make([]*schema.AgenticMessage, 0, len(msgs)-(contCount-1))
dropped := 0
for i, msg := range msgs {
if isAgenticContinuationUserMessage(msg) && i != lastIdx {
dropped++
continue
}
out = append(out, msg)
}
return out, dropped
}
func isAgenticContinuationUserMessage(msg *schema.AgenticMessage) bool {
if msg == nil || msg.Role != schema.AgenticRoleTypeUser {
return false
}
return strings.Contains(agenticMessageText(msg), continuationSessionMarker)
}
func agenticMessageText(msg *schema.AgenticMessage) string {
if msg == nil {
return ""
}
var b strings.Builder
for _, block := range msg.ContentBlocks {
if block == nil {
continue
}
switch {
case block.UserInputText != nil:
if s := strings.TrimSpace(block.UserInputText.Text); s != "" {
if b.Len() > 0 {
b.WriteByte('\n')
}
b.WriteString(s)
}
case block.AssistantGenText != nil:
if s := strings.TrimSpace(block.AssistantGenText.Text); s != "" {
if b.Len() > 0 {
b.WriteByte('\n')
}
b.WriteString(s)
}
}
}
return b.String()
}
@@ -0,0 +1,112 @@
package multiagent
import (
"context"
"strings"
"testing"
"github.com/cloudwego/eino/adk"
"github.com/cloudwego/eino/schema"
)
func TestAgenticSystemMessageNormalizerMiddlewareMergesDuplicates(t *testing.T) {
t.Parallel()
mw := newAgenticSystemMessageNormalizerMiddleware(nil, "test")
state := &adk.TypedChatModelAgentState[*schema.AgenticMessage]{
Messages: []*schema.AgenticMessage{
schema.SystemAgenticMessage("first"),
schema.UserAgenticMessage("hello"),
schema.SystemAgenticMessage("second"),
},
}
_, out, err := mw.BeforeModelRewriteState(context.Background(), state, nil)
if err != nil {
t.Fatalf("BeforeModelRewriteState: %v", err)
}
if out == state {
t.Fatal("expected rewritten state")
}
if got := countAgenticSystemMessages(out.Messages); got != 1 {
t.Fatalf("system messages = %d, want 1", got)
}
if out.Messages[0].Role != schema.AgenticRoleTypeSystem {
t.Fatalf("first role = %s, want system", out.Messages[0].Role)
}
text := agenticMessageText(out.Messages[0])
if !strings.Contains(text, "first") || !strings.Contains(text, "second") {
t.Fatalf("merged system text = %q", text)
}
if len(out.Messages) != 2 || agenticMessageText(out.Messages[1]) != "hello" {
t.Fatalf("normalized messages = %#v", out.Messages)
}
}
func TestAgenticContinuationUserDedupMiddlewareKeepsLatest(t *testing.T) {
t.Parallel()
mw := newAgenticContinuationUserDedupMiddleware(nil, "test")
state := &adk.TypedChatModelAgentState[*schema.AgenticMessage]{
Messages: []*schema.AgenticMessage{
schema.UserAgenticMessage(continuationSessionMarker + "\nold"),
schema.UserAgenticMessage("real user request"),
schema.UserAgenticMessage(continuationSessionMarker + "\nnew"),
},
}
_, out, err := mw.BeforeModelRewriteState(context.Background(), state, nil)
if err != nil {
t.Fatalf("BeforeModelRewriteState: %v", err)
}
if out == state {
t.Fatal("expected rewritten state")
}
if len(out.Messages) != 2 {
t.Fatalf("messages = %d, want 2", len(out.Messages))
}
if strings.Contains(agenticMessageText(out.Messages[0]), continuationSessionMarker) {
t.Fatalf("old continuation was not dropped: %#v", out.Messages)
}
if !strings.Contains(agenticMessageText(out.Messages[1]), "new") {
t.Fatalf("latest continuation not retained: %#v", out.Messages)
}
}
func TestAgenticModelFacingTraceMiddlewareStoresClassicTrace(t *testing.T) {
t.Parallel()
holder := newModelFacingTraceHolder()
mw := newAgenticModelFacingTraceMiddleware(holder)
state := &adk.TypedChatModelAgentState[*schema.AgenticMessage]{
Messages: []*schema.AgenticMessage{
schema.SystemAgenticMessage("instruction"),
{
Role: schema.AgenticRoleTypeAssistant,
ContentBlocks: []*schema.ContentBlock{
schema.NewContentBlock(&schema.AssistantGenText{Text: "answer"}),
},
},
},
}
if _, _, err := mw.BeforeModelRewriteState(context.Background(), state, nil); err != nil {
t.Fatalf("BeforeModelRewriteState: %v", err)
}
got := holder.Snapshot()
if len(got) != 2 {
t.Fatalf("trace len = %d, want 2", len(got))
}
if got[0].Role != schema.System || got[0].Content != "instruction" {
t.Fatalf("system trace = %#v", got[0])
}
if got[1].Role != schema.Assistant || got[1].Content != "answer" {
t.Fatalf("assistant trace = %#v", got[1])
}
}
func TestAppendEinoAgenticChatModelTailMiddlewares(t *testing.T) {
t.Parallel()
holder := newModelFacingTraceHolder()
handlers := appendEinoAgenticChatModelTailMiddlewares(nil, einoChatModelTailConfig{
phase: "agentic",
trace: holder,
})
if len(handlers) != 3 {
t.Fatalf("handlers = %d, want system + continuation + trace", len(handlers))
}
}
@@ -0,0 +1,106 @@
package multiagent
import (
"io"
"github.com/cloudwego/eino/adk"
"github.com/cloudwego/eino/schema"
)
// adaptAgenticEventToEinoEvents converts typed AgenticMessage ADK events into
// the classic schema.Message events consumed by the existing SSE/MCP drain.
func adaptAgenticEventToEinoEvents(ev *adk.TypedAgentEvent[*schema.AgenticMessage]) []*adk.AgentEvent {
if ev == nil {
return nil
}
base := func(output *adk.AgentOutput) *adk.AgentEvent {
return &adk.AgentEvent{
AgentName: ev.AgentName,
RunPath: append([]adk.RunStep(nil), ev.RunPath...),
Output: output,
Action: ev.Action,
Err: ev.Err,
}
}
if ev.Output == nil {
return []*adk.AgentEvent{base(nil)}
}
customized := ev.Output.CustomizedOutput
mv := ev.Output.MessageOutput
if mv == nil {
return []*adk.AgentEvent{base(&adk.AgentOutput{CustomizedOutput: customized})}
}
if mv.IsStreaming {
return []*adk.AgentEvent{base(&adk.AgentOutput{
MessageOutput: &adk.MessageVariant{
IsStreaming: true,
MessageStream: agenticStreamToEinoStream(mv.MessageStream),
Role: agenticVariantRole(mv),
},
CustomizedOutput: customized,
})}
}
msgs := AgenticMessageToEino(mv.Message)
if len(msgs) == 0 {
return []*adk.AgentEvent{base(&adk.AgentOutput{CustomizedOutput: customized})}
}
out := make([]*adk.AgentEvent, 0, len(msgs))
for i, msg := range msgs {
eventCustomized := any(nil)
if i == 0 {
eventCustomized = customized
}
out = append(out, base(&adk.AgentOutput{
MessageOutput: &adk.MessageVariant{
Message: msg,
Role: msg.Role,
ToolName: msg.ToolName,
},
CustomizedOutput: eventCustomized,
}))
}
return out
}
func agenticStreamToEinoStream(sr *schema.StreamReader[*schema.AgenticMessage]) *schema.StreamReader[*schema.Message] {
out, writer := schema.Pipe[*schema.Message](8)
go func() {
defer writer.Close()
if sr == nil {
return
}
defer sr.Close()
for {
chunk, err := sr.Recv()
if err != nil {
if err != io.EOF {
writer.Send(nil, err)
}
return
}
for _, msg := range AgenticMessageToEino(chunk) {
if msg != nil && writer.Send(msg, nil) {
return
}
}
}
}()
return out
}
func agenticVariantRole(mv *adk.TypedMessageVariant[*schema.AgenticMessage]) schema.RoleType {
if mv == nil {
return schema.Assistant
}
switch mv.AgenticRole {
case schema.AgenticRoleTypeSystem:
return schema.System
case schema.AgenticRoleTypeUser:
// In Agentic ReAct output, user-role events from the graph are local
// FunctionToolResult messages emitted by AgenticToolsNode.
return schema.Tool
default:
return schema.Assistant
}
}
@@ -0,0 +1,249 @@
package multiagent
import (
"errors"
"io"
"testing"
"github.com/cloudwego/eino/adk"
"github.com/cloudwego/eino/schema"
)
func TestAdaptAgenticEventToEinoEventsAssistantMessage(t *testing.T) {
usage := &schema.TokenUsage{PromptTokens: 10, CompletionTokens: 5, TotalTokens: 15}
ev := &adk.TypedAgentEvent[*schema.AgenticMessage]{
AgentName: "agentic",
Output: &adk.TypedAgentOutput[*schema.AgenticMessage]{
MessageOutput: &adk.TypedMessageVariant[*schema.AgenticMessage]{
Message: &schema.AgenticMessage{
Role: schema.AgenticRoleTypeAssistant,
ResponseMeta: &schema.AgenticResponseMeta{TokenUsage: usage},
ContentBlocks: []*schema.ContentBlock{
schema.NewContentBlock(&schema.Reasoning{Text: "think"}),
schema.NewContentBlock(&schema.AssistantGenText{Text: "calling"}),
schema.NewContentBlock(&schema.FunctionToolCall{CallID: "call-1", Name: "scan", Arguments: `{"host":"127.0.0.1"}`}),
},
},
},
CustomizedOutput: "custom",
},
}
got := adaptAgenticEventToEinoEvents(ev)
if len(got) != 1 {
t.Fatalf("events = %d, want 1", len(got))
}
mv := got[0].Output.MessageOutput
if got[0].AgentName != "agentic" || got[0].Output.CustomizedOutput != "custom" {
t.Fatalf("event metadata = %#v", got[0])
}
if mv.Role != schema.Assistant || mv.Message.Role != schema.Assistant {
t.Fatalf("role = %q/%q, want assistant", mv.Role, mv.Message.Role)
}
if mv.Message.Content != "calling" || mv.Message.ReasoningContent != "think" {
t.Fatalf("message text = %#v", mv.Message)
}
if len(mv.Message.ToolCalls) != 1 || mv.Message.ToolCalls[0].ID != "call-1" || mv.Message.ToolCalls[0].Function.Name != "scan" {
t.Fatalf("tool calls = %#v", mv.Message.ToolCalls)
}
if mv.Message.ResponseMeta == nil || mv.Message.ResponseMeta.Usage != usage {
t.Fatalf("usage = %#v, want original usage", mv.Message.ResponseMeta)
}
}
func TestAdaptAgenticEventToEinoEventsPureToolResult(t *testing.T) {
ev := &adk.TypedAgentEvent[*schema.AgenticMessage]{
Output: &adk.TypedAgentOutput[*schema.AgenticMessage]{
MessageOutput: &adk.TypedMessageVariant[*schema.AgenticMessage]{
Message: &schema.AgenticMessage{
Role: schema.AgenticRoleTypeUser,
ContentBlocks: []*schema.ContentBlock{
schema.NewContentBlock(&schema.FunctionToolResult{
CallID: "call-2",
Name: "execute",
Content: []*schema.FunctionToolResultContentBlock{{
Type: schema.FunctionToolResultContentBlockTypeText,
Text: &schema.UserInputText{Text: "done"},
}},
}),
},
},
},
},
}
got := adaptAgenticEventToEinoEvents(ev)
if len(got) != 1 {
t.Fatalf("events = %d, want 1", len(got))
}
msg := got[0].Output.MessageOutput.Message
if got[0].Output.MessageOutput.Role != schema.Tool || msg.Role != schema.Tool || msg.ToolName != "execute" || msg.ToolCallID != "call-2" || msg.Content != "done" {
t.Fatalf("tool event = %#v message=%#v", got[0].Output.MessageOutput, msg)
}
}
func TestAdaptAgenticEventToEinoEventsSplitsMixedToolResult(t *testing.T) {
ev := &adk.TypedAgentEvent[*schema.AgenticMessage]{
Output: &adk.TypedAgentOutput[*schema.AgenticMessage]{
MessageOutput: &adk.TypedMessageVariant[*schema.AgenticMessage]{
Message: &schema.AgenticMessage{
Role: schema.AgenticRoleTypeAssistant,
ContentBlocks: []*schema.ContentBlock{
schema.NewContentBlock(&schema.AssistantGenText{Text: "text"}),
schema.NewContentBlock(&schema.FunctionToolResult{
CallID: "call-3",
Name: "grep",
Content: []*schema.FunctionToolResultContentBlock{{
Type: schema.FunctionToolResultContentBlockTypeText,
Text: &schema.UserInputText{Text: "match"},
}},
}),
},
},
},
},
}
got := adaptAgenticEventToEinoEvents(ev)
if len(got) != 2 {
t.Fatalf("events = %d, want assistant + tool", len(got))
}
if got[0].Output.MessageOutput.Role != schema.Assistant || got[0].Output.MessageOutput.Message.Content != "text" {
t.Fatalf("assistant event = %#v", got[0].Output.MessageOutput)
}
if got[1].Output.MessageOutput.Role != schema.Tool || got[1].Output.MessageOutput.Message.ToolName != "grep" {
t.Fatalf("tool event = %#v", got[1].Output.MessageOutput)
}
}
func TestAdaptAgenticEventToEinoEventsStreamingAssistant(t *testing.T) {
stream := schema.StreamReaderFromArray([]*schema.AgenticMessage{
{
Role: schema.AgenticRoleTypeAssistant,
ContentBlocks: []*schema.ContentBlock{
schema.NewContentBlock(&schema.AssistantGenText{Text: "hel"}),
},
},
{
Role: schema.AgenticRoleTypeAssistant,
ContentBlocks: []*schema.ContentBlock{
schema.NewContentBlock(&schema.AssistantGenText{Text: "lo"}),
},
},
})
ev := &adk.TypedAgentEvent[*schema.AgenticMessage]{
Output: &adk.TypedAgentOutput[*schema.AgenticMessage]{
MessageOutput: &adk.TypedMessageVariant[*schema.AgenticMessage]{
IsStreaming: true,
MessageStream: stream,
AgenticRole: schema.AgenticRoleTypeAssistant,
},
},
}
got := adaptAgenticEventToEinoEvents(ev)
if len(got) != 1 {
t.Fatalf("events = %d, want 1", len(got))
}
mv := got[0].Output.MessageOutput
if !mv.IsStreaming || mv.Role != schema.Assistant {
t.Fatalf("stream variant = %#v", mv)
}
first, err := mv.MessageStream.Recv()
if err != nil || first.Content != "hel" {
t.Fatalf("first = %#v err=%v", first, err)
}
second, err := mv.MessageStream.Recv()
if err != nil || second.Content != "lo" {
t.Fatalf("second = %#v err=%v", second, err)
}
_, err = mv.MessageStream.Recv()
if !errors.Is(err, io.EOF) {
t.Fatalf("final err = %v, want EOF", err)
}
}
func TestAdaptAgenticStreamingToolResultFeedsClassicToolResultHandler(t *testing.T) {
stream := schema.StreamReaderFromArray([]*schema.AgenticMessage{
{
Role: schema.AgenticRoleTypeUser,
ContentBlocks: []*schema.ContentBlock{
schema.NewContentBlock(&schema.FunctionToolResult{
CallID: "call-agentic-stream",
Name: "execute",
Content: []*schema.FunctionToolResultContentBlock{{
Type: schema.FunctionToolResultContentBlockTypeText,
Text: &schema.UserInputText{Text: "partial "},
}},
}),
},
},
{
Role: schema.AgenticRoleTypeUser,
ContentBlocks: []*schema.ContentBlock{
schema.NewContentBlock(&schema.FunctionToolResult{
CallID: "call-agentic-stream",
Name: "execute",
Content: []*schema.FunctionToolResultContentBlock{{
Type: schema.FunctionToolResultContentBlockTypeText,
Text: &schema.UserInputText{Text: "done"},
}},
}),
},
},
})
ev := &adk.TypedAgentEvent[*schema.AgenticMessage]{
AgentName: "agentic",
Output: &adk.TypedAgentOutput[*schema.AgenticMessage]{
MessageOutput: &adk.TypedMessageVariant[*schema.AgenticMessage]{
IsStreaming: true,
MessageStream: stream,
AgenticRole: schema.AgenticRoleTypeUser,
},
},
}
got := adaptAgenticEventToEinoEvents(ev)
if len(got) != 1 || got[0].Output == nil || got[0].Output.MessageOutput == nil {
t.Fatalf("events = %#v", got)
}
mv := got[0].Output.MessageOutput
if !mv.IsStreaming || mv.Role != schema.Tool {
t.Fatalf("streaming variant = %#v, want tool stream", mv)
}
var event map[string]interface{}
runMessages := newEinoRunMessageAccumulator(nil)
emitter := newEinoToolResultProgressEmitter(einoToolResultProgressEmitterConfig{
ConversationID: "conv-agentic",
Progress: func(eventType, _ string, data interface{}) {
if eventType == "tool_result" {
event, _ = data.(map[string]interface{})
}
},
})
handler := newEinoToolResultEventHandler(einoToolResultEventHandlerConfig{
RunMessages: runMessages,
Emitter: emitter,
})
if !handler.HandleStreaming(mv, "agentic") {
t.Fatal("agentic streaming tool result was not handled")
}
if event["toolName"] != "execute" || event["toolCallId"] != "call-agentic-stream" || event["result"] != "partial done" {
t.Fatalf("tool result event = %#v", event)
}
msgs := runMessages.Messages()
if len(msgs) != 1 || msgs[0].ToolName != "execute" || msgs[0].ToolCallID != "call-agentic-stream" || msgs[0].Content != "partial done" {
t.Fatalf("run messages = %#v", msgs)
}
}
func TestAdaptAgenticEventToEinoEventsPreservesErrorOnlyEvent(t *testing.T) {
wantErr := errors.New("boom")
ev := &adk.TypedAgentEvent[*schema.AgenticMessage]{AgentName: "agentic", Err: wantErr}
got := adaptAgenticEventToEinoEvents(ev)
if len(got) != 1 || got[0].AgentName != "agentic" || !errors.Is(got[0].Err, wantErr) {
t.Fatalf("events = %#v", got)
}
}
+184
View File
@@ -0,0 +1,184 @@
package multiagent
import (
"strings"
"github.com/cloudwego/eino/schema"
)
// EinoMessagesToAgentic converts the project's current ADK message history to
// Eino's native AgenticMessage shape. It intentionally covers the text,
// reasoning, function tool-call, and function tool-result channels used by the
// agent runtime today; unsupported multimodal/provider-specific fields stay in
// schema.Message until a real AgenticModel backend is wired.
func EinoMessagesToAgentic(msgs []*schema.Message) []*schema.AgenticMessage {
if len(msgs) == 0 {
return nil
}
out := make([]*schema.AgenticMessage, 0, len(msgs))
for _, msg := range msgs {
if msg == nil {
continue
}
out = append(out, EinoMessageToAgentic(msg))
}
return out
}
func EinoMessageToAgentic(msg *schema.Message) *schema.AgenticMessage {
if msg == nil {
return nil
}
out := &schema.AgenticMessage{
Role: messageRoleToAgentic(msg.Role),
Extra: cloneAnyMap(msg.Extra),
}
if msg.ResponseMeta != nil {
out.ResponseMeta = &schema.AgenticResponseMeta{TokenUsage: msg.ResponseMeta.Usage}
}
if text := strings.TrimSpace(msg.ReasoningContent); text != "" {
out.ContentBlocks = append(out.ContentBlocks, schema.NewContentBlock(&schema.Reasoning{Text: msg.ReasoningContent}))
}
switch msg.Role {
case schema.Assistant:
if msg.Content != "" {
out.ContentBlocks = append(out.ContentBlocks, schema.NewContentBlock(&schema.AssistantGenText{Text: msg.Content}))
}
for _, tc := range msg.ToolCalls {
out.ContentBlocks = append(out.ContentBlocks, schema.NewContentBlock(&schema.FunctionToolCall{
CallID: tc.ID,
Name: tc.Function.Name,
Arguments: tc.Function.Arguments,
}))
}
case schema.Tool:
out.Role = schema.AgenticRoleTypeUser
out.ContentBlocks = append(out.ContentBlocks, schema.NewContentBlock(&schema.FunctionToolResult{
CallID: msg.ToolCallID,
Name: msg.ToolName,
Content: []*schema.FunctionToolResultContentBlock{{
Type: schema.FunctionToolResultContentBlockTypeText,
Text: &schema.UserInputText{Text: msg.Content},
}},
}))
default:
if msg.Content != "" {
out.ContentBlocks = append(out.ContentBlocks, schema.NewContentBlock(&schema.UserInputText{Text: msg.Content}))
}
}
return out
}
// AgenticMessagesToEino converts AgenticMessage values back into the classic
// schema.Message form used by the existing ADK event drain and persistence code.
func AgenticMessagesToEino(msgs []*schema.AgenticMessage) []*schema.Message {
if len(msgs) == 0 {
return nil
}
out := make([]*schema.Message, 0, len(msgs))
for _, msg := range msgs {
if msg == nil {
continue
}
out = append(out, AgenticMessageToEino(msg)...)
}
return out
}
func AgenticMessageToEino(msg *schema.AgenticMessage) []*schema.Message {
if msg == nil {
return nil
}
base := &schema.Message{
Role: agenticRoleToMessage(msg.Role),
Extra: cloneAnyMap(msg.Extra),
}
if msg.ResponseMeta != nil {
base.ResponseMeta = &schema.ResponseMeta{Usage: msg.ResponseMeta.TokenUsage}
}
var toolResults []*schema.Message
for _, block := range msg.ContentBlocks {
if block == nil {
continue
}
switch {
case block.Reasoning != nil:
base.ReasoningContent += block.Reasoning.Text
case block.UserInputText != nil:
base.Content += block.UserInputText.Text
case block.AssistantGenText != nil:
base.Role = schema.Assistant
base.Content += block.AssistantGenText.Text
case block.FunctionToolCall != nil:
base.Role = schema.Assistant
base.ToolCalls = append(base.ToolCalls, schema.ToolCall{
ID: block.FunctionToolCall.CallID,
Type: "function",
Function: schema.FunctionCall{
Name: block.FunctionToolCall.Name,
Arguments: block.FunctionToolCall.Arguments,
},
})
case block.FunctionToolResult != nil:
toolResults = append(toolResults, functionToolResultToMessage(block.FunctionToolResult))
}
}
if len(toolResults) > 0 && base.Content == "" && base.ReasoningContent == "" && len(base.ToolCalls) == 0 {
return toolResults
}
out := []*schema.Message{base}
out = append(out, toolResults...)
return out
}
func messageRoleToAgentic(role schema.RoleType) schema.AgenticRoleType {
switch role {
case schema.System:
return schema.AgenticRoleTypeSystem
case schema.Assistant:
return schema.AgenticRoleTypeAssistant
default:
return schema.AgenticRoleTypeUser
}
}
func agenticRoleToMessage(role schema.AgenticRoleType) schema.RoleType {
switch role {
case schema.AgenticRoleTypeSystem:
return schema.System
case schema.AgenticRoleTypeAssistant:
return schema.Assistant
default:
return schema.User
}
}
func functionToolResultToMessage(result *schema.FunctionToolResult) *schema.Message {
if result == nil {
return nil
}
parts := make([]string, 0, len(result.Content))
for _, block := range result.Content {
if block == nil || block.Text == nil {
continue
}
parts = append(parts, block.Text.Text)
}
return &schema.Message{
Role: schema.Tool,
Content: strings.Join(parts, ""),
ToolCallID: result.CallID,
ToolName: result.Name,
}
}
func cloneAnyMap(in map[string]any) map[string]any {
if len(in) == 0 {
return nil
}
out := make(map[string]any, len(in))
for k, v := range in {
out[k] = v
}
return out
}
@@ -0,0 +1,154 @@
package multiagent
import (
"testing"
"github.com/cloudwego/eino/schema"
)
func TestEinoMessageToAgenticPreservesAssistantToolCalls(t *testing.T) {
msg := &schema.Message{
Role: schema.Assistant,
Content: "I will scan it.",
ReasoningContent: "Need enumerate first.",
ToolCalls: []schema.ToolCall{{
ID: "call-1",
Type: "function",
Function: schema.FunctionCall{
Name: "nmap",
Arguments: `{"target":"127.0.0.1"}`,
},
}},
Extra: map[string]any{"trace": "kept"},
}
got := EinoMessageToAgentic(msg)
if got.Role != schema.AgenticRoleTypeAssistant {
t.Fatalf("role = %q, want assistant", got.Role)
}
if len(got.ContentBlocks) != 3 {
t.Fatalf("blocks = %d, want 3", len(got.ContentBlocks))
}
if got.ContentBlocks[0].Reasoning == nil || got.ContentBlocks[0].Reasoning.Text != msg.ReasoningContent {
t.Fatalf("reasoning block = %#v", got.ContentBlocks[0])
}
if got.ContentBlocks[1].AssistantGenText == nil || got.ContentBlocks[1].AssistantGenText.Text != msg.Content {
t.Fatalf("assistant text block = %#v", got.ContentBlocks[1])
}
call := got.ContentBlocks[2].FunctionToolCall
if call == nil || call.CallID != "call-1" || call.Name != "nmap" || call.Arguments != `{"target":"127.0.0.1"}` {
t.Fatalf("tool call block = %#v", got.ContentBlocks[2])
}
if got.Extra["trace"] != "kept" {
t.Fatalf("extra = %#v", got.Extra)
}
}
func TestEinoMessageToAgenticMapsToolResultAsUserFunctionResult(t *testing.T) {
msg := &schema.Message{
Role: schema.Tool,
Content: "22/tcp open ssh",
ToolCallID: "call-ssh",
ToolName: "nmap",
}
got := EinoMessageToAgentic(msg)
if got.Role != schema.AgenticRoleTypeUser {
t.Fatalf("role = %q, want user", got.Role)
}
if len(got.ContentBlocks) != 1 || got.ContentBlocks[0].FunctionToolResult == nil {
t.Fatalf("blocks = %#v", got.ContentBlocks)
}
result := got.ContentBlocks[0].FunctionToolResult
if result.CallID != "call-ssh" || result.Name != "nmap" {
t.Fatalf("tool result metadata = %#v", result)
}
if len(result.Content) != 1 || result.Content[0].Text == nil || result.Content[0].Text.Text != "22/tcp open ssh" {
t.Fatalf("tool result content = %#v", result.Content)
}
}
func TestAgenticMessageToEinoPreservesAssistantBlocks(t *testing.T) {
msg := &schema.AgenticMessage{
Role: schema.AgenticRoleTypeAssistant,
ContentBlocks: []*schema.ContentBlock{
schema.NewContentBlock(&schema.Reasoning{Text: "Think first."}),
schema.NewContentBlock(&schema.AssistantGenText{Text: "Calling scanner."}),
schema.NewContentBlock(&schema.FunctionToolCall{
CallID: "call-2",
Name: "scan",
Arguments: `{"host":"example.com"}`,
}),
},
}
got := AgenticMessageToEino(msg)
if len(got) != 1 {
t.Fatalf("messages = %d, want 1", len(got))
}
if got[0].Role != schema.Assistant || got[0].Content != "Calling scanner." || got[0].ReasoningContent != "Think first." {
t.Fatalf("assistant message = %#v", got[0])
}
if len(got[0].ToolCalls) != 1 || got[0].ToolCalls[0].ID != "call-2" || got[0].ToolCalls[0].Function.Name != "scan" {
t.Fatalf("tool calls = %#v", got[0].ToolCalls)
}
}
func TestAgenticMessageToEinoSplitsPureToolResult(t *testing.T) {
msg := &schema.AgenticMessage{
Role: schema.AgenticRoleTypeUser,
ContentBlocks: []*schema.ContentBlock{
schema.NewContentBlock(&schema.FunctionToolResult{
CallID: "call-3",
Name: "execute",
Content: []*schema.FunctionToolResultContentBlock{{
Type: schema.FunctionToolResultContentBlockTypeText,
Text: &schema.UserInputText{Text: "done"},
}},
}),
},
}
got := AgenticMessageToEino(msg)
if len(got) != 1 {
t.Fatalf("messages = %d, want 1", len(got))
}
if got[0].Role != schema.Tool || got[0].ToolCallID != "call-3" || got[0].ToolName != "execute" || got[0].Content != "done" {
t.Fatalf("tool message = %#v", got[0])
}
}
func TestEinoAgenticRoundTripForSupportedFields(t *testing.T) {
msgs := []*schema.Message{
schema.SystemMessage("system"),
schema.UserMessage("user"),
{
Role: schema.Assistant,
Content: "assistant",
ToolCalls: []schema.ToolCall{{
ID: "call-4",
Type: "function",
Function: schema.FunctionCall{Name: "grep", Arguments: `{"q":"token"}`},
}},
},
{
Role: schema.Tool,
Content: "match",
ToolCallID: "call-4",
ToolName: "grep",
},
}
got := AgenticMessagesToEino(EinoMessagesToAgentic(msgs))
if len(got) != len(msgs) {
t.Fatalf("round trip messages = %d, want %d: %#v", len(got), len(msgs), got)
}
for i := range msgs {
if got[i].Role != msgs[i].Role || got[i].Content != msgs[i].Content || got[i].ToolCallID != msgs[i].ToolCallID || got[i].ToolName != msgs[i].ToolName {
t.Fatalf("message[%d] = %#v, want %#v", i, got[i], msgs[i])
}
if len(got[i].ToolCalls) != len(msgs[i].ToolCalls) {
t.Fatalf("message[%d] tool calls = %#v, want %#v", i, got[i].ToolCalls, msgs[i].ToolCalls)
}
}
}
@@ -0,0 +1,109 @@
package multiagent
import (
"context"
"strings"
"github.com/cloudwego/eino/components/model"
"go.uber.org/zap"
)
type einoAgenticModelFactory func(context.Context) (model.AgenticModel, error)
type einoAgenticRuntimeSupport struct {
TypedRunner bool
Streaming bool
CancelMonitoring bool
ModelRetry bool
ModelFailover bool
ToolResultObservation bool
MCPExecutionAudit bool
}
type einoAgenticModelGate struct {
Ready bool
Reason string
Missing []string
}
// Eino v0.9.14 wires AgenticMessage through the same generic TypedRunner,
// stream cancel monitoring, model retry, and model failover wrappers used by
// schema.Message. Keep this matrix explicit so future upgrades are audited
// deliberately instead of flipping the AgenticModel path by accident.
func einoAgenticRuntimeSupportV0914() einoAgenticRuntimeSupport {
return einoAgenticRuntimeSupport{
TypedRunner: true,
Streaming: true,
CancelMonitoring: true,
ModelRetry: true,
ModelFailover: true,
ToolResultObservation: true,
MCPExecutionAudit: true,
}
}
func evaluateEinoAgenticModelGate(factory einoAgenticModelFactory, support einoAgenticRuntimeSupport) einoAgenticModelGate {
missing := make([]string, 0, 8)
if factory == nil {
missing = append(missing, "model.AgenticModel backend")
} else {
if m, err := factory(context.Background()); err != nil || m == nil {
missing = append(missing, "model.AgenticModel backend")
}
}
if !support.TypedRunner {
missing = append(missing, "adk.TypedRunner[*schema.AgenticMessage]")
}
if !support.Streaming {
missing = append(missing, "AgenticMessage streaming")
}
if !support.CancelMonitoring {
missing = append(missing, "AgenticMessage model-stream cancel monitoring")
}
if !support.ModelRetry {
missing = append(missing, "AgenticMessage ModelRetry")
}
if !support.ModelFailover {
missing = append(missing, "AgenticMessage ModelFailover")
}
if !support.ToolResultObservation {
missing = append(missing, "AgenticMessage tool-result observation")
}
if !support.MCPExecutionAudit {
missing = append(missing, "AgenticMessage MCP execution audit")
}
if len(missing) == 0 {
return einoAgenticModelGate{Ready: true, Reason: "ready"}
}
return einoAgenticModelGate{
Reason: "agentic_model_not_ready: " + strings.Join(missing, ", "),
Missing: missing,
}
}
func logEinoAgenticModelGate(logger *zap.Logger, scope, orchestration string, gate einoAgenticModelGate) {
if logger == nil {
return
}
fields := []zap.Field{
zap.String("scope", scope),
zap.String("orchestration", orchestration),
zap.Bool("ready", gate.Ready),
zap.String("reason", gate.Reason),
zap.Strings("missing", gate.Missing),
}
if gate.Ready {
logger.Info("eino agentic model gate ready", fields...)
return
}
logger.Info("eino agentic model gate disabled", fields...)
}
func agenticTextModelFactory(m model.AgenticModel) einoAgenticModelFactory {
if m == nil {
return nil
}
return func(context.Context) (model.AgenticModel, error) {
return m, nil
}
}
@@ -0,0 +1,93 @@
package multiagent
import (
"context"
"errors"
"testing"
"github.com/cloudwego/eino/components/model"
"github.com/cloudwego/eino/schema"
)
type fakeAgenticGateModel struct{}
func (m *fakeAgenticGateModel) Generate(context.Context, []*schema.AgenticMessage, ...model.Option) (*schema.AgenticMessage, error) {
return &schema.AgenticMessage{Role: schema.AgenticRoleTypeAssistant}, nil
}
func (m *fakeAgenticGateModel) Stream(context.Context, []*schema.AgenticMessage, ...model.Option) (*schema.StreamReader[*schema.AgenticMessage], error) {
return schema.StreamReaderFromArray([]*schema.AgenticMessage{{Role: schema.AgenticRoleTypeAssistant}}), nil
}
func TestEinoAgenticModelGateV0914WaitsOnlyForBackend(t *testing.T) {
gate := evaluateEinoAgenticModelGate(nil, einoAgenticRuntimeSupportV0914())
if gate.Ready {
t.Fatal("v0.9.14 gate should stay disabled without an AgenticModel backend")
}
if !containsString(gate.Missing, "model.AgenticModel backend") {
t.Fatalf("missing = %#v, want backend reason", gate.Missing)
}
for _, unexpected := range []string{
"AgenticMessage model-stream cancel monitoring",
"AgenticMessage ModelRetry",
"AgenticMessage ModelFailover",
"AgenticMessage tool-result observation",
"AgenticMessage MCP execution audit",
} {
if containsString(gate.Missing, unexpected) {
t.Fatalf("missing = %#v, should not include %q for v0.9.14 runtime support", gate.Missing, unexpected)
}
}
}
func TestEinoAgenticModelGateV0914ReadyWithBackend(t *testing.T) {
gate := evaluateEinoAgenticModelGate(agenticTextModelFactory(&fakeAgenticGateModel{}), einoAgenticRuntimeSupportV0914())
if !gate.Ready {
t.Fatalf("gate = %#v, want ready when v0.9.14 runtime support has a backend", gate)
}
if gate.Reason != "ready" || len(gate.Missing) != 0 {
t.Fatalf("gate details = %#v", gate)
}
}
func TestEinoAgenticModelGateReadyWhenBackendAndRuntimeParityExist(t *testing.T) {
gate := evaluateEinoAgenticModelGate(agenticTextModelFactory(&fakeAgenticGateModel{}), einoAgenticRuntimeSupport{
TypedRunner: true,
Streaming: true,
CancelMonitoring: true,
ModelRetry: true,
ModelFailover: true,
ToolResultObservation: true,
MCPExecutionAudit: true,
})
if !gate.Ready {
t.Fatalf("gate = %#v, want ready", gate)
}
if gate.Reason != "ready" || len(gate.Missing) != 0 {
t.Fatalf("gate details = %#v", gate)
}
}
func TestEinoAgenticModelGateTreatsFactoryErrorAsMissingBackend(t *testing.T) {
gate := evaluateEinoAgenticModelGate(func(context.Context) (model.AgenticModel, error) {
return nil, errors.New("not implemented")
}, einoAgenticRuntimeSupport{
TypedRunner: true,
Streaming: true,
CancelMonitoring: true,
ModelRetry: true,
ModelFailover: true,
ToolResultObservation: true,
MCPExecutionAudit: true,
})
if gate.Ready {
t.Fatal("factory error should disable gate")
}
if !containsString(gate.Missing, "model.AgenticModel backend") {
t.Fatalf("missing = %#v, want backend reason", gate.Missing)
}
}
@@ -0,0 +1,278 @@
package multiagent
import (
"context"
"fmt"
"os"
"path/filepath"
"strings"
"cyberstrike-ai/internal/config"
"cyberstrike-ai/internal/database"
einoopenai "github.com/cloudwego/eino-ext/components/model/openai"
"github.com/cloudwego/eino/adk"
"github.com/cloudwego/eino/adk/middlewares/summarization"
"github.com/cloudwego/eino/components/model"
"github.com/cloudwego/eino/schema"
"go.uber.org/zap"
)
// newEinoAgenticSummarizationMiddleware wires the project's domain-specific
// compaction policy into Eino's native typed AgenticMessage summarization.
func newEinoAgenticSummarizationMiddleware(
ctx context.Context,
summaryModel model.BaseModel[*schema.AgenticMessage],
appCfg *config.Config,
mwCfg *config.MultiAgentEinoMiddlewareConfig,
conversationID string,
db *database.DB,
projectID string,
logger *zap.Logger,
) (adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage], error) {
if summaryModel == nil || appCfg == nil {
return nil, fmt.Errorf("multiagent: agentic summarization 需要 model 与配置")
}
maxTotal := appCfg.OpenAI.MaxTotalTokens
if maxTotal <= 0 {
maxTotal = 120000
}
triggerRatio := 0.8
emitInternalEvents := true
outputReserve := config.DefaultSummarizationOutputReserveTokens
userLedgerMaxRunes := config.DefaultSummarizationUserIntentLedgerMaxRunes
userLedgerEntryMaxRunes := config.DefaultSummarizationUserIntentLedgerEntryMaxRunes
toolMaxBytes := config.MultiAgentEinoMiddlewareConfig{}.ReductionMaxLengthForTruncEffective()
if mwCfg != nil {
triggerRatio = mwCfg.SummarizationTriggerRatioEffective()
emitInternalEvents = mwCfg.SummarizationEmitInternalEventsEffective()
outputReserve = mwCfg.SummarizationOutputReserveTokensEffective()
userLedgerMaxRunes = mwCfg.SummarizationUserIntentLedgerMaxRunesEffective()
userLedgerEntryMaxRunes = mwCfg.SummarizationUserIntentLedgerEntryMaxRunesEffective()
toolMaxBytes = mwCfg.ReductionMaxLengthForTruncEffective()
}
ledgerWindowCap := modelFacingRuneBudget(maxTotal, 0.20)
userLedgerMaxRunes = minPositiveInt(userLedgerMaxRunes, ledgerWindowCap)
userLedgerEntryMaxRunes = minPositiveInt(userLedgerEntryMaxRunes, userLedgerMaxRunes)
trigger := int(float64(maxTotal) * triggerRatio)
if trigger < 4096 {
trigger = maxTotal
if trigger < 4096 {
trigger = 4096
}
}
modelName := strings.TrimSpace(appCfg.OpenAI.Model)
if modelName == "" {
modelName = "gpt-4o"
}
classicTokenCounter := einoSummarizationTokenCounter(modelName)
agenticTokenCounter := func(ctx context.Context, input *summarization.TypedTokenCounterInput[*schema.AgenticMessage]) (int, error) {
if input == nil {
return 0, nil
}
return classicTokenCounter(ctx, &summarization.TokenCounterInput{
Messages: AgenticMessagesToEino(input.Messages),
Tools: input.Tools,
})
}
recentTrailMax := trigger / 4
if recentTrailMax < 2048 {
recentTrailMax = 2048
}
if recentTrailMax > trigger/2 {
recentTrailMax = trigger / 2
}
summaryInputMax := trigger - outputReserve
if summaryInputMax < 4096 {
summaryInputMax = trigger * 80 / 100
}
if summaryInputMax < 4096 {
summaryInputMax = 4096
}
transcriptPath := ""
if conv := strings.TrimSpace(conversationID); conv != "" {
baseRoot := filepath.Join(os.TempDir(), "cyberstrike-summarization")
if dbPath := strings.TrimSpace(appCfg.Database.Path); dbPath != "" {
baseRoot = filepath.Join(filepath.Dir(dbPath), "conversation_artifacts", sanitizeEinoPathSegment(conv), "summarization")
}
base := baseRoot
if abs, err := filepath.Abs(base); err == nil {
base = abs
}
if mkErr := os.MkdirAll(base, 0o755); mkErr == nil {
transcriptPath = filepath.Join(base, "transcript.txt")
}
}
retryPolicy := einoTransientRunRetryPolicyFromMW(mwCfg)
retryMax := retryPolicy.maxAttempts
var summaryOverflowRetries int
summaryModelOpts := []model.Option{
einoopenai.WithMaxCompletionTokens(outputReserve),
}
mw, err := summarization.NewTyped[*schema.AgenticMessage](ctx, &summarization.TypedConfig[*schema.AgenticMessage]{
Model: summaryModel,
ModelOptions: summaryModelOpts,
GenModelInput: func(ctx context.Context, sysInstruction, userInstruction *schema.AgenticMessage, originalMsgs []*schema.AgenticMessage) ([]*schema.AgenticMessage, error) {
classicOriginal := AgenticMessagesToEino(originalMsgs)
if transcriptPath != "" && len(classicOriginal) > 0 {
if werr := writeSummarizationTranscript(transcriptPath, classicOriginal); werr != nil && logger != nil {
logger.Warn("eino agentic summarization transcript preflight 写入失败",
zap.String("path", transcriptPath), zap.Error(werr))
}
}
budget := summaryInputMax
aggressive := summaryOverflowRetries > 0
if aggressive {
budget = summaryInputMax * 70 / 100
if budget < 4096 {
budget = 4096
}
}
input, dropped, berr := buildBudgetedSummarizationModelInput(
ctx,
agenticInstructionToClassic(sysInstruction, schema.System),
agenticInstructionToClassic(userInstruction, schema.User),
classicOriginal,
classicTokenCounter,
budget,
summarizationInputBudgetOpts{
toolMaxBytes: toolMaxBytes,
spillRef: transcriptPath,
aggressive: aggressive,
},
)
if logger != nil && (berr != nil || dropped > 0 || aggressive) {
fields := []zap.Field{
zap.Int("max_input_tokens", budget),
zap.Int("trigger_context_tokens", trigger),
zap.Int("output_reserve_tokens", outputReserve),
zap.Int("dropped_rounds", dropped),
zap.Bool("aggressive", aggressive),
}
if berr != nil {
fields = append(fields, zap.Error(berr))
logger.Warn("eino agentic summarization input budget failed", fields...)
} else {
logger.Info("eino agentic summarization input bounded", fields...)
}
}
return EinoMessagesToAgentic(input), berr
},
Trigger: &summarization.TriggerCondition{
ContextTokens: trigger,
},
TokenCounter: agenticTokenCounter,
UserInstruction: einoSummarizeUserInstruction,
EmitInternalEvents: emitInternalEvents,
TranscriptFilePath: transcriptPath,
Retry: &summarization.TypedRetryConfig[*schema.AgenticMessage]{
MaxRetries: &retryMax,
ShouldRetry: func(_ context.Context, _ *schema.AgenticMessage, err error) bool {
if isEinoContextOverflowError(err) && summaryOverflowRetries < 1 {
summaryOverflowRetries++
if logger != nil {
logger.Warn("eino agentic summarization context overflow, retrying with aggressive compaction",
zap.Error(err),
)
}
return true
}
retry := isEinoTransientRunError(err)
if retry && logger != nil {
logger.Warn("eino agentic summarization generate transient error, will retry if attempts remain",
zap.Error(err),
zap.Int("max_retries", retryMax),
)
}
return retry
},
},
Finalize: func(ctx context.Context, originalMessages []*schema.AgenticMessage, summary *schema.AgenticMessage) ([]*schema.AgenticMessage, error) {
classicOriginal := AgenticMessagesToEino(originalMessages)
classicSummary := agenticSummaryToClassicMessage(summary)
if classicSummary == nil {
return nil, fmt.Errorf("agentic summarization returned empty summary")
}
compactionMessages := stripOriginalUserIntentLedgerFromMessages(classicOriginal)
defaultFinalized, derr := summarization.DefaultFinalize(ctx, compactionMessages, classicSummary)
if derr != nil {
return nil, derr
}
if len(defaultFinalized) == 0 {
return nil, fmt.Errorf("agentic summarization default finalize returned no messages")
}
summaryMsg := appendTranscriptPathToSummarizationMessage(defaultFinalized[len(defaultFinalized)-1], transcriptPath)
summaryMsg = stripAnalysisFromSummarizationMessage(summaryMsg)
userLedger := buildOriginalUserIntentLedgerMessage(classicOriginal, userLedgerMaxRunes, userLedgerEntryMaxRunes)
out, ferr := summarizeFinalizeWithRecentAssistantToolTrail(ctx, compactionMessages, summaryMsg, classicTokenCounter, recentTrailMax)
if ferr != nil {
return nil, ferr
}
out = mergeMessageIntoLeadingSystem(out, userLedger)
if appCfg != nil {
out = refreshFactIndexInMessages(out, db, projectID, appCfg.Project, logger)
}
return EinoMessagesToAgentic(out), nil
},
Callback: func(ctx context.Context, before, after adk.TypedChatModelAgentState[*schema.AgenticMessage]) error {
classicBefore := AgenticMessagesToEino(before.Messages)
classicAfter := AgenticMessagesToEino(after.Messages)
if transcriptPath != "" && len(classicBefore) > 0 {
if werr := writeSummarizationTranscript(transcriptPath, classicBefore); werr != nil && logger != nil {
logger.Warn("eino agentic summarization transcript 写入失败",
zap.String("path", transcriptPath),
zap.Error(werr),
)
}
}
if logger != nil {
beforeTokens, _ := classicTokenCounter(ctx, &summarization.TokenCounterInput{Messages: classicBefore})
afterTokens, _ := classicTokenCounter(ctx, &summarization.TokenCounterInput{Messages: classicAfter})
logger.Info("eino agentic summarization 已压缩上下文",
zap.Int("messages_before", len(before.Messages)),
zap.Int("messages_after", len(after.Messages)),
zap.Int("tokens_before_estimated", beforeTokens),
zap.Int("tokens_after_estimated", afterTokens),
zap.Int("max_total_tokens", maxTotal),
zap.Int("trigger_context_tokens", trigger),
zap.String("transcript_file", transcriptPath),
)
}
return nil
},
})
if err != nil {
return nil, fmt.Errorf("summarization.NewTyped[AgenticMessage]: %w", err)
}
return mw, nil
}
func agenticInstructionToClassic(msg *schema.AgenticMessage, fallbackRole schema.RoleType) *schema.Message {
msgs := AgenticMessageToEino(msg)
if len(msgs) > 0 && msgs[0] != nil {
return msgs[0]
}
return &schema.Message{Role: fallbackRole}
}
func agenticSummaryToClassicMessage(msg *schema.AgenticMessage) *schema.Message {
msgs := AgenticMessageToEino(msg)
for _, m := range msgs {
if m == nil {
continue
}
if m.Role == schema.Assistant || strings.TrimSpace(m.Content) != "" || m.ReasoningContent != "" {
if m.Role != schema.Assistant {
cp := *m
cp.Role = schema.Assistant
return &cp
}
return m
}
}
return nil
}
@@ -0,0 +1,210 @@
package multiagent
import (
"context"
"path/filepath"
"strings"
"testing"
"cyberstrike-ai/internal/config"
"github.com/cloudwego/eino/adk"
"github.com/cloudwego/eino/schema"
)
func TestNewEinoAgenticSummarizationMiddlewareCompactsWithNativeTypedMiddleware(t *testing.T) {
t.Parallel()
ctx := context.Background()
emit := false
summaryModel := &capturingAgenticChatModel{
output: &schema.AgenticMessage{
Role: schema.AgenticRoleTypeAssistant,
ContentBlocks: []*schema.ContentBlock{schema.NewContentBlock(&schema.AssistantGenText{Text: `<analysis>检查历史</analysis>
<summary>
## 1. 授权范围与约束
- 仅测试 example.com
## 7. 当前进度、策略决策与下一步
- 继续验证 SQL 注入路径
</summary>`})},
},
}
appCfg := &config.Config{}
appCfg.OpenAI.Model = "gpt-4o"
appCfg.OpenAI.MaxTotalTokens = 5000
appCfg.Database.Path = filepath.Join(t.TempDir(), "cyberstrike.db")
mwCfg := &config.MultiAgentEinoMiddlewareConfig{
SummarizationEmitInternalEvents: &emit,
SummarizationOutputReserveTokens: 1024,
}
mw, err := newEinoAgenticSummarizationMiddleware(ctx, summaryModel, appCfg, mwCfg, "conv-agentic", nil, "", nil)
if err != nil {
t.Fatalf("newEinoAgenticSummarizationMiddleware: %v", err)
}
state := &adk.TypedChatModelAgentState[*schema.AgenticMessage]{
Messages: []*schema.AgenticMessage{
schema.SystemAgenticMessage("system root"),
schema.UserAgenticMessage("授权范围 example.com\n" + strings.Repeat("历史扫描输出 ", 12000)),
agenticAssistantTextMessage("已记录范围"),
schema.UserAgenticMessage("继续验证 SQL 注入路径"),
},
}
_, after, err := mw.BeforeModelRewriteState(ctx, state, nil)
if err != nil {
t.Fatalf("BeforeModelRewriteState: %v", err)
}
inputs := summaryModel.snapshotInputs()
if len(inputs) != 1 || len(inputs[0]) == 0 {
t.Fatalf("summary model inputs = %#v, want one typed AgenticMessage call", inputs)
}
if after == nil {
t.Fatal("after state is nil")
}
classicAfter := AgenticMessagesToEino(after.Messages)
joined := joinClassicMessageContent(classicAfter)
if strings.Contains(joined, "<analysis>") {
t.Fatalf("analysis block leaked into compacted context: %s", joined)
}
for _, want := range []string{"继续验证 SQL 注入路径", "原始用户输入与约束账本", "完整的对话记录位于"} {
if !strings.Contains(joined, want) {
t.Fatalf("compacted context missing %q:\n%s", want, joined)
}
}
}
func TestEinoAgenticChatModelAgentCompactsContextBeforeBusinessModel(t *testing.T) {
t.Parallel()
ctx := context.Background()
emit := false
summaryModel := &capturingAgenticChatModel{
output: agenticAssistantTextMessage(`<analysis>internal scratchpad</analysis>
<summary>
## 1. 授权范围与约束
- 仅测试 example.com
## 7. 当前进度、策略决策与下一步
- 继续验证 SQL 注入路径
</summary>`),
}
businessModel := &capturingAgenticChatModel{
output: agenticAssistantTextMessage("business answer after compaction"),
}
appCfg := &config.Config{}
appCfg.OpenAI.Model = "gpt-4o"
appCfg.OpenAI.MaxTotalTokens = 5000
appCfg.Database.Path = filepath.Join(t.TempDir(), "cyberstrike.db")
mwCfg := &config.MultiAgentEinoMiddlewareConfig{
SummarizationEmitInternalEvents: &emit,
SummarizationOutputReserveTokens: 1024,
}
sumMw, err := newEinoAgenticSummarizationMiddleware(ctx, summaryModel, appCfg, mwCfg, "conv-agentic-e2e", nil, "", nil)
if err != nil {
t.Fatalf("newEinoAgenticSummarizationMiddleware: %v", err)
}
trace := newModelFacingTraceHolder()
agent, err := newEinoAgenticChatModelAgentAdapter(ctx, einoAgenticChatModelAgentConfig{
Name: "agentic",
Description: "agentic compaction e2e test",
Instruction: "system root",
Model: businessModel,
Handlers: appendEinoAgenticChatModelTailMiddlewares(nil, einoChatModelTailConfig{
phase: "agentic",
agenticSummarization: sumMw,
trace: trace,
}),
})
if err != nil {
t.Fatalf("newEinoAgenticChatModelAgentAdapter: %v", err)
}
rawHistory := "授权范围 example.com\n" + strings.Repeat("原始扫描输出SHOULD_NOT_REACH_BUSINESS_MODEL ", 12000)
iter := agent.Run(ctx, &adk.AgentInput{
Messages: []*schema.Message{
schema.UserMessage(rawHistory),
schema.AssistantMessage("已记录范围", nil),
schema.UserMessage("继续验证 SQL 注入路径"),
},
})
var last *adk.AgentEvent
for {
ev, ok := iter.Next()
if !ok {
break
}
if ev.Err != nil {
t.Fatalf("agent event error: %v", ev.Err)
}
last = ev
}
if last == nil || last.Output == nil || last.Output.MessageOutput == nil {
t.Fatalf("last event = %#v, want message output", last)
}
if got := last.Output.MessageOutput.Message.Content; got != "business answer after compaction" {
t.Fatalf("business output = %q", got)
}
if inputs := summaryModel.snapshotInputs(); len(inputs) != 1 {
t.Fatalf("summary model calls = %d, want 1", len(inputs))
}
businessInputs := businessModel.snapshotInputs()
if len(businessInputs) != 1 {
t.Fatalf("business model calls = %d, want 1", len(businessInputs))
}
finalClassicInput := AgenticMessagesToEino(businessInputs[0])
joined := joinClassicMessageContent(finalClassicInput)
for _, want := range []string{"继续验证 SQL 注入路径", "原始用户输入与约束账本", "完整的对话记录位于"} {
if !strings.Contains(joined, want) {
t.Fatalf("business model input missing %q:\n%s", want, joined)
}
}
if strings.Contains(joined, "<analysis>") {
t.Fatalf("analysis leaked to business model input:\n%s", joined)
}
if strings.Count(joined, "原始扫描输出SHOULD_NOT_REACH_BUSINESS_MODEL") > 3 {
t.Fatalf("raw oversized history leaked to business model input, count=%d", strings.Count(joined, "原始扫描输出SHOULD_NOT_REACH_BUSINESS_MODEL"))
}
traceJoined := joinClassicMessageContent(trace.Snapshot())
if !strings.Contains(traceJoined, "继续验证 SQL 注入路径") || strings.Count(traceJoined, "原始扫描输出SHOULD_NOT_REACH_BUSINESS_MODEL") > 3 {
t.Fatalf("model-facing trace not compacted:\n%s", traceJoined)
}
}
func TestAppendEinoAgenticChatModelTailMiddlewaresIncludesTypedSummarization(t *testing.T) {
t.Parallel()
mw := newAgenticSystemMessageNormalizerMiddleware(nil, "summary")
handlers := appendEinoAgenticChatModelTailMiddlewares(nil, einoChatModelTailConfig{
agenticSummarization: mw,
skipTrace: true,
})
found := false
for _, h := range handlers {
if h == mw {
found = true
break
}
}
if !found {
t.Fatal("agentic summarization middleware was not appended")
}
}
func agenticAssistantTextMessage(text string) *schema.AgenticMessage {
return &schema.AgenticMessage{
Role: schema.AgenticRoleTypeAssistant,
ContentBlocks: []*schema.ContentBlock{schema.NewContentBlock(&schema.AssistantGenText{Text: text})},
}
}
func joinClassicMessageContent(msgs []*schema.Message) string {
var b strings.Builder
for _, msg := range msgs {
if msg == nil {
continue
}
b.WriteString(msg.Content)
b.WriteByte('\n')
}
return b.String()
}
@@ -0,0 +1,42 @@
package multiagent
import "strings"
type einoAssistantOutputAccumulator struct {
orchMode string
lastAssistant string
lastPlanExecuteExecutor string
}
func newEinoAssistantOutputAccumulator(orchMode string) *einoAssistantOutputAccumulator {
return &einoAssistantOutputAccumulator{orchMode: orchMode}
}
func (a *einoAssistantOutputAccumulator) RecordMainAssistant(agentName, content string) bool {
if a == nil {
return false
}
content = strings.TrimSpace(content)
if content == "" {
return false
}
a.lastAssistant = content
if a.orchMode == "plan_execute" && strings.EqualFold(strings.TrimSpace(agentName), "executor") {
a.lastPlanExecuteExecutor = UnwrapPlanExecuteUserText(content)
}
return true
}
func (a *einoAssistantOutputAccumulator) LastAssistant() string {
if a == nil {
return ""
}
return a.lastAssistant
}
func (a *einoAssistantOutputAccumulator) LastPlanExecuteExecutor() string {
if a == nil {
return ""
}
return a.lastPlanExecuteExecutor
}
@@ -0,0 +1,52 @@
package multiagent
import "testing"
func TestEinoAssistantOutputAccumulatorRecordsMainAssistant(t *testing.T) {
acc := newEinoAssistantOutputAccumulator("deep")
if acc.RecordMainAssistant("lead", " hello ") != true {
t.Fatal("expected record")
}
if got := acc.LastAssistant(); got != "hello" {
t.Fatalf("last assistant = %q, want hello", got)
}
if got := acc.LastPlanExecuteExecutor(); got != "" {
t.Fatalf("plan execute executor = %q, want empty", got)
}
if acc.RecordMainAssistant("lead", " ") {
t.Fatal("blank content should not record")
}
if got := acc.LastAssistant(); got != "hello" {
t.Fatalf("blank content changed last assistant to %q", got)
}
}
func TestEinoAssistantOutputAccumulatorPlanExecuteExecutor(t *testing.T) {
acc := newEinoAssistantOutputAccumulator("plan_execute")
raw := `{"response":"给用户看的正文","scratchpad":"internal"}`
acc.RecordMainAssistant("executor", raw)
if got := acc.LastAssistant(); got != raw {
t.Fatalf("last assistant = %q, want raw", got)
}
if got := acc.LastPlanExecuteExecutor(); got != "给用户看的正文" {
t.Fatalf("executor output = %q", got)
}
acc.RecordMainAssistant("planner", "planner note")
if got := acc.LastAssistant(); got != "planner note" {
t.Fatalf("last assistant after planner = %q", got)
}
if got := acc.LastPlanExecuteExecutor(); got != "给用户看的正文" {
t.Fatalf("planner should not overwrite executor output, got %q", got)
}
}
func TestEinoAssistantOutputAccumulatorNilSafe(t *testing.T) {
var acc *einoAssistantOutputAccumulator
if acc.RecordMainAssistant("agent", "hello") {
t.Fatal("nil accumulator should not record")
}
if acc.LastAssistant() != "" || acc.LastPlanExecuteExecutor() != "" {
t.Fatal("nil accumulator should return empty values")
}
}
@@ -0,0 +1,167 @@
package multiagent
import (
"context"
"errors"
"github.com/cloudwego/eino/adk"
"github.com/cloudwego/eino/schema"
"go.uber.org/zap"
)
type einoAssistantStreamEventHandlerConfig struct {
Context context.Context
ConversationID string
OrchMode string
Progress func(eventType, message string, data interface{})
Logger *zap.Logger
SnapshotMCPIDs func() []string
StreamsMainAssistant func(agent string) bool
EinoRoleTag func(agent string) string
RunProgress *einoRunProgressTracker
StdoutSuppressor *einoExecuteStdoutSuppressor
AssistantOutput *einoAssistantOutputAccumulator
RunMessages *einoRunMessageAccumulator
Usage *einoRunUsageAccumulator
ToolCallCompletion *einoStreamToolCallCompletionHandler
NextMainStreamID func() string
NextReasoningStreamID func() string
NextSubAgentReplyStreamID func() string
}
type einoAssistantStreamEventHandler struct {
ctx context.Context
conversationID string
orchMode string
progress func(eventType, message string, data interface{})
logger *zap.Logger
snapshotMCPIDs func() []string
streamsMainAssistant func(agent string) bool
einoRoleTag func(agent string) string
runProgress *einoRunProgressTracker
stdoutSuppressor *einoExecuteStdoutSuppressor
assistantOutput *einoAssistantOutputAccumulator
runMessages *einoRunMessageAccumulator
usage *einoRunUsageAccumulator
toolCallCompletion *einoStreamToolCallCompletionHandler
nextMainStreamID func() string
nextReasoningStreamID func() string
nextSubAgentReplyStreamID func() string
}
func newEinoAssistantStreamEventHandler(cfg einoAssistantStreamEventHandlerConfig) *einoAssistantStreamEventHandler {
if cfg.Context == nil {
cfg.Context = context.Background()
}
if cfg.SnapshotMCPIDs == nil {
cfg.SnapshotMCPIDs = func() []string { return nil }
}
if cfg.StreamsMainAssistant == nil {
cfg.StreamsMainAssistant = func(string) bool { return true }
}
if cfg.EinoRoleTag == nil {
cfg.EinoRoleTag = func(string) string { return "" }
}
if cfg.NextMainStreamID == nil {
cfg.NextMainStreamID = func() string { return "eino-main" }
}
if cfg.NextReasoningStreamID == nil {
cfg.NextReasoningStreamID = func() string { return "eino-reasoning" }
}
if cfg.NextSubAgentReplyStreamID == nil {
cfg.NextSubAgentReplyStreamID = func() string { return "eino-sub-reply" }
}
return &einoAssistantStreamEventHandler{
ctx: cfg.Context,
conversationID: cfg.ConversationID,
orchMode: cfg.OrchMode,
progress: cfg.Progress,
logger: cfg.Logger,
snapshotMCPIDs: cfg.SnapshotMCPIDs,
streamsMainAssistant: cfg.StreamsMainAssistant,
einoRoleTag: cfg.EinoRoleTag,
runProgress: cfg.RunProgress,
stdoutSuppressor: cfg.StdoutSuppressor,
assistantOutput: cfg.AssistantOutput,
runMessages: cfg.RunMessages,
usage: cfg.Usage,
toolCallCompletion: cfg.ToolCallCompletion,
nextMainStreamID: cfg.NextMainStreamID,
nextReasoningStreamID: cfg.NextReasoningStreamID,
nextSubAgentReplyStreamID: cfg.NextSubAgentReplyStreamID,
}
}
func (h *einoAssistantStreamEventHandler) Handle(mv *adk.MessageVariant, agentName string) (handled bool, recvErr error) {
if h == nil || mv == nil || !mv.IsStreaming || mv.MessageStream == nil || mv.Role == schema.Tool {
return false, nil
}
mainStreamID := h.nextMainStreamID()
mainEmitter := newEinoMainResponseStreamEmitter(
h.conversationID, h.orchMode, agentName, mainStreamID, h.mainIteration(agentName), h.progress, h.snapshotMCPIDs,
)
reasoningEmitter := newEinoReasoningStreamEmitter(
h.conversationID,
h.orchMode,
agentName,
h.einoRoleTag(agentName),
h.progress,
h.nextReasoningStreamID,
)
var toolStreamFragments []schema.ToolCall
var streamUsage *schema.TokenUsage
subReplyEmitter := newEinoSubAgentReplyEmitter(
h.conversationID,
agentName,
h.progress,
h.nextSubAgentReplyStreamID,
)
mainAssistantStream := newEinoMainAssistantStreamHandler(einoMainAssistantStreamHandlerConfig{
AgentName: agentName,
Emitter: mainEmitter,
StdoutSuppressor: h.stdoutSuppressor,
AssistantOutput: h.assistantOutput,
RunMessages: h.runMessages,
})
recvErr = recvEinoSchemaMessageStreamWithContext(h.ctx, mv.MessageStream, 8, func(chunk *schema.Message) {
reasoningEmitter.EmitDelta(chunk.ReasoningContent)
if chunk.Content != "" {
if h.streamsMainAssistant(agentName) {
mainAssistantStream.EmitDelta(chunk.Content)
} else if !h.streamsMainAssistant(agentName) {
subReplyEmitter.EmitDelta(chunk.Content)
}
}
if len(chunk.ToolCalls) > 0 {
toolStreamFragments = append(toolStreamFragments, chunk.ToolCalls...)
}
if chunk.ResponseMeta != nil && chunk.ResponseMeta.Usage != nil {
streamUsage = maxEinoTokenUsage(streamUsage, chunk.ResponseMeta.Usage)
}
})
if recvErr != nil && !errors.Is(recvErr, context.Canceled) && h.logger != nil {
h.logger.Warn("eino stream recv error, flushing incomplete stream",
zap.Error(recvErr),
zap.String("agent", agentName),
zap.Int("toolFragments", len(toolStreamFragments)))
}
reasoningEmitter.Finish()
if h.streamsMainAssistant(agentName) {
mainAssistantStream.Finish()
}
subReplyEmitter.Finish()
if h.toolCallCompletion != nil {
h.toolCallCompletion.Complete(toolStreamFragments, agentName)
}
if h.usage != nil {
h.usage.AddUsage(streamUsage)
}
return true, recvErr
}
func (h *einoAssistantStreamEventHandler) mainIteration(agentName string) int {
if h == nil || h.runProgress == nil {
return 0
}
return h.runProgress.MainIteration(agentName)
}
@@ -0,0 +1,148 @@
package multiagent
import (
"testing"
"github.com/cloudwego/eino/adk"
"github.com/cloudwego/eino/schema"
)
func TestEinoAssistantStreamEventHandlerHandlesMainAssistantStream(t *testing.T) {
var events []string
runMessages := newEinoRunMessageAccumulator(nil)
assistantOutput := newEinoAssistantOutputAccumulator("deep")
usage := newEinoRunUsageAccumulator()
handler := newEinoAssistantStreamEventHandler(einoAssistantStreamEventHandlerConfig{
ConversationID: "conv-1",
OrchMode: "deep",
RunMessages: runMessages,
Usage: usage,
AssistantOutput: assistantOutput,
StreamsMainAssistant: func(agent string) bool { return agent == "lead" },
EinoRoleTag: func(string) string { return "orchestrator" },
NextMainStreamID: func() string { return "main-stream-1" },
Progress: func(eventType, _ string, _ interface{}) {
events = append(events, eventType)
},
})
mv := &adk.MessageVariant{
IsStreaming: true,
Role: schema.Assistant,
MessageStream: schema.StreamReaderFromArray([]*schema.Message{
{Role: schema.Assistant, Content: "he", ResponseMeta: &schema.ResponseMeta{Usage: &schema.TokenUsage{PromptTokens: 10, CompletionTokens: 1, TotalTokens: 11}}},
{Role: schema.Assistant, Content: "hello", ResponseMeta: &schema.ResponseMeta{Usage: &schema.TokenUsage{PromptTokens: 10, CompletionTokens: 5, TotalTokens: 15}}},
}),
}
handled, err := handler.Handle(mv, "lead")
if !handled || err != nil {
t.Fatalf("handled=%v err=%v", handled, err)
}
if assistantOutput.LastAssistant() != "hello" {
t.Fatalf("last assistant = %q", assistantOutput.LastAssistant())
}
if msgs := runMessages.Messages(); len(msgs) != 1 || msgs[0].Content != "hello" {
t.Fatalf("run messages = %#v", msgs)
}
if got := usage.Summary(); got.ModelCalls != 1 || got.PromptTokens != 10 || got.CompletionTokens != 5 || got.TotalTokens != 15 {
t.Fatalf("usage = %#v, want one stream model call", got)
}
if !containsString(events, "response_start") || !containsString(events, "response_delta") {
t.Fatalf("events = %#v, want response stream events", events)
}
}
func TestEinoAssistantStreamEventHandlerHandlesSubAgentStream(t *testing.T) {
var events []string
runMessages := newEinoRunMessageAccumulator(nil)
assistantOutput := newEinoAssistantOutputAccumulator("deep")
handler := newEinoAssistantStreamEventHandler(einoAssistantStreamEventHandlerConfig{
ConversationID: "conv-1",
OrchMode: "deep",
RunMessages: runMessages,
AssistantOutput: assistantOutput,
StreamsMainAssistant: func(agent string) bool { return agent == "lead" },
EinoRoleTag: func(string) string { return "sub" },
NextSubAgentReplyStreamID: func() string {
return "sub-stream-1"
},
Progress: func(eventType, _ string, _ interface{}) {
events = append(events, eventType)
},
})
mv := &adk.MessageVariant{
IsStreaming: true,
Role: schema.Assistant,
MessageStream: schema.StreamReaderFromArray([]*schema.Message{{Role: schema.Assistant, Content: "sub reply"}}),
}
handled, err := handler.Handle(mv, "worker")
if !handled || err != nil {
t.Fatalf("handled=%v err=%v", handled, err)
}
if len(runMessages.Messages()) != 0 {
t.Fatalf("sub stream should not append main run text, got %#v", runMessages.Messages())
}
if assistantOutput.LastAssistant() != "" {
t.Fatalf("sub stream should not record main assistant, got %q", assistantOutput.LastAssistant())
}
if !containsString(events, "eino_agent_reply_stream_start") ||
!containsString(events, "eino_agent_reply_stream_delta") ||
!containsString(events, "eino_agent_reply_stream_end") {
t.Fatalf("events = %#v, want sub reply stream events", events)
}
}
func TestEinoAssistantStreamEventHandlerCompletesToolFragments(t *testing.T) {
idx := 0
var events []string
runMessages := newEinoRunMessageAccumulator(nil)
runProgress := newEinoRunProgressTracker(
"deep", "lead", "conv-1",
func(eventType, _ string, _ interface{}) { events = append(events, eventType) },
func(agent string) bool { return agent == "lead" },
nil,
)
completion := newEinoStreamToolCallCompletionHandler(einoStreamToolCallCompletionHandlerConfig{
ConversationID: "conv-1",
OrchMode: "deep",
Progress: func(eventType, _ string, _ interface{}) { events = append(events, eventType) },
RunProgress: runProgress,
RunMessages: runMessages,
})
handler := newEinoAssistantStreamEventHandler(einoAssistantStreamEventHandlerConfig{
ConversationID: "conv-1",
OrchMode: "deep",
RunMessages: runMessages,
StreamsMainAssistant: func(string) bool { return true },
ToolCallCompletion: completion,
})
mv := &adk.MessageVariant{
IsStreaming: true,
Role: schema.Assistant,
MessageStream: schema.StreamReaderFromArray([]*schema.Message{
{Role: schema.Assistant, ToolCalls: []schema.ToolCall{{ID: "call-1", Index: &idx, Type: "function", Function: schema.FunctionCall{Name: "execute", Arguments: `{"command":`}}}},
{Role: schema.Assistant, ToolCalls: []schema.ToolCall{{Index: &idx, Function: schema.FunctionCall{Arguments: `"pwd"}`}}}},
}),
}
handled, err := handler.Handle(mv, "lead")
if !handled || err != nil {
t.Fatalf("handled=%v err=%v", handled, err)
}
msgs := runMessages.Messages()
if len(msgs) != 1 || len(msgs[0].ToolCalls) != 1 || msgs[0].ToolCalls[0].Function.Arguments != `{"command":"pwd"}` {
t.Fatalf("run messages = %#v, want merged tool call", msgs)
}
if !containsString(events, "tool_call") {
t.Fatalf("events = %#v, want tool_call", events)
}
}
func TestEinoAssistantStreamEventHandlerIgnoresToolStream(t *testing.T) {
handler := newEinoAssistantStreamEventHandler(einoAssistantStreamEventHandlerConfig{})
handled, err := handler.Handle(&adk.MessageVariant{IsStreaming: true, Role: schema.Tool, MessageStream: schema.StreamReaderFromArray([]*schema.Message{})}, "lead")
if handled || err != nil {
t.Fatalf("handled=%v err=%v, want ignored", handled, err)
}
}
@@ -4,6 +4,7 @@ import (
"cyberstrike-ai/internal/config"
"github.com/cloudwego/eino/adk"
"github.com/cloudwego/eino/schema"
"go.uber.org/zap"
)
@@ -24,18 +25,19 @@ import (
// 11. telemetry
// 12. model-facing trace snapshot
type einoChatModelTailConfig struct {
logger *zap.Logger
phase string
summarization adk.ChatModelAgentMiddleware
modelName string
maxTotalTokens int
toolMaxBytes int
conversationID string
trace *modelFacingTraceHolder
middlewareConfig *config.MultiAgentEinoMiddlewareConfig
skipOrphanPruner bool
skipTelemetry bool
skipTrace bool
logger *zap.Logger
phase string
summarization adk.ChatModelAgentMiddleware
agenticSummarization adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage]
modelName string
maxTotalTokens int
toolMaxBytes int
conversationID string
trace *modelFacingTraceHolder
middlewareConfig *config.MultiAgentEinoMiddlewareConfig
skipOrphanPruner bool
skipTelemetry bool
skipTrace bool
}
func appendEinoChatModelTailMiddlewares(handlers []adk.ChatModelAgentMiddleware, cfg einoChatModelTailConfig) []adk.ChatModelAgentMiddleware {
@@ -65,7 +67,6 @@ func appendEinoChatModelTailMiddlewares(handlers []adk.ChatModelAgentMiddleware,
handlers = append(handlers, capMw)
}
}
handlers = append(handlers, newModelOutputGuardMiddleware(cfg.middlewareConfig, cfg.logger, cfg.phase))
return handlers
}
@@ -0,0 +1,71 @@
package multiagent
import (
"context"
"github.com/cloudwego/eino/adk"
"go.uber.org/zap"
)
type einoCheckpointResumeHandlerConfig struct {
Context context.Context
ConversationID string
OrchMode string
Progress func(eventType, message string, data interface{})
Logger *zap.Logger
Store *fileCheckPointStore
CheckPointID string
Resume func(checkPointID string) (*adk.AsyncIterator[*adk.AgentEvent], error)
}
type einoCheckpointResumeHandler struct {
cfg einoCheckpointResumeHandlerConfig
}
func newEinoCheckpointResumeHandler(cfg einoCheckpointResumeHandlerConfig) *einoCheckpointResumeHandler {
if cfg.Context == nil {
cfg.Context = context.Background()
}
return &einoCheckpointResumeHandler{cfg: cfg}
}
func (h *einoCheckpointResumeHandler) TryResume() *adk.AsyncIterator[*adk.AgentEvent] {
if h == nil || h.cfg.Store == nil || h.cfg.CheckPointID == "" || h.cfg.Resume == nil {
return nil
}
if _, existed, err := h.cfg.Store.Get(h.cfg.Context, h.cfg.CheckPointID); err != nil {
if h.cfg.Logger != nil {
h.cfg.Logger.Warn("eino checkpoint preflight get failed", zap.String("checkPointID", h.cfg.CheckPointID), zap.Error(err))
}
return nil
} else if !existed {
return nil
}
h.emitProgress("检测到断点,正在从中断节点恢复执行...")
if h.cfg.Logger != nil {
h.cfg.Logger.Info("eino runner: resume from checkpoint", zap.String("checkPointID", h.cfg.CheckPointID))
}
iter, err := h.cfg.Resume(h.cfg.CheckPointID)
if err == nil {
return iter
}
if h.cfg.Logger != nil {
h.cfg.Logger.Warn("eino runner: resume failed, fallback to fresh run",
zap.String("checkPointID", h.cfg.CheckPointID),
zap.Error(err))
}
h.emitProgress("断点恢复失败,已回退为全新执行。")
return nil
}
func (h *einoCheckpointResumeHandler) emitProgress(message string) {
if h == nil || h.cfg.Progress == nil {
return
}
h.cfg.Progress("progress", message, map[string]interface{}{
"conversationId": h.cfg.ConversationID,
"source": "eino",
"orchestration": h.cfg.OrchMode,
"checkPointID": h.cfg.CheckPointID,
})
}
@@ -0,0 +1,138 @@
package multiagent
import (
"context"
"errors"
"testing"
"github.com/cloudwego/eino/adk"
"go.uber.org/zap"
"go.uber.org/zap/zaptest/observer"
)
func TestEinoCheckpointResumeHandlerSkipsWithoutCheckpoint(t *testing.T) {
called := false
handler := newEinoCheckpointResumeHandler(einoCheckpointResumeHandlerConfig{
Resume: func(string) (*adk.AsyncIterator[*adk.AgentEvent], error) {
called = true
return nil, nil
},
})
if iter := handler.TryResume(); iter != nil {
t.Fatalf("iter = %#v, want nil", iter)
}
if called {
t.Fatal("resume should not be called without checkpoint state")
}
}
func TestEinoCheckpointResumeHandlerResumesExistingCheckpoint(t *testing.T) {
store, err := newFileCheckPointStore(t.TempDir())
if err != nil {
t.Fatal(err)
}
if err := store.Set(context.Background(), "cp-1", []byte("checkpoint")); err != nil {
t.Fatal(err)
}
var progressMessages []string
var resumedID string
core, logs := observer.New(zap.InfoLevel)
wantIter, gen := adk.NewAsyncIteratorPair[*adk.AgentEvent]()
defer gen.Close()
handler := newEinoCheckpointResumeHandler(einoCheckpointResumeHandlerConfig{
Context: context.Background(),
ConversationID: "conv-1",
OrchMode: "deep",
Store: store,
CheckPointID: "cp-1",
Logger: zap.New(core),
Progress: func(eventType, message string, data interface{}) {
if eventType != "progress" {
return
}
progressMessages = append(progressMessages, message)
m, _ := data.(map[string]interface{})
if m["conversationId"] != "conv-1" || m["orchestration"] != "deep" || m["checkPointID"] != "cp-1" {
t.Fatalf("progress data = %#v", m)
}
},
Resume: func(checkPointID string) (*adk.AsyncIterator[*adk.AgentEvent], error) {
resumedID = checkPointID
return wantIter, nil
},
})
got := handler.TryResume()
if got != wantIter {
t.Fatalf("iter = %#v, want resume iterator", got)
}
if resumedID != "cp-1" {
t.Fatalf("resumed id = %q", resumedID)
}
if len(progressMessages) != 1 || progressMessages[0] != "检测到断点,正在从中断节点恢复执行..." {
t.Fatalf("progress messages = %#v", progressMessages)
}
if logs.FilterMessage("eino runner: resume from checkpoint").Len() != 1 {
t.Fatalf("expected resume log, got %d", logs.Len())
}
}
func TestEinoCheckpointResumeHandlerFallsBackOnResumeError(t *testing.T) {
store, err := newFileCheckPointStore(t.TempDir())
if err != nil {
t.Fatal(err)
}
if err := store.Set(context.Background(), "cp-1", []byte("checkpoint")); err != nil {
t.Fatal(err)
}
var progressMessages []string
core, logs := observer.New(zap.WarnLevel)
handler := newEinoCheckpointResumeHandler(einoCheckpointResumeHandlerConfig{
Context: context.Background(),
Store: store,
CheckPointID: "cp-1",
Logger: zap.New(core),
Progress: func(eventType, message string, _ interface{}) {
if eventType == "progress" {
progressMessages = append(progressMessages, message)
}
},
Resume: func(string) (*adk.AsyncIterator[*adk.AgentEvent], error) {
return nil, errors.New("resume failed")
},
})
if iter := handler.TryResume(); iter != nil {
t.Fatalf("iter = %#v, want nil fallback", iter)
}
if len(progressMessages) != 2 || progressMessages[1] != "断点恢复失败,已回退为全新执行。" {
t.Fatalf("progress messages = %#v", progressMessages)
}
if logs.FilterMessage("eino runner: resume failed, fallback to fresh run").Len() != 1 {
t.Fatalf("expected fallback log, got %d", logs.Len())
}
}
func TestEinoCheckpointResumeHandlerLogsPreflightError(t *testing.T) {
store, err := newFileCheckPointStore(t.TempDir())
if err != nil {
t.Fatal(err)
}
core, logs := observer.New(zap.WarnLevel)
handler := newEinoCheckpointResumeHandler(einoCheckpointResumeHandlerConfig{
Context: context.Background(),
Store: store,
CheckPointID: "bad/id",
Logger: zap.New(core),
Resume: func(string) (*adk.AsyncIterator[*adk.AgentEvent], error) {
t.Fatal("resume should not be called after preflight error")
return nil, nil
},
})
if iter := handler.TryResume(); iter != nil {
t.Fatalf("iter = %#v, want nil", iter)
}
if logs.FilterMessage("eino checkpoint preflight get failed").Len() != 1 {
t.Fatalf("expected preflight warning, got %d", logs.Len())
}
}
@@ -0,0 +1,38 @@
package multiagent
import (
"path/filepath"
"strings"
"go.uber.org/zap"
)
type einoCheckpointRuntime struct {
Store *fileCheckPointStore
CheckPointID string
}
func newEinoCheckpointRuntime(checkpointDir, conversationID, orchMode string, logger *zap.Logger) *einoCheckpointRuntime {
checkpointDir = strings.TrimSpace(checkpointDir)
if checkpointDir == "" {
return nil
}
cpDir := filepath.Join(checkpointDir, sanitizeEinoPathSegment(conversationID))
store, err := newFileCheckPointStore(cpDir)
if err != nil {
if logger != nil {
logger.Warn("eino checkpoint store disabled", zap.String("dir", cpDir), zap.Error(err))
}
return nil
}
checkPointID := buildEinoCheckpointID(orchMode)
if logger != nil {
logger.Info("eino runner: checkpoint store enabled",
zap.String("dir", cpDir),
zap.String("checkPointID", checkPointID))
}
return &einoCheckpointRuntime{
Store: store,
CheckPointID: checkPointID,
}
}
@@ -0,0 +1,48 @@
package multiagent
import (
"os"
"strings"
"testing"
"go.uber.org/zap"
"go.uber.org/zap/zaptest/observer"
)
func TestNewEinoCheckpointRuntimeDisabledWithoutDir(t *testing.T) {
if got := newEinoCheckpointRuntime(" ", "conv-1", "deep", nil); got != nil {
t.Fatalf("runtime = %#v, want nil", got)
}
}
func TestNewEinoCheckpointRuntimeCreatesStore(t *testing.T) {
core, logs := observer.New(zap.InfoLevel)
runtime := newEinoCheckpointRuntime(t.TempDir(), "conv/1", "deep", zap.New(core))
if runtime == nil || runtime.Store == nil {
t.Fatal("expected checkpoint runtime with store")
}
if runtime.CheckPointID != buildEinoCheckpointID("deep") {
t.Fatalf("checkpoint id = %q", runtime.CheckPointID)
}
if !strings.Contains(runtime.Store.dir, sanitizeEinoPathSegment("conv/1")) {
t.Fatalf("store dir = %q, want sanitized conversation segment", runtime.Store.dir)
}
if logs.FilterMessage("eino runner: checkpoint store enabled").Len() != 1 {
t.Fatalf("expected enabled log, got %d", logs.Len())
}
}
func TestNewEinoCheckpointRuntimeLogsCreateFailure(t *testing.T) {
filePath := t.TempDir() + "/not-a-dir"
if err := os.WriteFile(filePath, []byte("x"), 0o600); err != nil {
t.Fatal(err)
}
core, logs := observer.New(zap.WarnLevel)
runtime := newEinoCheckpointRuntime(filePath, "conv-1", "deep", zap.New(core))
if runtime != nil {
t.Fatalf("runtime = %#v, want nil", runtime)
}
if logs.FilterMessage("eino checkpoint store disabled").Len() != 1 {
t.Fatalf("expected disabled log, got %d", logs.Len())
}
}
@@ -0,0 +1,90 @@
package multiagent
import (
"context"
"github.com/cloudwego/eino/adk"
"go.uber.org/zap"
)
type einoContextOverflowRetryConfig struct {
Context context.Context
ConversationID string
OrchMode string
Args *einoADKRunLoopArgs
BaseMsgs []adk.Message
Progress func(eventType, message string, data interface{})
Logger *zap.Logger
}
type einoContextOverflowRetryResult struct {
Handled bool
RestartMsgs []adk.Message
ContextSrc einoRunRestartContextSource
}
type einoContextOverflowRetryHandler struct {
cfg einoContextOverflowRetryConfig
retried bool
}
func newEinoContextOverflowRetryHandler(cfg einoContextOverflowRetryConfig) *einoContextOverflowRetryHandler {
if cfg.Context == nil {
cfg.Context = context.Background()
}
if cfg.Args == nil {
cfg.Args = &einoADKRunLoopArgs{}
}
return &einoContextOverflowRetryHandler{cfg: cfg}
}
func (h *einoContextOverflowRetryHandler) Prepare(
runErr error,
accumulated []adk.Message,
baseCount int,
) einoContextOverflowRetryResult {
if h == nil || !isEinoContextOverflowError(runErr) || h.retried {
return einoContextOverflowRetryResult{}
}
h.retried = true
restartMsgs, ctxSource := einoMessagesForRunRestart(h.cfg.Args, h.cfg.BaseMsgs, accumulated, baseCount)
restartMsgs = aggressiveCompactMessagesForOverflow(
h.cfg.Context,
restartMsgs,
h.cfg.Args.MaxTotalTokens,
h.cfg.Args.ModelName,
h.cfg.Args.ToolMaxBytes,
h.cfg.OrchMode,
h.cfg.Logger,
)
if h.cfg.Logger != nil {
h.cfg.Logger.Warn("eino context overflow, retrying with aggressive compaction",
zap.Error(runErr),
zap.String("orchestration", h.cfg.OrchMode),
zap.String("contextSource", string(ctxSource)),
)
}
emitEinoContextOverflowRetryProgress(h.cfg.Progress, h.cfg.ConversationID, h.cfg.OrchMode, ctxSource)
return einoContextOverflowRetryResult{
Handled: true,
RestartMsgs: restartMsgs,
ContextSrc: ctxSource,
}
}
func emitEinoContextOverflowRetryProgress(
progress func(eventType, message string, data interface{}),
conversationID, orchMode string,
ctxSource einoRunRestartContextSource,
) bool {
if progress == nil {
return false
}
progress("eino_context_overflow_retry", "上下文超限,正在激进压缩后重试…", map[string]interface{}{
"conversationId": conversationID,
"source": "eino",
"orchestration": orchMode,
"contextSource": string(ctxSource),
})
return true
}
@@ -0,0 +1,90 @@
package multiagent
import (
"context"
"errors"
"testing"
"github.com/cloudwego/eino/adk"
"github.com/cloudwego/eino/schema"
"go.uber.org/zap"
"go.uber.org/zap/zaptest/observer"
)
func TestEinoContextOverflowRetryHandlerPreparesOnce(t *testing.T) {
baseMsgs := []adk.Message{
schema.UserMessage("base"),
}
accumulated := []adk.Message{
schema.UserMessage("base"),
schema.AssistantMessage("partial", nil),
}
var gotType, gotMessage string
var gotData map[string]interface{}
core, logs := observer.New(zap.WarnLevel)
handler := newEinoContextOverflowRetryHandler(einoContextOverflowRetryConfig{
Context: context.Background(),
ConversationID: "conv-1",
OrchMode: "deep_agent",
Args: &einoADKRunLoopArgs{},
BaseMsgs: baseMsgs,
Progress: func(eventType, message string, data interface{}) {
gotType = eventType
gotMessage = message
var ok bool
gotData, ok = data.(map[string]interface{})
if !ok {
t.Fatalf("progress data type = %T, want map[string]interface{}", data)
}
},
Logger: zap.New(core),
})
result := handler.Prepare(errors.New("context length exceeded"), accumulated, len(baseMsgs))
if !result.Handled {
t.Fatal("handled = false, want true")
}
if result.ContextSrc != einoRestartContextAccumulated {
t.Fatalf("context source = %q, want %q", result.ContextSrc, einoRestartContextAccumulated)
}
if len(result.RestartMsgs) != len(accumulated) {
t.Fatalf("restart message count = %d, want %d", len(result.RestartMsgs), len(accumulated))
}
if gotType != "eino_context_overflow_retry" {
t.Fatalf("event type = %q, want eino_context_overflow_retry", gotType)
}
if gotMessage != "上下文超限,正在激进压缩后重试…" {
t.Fatalf("message = %q", gotMessage)
}
assertContextOverflowMapValue(t, gotData, "conversationId", "conv-1")
assertContextOverflowMapValue(t, gotData, "source", "eino")
assertContextOverflowMapValue(t, gotData, "orchestration", "deep_agent")
assertContextOverflowMapValue(t, gotData, "contextSource", string(einoRestartContextAccumulated))
if logs.FilterMessage("eino context overflow, retrying with aggressive compaction").Len() != 1 {
t.Fatalf("expected one context overflow retry log, got %d", logs.Len())
}
second := handler.Prepare(errors.New("maximum context length"), accumulated, len(baseMsgs))
if second.Handled {
t.Fatalf("second result = %+v, want unhandled after first retry", second)
}
}
func TestEinoContextOverflowRetryHandlerIgnoresOtherErrors(t *testing.T) {
handler := newEinoContextOverflowRetryHandler(einoContextOverflowRetryConfig{
Context: context.Background(),
Args: &einoADKRunLoopArgs{},
BaseMsgs: []adk.Message{schema.UserMessage("base")},
})
result := handler.Prepare(errors.New("HTTP 429 Too Many Requests"), nil, 0)
if result.Handled {
t.Fatalf("result = %+v, want unhandled", result)
}
}
func assertContextOverflowMapValue(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)
}
}
@@ -0,0 +1,57 @@
package multiagent
import (
"strings"
"sync"
)
type einoExecuteStdoutSuppressor struct {
mu sync.Mutex
pending string
}
func newEinoExecuteStdoutSuppressor() *einoExecuteStdoutSuppressor {
return &einoExecuteStdoutSuppressor{}
}
func (s *einoExecuteStdoutSuppressor) Record(toolName, stdout string, isErr bool) {
if s == nil || isErr || !strings.EqualFold(strings.TrimSpace(toolName), "execute") {
return
}
t := strings.TrimSpace(stdout)
if t == "" {
return
}
s.mu.Lock()
s.pending = t
s.mu.Unlock()
}
func (s *einoExecuteStdoutSuppressor) Peek() string {
if s == nil {
return ""
}
s.mu.Lock()
defer s.mu.Unlock()
return s.pending
}
func (s *einoExecuteStdoutSuppressor) Consume() string {
if s == nil {
return ""
}
s.mu.Lock()
defer s.mu.Unlock()
out := s.pending
s.pending = ""
return out
}
func (s *einoExecuteStdoutSuppressor) Clear() {
if s == nil {
return
}
s.mu.Lock()
s.pending = ""
s.mu.Unlock()
}
@@ -0,0 +1,42 @@
package multiagent
import "testing"
func TestEinoExecuteStdoutSuppressorRecordsOnlySuccessfulExecute(t *testing.T) {
s := newEinoExecuteStdoutSuppressor()
s.Record("read_file", "file body", false)
if got := s.Peek(); got != "" {
t.Fatalf("non-execute should not be recorded, got %q", got)
}
s.Record("execute", "failed", true)
if got := s.Peek(); got != "" {
t.Fatalf("failed execute should not be recorded, got %q", got)
}
s.Record(" execute ", " hello\n", false)
if got := s.Peek(); got != "hello" {
t.Fatalf("Peek = %q, want hello", got)
}
}
func TestEinoExecuteStdoutSuppressorConsumeAndClear(t *testing.T) {
s := newEinoExecuteStdoutSuppressor()
s.Record("execute", "stdout", false)
if got := s.Peek(); got != "stdout" {
t.Fatalf("Peek = %q, want stdout", got)
}
if got := s.Peek(); got != "stdout" {
t.Fatalf("Peek should not clear, got %q", got)
}
if got := s.Consume(); got != "stdout" {
t.Fatalf("Consume = %q, want stdout", got)
}
if got := s.Peek(); got != "" {
t.Fatalf("Consume should clear, got %q", got)
}
s.Record("execute", "again", false)
s.Clear()
if got := s.Consume(); got != "" {
t.Fatalf("Clear should remove pending value, got %q", got)
}
}
@@ -0,0 +1,82 @@
package multiagent
import (
"context"
"testing"
"cyberstrike-ai/internal/agent"
"cyberstrike-ai/internal/config"
"cyberstrike-ai/internal/einomcp"
"cyberstrike-ai/internal/mcp"
"github.com/cloudwego/eino/adk"
"github.com/cloudwego/eino/schema"
"go.uber.org/zap"
)
func TestEinoADKFilesystemToolMonitorBindsFinishesAndUpdatesDisplayResult(t *testing.T) {
t.Parallel()
ctx := context.Background()
logger := zap.NewNop()
server := mcp.NewServer(logger)
ag := agent.NewAgent(&config.OpenAIConfig{}, &config.AgentConfig{}, server, nil, logger, 1)
binder := NewMCPExecutionBinder()
var recorded []string
rec := einomcp.ExecutionRecorder(func(executionID, toolCallID string) {
recorded = append(recorded, executionID+"|"+toolCallID)
})
beginEinoADKFilesystemToolMonitor(ctx, ag, rec, binder, "call-read", "read_file")
execID := binder.ExecutionID("call-read")
if execID == "" {
t.Fatal("expected begin to bind execution id")
}
exec, ok := server.GetExecution(execID)
if !ok || exec == nil || exec.Status != "running" || exec.ToolName != "eino_fs::read_file" {
t.Fatalf("begin execution = %#v ok=%v", exec, ok)
}
if len(recorded) != 1 || recorded[0] != execID+"|call-read" {
t.Fatalf("recorded begin ids = %#v", recorded)
}
runMessages := newEinoRunMessageAccumulator([]adk.Message{
&schema.Message{
Role: schema.Assistant,
ToolCalls: []schema.ToolCall{{
ID: "call-read",
Type: "function",
Function: schema.FunctionCall{
Name: "read_file",
Arguments: `{"path":"/tmp/secret.txt"}`,
},
}},
},
})
emitter := newEinoToolResultProgressEmitter(einoToolResultProgressEmitterConfig{
ConversationID: "conv-1",
RunMessages: runMessages,
FilesystemMonitorAgent: ag,
FilesystemMonitorRecord: rec,
MCPExecutionBinder: binder,
})
if !emitter.Emit(ctx, "read_file", "model-facing truncated body", "call-read", false, "lead") {
t.Fatal("expected tool_result emit")
}
exec, ok = server.GetExecution(execID)
if !ok || exec == nil {
t.Fatalf("finished execution missing: ok=%v exec=%#v", ok, exec)
}
if exec.Status != "completed" || exec.ToolName != "eino_fs::read_file" {
t.Fatalf("finished execution status/name = %#v", exec)
}
if got, _ := exec.Arguments["path"].(string); got != "/tmp/secret.txt" {
t.Fatalf("execution args = %#v", exec.Arguments)
}
if exec.Result == nil || len(exec.Result.Content) != 1 || exec.Result.Content[0].Text != "model-facing truncated body" {
t.Fatalf("execution display result = %#v", exec.Result)
}
if len(recorded) != 1 {
t.Fatalf("finish should reuse existing execution without recording a second id, got %#v", recorded)
}
}
@@ -0,0 +1,54 @@
package multiagent
import "github.com/cloudwego/eino/adk"
type einoAgentEventIteratorStarter func([]adk.Message) *adk.AsyncIterator[*adk.AgentEvent]
type einoInitialIteratorStartHandlerConfig struct {
ConversationID string
OrchMode string
Progress func(eventType, message string, data interface{})
UseTurnLoop bool
StartRunner einoAgentEventIteratorStarter
StartTurnLoop einoAgentEventIteratorStarter
}
type einoInitialIteratorStartHandler struct {
cfg einoInitialIteratorStartHandlerConfig
}
func newEinoInitialIteratorStartHandler(cfg einoInitialIteratorStartHandlerConfig) *einoInitialIteratorStartHandler {
return &einoInitialIteratorStartHandler{cfg: cfg}
}
func (h *einoInitialIteratorStartHandler) StartIfNeeded(existing *adk.AsyncIterator[*adk.AgentEvent], msgs []adk.Message) *adk.AsyncIterator[*adk.AgentEvent] {
if existing != nil {
return existing
}
if h == nil {
return nil
}
if h.cfg.UseTurnLoop {
h.emitTurnLoopTakeover()
if h.cfg.StartTurnLoop == nil {
return nil
}
return h.cfg.StartTurnLoop(msgs)
}
if h.cfg.StartRunner == nil {
return nil
}
return h.cfg.StartRunner(msgs)
}
func (h *einoInitialIteratorStartHandler) emitTurnLoopTakeover() {
if h == nil || h.cfg.Progress == nil {
return
}
h.cfg.Progress("progress", "Eino TurnLoop 常驻多轮 runtime 已接管本轮会话。", map[string]interface{}{
"conversationId": h.cfg.ConversationID,
"source": "eino",
"orchestration": h.cfg.OrchMode,
"kind": "turn_loop_takeover",
})
}
@@ -0,0 +1,111 @@
package multiagent
import (
"testing"
"github.com/cloudwego/eino/adk"
)
func TestEinoInitialIteratorStartHandlerKeepsExistingIterator(t *testing.T) {
existing, gen := adk.NewAsyncIteratorPair[*adk.AgentEvent]()
defer gen.Close()
var started bool
got := newEinoInitialIteratorStartHandler(einoInitialIteratorStartHandlerConfig{
UseTurnLoop: true,
StartTurnLoop: func([]adk.Message) *adk.AsyncIterator[*adk.AgentEvent] {
started = true
iter, iterGen := adk.NewAsyncIteratorPair[*adk.AgentEvent]()
iterGen.Close()
return iter
},
Progress: func(string, string, interface{}) {
t.Fatal("progress should not be emitted when an iterator already exists")
},
}).StartIfNeeded(existing, nil)
if got != existing {
t.Fatal("existing iterator should be preserved")
}
if started {
t.Fatal("start function should not be called when an iterator already exists")
}
}
func TestEinoInitialIteratorStartHandlerStartsRunner(t *testing.T) {
wantIter, gen := adk.NewAsyncIteratorPair[*adk.AgentEvent]()
defer gen.Close()
var runnerStarted bool
got := newEinoInitialIteratorStartHandler(einoInitialIteratorStartHandlerConfig{
StartRunner: func(msgs []adk.Message) *adk.AsyncIterator[*adk.AgentEvent] {
runnerStarted = true
if msgs == nil {
t.Fatal("msgs should be forwarded")
}
return wantIter
},
StartTurnLoop: func([]adk.Message) *adk.AsyncIterator[*adk.AgentEvent] {
t.Fatal("turn loop should not start when UseTurnLoop is false")
return nil
},
Progress: func(string, string, interface{}) {
t.Fatal("runner start should not emit TurnLoop takeover progress")
},
}).StartIfNeeded(nil, []adk.Message{})
if !runnerStarted {
t.Fatal("runner start was not called")
}
if got != wantIter {
t.Fatal("runner iterator should be returned")
}
}
func TestEinoInitialIteratorStartHandlerStartsTurnLoopWithTakeoverProgress(t *testing.T) {
wantIter, gen := adk.NewAsyncIteratorPair[*adk.AgentEvent]()
defer gen.Close()
var turnLoopStarted bool
var gotType, gotMessage string
var gotData map[string]interface{}
got := newEinoInitialIteratorStartHandler(einoInitialIteratorStartHandlerConfig{
ConversationID: "conv-1",
OrchMode: "deep",
UseTurnLoop: true,
StartRunner: func([]adk.Message) *adk.AsyncIterator[*adk.AgentEvent] {
t.Fatal("runner should not start when UseTurnLoop is true")
return nil
},
StartTurnLoop: func(msgs []adk.Message) *adk.AsyncIterator[*adk.AgentEvent] {
turnLoopStarted = true
if msgs == nil {
t.Fatal("msgs should be forwarded")
}
return wantIter
},
Progress: func(eventType, message string, data interface{}) {
gotType = eventType
gotMessage = message
if m, ok := data.(map[string]interface{}); ok {
gotData = m
}
},
}).StartIfNeeded(nil, []adk.Message{})
if !turnLoopStarted {
t.Fatal("turn loop start was not called")
}
if got != wantIter {
t.Fatal("turn loop iterator should be returned")
}
if gotType != "progress" {
t.Fatalf("progress type = %q, want progress", gotType)
}
if gotMessage != "Eino TurnLoop 常驻多轮 runtime 已接管本轮会话。" {
t.Fatalf("progress message = %q", gotMessage)
}
if gotData["conversationId"] != "conv-1" || gotData["source"] != "eino" || gotData["orchestration"] != "deep" {
t.Fatalf("progress data = %#v", gotData)
}
}
@@ -0,0 +1,49 @@
package multiagent
import "strings"
type einoMainAssistantCompleteHandler struct {
agentName string
emitter *einoMainResponseStreamEmitter
stdoutSuppressor *einoExecuteStdoutSuppressor
assistantOutput *einoAssistantOutputAccumulator
}
type einoMainAssistantCompleteHandlerConfig struct {
AgentName string
Emitter *einoMainResponseStreamEmitter
StdoutSuppressor *einoExecuteStdoutSuppressor
AssistantOutput *einoAssistantOutputAccumulator
}
func newEinoMainAssistantCompleteHandler(cfg einoMainAssistantCompleteHandlerConfig) *einoMainAssistantCompleteHandler {
return &einoMainAssistantCompleteHandler{
agentName: cfg.AgentName,
emitter: cfg.Emitter,
stdoutSuppressor: cfg.StdoutSuppressor,
assistantOutput: cfg.AssistantOutput,
}
}
func (h *einoMainAssistantCompleteHandler) EmitComplete(content string) bool {
if h == nil {
return false
}
body := strings.TrimSpace(content)
if body == "" {
return false
}
if h.stdoutSuppressor != nil {
if dup := h.stdoutSuppressor.Consume(); dup != "" && body == dup {
if h.assistantOutput != nil {
h.assistantOutput.RecordMainAssistant(h.agentName, body)
}
return false
}
}
emitted := h.emitter.EmitDelta(body, body)
if h.assistantOutput != nil {
h.assistantOutput.RecordMainAssistant(h.agentName, body)
}
return emitted
}
@@ -0,0 +1,76 @@
package multiagent
import "testing"
func TestEinoMainAssistantCompleteHandlerEmitsAndRecords(t *testing.T) {
var eventTypes []string
var messages []string
progress := func(eventType, message string, _ interface{}) {
eventTypes = append(eventTypes, eventType)
messages = append(messages, message)
}
out := newEinoAssistantOutputAccumulator("deep")
handler := newEinoMainAssistantCompleteHandler(einoMainAssistantCompleteHandlerConfig{
AgentName: "lead",
Emitter: newEinoMainResponseStreamEmitter("conv-1", "deep", "lead", "stream-1", 2, progress, nil),
AssistantOutput: out,
})
if !handler.EmitComplete(" hello ") {
t.Fatal("complete assistant should emit")
}
if len(eventTypes) != 2 || eventTypes[0] != "response_start" || eventTypes[1] != "response_delta" {
t.Fatalf("events = %#v", eventTypes)
}
if messages[1] != "hello" {
t.Fatalf("delta message = %q", messages[1])
}
if out.LastAssistant() != "hello" {
t.Fatalf("last assistant = %q", out.LastAssistant())
}
}
func TestEinoMainAssistantCompleteHandlerSuppressesDuplicateExecuteStdout(t *testing.T) {
var eventTypes []string
progress := func(eventType, _ string, _ interface{}) {
eventTypes = append(eventTypes, eventType)
}
stdoutDup := newEinoExecuteStdoutSuppressor()
stdoutDup.Record("execute", "hello", false)
out := newEinoAssistantOutputAccumulator("deep")
handler := newEinoMainAssistantCompleteHandler(einoMainAssistantCompleteHandlerConfig{
AgentName: "lead",
Emitter: newEinoMainResponseStreamEmitter("conv-1", "deep", "lead", "stream-1", 1, progress, nil),
StdoutSuppressor: stdoutDup,
AssistantOutput: out,
})
if handler.EmitComplete("hello") {
t.Fatal("duplicate execute stdout should not emit")
}
if len(eventTypes) != 0 {
t.Fatalf("events = %#v, want none", eventTypes)
}
if out.LastAssistant() != "hello" {
t.Fatalf("last assistant = %q", out.LastAssistant())
}
if stdoutDup.Peek() != "" {
t.Fatal("duplicate target should be consumed")
}
}
func TestEinoMainAssistantCompleteHandlerRecordsWithoutProgress(t *testing.T) {
out := newEinoAssistantOutputAccumulator("plan_execute")
handler := newEinoMainAssistantCompleteHandler(einoMainAssistantCompleteHandlerConfig{
AgentName: "executor",
Emitter: newEinoMainResponseStreamEmitter("conv-1", "plan_execute", "executor", "stream-1", 1, nil, nil),
AssistantOutput: out,
})
if handler.EmitComplete(`{"response":"done"}`) {
t.Fatal("nil progress should not emit")
}
if out.LastPlanExecuteExecutor() != "done" {
t.Fatalf("executor output = %q", out.LastPlanExecuteExecutor())
}
}
@@ -0,0 +1,77 @@
package multiagent
import "strings"
type einoMainAssistantStreamHandler struct {
agentName string
emitter *einoMainResponseStreamEmitter
stdoutSuppressor *einoExecuteStdoutSuppressor
assistantOutput *einoAssistantOutputAccumulator
runMessages *einoRunMessageAccumulator
buf string
dupTarget string
}
type einoMainAssistantStreamHandlerConfig struct {
AgentName string
Emitter *einoMainResponseStreamEmitter
StdoutSuppressor *einoExecuteStdoutSuppressor
AssistantOutput *einoAssistantOutputAccumulator
RunMessages *einoRunMessageAccumulator
}
func newEinoMainAssistantStreamHandler(cfg einoMainAssistantStreamHandlerConfig) *einoMainAssistantStreamHandler {
return &einoMainAssistantStreamHandler{
agentName: cfg.AgentName,
emitter: cfg.Emitter,
stdoutSuppressor: cfg.StdoutSuppressor,
assistantOutput: cfg.AssistantOutput,
runMessages: cfg.RunMessages,
}
}
func (h *einoMainAssistantStreamHandler) EmitDelta(content string) bool {
if h == nil || content == "" {
return false
}
var delta string
h.buf, delta = normalizeStreamingDelta(h.buf, content)
if delta == "" {
return false
}
if h.dupTarget == "" && h.stdoutSuppressor != nil {
h.dupTarget = h.stdoutSuppressor.Peek()
}
if h.dupTarget != "" {
return false
}
return h.emitter.EmitDelta(delta, h.buf)
}
func (h *einoMainAssistantStreamHandler) Finish() string {
if h == nil {
return ""
}
body := strings.TrimSpace(h.buf)
if body == "" {
return ""
}
if h.dupTarget != "" {
if h.stdoutSuppressor != nil {
h.stdoutSuppressor.Clear()
}
if body != h.dupTarget {
h.emitter.EmitTailFromFull(h.buf)
}
} else {
h.emitter.EmitTailFromFull(h.buf)
}
if h.assistantOutput != nil {
h.assistantOutput.RecordMainAssistant(h.agentName, body)
}
if h.runMessages != nil {
h.runMessages.AppendAssistantText(body)
}
return body
}
@@ -0,0 +1,103 @@
package multiagent
import "testing"
func TestEinoMainAssistantStreamHandlerEmitsAndRecords(t *testing.T) {
var eventTypes []string
var messages []string
progress := func(eventType, message string, _ interface{}) {
eventTypes = append(eventTypes, eventType)
messages = append(messages, message)
}
out := newEinoAssistantOutputAccumulator("deep")
runMsgs := newEinoRunMessageAccumulator(nil)
emitter := newEinoMainResponseStreamEmitter("conv-1", "deep", "lead", "stream-1", 2, progress, nil)
handler := newEinoMainAssistantStreamHandler(einoMainAssistantStreamHandlerConfig{
AgentName: "lead",
Emitter: emitter,
AssistantOutput: out,
RunMessages: runMsgs,
})
if !handler.EmitDelta("he") {
t.Fatal("first delta should emit")
}
if !handler.EmitDelta("hello") {
t.Fatal("cumulative chunk should emit tail")
}
if got := handler.Finish(); got != "hello" {
t.Fatalf("finish = %q, want hello", got)
}
if len(eventTypes) != 3 || eventTypes[0] != "response_start" || eventTypes[1] != "response_delta" || eventTypes[2] != "response_delta" {
t.Fatalf("events = %#v", eventTypes)
}
if messages[1] != "he" || messages[2] != "llo" {
t.Fatalf("delta messages = %#v", messages)
}
if out.LastAssistant() != "hello" {
t.Fatalf("last assistant = %q", out.LastAssistant())
}
if msgs := runMsgs.Messages(); len(msgs) != 1 || msgs[0].Content != "hello" {
t.Fatalf("run messages = %#v", msgs)
}
}
func TestEinoMainAssistantStreamHandlerSuppressesDuplicateExecuteStdout(t *testing.T) {
var eventTypes []string
progress := func(eventType, _ string, _ interface{}) {
eventTypes = append(eventTypes, eventType)
}
stdoutDup := newEinoExecuteStdoutSuppressor()
stdoutDup.Record("execute", "hello", false)
out := newEinoAssistantOutputAccumulator("deep")
runMsgs := newEinoRunMessageAccumulator(nil)
handler := newEinoMainAssistantStreamHandler(einoMainAssistantStreamHandlerConfig{
AgentName: "lead",
Emitter: newEinoMainResponseStreamEmitter("conv-1", "deep", "lead", "stream-1", 1, progress, nil),
StdoutSuppressor: stdoutDup,
AssistantOutput: out,
RunMessages: runMsgs,
})
if handler.EmitDelta("hello") {
t.Fatal("duplicate execute stdout should not emit delta")
}
if got := handler.Finish(); got != "hello" {
t.Fatalf("finish = %q, want hello", got)
}
if len(eventTypes) != 0 {
t.Fatalf("events = %#v, want none", eventTypes)
}
if stdoutDup.Peek() != "" {
t.Fatal("duplicate target should be cleared on finish")
}
if out.LastAssistant() != "hello" {
t.Fatalf("last assistant = %q", out.LastAssistant())
}
if msgs := runMsgs.Messages(); len(msgs) != 1 || msgs[0].Content != "hello" {
t.Fatalf("run messages = %#v", msgs)
}
}
func TestEinoMainAssistantStreamHandlerRecordsWithoutProgress(t *testing.T) {
out := newEinoAssistantOutputAccumulator("plan_execute")
runMsgs := newEinoRunMessageAccumulator(nil)
handler := newEinoMainAssistantStreamHandler(einoMainAssistantStreamHandlerConfig{
AgentName: "executor",
Emitter: newEinoMainResponseStreamEmitter("conv-1", "plan_execute", "executor", "stream-1", 1, nil, nil),
AssistantOutput: out,
RunMessages: runMsgs,
})
handler.EmitDelta(`{"response":"done"}`)
if got := handler.Finish(); got != `{"response":"done"}` {
t.Fatalf("finish = %q", got)
}
if out.LastPlanExecuteExecutor() != "done" {
t.Fatalf("executor output = %q", out.LastPlanExecuteExecutor())
}
if msgs := runMsgs.Messages(); len(msgs) != 1 || msgs[0].Content != `{"response":"done"}` {
t.Fatalf("run messages = %#v", msgs)
}
}
@@ -0,0 +1,85 @@
package multiagent
import "cyberstrike-ai/internal/openai"
type einoMainResponseStreamEmitter struct {
progress func(eventType, message string, data interface{})
snapshotMCPIDs func() []string
conversationID string
orchMode string
agentName string
streamID string
iteration int
headerSent bool
wireAccum string
}
func newEinoMainResponseStreamEmitter(
conversationID, orchMode, agentName, streamID string,
iteration int,
progress func(eventType, message string, data interface{}),
snapshotMCPIDs func() []string,
) *einoMainResponseStreamEmitter {
if snapshotMCPIDs == nil {
snapshotMCPIDs = func() []string { return nil }
}
return &einoMainResponseStreamEmitter{
progress: progress,
snapshotMCPIDs: snapshotMCPIDs,
conversationID: conversationID,
orchMode: orchMode,
agentName: agentName,
streamID: streamID,
iteration: iteration,
}
}
func (e *einoMainResponseStreamEmitter) EmitDelta(delta, accumulated string) bool {
if e == nil || e.progress == nil || delta == "" {
return false
}
e.emitStart()
e.progress("response_delta", delta, openai.WithSSEAccumulated(e.responseData(), accumulated))
e.wireAccum, _ = normalizeStreamingDelta(e.wireAccum, delta)
return true
}
func (e *einoMainResponseStreamEmitter) EmitTailFromFull(full string) bool {
if e == nil || full == "" {
return false
}
_, tail := normalizeStreamingDelta(e.wireAccum, full)
if tail == "" {
return false
}
return e.EmitDelta(tail, full)
}
func (e *einoMainResponseStreamEmitter) emitStart() {
if e.headerSent || e.progress == nil {
return
}
e.progress("response_start", "", map[string]interface{}{
"conversationId": e.conversationID,
"mcpExecutionIds": e.snapshotMCPIDs(),
"messageGeneratedBy": "eino:" + e.agentName,
"einoRole": "orchestrator",
"einoAgent": e.agentName,
"orchestration": e.orchMode,
"iteration": e.iteration,
"streamId": e.streamID,
})
e.headerSent = true
}
func (e *einoMainResponseStreamEmitter) responseData() map[string]interface{} {
return map[string]interface{}{
"conversationId": e.conversationID,
"mcpExecutionIds": e.snapshotMCPIDs(),
"einoRole": "orchestrator",
"einoAgent": e.agentName,
"orchestration": e.orchMode,
"iteration": e.iteration,
"streamId": e.streamID,
}
}
@@ -0,0 +1,65 @@
package multiagent
import (
"testing"
"cyberstrike-ai/internal/openai"
)
func TestEinoMainResponseStreamEmitterEmitsStartOnceAndTail(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 := newEinoMainResponseStreamEmitter(
"conv-1", "supervisor", "lead", "stream-1", 3, progress, func() []string { return []string{"mcp-1"} },
)
if !emitter.EmitDelta("he", "he") {
t.Fatal("first delta should be emitted")
}
if !emitter.EmitTailFromFull("hello") {
t.Fatal("tail should be emitted")
}
if emitter.EmitTailFromFull("hello") {
t.Fatal("duplicate tail should not be emitted")
}
if len(events) != 3 {
t.Fatalf("events = %#v, want start + 2 deltas", events)
}
if events[0].eventType != "response_start" {
t.Fatalf("event[0] = %s, want response_start", events[0].eventType)
}
if events[1].eventType != "response_delta" || events[1].message != "he" {
t.Fatalf("event[1] = %#v, want first delta", events[1])
}
if events[2].eventType != "response_delta" || events[2].message != "llo" {
t.Fatalf("event[2] = %#v, want tail delta", events[2])
}
if got := events[2].data[openai.SSEAccumulatedKey]; got != "hello" {
t.Fatalf("accumulated = %#v, want hello", got)
}
if got := events[0].data["messageGeneratedBy"]; got != "eino:lead" {
t.Fatalf("messageGeneratedBy = %#v", got)
}
if got := events[0].data["iteration"]; got != 3 {
t.Fatalf("iteration = %#v", got)
}
}
func TestEinoMainResponseStreamEmitterNoProgress(t *testing.T) {
emitter := newEinoMainResponseStreamEmitter("conv", "deep", "agent", "stream", 1, nil, nil)
if emitter.EmitDelta("hello", "hello") {
t.Fatal("nil progress should not emit")
}
if emitter.EmitTailFromFull("hello") {
t.Fatal("nil progress should not emit tail")
}
}
@@ -0,0 +1,115 @@
package multiagent
import (
"strings"
"github.com/cloudwego/eino/adk"
"github.com/cloudwego/eino/schema"
)
type einoMaterializedMessageEventHandlerConfig struct {
ConversationID string
OrchMode string
Progress func(eventType, message string, data interface{})
SnapshotMCPIDs func() []string
StreamsMainAssistant func(agent string) bool
EinoRoleTag func(agent string) string
RunProgress *einoRunProgressTracker
StdoutSuppressor *einoExecuteStdoutSuppressor
AssistantOutput *einoAssistantOutputAccumulator
RunMessages *einoRunMessageAccumulator
Usage *einoRunUsageAccumulator
ToolResultHandler *einoToolResultEventHandler
MarkPending func(toolCallPendingInfo)
NextMainStreamID func() string
}
type einoMaterializedMessageEventHandler struct {
conversationID string
orchMode string
progress func(eventType, message string, data interface{})
snapshotMCPIDs func() []string
streamsMainAssistant func(agent string) bool
einoRoleTag func(agent string) string
runProgress *einoRunProgressTracker
stdoutSuppressor *einoExecuteStdoutSuppressor
assistantOutput *einoAssistantOutputAccumulator
runMessages *einoRunMessageAccumulator
usage *einoRunUsageAccumulator
toolResultHandler *einoToolResultEventHandler
markPending func(toolCallPendingInfo)
nextMainStreamID func() string
}
func newEinoMaterializedMessageEventHandler(cfg einoMaterializedMessageEventHandlerConfig) *einoMaterializedMessageEventHandler {
if cfg.SnapshotMCPIDs == nil {
cfg.SnapshotMCPIDs = func() []string { return nil }
}
if cfg.StreamsMainAssistant == nil {
cfg.StreamsMainAssistant = func(string) bool { return true }
}
if cfg.EinoRoleTag == nil {
cfg.EinoRoleTag = func(string) string { return "" }
}
if cfg.NextMainStreamID == nil {
cfg.NextMainStreamID = func() string { return "eino-main" }
}
return &einoMaterializedMessageEventHandler{
conversationID: cfg.ConversationID,
orchMode: cfg.OrchMode,
progress: cfg.Progress,
snapshotMCPIDs: cfg.SnapshotMCPIDs,
streamsMainAssistant: cfg.StreamsMainAssistant,
einoRoleTag: cfg.EinoRoleTag,
runProgress: cfg.RunProgress,
stdoutSuppressor: cfg.StdoutSuppressor,
assistantOutput: cfg.AssistantOutput,
runMessages: cfg.RunMessages,
usage: cfg.Usage,
toolResultHandler: cfg.ToolResultHandler,
markPending: cfg.MarkPending,
nextMainStreamID: cfg.NextMainStreamID,
}
}
func (h *einoMaterializedMessageEventHandler) Handle(mv *adk.MessageVariant, msg adk.Message, agentName string) bool {
if h == nil || mv == nil || msg == nil {
return false
}
if h.runMessages != nil {
h.runMessages.Append(msg)
}
if msg.Role == schema.Assistant && h.usage != nil {
h.usage.AddMessage(msg)
}
if h.runProgress != nil {
h.runProgress.EmitToolCalls(mergeMessageToolCalls(msg), agentName, h.markPending)
}
if mv.Role == schema.Assistant {
newEinoReasoningStreamEmitter(h.conversationID, h.orchMode, agentName, h.einoRoleTag(agentName), h.progress, nil).EmitComplete(msg.ReasoningContent)
body := strings.TrimSpace(msg.Content)
if body != "" {
if h.streamsMainAssistant(agentName) {
newEinoMainAssistantCompleteHandler(einoMainAssistantCompleteHandlerConfig{
AgentName: agentName,
Emitter: newEinoMainResponseStreamEmitter(h.conversationID, h.orchMode, agentName, h.nextMainStreamID(), h.mainIteration(agentName), h.progress, h.snapshotMCPIDs),
StdoutSuppressor: h.stdoutSuppressor,
AssistantOutput: h.assistantOutput,
}).EmitComplete(body)
} else {
newEinoSubAgentReplyEmitter(h.conversationID, agentName, h.progress, nil).EmitComplete(body)
}
}
}
if h.toolResultHandler != nil {
h.toolResultHandler.HandleMaterialized(mv, msg, agentName)
}
return true
}
func (h *einoMaterializedMessageEventHandler) mainIteration(agentName string) int {
if h == nil || h.runProgress == nil {
return 0
}
return h.runProgress.MainIteration(agentName)
}
@@ -0,0 +1,151 @@
package multiagent
import (
"testing"
"cyberstrike-ai/internal/einomcp"
"github.com/cloudwego/eino/adk"
"github.com/cloudwego/eino/schema"
)
func TestEinoMaterializedMessageEventHandlerHandlesMainAssistant(t *testing.T) {
var events []string
runMessages := newEinoRunMessageAccumulator(nil)
assistantOutput := newEinoAssistantOutputAccumulator("deep")
usage := newEinoRunUsageAccumulator()
runProgress := newEinoRunProgressTracker(
"deep", "lead", "conv-1",
func(eventType, _ string, _ interface{}) { events = append(events, eventType) },
func(agent string) bool { return agent == "lead" },
nil,
)
handler := newEinoMaterializedMessageEventHandler(einoMaterializedMessageEventHandlerConfig{
ConversationID: "conv-1",
OrchMode: "deep",
Progress: func(eventType, _ string, _ interface{}) { events = append(events, eventType) },
RunMessages: runMessages,
Usage: usage,
AssistantOutput: assistantOutput,
RunProgress: runProgress,
StreamsMainAssistant: func(agent string) bool { return agent == "lead" },
EinoRoleTag: func(string) string { return "orchestrator" },
NextMainStreamID: func() string { return "main-complete-1" },
})
msg := schema.AssistantMessage(" done ", nil)
msg.ReasoningContent = "thought"
msg.ResponseMeta = &schema.ResponseMeta{Usage: &schema.TokenUsage{
PromptTokens: 11,
CompletionTokens: 7,
TotalTokens: 18,
}}
mv := &adk.MessageVariant{Role: schema.Assistant}
if !handler.Handle(mv, msg, "lead") {
t.Fatal("main assistant message was not handled")
}
if assistantOutput.LastAssistant() != "done" {
t.Fatalf("last assistant = %q", assistantOutput.LastAssistant())
}
if msgs := runMessages.Messages(); len(msgs) != 1 || msgs[0].Content != " done " {
t.Fatalf("run messages = %#v", msgs)
}
if got := usage.Summary(); got.ModelCalls != 1 || got.TotalTokens != 18 {
t.Fatalf("usage = %#v, want one assistant model call", got)
}
if !containsString(events, "reasoning_chain") || !containsString(events, "response_start") || !containsString(events, "response_delta") {
t.Fatalf("events = %#v, want reasoning and response events", events)
}
}
func TestEinoMaterializedMessageEventHandlerHandlesSubAssistant(t *testing.T) {
var events []string
runMessages := newEinoRunMessageAccumulator(nil)
assistantOutput := newEinoAssistantOutputAccumulator("deep")
handler := newEinoMaterializedMessageEventHandler(einoMaterializedMessageEventHandlerConfig{
ConversationID: "conv-1",
OrchMode: "deep",
Progress: func(eventType, _ string, _ interface{}) { events = append(events, eventType) },
RunMessages: runMessages,
AssistantOutput: assistantOutput,
StreamsMainAssistant: func(agent string) bool { return agent == "lead" },
EinoRoleTag: func(string) string { return "sub" },
})
if !handler.Handle(&adk.MessageVariant{Role: schema.Assistant}, schema.AssistantMessage("sub done", nil), "worker") {
t.Fatal("sub assistant message was not handled")
}
if assistantOutput.LastAssistant() != "" {
t.Fatalf("sub assistant should not update main output, got %q", assistantOutput.LastAssistant())
}
if len(runMessages.Messages()) != 1 {
t.Fatalf("run messages = %#v, want appended original message", runMessages.Messages())
}
if !containsString(events, "eino_agent_reply") {
t.Fatalf("events = %#v, want sub reply event", events)
}
}
func TestEinoMaterializedMessageEventHandlerHandlesToolCallsAndToolResult(t *testing.T) {
var events []string
var marked []toolCallPendingInfo
runMessages := newEinoRunMessageAccumulator(nil)
runProgress := newEinoRunProgressTracker(
"deep", "lead", "conv-1",
func(eventType, _ string, _ interface{}) { events = append(events, eventType) },
func(agent string) bool { return agent == "lead" },
nil,
)
toolResultEmitter := newEinoToolResultProgressEmitter(einoToolResultProgressEmitterConfig{
ConversationID: "conv-1",
Progress: func(eventType, _ string, _ interface{}) { events = append(events, eventType) },
})
toolResultHandler := newEinoToolResultEventHandler(einoToolResultEventHandlerConfig{Emitter: toolResultEmitter})
handler := newEinoMaterializedMessageEventHandler(einoMaterializedMessageEventHandlerConfig{
ConversationID: "conv-1",
OrchMode: "deep",
Progress: func(eventType, _ string, _ interface{}) { events = append(events, eventType) },
RunMessages: runMessages,
RunProgress: runProgress,
ToolResultHandler: toolResultHandler,
MarkPending: func(info toolCallPendingInfo) {
marked = append(marked, info)
},
})
toolCallMsg := &schema.Message{
Role: schema.Assistant,
ToolCalls: []schema.ToolCall{{
ID: "call-1",
Type: "function",
Function: schema.FunctionCall{
Name: "execute",
Arguments: `{"command":`,
},
}},
}
if !handler.Handle(&adk.MessageVariant{Role: schema.Assistant}, toolCallMsg, "lead") {
t.Fatal("tool call message was not handled")
}
toolMsg := schema.ToolMessage(einomcp.ToolErrorPrefix+"bad command", "call-1", schema.WithToolName("execute"))
if !handler.Handle(&adk.MessageVariant{Role: schema.Tool}, toolMsg, "lead") {
t.Fatal("tool message was not handled")
}
if !containsString(events, "tool_call") || !containsString(events, "tool_result") || containsString(events, "model_output_rejected") {
t.Fatalf("events = %#v, want real tool_call and tool_result without model-output recovery", events)
}
if len(marked) != 1 || marked[0].ToolCallID != "call-1" || marked[0].ToolName != "execute" {
t.Fatalf("marked pending = %#v", marked)
}
if len(runMessages.Messages()) != 2 {
t.Fatalf("run messages = %#v, want assistant and tool messages", runMessages.Messages())
}
}
func TestEinoMaterializedMessageEventHandlerIgnoresNil(t *testing.T) {
handler := newEinoMaterializedMessageEventHandler(einoMaterializedMessageEventHandlerConfig{})
if handler.Handle(nil, nil, "lead") {
t.Fatal("nil message should be ignored")
}
}
@@ -0,0 +1,61 @@
package multiagent
import (
"context"
"errors"
"io"
"github.com/cloudwego/eino/schema"
)
// recvEinoSchemaMessageStreamWithContext consumes an Eino schema.Message stream
// and stops promptly when ctx is canceled. EOF and nil chunks are treated as a
// normal stream boundary.
func recvEinoSchemaMessageStreamWithContext(
ctx context.Context,
stream *schema.StreamReader[*schema.Message],
buffer int,
onChunk func(*schema.Message),
) error {
if stream == nil {
return nil
}
if buffer <= 0 {
buffer = 1
}
type streamMsg struct {
chunk *schema.Message
err error
}
recvCh := make(chan streamMsg, buffer)
go func() {
defer close(recvCh)
for {
ch, rerr := stream.Recv()
recvCh <- streamMsg{chunk: ch, err: rerr}
if rerr != nil {
return
}
}
}()
for {
select {
case <-ctx.Done():
return ctx.Err()
case sm, ok := <-recvCh:
if !ok {
return nil
}
if errors.Is(sm.err, io.EOF) {
return nil
}
if sm.err != nil {
return sm.err
}
if sm.chunk == nil || onChunk == nil {
continue
}
onChunk(sm.chunk)
}
}
}
+160
View File
@@ -17,6 +17,7 @@ import (
"github.com/cloudwego/eino/adk/middlewares/plantask"
"github.com/cloudwego/eino/adk/middlewares/reduction"
"github.com/cloudwego/eino/components/tool"
"github.com/cloudwego/eino/schema"
"go.uber.org/zap"
)
@@ -149,6 +150,43 @@ func buildReductionMiddleware(ctx context.Context, mw config.MultiAgentEinoMiddl
return redMW, nil
}
func buildAgenticReductionMiddleware(
ctx context.Context,
mw config.MultiAgentEinoMiddlewareConfig,
projectID, convID string,
loc *localbk.Local,
logger *zap.Logger,
) (adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage], error) {
if loc == nil {
return nil, fmt.Errorf("agentic reduction: local backend nil")
}
root := reductionCacheRootDir(mw.ReductionRootDir, projectID, convID)
if err := os.MkdirAll(root, 0o755); err != nil {
return nil, fmt.Errorf("agentic reduction root: %w", err)
}
excl := append([]string(nil), mw.ReductionClearExclude...)
defaultExcl := []string{
"task", "transfer_to_agent", "exit", "write_todos", "skill", "tool_search",
"TaskCreate", "TaskGet", "TaskUpdate", "TaskList",
}
excl = append(excl, defaultExcl...)
redMW, err := reduction.NewTyped[*schema.AgenticMessage](ctx, &reduction.TypedConfig[*schema.AgenticMessage]{
Backend: loc,
RootDir: root,
ReadFileToolName: "read_file",
ClearExcludeTools: excl,
MaxLengthForTrunc: mw.ReductionMaxLengthForTruncEffective(),
MaxTokensForClear: int64(mw.ReductionMaxTokensForClearEffective()),
})
if err != nil {
return nil, err
}
if logger != nil {
logger.Info("eino middleware: agentic reduction enabled", zap.String("root", root))
}
return redMW, nil
}
// prependEinoMiddlewares returns handlers to prepend (outermost first) and optionally replaces tools when tool_search is used.
// toolSearchActive is true when the toolsearch middleware was mounted (dynamic tools split off); callers should pass this to
// injectToolNamesOnlyInstruction — tool_search is not part of the pre-middleware tools list, so name-scanning alone cannot detect it.
@@ -243,6 +281,97 @@ func prependEinoMiddlewares(
return outTools, extraHandlers, toolSearchActive, nil
}
func prependEinoAgenticMiddlewares(
ctx context.Context,
mw *config.MultiAgentEinoMiddlewareConfig,
place einoMWPlacement,
tools []tool.BaseTool,
einoLoc *localbk.Local,
skillsRoot string,
conversationID string,
projectID string,
logger *zap.Logger,
) (outTools []tool.BaseTool, extraHandlers []adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage], toolSearchActive bool, err error) {
if mw == nil {
return tools, nil, false, nil
}
outTools = tools
if mw.PatchToolCallsEffective() {
patchMW, perr := patchtoolcalls.NewTyped[*schema.AgenticMessage](ctx, &patchtoolcalls.Config{})
if perr != nil {
return nil, nil, false, fmt.Errorf("agentic patchtoolcalls: %w", perr)
}
extraHandlers = append(extraHandlers, patchMW)
}
if mw.ReductionEnable && einoLoc != nil {
if place == einoMWSub && !mw.ReductionSubAgents {
// skip
} else {
redMW, rerr := buildAgenticReductionMiddleware(ctx, *mw, projectID, conversationID, einoLoc, logger)
if rerr != nil {
return nil, nil, false, rerr
}
extraHandlers = append(extraHandlers, redMW)
}
}
minTools := mw.ToolSearchMinTools
if minTools <= 0 {
minTools = 20
}
alwaysVis := mw.ToolSearchAlwaysVisible
if alwaysVis <= 0 {
alwaysVis = 12
}
if mw.ToolSearchEnable && len(tools) >= minTools {
static, dynamic, split := splitToolsForToolSearchByNames(tools, mergeAlwaysVisibleToolNames(mw.ToolSearchAlwaysVisibleTools), alwaysVis)
if split && len(dynamic) > 0 {
ts, terr := toolsearch.NewTyped[*schema.AgenticMessage](ctx, &toolsearch.Config{DynamicTools: dynamic})
if terr != nil {
return nil, nil, false, fmt.Errorf("agentic toolsearch: %w", terr)
}
extraHandlers = append(extraHandlers, ts)
outTools = static
toolSearchActive = true
if logger != nil {
logger.Info("eino middleware: agentic tool_search enabled",
zap.Int("static_tools", len(static)),
zap.Int("dynamic_tools", len(dynamic)))
}
}
}
if place == einoMWMain && mw.PlantaskEnable {
if einoLoc == nil || strings.TrimSpace(skillsRoot) == "" {
if logger != nil {
logger.Warn("eino middleware: agentic plantask_enable ignored (need eino_skills + skills_dir)")
}
} else {
rel := strings.TrimSpace(mw.PlantaskRelDir)
if rel == "" {
rel = ".eino/plantask"
}
baseDir := filepath.Join(skillsRoot, rel, sanitizeEinoPathSegment(conversationID))
if mk := os.MkdirAll(baseDir, 0o755); mk != nil {
return nil, nil, toolSearchActive, fmt.Errorf("agentic plantask mkdir: %w", mk)
}
ptBE := newLocalPlantaskBackend(einoLoc)
pt, perr := plantask.NewTyped[*schema.AgenticMessage](ctx, &plantask.Config{Backend: ptBE, BaseDir: baseDir})
if perr != nil {
return nil, nil, toolSearchActive, fmt.Errorf("agentic plantask: %w", perr)
}
extraHandlers = append(extraHandlers, pt)
if logger != nil {
logger.Info("eino middleware: agentic plantask enabled", zap.String("baseDir", baseDir))
}
}
}
return outTools, extraHandlers, toolSearchActive, nil
}
func deepExtrasFromConfig(ma *config.MultiAgentConfig) (outputKey string, taskDesc func(context.Context, []adk.Agent) (string, error)) {
if ma == nil {
return "", nil
@@ -273,3 +402,34 @@ func deepExtrasFromConfig(ma *config.MultiAgentConfig) (outputKey string, taskDe
}
return outputKey, taskDesc
}
func deepAgenticExtrasFromConfig(ma *config.MultiAgentConfig) (outputKey string, taskDesc func(context.Context, []adk.TypedAgent[*schema.AgenticMessage]) (string, error)) {
if ma == nil {
return "", nil
}
mw := ma.EinoMiddleware
if k := strings.TrimSpace(mw.DeepOutputKey); k != "" {
outputKey = k
}
prefix := strings.TrimSpace(mw.TaskToolDescriptionPrefix)
if prefix != "" {
taskDesc = func(ctx context.Context, agents []adk.TypedAgent[*schema.AgenticMessage]) (string, error) {
_ = ctx
var names []string
for _, a := range agents {
if a == nil {
continue
}
n := strings.TrimSpace(a.Name(ctx))
if n != "" {
names = append(names, n)
}
}
if len(names) == 0 {
return prefix, nil
}
return prefix + "\n可用子代理(按名称 transfer / task 调用):" + strings.Join(names, "、"), nil
}
}
return outputKey, taskDesc
}
+167
View File
@@ -7,6 +7,10 @@ import (
"strings"
"testing"
"cyberstrike-ai/internal/config"
localbk "github.com/cloudwego/eino-ext/adk/backend/local"
"github.com/cloudwego/eino/adk"
"github.com/cloudwego/eino/components/tool"
"github.com/cloudwego/eino/schema"
)
@@ -28,6 +32,169 @@ func TestReductionCacheRootDir(t *testing.T) {
}
}
func TestBuildAgenticReductionMiddlewareClearsOldAgenticToolResult(t *testing.T) {
ctx := context.Background()
loc, err := localbk.NewBackend(ctx, &localbk.Config{})
if err != nil {
t.Fatalf("NewBackend: %v", err)
}
root := t.TempDir()
mw, err := buildAgenticReductionMiddleware(ctx, config.MultiAgentEinoMiddlewareConfig{
ReductionRootDir: root,
ReductionMaxTokensForClear: 1,
}, "", "conv-1", loc, nil)
if err != nil {
t.Fatalf("buildAgenticReductionMiddleware: %v", err)
}
oldText := strings.Repeat("old-tool-output-", 20)
newText := strings.Repeat("new-tool-output-", 20)
state := &adk.TypedChatModelAgentState[*schema.AgenticMessage]{
Messages: []*schema.AgenticMessage{
agenticAssistantToolCall("old-call", "execute", `{"command":"old"}`),
agenticToolResult("old-call", "execute", oldText),
agenticAssistantToolCall("new-call", "execute", `{"command":"new"}`),
agenticToolResult("new-call", "execute", newText),
},
}
_, out, err := mw.BeforeModelRewriteState(ctx, state, nil)
if err != nil {
t.Fatalf("BeforeModelRewriteState: %v", err)
}
oldGot := out.Messages[1].ContentBlocks[0].FunctionToolResult.Content[0].Text.Text
newGot := out.Messages[3].ContentBlocks[0].FunctionToolResult.Content[0].Text.Text
if oldGot == oldText {
t.Fatal("agentic reduction did not clear old oversized tool result")
}
if !strings.Contains(oldGot, "read_file") {
t.Fatalf("cleared content should mention read_file, got %q", oldGot)
}
if newGot != newText {
t.Fatalf("latest tool result should be retained, got %q", newGot)
}
}
func agenticAssistantToolCall(callID, name, arguments string) *schema.AgenticMessage {
return &schema.AgenticMessage{
Role: schema.AgenticRoleTypeAssistant,
ContentBlocks: []*schema.ContentBlock{schema.NewContentBlock(&schema.FunctionToolCall{
CallID: callID,
Name: name,
Arguments: arguments,
})},
}
}
func agenticToolResult(callID, name, text string) *schema.AgenticMessage {
return &schema.AgenticMessage{
Role: schema.AgenticRoleTypeUser,
ContentBlocks: []*schema.ContentBlock{schema.NewContentBlock(&schema.FunctionToolResult{
CallID: callID,
Name: name,
Content: []*schema.FunctionToolResultContentBlock{{
Type: schema.FunctionToolResultContentBlockTypeText,
Text: &schema.UserInputText{Text: text},
}},
})},
}
}
func TestBuildAgenticReductionMiddlewareHandlesSingleAgenticToolResult(t *testing.T) {
ctx := context.Background()
loc, err := localbk.NewBackend(ctx, &localbk.Config{})
if err != nil {
t.Fatalf("NewBackend: %v", err)
}
mw, err := buildAgenticReductionMiddleware(ctx, config.MultiAgentEinoMiddlewareConfig{
ReductionRootDir: t.TempDir(),
ReductionMaxTokensForClear: 1,
}, "", "conv-1", loc, nil)
if err != nil {
t.Fatalf("buildAgenticReductionMiddleware: %v", err)
}
state := &adk.TypedChatModelAgentState[*schema.AgenticMessage]{
Messages: []*schema.AgenticMessage{
{
Role: schema.AgenticRoleTypeUser,
ContentBlocks: []*schema.ContentBlock{schema.NewContentBlock(&schema.FunctionToolResult{
CallID: "call-1",
Name: "execute",
Content: []*schema.FunctionToolResultContentBlock{{
Type: schema.FunctionToolResultContentBlockTypeText,
Text: &schema.UserInputText{Text: strings.Repeat("tool-output-", 20)},
}},
})},
},
},
}
_, out, err := mw.BeforeModelRewriteState(ctx, state, nil)
if err != nil {
t.Fatalf("BeforeModelRewriteState: %v", err)
}
got := out.Messages[0].ContentBlocks[0].FunctionToolResult.Content[0].Text.Text
if got != strings.Repeat("tool-output-", 20) {
t.Fatalf("single retained tool result should not be cleared, got %q", got)
}
}
func TestPrependEinoAgenticMiddlewaresRespectsReductionPlacement(t *testing.T) {
ctx := context.Background()
loc, err := localbk.NewBackend(ctx, &localbk.Config{})
if err != nil {
t.Fatalf("NewBackend: %v", err)
}
patchToolCalls := false
mw := &config.MultiAgentEinoMiddlewareConfig{
ReductionEnable: true,
ReductionRootDir: t.TempDir(),
ReductionMaxTokensForClear: 100,
PatchToolCalls: &patchToolCalls,
}
_, mainHandlers, _, err := prependEinoAgenticMiddlewares(ctx, mw, einoMWMain, nil, loc, "", "conv-1", "", nil)
if err != nil {
t.Fatalf("prepend main: %v", err)
}
if len(mainHandlers) != 1 {
t.Fatalf("main handlers = %d, want reduction", len(mainHandlers))
}
_, subHandlers, _, err := prependEinoAgenticMiddlewares(ctx, mw, einoMWSub, nil, loc, "", "conv-1", "", nil)
if err != nil {
t.Fatalf("prepend sub: %v", err)
}
if len(subHandlers) != 0 {
t.Fatalf("sub handlers = %d, want skipped when reduction_sub_agents=false", len(subHandlers))
}
mw.ReductionSubAgents = true
_, subHandlers, _, err = prependEinoAgenticMiddlewares(ctx, mw, einoMWSub, nil, loc, "", "conv-1", "", nil)
if err != nil {
t.Fatalf("prepend sub enabled: %v", err)
}
if len(subHandlers) != 1 {
t.Fatalf("sub handlers = %d, want reduction when reduction_sub_agents=true", len(subHandlers))
}
}
func TestPrependEinoAgenticMiddlewaresMountsToolSearchAndPatchToolCalls(t *testing.T) {
ctx := context.Background()
mw := &config.MultiAgentEinoMiddlewareConfig{
ToolSearchEnable: true,
ToolSearchMinTools: 20,
ToolSearchAlwaysVisible: 5,
}
outTools, handlers, toolSearchActive, err := prependEinoAgenticMiddlewares(ctx, mw, einoMWMain, stubTools(25), nil, "", "conv-test", "", nil)
if err != nil {
t.Fatalf("prependEinoAgenticMiddlewares: %v", err)
}
if !toolSearchActive {
t.Fatal("agentic tool_search should be active")
}
if len(outTools) != 5 {
t.Fatalf("mounted tools = %d, want static visible tools only", len(outTools))
}
if len(handlers) != 2 {
t.Fatalf("handlers = %d, want patchtoolcalls + toolsearch", len(handlers))
}
}
type stubTool struct{ name string }
func (s stubTool) Info(_ context.Context) (*schema.ToolInfo, error) {
@@ -6,6 +6,7 @@ import (
"sync"
"github.com/cloudwego/eino/adk"
"github.com/cloudwego/eino/schema"
)
// modelFacingTraceHolder 保存「即将送入 ChatModel」的消息快照(已走 summarization / reduction / orphan 修剪等),
@@ -43,6 +44,19 @@ func (h *modelFacingTraceHolder) storeFromState(state *adk.ChatModelAgentState)
h.mu.Unlock()
}
func (h *modelFacingTraceHolder) storeFromAgenticState(state *adk.TypedChatModelAgentState[*schema.AgenticMessage]) {
if h == nil || state == nil || len(state.Messages) == 0 {
return
}
cloned := cloneADKMessagesForTrace(AgenticMessagesToEino(state.Messages))
if len(cloned) == 0 {
return
}
h.mu.Lock()
h.msgs = cloned
h.mu.Unlock()
}
func cloneADKMessagesForTrace(msgs []adk.Message) []adk.Message {
if len(msgs) == 0 {
return nil
@@ -82,3 +96,29 @@ func (m *modelFacingTraceMiddleware) BeforeModelRewriteState(
}
return ctx, state, nil
}
type agenticModelFacingTraceMiddleware struct {
*adk.TypedBaseChatModelAgentMiddleware[*schema.AgenticMessage]
holder *modelFacingTraceHolder
}
func newAgenticModelFacingTraceMiddleware(holder *modelFacingTraceHolder) adk.TypedChatModelAgentMiddleware[*schema.AgenticMessage] {
if holder == nil {
return nil
}
return &agenticModelFacingTraceMiddleware{
TypedBaseChatModelAgentMiddleware: &adk.TypedBaseChatModelAgentMiddleware[*schema.AgenticMessage]{},
holder: holder,
}
}
func (m *agenticModelFacingTraceMiddleware) BeforeModelRewriteState(
ctx context.Context,
state *adk.TypedChatModelAgentState[*schema.AgenticMessage],
mc *adk.TypedModelContext[*schema.AgenticMessage],
) (context.Context, *adk.TypedChatModelAgentState[*schema.AgenticMessage], error) {
if m.holder != nil && state != nil {
m.holder.storeFromAgenticState(state)
}
return ctx, state, nil
}
@@ -0,0 +1,500 @@
package multiagent
import (
"context"
"errors"
"fmt"
"net"
"net/http"
"strings"
"sync"
"time"
"cyberstrike-ai/internal/config"
"cyberstrike-ai/internal/openai"
"cyberstrike-ai/internal/reasoning"
agenticopenai "github.com/cloudwego/eino-ext/components/model/agenticopenai"
einoopenai "github.com/cloudwego/eino-ext/components/model/openai"
"github.com/cloudwego/eino/adk"
"github.com/cloudwego/eino/components/model"
"github.com/cloudwego/eino/schema"
"go.uber.org/zap"
)
type einoModelMode string
const (
einoModelModeNormal einoModelMode = "normal"
einoModelModePlanner einoModelMode = "planner"
)
type einoModelFactory func(ctx context.Context, oa config.OpenAIConfig, mode einoModelMode) (model.ToolCallingChatModel, error)
type einoAgenticModelConfigFactory func(ctx context.Context, oa config.OpenAIConfig, mode einoModelMode) (model.AgenticModel, error)
func newEinoBaseHTTPClient() *http.Client {
return &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,
},
}
}
func newEinoOpenAIChatModelFactory(
baseHTTPClient *http.Client,
reasoningClient *reasoning.ClientIntent,
logger *zap.Logger,
) einoModelFactory {
if baseHTTPClient == nil {
baseHTTPClient = newEinoBaseHTTPClient()
}
return func(ctx context.Context, oa config.OpenAIConfig, mode einoModelMode) (model.ToolCallingChatModel, error) {
httpClient := openai.NewEinoHTTPClient(&oa, baseHTTPClient)
openai.AttachSummarizationDiagTransport(httpClient, logger)
maxCompletionTokens := oa.MaxCompletionTokensEffective()
modelCfg := &einoopenai.ChatModelConfig{
APIKey: oa.APIKey,
BaseURL: strings.TrimSuffix(oa.BaseURL, "/"),
Model: oa.Model,
HTTPClient: httpClient,
MaxCompletionTokens: &maxCompletionTokens,
}
if mode == einoModelModePlanner {
reasoning.ApplyPlanExecutePlannerModelConfig(modelCfg, &oa)
} else {
reasoning.ApplyToEinoChatModelConfig(modelCfg, &oa, reasoningClient)
}
baseModel, err := einoopenai.NewChatModel(ctx, modelCfg)
if err != nil {
return nil, err
}
return newStreamToolCallIndexRepairModel(baseModel), nil
}
}
func newEinoOpenAIAgenticChatModelFactory(
baseHTTPClient *http.Client,
reasoningClient *reasoning.ClientIntent,
logger *zap.Logger,
) einoAgenticModelConfigFactory {
if baseHTTPClient == nil {
baseHTTPClient = newEinoBaseHTTPClient()
}
return func(ctx context.Context, oa config.OpenAIConfig, mode einoModelMode) (model.AgenticModel, error) {
if !supportsEinoAgenticOpenAIBackend(oa) {
return nil, fmt.Errorf("eino agentic model: provider %q is not enabled for agenticopenai backend", strings.TrimSpace(oa.Provider))
}
httpClient := openai.NewEinoHTTPClient(&oa, baseHTTPClient)
openai.AttachSummarizationDiagTransport(httpClient, logger)
maxCompletionTokens := oa.MaxCompletionTokensEffective()
modelCfg := &agenticopenai.ChatConfig{
APIKey: oa.APIKey,
BaseURL: strings.TrimSuffix(oa.BaseURL, "/"),
Model: oa.Model,
HTTPClient: httpClient,
MaxCompletionTokens: &maxCompletionTokens,
ExtraFields: reasoning.AgenticOpenAIExtraFields(&oa, reasoningClient),
}
if mode == einoModelModePlanner {
modelCfg.ExtraFields = reasoning.AgenticOpenAIPlannerExtraFields(&oa)
}
return agenticopenai.NewChatModel(ctx, modelCfg)
}
}
func supportsEinoAgenticOpenAIBackend(oa config.OpenAIConfig) bool {
provider := strings.ToLower(strings.TrimSpace(oa.Provider))
return provider == "" || provider == "openai" || provider == "openai_compatible"
}
func agenticModelGateFactory(factory einoAgenticModelConfigFactory, oa config.OpenAIConfig, mode einoModelMode) einoAgenticModelFactory {
if factory == nil {
return nil
}
return func(ctx context.Context) (model.AgenticModel, error) {
return factory(ctx, oa, mode)
}
}
func newEinoModelRetryConfig(
mw *config.MultiAgentEinoMiddlewareConfig,
logger *zap.Logger,
scope string,
) *adk.ModelRetryConfig {
maxRetries := RunRetryMaxAttemptsFromConfig(mw)
maxBackoff := einoRunRetryMaxBackoffFromConfig(mw)
return &adk.ModelRetryConfig{
MaxRetries: maxRetries,
BackoffFunc: func(_ context.Context, attempt int) time.Duration {
return einoTransientRetryBackoff(attempt-1, maxBackoff)
},
ShouldRetry: func(ctx context.Context, retryCtx *adk.RetryContext) *adk.RetryDecision {
if retryCtx == nil || ctx.Err() != nil {
return &adk.RetryDecision{}
}
if retryCtx.Err != nil {
if !isEinoTransientRunError(retryCtx.Err) {
return &adk.RetryDecision{}
}
if logger != nil {
kind, summary := einoTransientRunErrorUserDetail(retryCtx.Err)
logger.Warn("eino native model retry",
zap.String("scope", scope),
zap.Int("attempt", retryCtx.RetryAttempt),
zap.Int("maxRetries", maxRetries),
zap.String("errorKind", kind),
zap.String("errorSummary", summary),
)
}
return &adk.RetryDecision{Retry: true, RejectReason: "transient_model_error"}
}
if isRetryableEmptyModelOutput(retryCtx.OutputMessage) {
if logger != nil {
logger.Warn("eino native model retry: empty model output",
zap.String("scope", scope),
zap.Int("attempt", retryCtx.RetryAttempt),
zap.Int("maxRetries", maxRetries),
)
}
return &adk.RetryDecision{Retry: true, RejectReason: "empty_model_output"}
}
return &adk.RetryDecision{}
},
}
}
func newEinoAgenticModelRetryConfig(
mw *config.MultiAgentEinoMiddlewareConfig,
logger *zap.Logger,
scope string,
) *adk.TypedModelRetryConfig[*schema.AgenticMessage] {
maxRetries := RunRetryMaxAttemptsFromConfig(mw)
maxBackoff := einoRunRetryMaxBackoffFromConfig(mw)
return &adk.TypedModelRetryConfig[*schema.AgenticMessage]{
MaxRetries: maxRetries,
BackoffFunc: func(_ context.Context, attempt int) time.Duration {
return einoTransientRetryBackoff(attempt-1, maxBackoff)
},
ShouldRetry: func(ctx context.Context, retryCtx *adk.TypedRetryContext[*schema.AgenticMessage]) *adk.TypedRetryDecision[*schema.AgenticMessage] {
if retryCtx == nil || ctx.Err() != nil {
return &adk.TypedRetryDecision[*schema.AgenticMessage]{}
}
if retryCtx.Err != nil {
if !isEinoTransientRunError(retryCtx.Err) {
return &adk.TypedRetryDecision[*schema.AgenticMessage]{}
}
if logger != nil {
kind, summary := einoTransientRunErrorUserDetail(retryCtx.Err)
logger.Warn("eino native agentic model retry",
zap.String("scope", scope),
zap.Int("attempt", retryCtx.RetryAttempt),
zap.Int("maxRetries", maxRetries),
zap.String("errorKind", kind),
zap.String("errorSummary", summary),
)
}
return &adk.TypedRetryDecision[*schema.AgenticMessage]{Retry: true, RejectReason: "transient_model_error"}
}
if isRetryableEmptyAgenticModelOutput(retryCtx.OutputMessage) {
if logger != nil {
logger.Warn("eino native agentic model retry: empty model output",
zap.String("scope", scope),
zap.Int("attempt", retryCtx.RetryAttempt),
zap.Int("maxRetries", maxRetries),
)
}
return &adk.TypedRetryDecision[*schema.AgenticMessage]{Retry: true, RejectReason: "empty_model_output"}
}
return &adk.TypedRetryDecision[*schema.AgenticMessage]{}
},
}
}
func newEinoModelFailoverConfig(
ctx context.Context,
appCfg *config.Config,
mw *config.MultiAgentEinoMiddlewareConfig,
mode einoModelMode,
factory einoModelFactory,
logger *zap.Logger,
scope string,
progress func(eventType, message string, data interface{}),
orchestration string,
conversationID string,
) (*adk.ModelFailoverConfig[*schema.Message], error) {
channels := resolveEinoFailoverChannels(appCfg, mw)
if len(channels) == 0 {
return nil, nil
}
if factory == nil {
return nil, fmt.Errorf("eino model failover: 模型工厂为空")
}
maxRetries := len(channels)
if mw != nil && mw.ModelFailoverMaxRetries > 0 && mw.ModelFailoverMaxRetries < maxRetries {
maxRetries = mw.ModelFailoverMaxRetries
}
channels = channels[:maxRetries]
cache := make(map[string]model.BaseModel[*schema.Message], len(channels))
var mu sync.Mutex
return &adk.ModelFailoverConfig[*schema.Message]{
MaxRetries: uint(maxRetries),
ShouldFailover: func(ctx context.Context, _ *schema.Message, err error) bool {
if ctx.Err() != nil || err == nil {
return false
}
err = unwrapEinoRetryExhausted(err)
return isEinoTransientRunError(err)
},
GetFailoverModel: func(ctx context.Context, failoverCtx *adk.FailoverContext[*schema.Message]) (model.BaseModel[*schema.Message], []*schema.Message, error) {
if failoverCtx == nil || failoverCtx.FailoverAttempt == 0 {
return nil, nil, fmt.Errorf("eino model failover: invalid failover attempt")
}
idx := int(failoverCtx.FailoverAttempt) - 1
if idx < 0 || idx >= len(channels) {
return nil, nil, fmt.Errorf("eino model failover: no channel for attempt %d", failoverCtx.FailoverAttempt)
}
ch := channels[idx]
mu.Lock()
cached := cache[ch.id]
mu.Unlock()
if cached != nil {
emitEinoModelFailoverEvent(progress, conversationID, orchestration, scope, ch.id, ch.cfg.Model, failoverCtx.FailoverAttempt)
if logger != nil {
logger.Warn("eino native model failover",
zap.String("scope", scope),
zap.String("channel", ch.id),
zap.String("model", ch.cfg.Model),
zap.Uint("attempt", failoverCtx.FailoverAttempt),
)
}
return cached, nil, nil
}
m, err := factory(ctx, ch.cfg, mode)
if err != nil {
return nil, nil, fmt.Errorf("eino model failover channel %q: %w", ch.id, err)
}
mu.Lock()
cache[ch.id] = m
mu.Unlock()
emitEinoModelFailoverEvent(progress, conversationID, orchestration, scope, ch.id, ch.cfg.Model, failoverCtx.FailoverAttempt)
if logger != nil {
logger.Warn("eino native model failover",
zap.String("scope", scope),
zap.String("channel", ch.id),
zap.String("model", ch.cfg.Model),
zap.Uint("attempt", failoverCtx.FailoverAttempt),
)
}
return m, nil, nil
},
}, nil
}
func newEinoAgenticModelFailoverConfig(
ctx context.Context,
appCfg *config.Config,
mw *config.MultiAgentEinoMiddlewareConfig,
mode einoModelMode,
factory einoAgenticModelConfigFactory,
logger *zap.Logger,
scope string,
progress func(eventType, message string, data interface{}),
orchestration string,
conversationID string,
) (*adk.ModelFailoverConfig[*schema.AgenticMessage], error) {
channels := resolveEinoFailoverChannels(appCfg, mw)
if len(channels) == 0 {
return nil, nil
}
if factory == nil {
return nil, fmt.Errorf("eino agentic model failover: 模型工厂为空")
}
maxRetries := len(channels)
if mw != nil && mw.ModelFailoverMaxRetries > 0 && mw.ModelFailoverMaxRetries < maxRetries {
maxRetries = mw.ModelFailoverMaxRetries
}
channels = channels[:maxRetries]
cache := make(map[string]model.BaseModel[*schema.AgenticMessage], len(channels))
var mu sync.Mutex
return &adk.ModelFailoverConfig[*schema.AgenticMessage]{
MaxRetries: uint(maxRetries),
ShouldFailover: func(ctx context.Context, _ *schema.AgenticMessage, err error) bool {
if ctx.Err() != nil || err == nil {
return false
}
err = unwrapEinoRetryExhausted(err)
return isEinoTransientRunError(err)
},
GetFailoverModel: func(ctx context.Context, failoverCtx *adk.FailoverContext[*schema.AgenticMessage]) (model.BaseModel[*schema.AgenticMessage], []*schema.AgenticMessage, error) {
if failoverCtx == nil || failoverCtx.FailoverAttempt == 0 {
return nil, nil, fmt.Errorf("eino agentic model failover: invalid failover attempt")
}
idx := int(failoverCtx.FailoverAttempt) - 1
if idx < 0 || idx >= len(channels) {
return nil, nil, fmt.Errorf("eino agentic model failover: no channel for attempt %d", failoverCtx.FailoverAttempt)
}
ch := channels[idx]
mu.Lock()
cached := cache[ch.id]
mu.Unlock()
if cached != nil {
emitEinoModelFailoverEvent(progress, conversationID, orchestration, scope, ch.id, ch.cfg.Model, failoverCtx.FailoverAttempt)
if logger != nil {
logger.Warn("eino native agentic model failover",
zap.String("scope", scope),
zap.String("channel", ch.id),
zap.String("model", ch.cfg.Model),
zap.Uint("attempt", failoverCtx.FailoverAttempt),
)
}
return cached, nil, nil
}
m, err := factory(ctx, ch.cfg, mode)
if err != nil {
return nil, nil, fmt.Errorf("eino agentic model failover channel %q: %w", ch.id, err)
}
mu.Lock()
cache[ch.id] = m
mu.Unlock()
emitEinoModelFailoverEvent(progress, conversationID, orchestration, scope, ch.id, ch.cfg.Model, failoverCtx.FailoverAttempt)
if logger != nil {
logger.Warn("eino native agentic model failover",
zap.String("scope", scope),
zap.String("channel", ch.id),
zap.String("model", ch.cfg.Model),
zap.Uint("attempt", failoverCtx.FailoverAttempt),
)
}
return m, nil, nil
},
}, nil
}
type resolvedEinoFailoverChannel struct {
id string
cfg config.OpenAIConfig
}
func resolveEinoFailoverChannels(appCfg *config.Config, mw *config.MultiAgentEinoMiddlewareConfig) []resolvedEinoFailoverChannel {
if appCfg == nil || mw == nil || len(mw.ModelFailoverChannels) == 0 {
return nil
}
primary := appCfg.OpenAI
seen := map[string]struct{}{}
out := make([]resolvedEinoFailoverChannel, 0, len(mw.ModelFailoverChannels))
for _, raw := range mw.ModelFailoverChannels {
id := config.NormalizeAIChannelID(raw)
if id == "" {
continue
}
if _, ok := seen[id]; ok {
continue
}
oa, resolvedID, ok := appCfg.AI.ResolveChannel(id)
if !ok {
continue
}
if sameOpenAIModelEndpoint(primary, oa) {
continue
}
seen[resolvedID] = struct{}{}
out = append(out, resolvedEinoFailoverChannel{id: resolvedID, cfg: oa})
}
return out
}
func sameOpenAIModelEndpoint(a, b config.OpenAIConfig) bool {
return strings.EqualFold(strings.TrimSpace(a.Provider), strings.TrimSpace(b.Provider)) &&
strings.TrimRight(strings.TrimSpace(a.BaseURL), "/") == strings.TrimRight(strings.TrimSpace(b.BaseURL), "/") &&
strings.TrimSpace(a.APIKey) == strings.TrimSpace(b.APIKey) &&
strings.TrimSpace(a.Model) == strings.TrimSpace(b.Model)
}
func isRetryableEmptyModelOutput(msg *schema.Message) bool {
if msg == nil {
return true
}
return strings.TrimSpace(msg.Content) == "" &&
strings.TrimSpace(msg.ReasoningContent) == "" &&
len(msg.ToolCalls) == 0 &&
len(msg.MultiContent) == 0 &&
len(msg.UserInputMultiContent) == 0 &&
len(msg.AssistantGenMultiContent) == 0
}
func isRetryableEmptyAgenticModelOutput(msg *schema.AgenticMessage) bool {
if msg == nil {
return true
}
for _, block := range msg.ContentBlocks {
if block == nil {
continue
}
switch {
case block.Reasoning != nil:
if strings.TrimSpace(block.Reasoning.Text) != "" {
return false
}
case block.UserInputText != nil:
if strings.TrimSpace(block.UserInputText.Text) != "" {
return false
}
case block.AssistantGenText != nil:
if strings.TrimSpace(block.AssistantGenText.Text) != "" {
return false
}
default:
return false
}
}
return true
}
func unwrapEinoRetryExhausted(err error) error {
var retryErr *adk.RetryExhaustedError
if errors.As(err, &retryErr) && retryErr.LastErr != nil {
return retryErr.LastErr
}
return err
}
func isEinoNativeWillRetry(err error) (*adk.WillRetryError, bool) {
var willRetry *adk.WillRetryError
if errors.As(err, &willRetry) {
return willRetry, true
}
return nil, false
}
func emitEinoModelFailoverEvent(
progress func(eventType, message string, data interface{}),
conversationID, orchestration, scope, channelID, modelName string,
attempt uint,
) {
if progress == nil {
return
}
msg := fmt.Sprintf("主模型重试耗尽,正在切换备用模型 %s。", modelName)
progress("eino_model_failover", msg, map[string]interface{}{
"conversationId": conversationID,
"source": "eino",
"orchestration": orchestration,
"scope": scope,
"channel": channelID,
"model": modelName,
"attempt": attempt,
})
}
@@ -0,0 +1,376 @@
package multiagent
import (
"context"
"errors"
"testing"
"time"
"cyberstrike-ai/internal/config"
"github.com/cloudwego/eino/adk"
"github.com/cloudwego/eino/components/model"
"github.com/cloudwego/eino/schema"
)
func TestNewEinoModelRetryConfigUsesNativeFieldsFirst(t *testing.T) {
t.Parallel()
mw := &config.MultiAgentEinoMiddlewareConfig{
ModelRetryMaxRetries: 2,
ModelRetryMaxBackoffSec: 7,
RunRetryMaxAttempts: 9,
RunRetryMaxBackoffSec: 11,
}
cfg := newEinoModelRetryConfig(mw, nil, "test")
if cfg.MaxRetries != 2 {
t.Fatalf("MaxRetries = %d, want 2", cfg.MaxRetries)
}
backoff := cfg.BackoffFunc(context.Background(), 1)
if backoff < 500*time.Millisecond || backoff > 2*time.Second {
t.Fatalf("attempt 1 backoff = %v, want first equal-jitter window", backoff)
}
if got := einoRunRetryMaxBackoffFromConfig(mw); got != 7*time.Second {
t.Fatalf("backoff from config = %v, want 7s", got)
}
}
func TestEinoModelRetryPolicyRetriesTransientAndEmptyOutput(t *testing.T) {
t.Parallel()
cfg := newEinoModelRetryConfig(&config.MultiAgentEinoMiddlewareConfig{ModelRetryMaxRetries: 1}, nil, "test")
if got := cfg.ShouldRetry(context.Background(), &adk.RetryContext{Err: errors.New("HTTP 429 Too Many Requests")}); got == nil || !got.Retry {
t.Fatal("transient model error should retry")
}
if got := cfg.ShouldRetry(context.Background(), &adk.RetryContext{OutputMessage: schema.AssistantMessage("", nil)}); got == nil || !got.Retry {
t.Fatal("empty assistant output should retry")
}
if got := cfg.ShouldRetry(context.Background(), &adk.RetryContext{OutputMessage: schema.AssistantMessage("", []schema.ToolCall{{ID: "call_1"}})}); got == nil || got.Retry {
t.Fatal("assistant tool call output should not be treated as empty")
}
if got := cfg.ShouldRetry(context.Background(), &adk.RetryContext{Err: errors.New("invalid api key")}); got == nil || got.Retry {
t.Fatal("permanent auth error should not retry")
}
}
func TestEinoAgenticModelRetryPolicyRetriesTransientAndEmptyOutput(t *testing.T) {
t.Parallel()
cfg := newEinoAgenticModelRetryConfig(&config.MultiAgentEinoMiddlewareConfig{ModelRetryMaxRetries: 1}, nil, "agentic")
if got := cfg.ShouldRetry(context.Background(), &adk.TypedRetryContext[*schema.AgenticMessage]{Err: errors.New("HTTP 429 Too Many Requests")}); got == nil || !got.Retry {
t.Fatal("transient agentic model error should retry")
}
if got := cfg.ShouldRetry(context.Background(), &adk.TypedRetryContext[*schema.AgenticMessage]{
OutputMessage: &schema.AgenticMessage{Role: schema.AgenticRoleTypeAssistant},
}); got == nil || !got.Retry {
t.Fatal("empty agentic assistant output should retry")
}
if got := cfg.ShouldRetry(context.Background(), &adk.TypedRetryContext[*schema.AgenticMessage]{
OutputMessage: &schema.AgenticMessage{
Role: schema.AgenticRoleTypeAssistant,
ContentBlocks: []*schema.ContentBlock{schema.NewContentBlock(&schema.AssistantGenText{Text: "ok"})},
},
}); got == nil || got.Retry {
t.Fatal("agentic assistant text should not be treated as empty")
}
if got := cfg.ShouldRetry(context.Background(), &adk.TypedRetryContext[*schema.AgenticMessage]{
OutputMessage: &schema.AgenticMessage{
Role: schema.AgenticRoleTypeAssistant,
ContentBlocks: []*schema.ContentBlock{schema.NewContentBlock(&schema.FunctionToolCall{
CallID: "call_1", Name: "search", Arguments: `{"q":"x"}`,
})},
},
}); got == nil || got.Retry {
t.Fatal("agentic tool call output should not be treated as empty")
}
if got := cfg.ShouldRetry(context.Background(), &adk.TypedRetryContext[*schema.AgenticMessage]{Err: errors.New("invalid api key")}); got == nil || got.Retry {
t.Fatal("permanent auth error should not retry")
}
}
func TestResolveEinoFailoverChannelsSkipsPrimaryDuplicateAndUnknown(t *testing.T) {
t.Parallel()
appCfg := &config.Config{
OpenAI: config.OpenAIConfig{Provider: "openai", APIKey: "k1", BaseURL: "https://api.example/v1", Model: "primary"},
AI: config.AIConfig{Channels: map[string]config.AIChannelConfig{
"same": {Provider: "openai", APIKey: "k1", BaseURL: "https://api.example/v1", Model: "primary"},
"fb1": {Provider: "openai", APIKey: "k2", BaseURL: "https://api.example/v1", Model: "fallback-1"},
"fb2": {Provider: "claude", APIKey: "k3", BaseURL: "https://api.anthropic.com/v1", Model: "claude-sonnet"},
}},
}
got := resolveEinoFailoverChannels(appCfg, &config.MultiAgentEinoMiddlewareConfig{
ModelFailoverChannels: []string{"same", "missing", "fb1", "fb1", "fb2"},
ModelFailoverMaxRetries: 1,
})
if len(got) != 2 {
t.Fatalf("resolved channels len = %d, want 2 before max cap is applied by config builder", len(got))
}
if got[0].id != "fb1" || got[1].id != "fb2" {
t.Fatalf("resolved channel order = %#v", got)
}
}
func TestNewEinoModelFailoverConfigBuildsDistinctFallbackModel(t *testing.T) {
t.Parallel()
appCfg := &config.Config{
OpenAI: config.OpenAIConfig{APIKey: "k1", BaseURL: "https://api.example/v1", Model: "primary"},
AI: config.AIConfig{Channels: map[string]config.AIChannelConfig{
"fb1": {APIKey: "k2", BaseURL: "https://api.example/v1", Model: "fallback-1"},
"fb2": {APIKey: "k3", BaseURL: "https://api.example/v1", Model: "fallback-2"},
}},
}
var built []string
cfg, err := newEinoModelFailoverConfig(
context.Background(),
appCfg,
&config.MultiAgentEinoMiddlewareConfig{
ModelFailoverChannels: []string{"fb1", "fb2"},
ModelFailoverMaxRetries: 1,
},
einoModelModeNormal,
func(_ context.Context, oa config.OpenAIConfig, _ einoModelMode) (model.ToolCallingChatModel, error) {
built = append(built, oa.Model)
return &streamToolCallIndexFakeModel{}, nil
},
nil,
"test",
nil,
"deep",
"conv-1",
)
if err != nil {
t.Fatalf("newEinoModelFailoverConfig: %v", err)
}
if cfg == nil || cfg.MaxRetries != 1 {
t.Fatalf("failover cfg = %#v, want max retries 1", cfg)
}
m, msgs, err := cfg.GetFailoverModel(context.Background(), &adk.FailoverContext[*schema.Message]{FailoverAttempt: 1})
if err != nil || m == nil || msgs != nil {
t.Fatalf("GetFailoverModel = (%v, %v, %v)", m, msgs, err)
}
if len(built) != 1 || built[0] != "fallback-1" {
t.Fatalf("built models = %v, want [fallback-1]", built)
}
if !cfg.ShouldFailover(context.Background(), nil, &adk.RetryExhaustedError{LastErr: errors.New("upstream returned 503"), TotalRetries: 4}) {
t.Fatal("retry-exhausted transient error should fail over")
}
if cfg.ShouldFailover(context.Background(), nil, &adk.RetryExhaustedError{LastErr: errors.New("invalid api key"), TotalRetries: 4}) {
t.Fatal("retry-exhausted permanent error should not fail over")
}
}
func TestNewEinoModelFailoverConfigEmitsProgressEvent(t *testing.T) {
t.Parallel()
appCfg := &config.Config{
OpenAI: config.OpenAIConfig{APIKey: "k1", BaseURL: "https://api.example/v1", Model: "primary"},
AI: config.AIConfig{Channels: map[string]config.AIChannelConfig{
"fb1": {APIKey: "k2", BaseURL: "https://api.example/v1", Model: "fallback-1"},
}},
}
var events []struct {
eventType string
message string
data interface{}
}
cfg, err := newEinoModelFailoverConfig(
context.Background(),
appCfg,
&config.MultiAgentEinoMiddlewareConfig{ModelFailoverChannels: []string{"fb1"}},
einoModelModeNormal,
func(_ context.Context, _ config.OpenAIConfig, _ einoModelMode) (model.ToolCallingChatModel, error) {
return &streamToolCallIndexFakeModel{}, nil
},
nil,
"test",
func(eventType, message string, data interface{}) {
events = append(events, struct {
eventType string
message string
data interface{}
}{eventType: eventType, message: message, data: data})
},
"deep",
"conv-1",
)
if err != nil {
t.Fatalf("newEinoModelFailoverConfig: %v", err)
}
if _, _, err := cfg.GetFailoverModel(context.Background(), &adk.FailoverContext[*schema.Message]{FailoverAttempt: 1}); err != nil {
t.Fatalf("GetFailoverModel: %v", err)
}
if len(events) != 1 || events[0].eventType != "eino_model_failover" {
t.Fatalf("events = %#v, want one eino_model_failover", events)
}
payload, ok := events[0].data.(map[string]interface{})
if !ok {
t.Fatalf("event payload type = %T", events[0].data)
}
if payload["conversationId"] != "conv-1" || payload["orchestration"] != "deep" || payload["channel"] != "fb1" || payload["model"] != "fallback-1" {
t.Fatalf("payload = %#v", payload)
}
}
func TestNewEinoAgenticModelFailoverConfigBuildsDistinctFallbackModel(t *testing.T) {
t.Parallel()
appCfg := &config.Config{
OpenAI: config.OpenAIConfig{Provider: "openai", APIKey: "k1", BaseURL: "https://api.example/v1", Model: "primary"},
AI: config.AIConfig{Channels: map[string]config.AIChannelConfig{
"fb1": {Provider: "openai", APIKey: "k2", BaseURL: "https://api.example/v1", Model: "fallback-1"},
"fb2": {Provider: "openai", APIKey: "k3", BaseURL: "https://api.example/v1", Model: "fallback-2"},
}},
}
var built []string
cfg, err := newEinoAgenticModelFailoverConfig(
context.Background(),
appCfg,
&config.MultiAgentEinoMiddlewareConfig{
ModelFailoverChannels: []string{"fb1", "fb2"},
ModelFailoverMaxRetries: 1,
},
einoModelModeNormal,
func(_ context.Context, oa config.OpenAIConfig, _ einoModelMode) (model.AgenticModel, error) {
built = append(built, oa.Model)
return &fakeAgenticGateModel{}, nil
},
nil,
"agentic",
nil,
"eino_single_agentic",
"conv-1",
)
if err != nil {
t.Fatalf("newEinoAgenticModelFailoverConfig: %v", err)
}
if cfg == nil || cfg.MaxRetries != 1 {
t.Fatalf("agentic failover cfg = %#v, want max retries 1", cfg)
}
m, msgs, err := cfg.GetFailoverModel(context.Background(), &adk.FailoverContext[*schema.AgenticMessage]{FailoverAttempt: 1})
if err != nil || m == nil || msgs != nil {
t.Fatalf("GetFailoverModel = (%v, %v, %v)", m, msgs, err)
}
if len(built) != 1 || built[0] != "fallback-1" {
t.Fatalf("built models = %v, want [fallback-1]", built)
}
if !cfg.ShouldFailover(context.Background(), nil, &adk.RetryExhaustedError{LastErr: errors.New("upstream returned 503"), TotalRetries: 4}) {
t.Fatal("retry-exhausted transient agentic error should fail over")
}
if cfg.ShouldFailover(context.Background(), nil, &adk.RetryExhaustedError{LastErr: errors.New("invalid api key"), TotalRetries: 4}) {
t.Fatal("retry-exhausted permanent agentic error should not fail over")
}
}
func TestNewEinoAgenticModelFailoverConfigEmitsProgressEvent(t *testing.T) {
t.Parallel()
appCfg := &config.Config{
OpenAI: config.OpenAIConfig{Provider: "openai", APIKey: "k1", BaseURL: "https://api.example/v1", Model: "primary"},
AI: config.AIConfig{Channels: map[string]config.AIChannelConfig{
"fb1": {Provider: "openai", APIKey: "k2", BaseURL: "https://api.example/v1", Model: "fallback-1"},
}},
}
var events []struct {
eventType string
message string
data interface{}
}
cfg, err := newEinoAgenticModelFailoverConfig(
context.Background(),
appCfg,
&config.MultiAgentEinoMiddlewareConfig{ModelFailoverChannels: []string{"fb1"}},
einoModelModeNormal,
func(_ context.Context, _ config.OpenAIConfig, _ einoModelMode) (model.AgenticModel, error) {
return &fakeAgenticGateModel{}, nil
},
nil,
"agentic",
func(eventType, message string, data interface{}) {
events = append(events, struct {
eventType string
message string
data interface{}
}{eventType: eventType, message: message, data: data})
},
"eino_single_agentic",
"conv-1",
)
if err != nil {
t.Fatalf("newEinoAgenticModelFailoverConfig: %v", err)
}
if _, _, err := cfg.GetFailoverModel(context.Background(), &adk.FailoverContext[*schema.AgenticMessage]{FailoverAttempt: 1}); err != nil {
t.Fatalf("GetFailoverModel: %v", err)
}
if len(events) != 1 || events[0].eventType != "eino_model_failover" {
t.Fatalf("events = %#v, want one eino_model_failover", events)
}
payload, ok := events[0].data.(map[string]interface{})
if !ok {
t.Fatalf("event payload type = %T", events[0].data)
}
if payload["conversationId"] != "conv-1" || payload["orchestration"] != "eino_single_agentic" || payload["channel"] != "fb1" || payload["model"] != "fallback-1" {
t.Fatalf("payload = %#v", payload)
}
}
func TestNewEinoOpenAIAgenticChatModelFactoryBuildsBackend(t *testing.T) {
t.Parallel()
factory := newEinoOpenAIAgenticChatModelFactory(newEinoBaseHTTPClient(), nil, nil)
m, err := factory(context.Background(), config.OpenAIConfig{
Provider: "openai",
APIKey: "test-key",
BaseURL: "https://api.example/v1",
Model: "gpt-4o-mini",
Reasoning: config.OpenAIReasoningConfig{
Profile: "openai_compat",
Mode: "on",
Effort: "high",
},
}, einoModelModeNormal)
if err != nil {
t.Fatalf("agentic factory: %v", err)
}
if m == nil {
t.Fatal("agentic factory returned nil model")
}
gate := evaluateEinoAgenticModelGate(agenticModelGateFactory(factory, config.OpenAIConfig{
Provider: "openai",
APIKey: "test-key",
BaseURL: "https://api.example/v1",
Model: "gpt-4o-mini",
}, einoModelModeNormal), einoAgenticRuntimeSupportV0914())
if !gate.Ready {
t.Fatalf("gate = %#v, want ready with buildable agentic backend", gate)
}
}
func TestNewEinoOpenAIAgenticChatModelFactoryRejectsUnsupportedProvider(t *testing.T) {
t.Parallel()
factory := newEinoOpenAIAgenticChatModelFactory(newEinoBaseHTTPClient(), nil, nil)
if _, err := factory(context.Background(), config.OpenAIConfig{
Provider: "claude",
APIKey: "test-key",
BaseURL: "https://api.anthropic.com/v1",
Model: "claude-sonnet-4",
}, einoModelModeNormal); err == nil {
t.Fatal("expected unsupported provider error")
}
gate := evaluateEinoAgenticModelGate(agenticModelGateFactory(factory, config.OpenAIConfig{
Provider: "claude",
APIKey: "test-key",
BaseURL: "https://api.anthropic.com/v1",
Model: "claude-sonnet-4",
}, einoModelModeNormal), einoAgenticRuntimeSupportV0914())
if gate.Ready || !containsString(gate.Missing, "model.AgenticModel backend") {
t.Fatalf("gate = %#v, want backend missing for unsupported provider", gate)
}
}
func TestEinoNativeRetryErrorsDoNotTriggerRunLevelTransientRetry(t *testing.T) {
t.Parallel()
err := &adk.WillRetryError{ErrStr: "HTTP 429 Too Many Requests", RetryAttempt: 1}
if isEinoTransientRunError(err) {
t.Fatal("WillRetryError should be observed, not treated as a run-level transient failure")
}
exhausted := &adk.RetryExhaustedError{LastErr: errors.New("HTTP 429 Too Many Requests"), TotalRetries: 4}
if isEinoTransientRunError(exhausted) {
t.Fatal("RetryExhaustedError should not trigger a second run-level retry layer")
}
if got := unwrapEinoRetryExhausted(exhausted); got == exhausted {
t.Fatal("unwrapEinoRetryExhausted should return the underlying model error")
}
}
@@ -0,0 +1,27 @@
package multiagent
import (
"context"
"testing"
)
func TestEinoNativeCancelOptionsByCause(t *testing.T) {
fullStopOpts, fullStopWait := einoNativeCancelOptions(context.Canceled)
if len(fullStopOpts) != 2 {
t.Fatalf("full stop options: got %d want 2", len(fullStopOpts))
}
if fullStopWait != einoNativeCancelImmediateWait {
t.Fatalf("full stop wait: got %v want %v", fullStopWait, einoNativeCancelImmediateWait)
}
interruptOpts, interruptWait := einoNativeCancelOptions(ErrInterruptContinue)
if len(interruptOpts) != 3 {
t.Fatalf("interrupt options: got %d want 3", len(interruptOpts))
}
if interruptWait != einoNativeCancelSafePointWait {
t.Fatalf("interrupt wait: got %v want %v", interruptWait, einoNativeCancelSafePointWait)
}
if interruptWait <= einoNativeCancelSafePointTTL {
t.Fatalf("interrupt wait must allow the Eino safe-point timeout to elapse: wait=%v ttl=%v", interruptWait, einoNativeCancelSafePointTTL)
}
}