From 9f3ae4f7e2d44d87f15333f2160499d287dbcf17 Mon Sep 17 00:00:00 2001 From: phantom5099 <1011668688@qq.com> Date: Sat, 4 Apr 2026 23:31:46 +0800 Subject: [PATCH 01/55] =?UTF-8?q?pref(provider):=E5=A2=9E=E5=8A=A0?= =?UTF-8?q?=E5=BC=BA=E6=9E=84=E9=80=A0=E5=99=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/provider/openai/openai.go | 33 +++------- internal/provider/openai/openai_test.go | 10 +-- internal/provider/types.go | 81 ++++++++++++++++++++++++- 3 files changed, 94 insertions(+), 30 deletions(-) diff --git a/internal/provider/openai/openai.go b/internal/provider/openai/openai.go index b1b3599c..d67ee294 100644 --- a/internal/provider/openai/openai.go +++ b/internal/provider/openai/openai.go @@ -53,11 +53,11 @@ func Driver() provider.DriverDefinition { return New(cfg, WithTransport(defaultRetryTransport())) }, Discover: func(ctx context.Context, cfg config.ResolvedProviderConfig) ([]config.ModelDescriptor, error) { - provider, err := New(cfg, WithTransport(defaultRetryTransport())) + p, err := New(cfg, WithTransport(defaultRetryTransport())) if err != nil { return nil, err } - return provider.DiscoverModels(ctx) + return p.DiscoverModels(ctx) }, } } @@ -331,34 +331,23 @@ func emitTextDelta(ctx context.Context, events chan<- provider.StreamEvent, text if text == "" { return nil } - return emitStreamEvent(ctx, events, provider.StreamEvent{ - Type: provider.StreamEventTextDelta, - Text: text, - }) + return emitStreamEvent(ctx, events, provider.NewTextDeltaStreamEvent(text)) } func emitToolCallStart(ctx context.Context, events chan<- provider.StreamEvent, index int, id, name string) error { if name == "" { return nil } - return emitStreamEvent(ctx, events, provider.StreamEvent{ - Type: provider.StreamEventToolCallStart, - ToolCallID: id, - ToolName: name, - ToolCallIndex: index, - }) + return emitStreamEvent(ctx, events, provider.NewToolCallStartStreamEvent(index, id, name)) } // emitToolCallDelta 发送工具调用参数增量事件。 -func emitToolCallDelta(ctx context.Context, events chan<- provider.StreamEvent, index int, argumentsDelta string) error { +// id 为工具调用 ID,由上游 mergeToolCallDelta 从累积状态中传入。 +func emitToolCallDelta(ctx context.Context, events chan<- provider.StreamEvent, index int, id, argumentsDelta string) error { if argumentsDelta == "" { return nil } - return emitStreamEvent(ctx, events, provider.StreamEvent{ - Type: provider.StreamEventToolCallDelta, - ToolCallIndex: index, - ToolArgumentsDelta: argumentsDelta, - }) + return emitStreamEvent(ctx, events, provider.NewToolCallDeltaStreamEvent(index, id, argumentsDelta)) } // emitMessageDone 发送消息完成事件。 @@ -366,11 +355,7 @@ func emitMessageDone(ctx context.Context, events chan<- provider.StreamEvent, fi if events == nil { return nil } - return emitStreamEvent(ctx, events, provider.StreamEvent{ - Type: provider.StreamEventMessageDone, - FinishReason: finishReason, - Usage: usage, - }) + return emitStreamEvent(ctx, events, provider.NewMessageDoneStreamEvent(finishReason, usage)) } // extractStreamUsage 从 OpenAI usage 响应提取并覆盖累积的 token 统计。 @@ -413,7 +398,7 @@ func mergeToolCallDelta(ctx context.Context, events chan<- provider.StreamEvent, // 发送参数增量事件(同一 chunk 可能同时携带 name 和 arguments) if args := delta.Function.Arguments; args != "" { call.Arguments += args - if err := emitToolCallDelta(ctx, events, delta.Index, args); err != nil { + if err := emitToolCallDelta(ctx, events, delta.Index, call.ID, args); err != nil { return err } } diff --git a/internal/provider/openai/openai_test.go b/internal/provider/openai/openai_test.go index 413e5b44..cdf62bac 100644 --- a/internal/provider/openai/openai_test.go +++ b/internal/provider/openai/openai_test.go @@ -92,7 +92,7 @@ func TestEmitToolCallDelta(t *testing.T) { t.Run("nil events guard", func(t *testing.T) { t.Parallel() - if err := emitToolCallDelta(context.Background(), nil, 0, "args"); err != nil { + if err := emitToolCallDelta(context.Background(), nil, 0, "", "args"); err != nil { t.Fatalf("expected nil events guard to return nil, got %v", err) } }) @@ -100,7 +100,7 @@ func TestEmitToolCallDelta(t *testing.T) { t.Run("empty arguments guard", func(t *testing.T) { t.Parallel() events := make(chan domain.StreamEvent, 1) - if err := emitToolCallDelta(context.Background(), events, 0, ""); err != nil { + if err := emitToolCallDelta(context.Background(), events, 0, "", ""); err != nil { t.Fatalf("expected empty arguments guard to return nil, got %v", err) } select { @@ -113,11 +113,11 @@ func TestEmitToolCallDelta(t *testing.T) { t.Run("normal send", func(t *testing.T) { t.Parallel() events := make(chan domain.StreamEvent, 1) - if err := emitToolCallDelta(context.Background(), events, 3, `{"path":"main.go"}`); err != nil { + if err := emitToolCallDelta(context.Background(), events, 3, "call_123", `{"path":"main.go"}`); err != nil { t.Fatalf("emitToolCallDelta() error = %v", err) } got := <-events - if got.Type != domain.StreamEventToolCallDelta || got.ToolCallIndex != 3 || got.ToolArgumentsDelta != `{"path":"main.go"}` { + if got.Type != domain.StreamEventToolCallDelta || got.ToolCallIndex != 3 || got.ToolArgumentsDelta != `{"path":"main.go"}` || got.ToolCallID != "call_123" { t.Fatalf("unexpected event: %+v", got) } }) @@ -126,7 +126,7 @@ func TestEmitToolCallDelta(t *testing.T) { t.Parallel() cancelledCtx, cancel := context.WithCancel(context.Background()) cancel() - if err := emitToolCallDelta(cancelledCtx, make(chan domain.StreamEvent), 0, "args"); err == nil { + if err := emitToolCallDelta(cancelledCtx, make(chan domain.StreamEvent), 0, "", "args"); err == nil { t.Fatal("expected cancellation error") } }) diff --git a/internal/provider/types.go b/internal/provider/types.go index eba95f19..294d31b1 100644 --- a/internal/provider/types.go +++ b/internal/provider/types.go @@ -68,8 +68,13 @@ const ( ) // StreamEvent 表示 provider 驱动层向 runtime 推送的流式事件。 +// 强制使用 NewXxxStreamEvent 构造器创建实例,禁止直接构造。 type StreamEvent struct { - Type StreamEventType + Type StreamEventType `json:"type"` + Payload interface{} `json:"payload,omitempty"` // 强类型载荷,使用类型断言访问 + + // --- 以下字段已弃用,保留仅用于 Phase 1 向后兼容 --- + // Phase 2 将由 Runtime 负责人移除,届时所有消费方应通过 Payload 类型断言访问。 // text_delta Text string `json:"text,omitempty"` // 文本片段 @@ -84,3 +89,77 @@ type StreamEvent struct { FinishReason string `json:"finish_reason,omitempty"` // 结束原因(仅 message_done 时有效) Usage *Usage `json:"usage,omitempty"` // 使用统计(仅 message_done 时有效) } + +// --- Payload 强类型定义 --- + +// TextDeltaPayload 表示文本增量事件的载荷。 +type TextDeltaPayload struct { + Text string `json:"text"` +} + +// ToolCallStartPayload 表示工具调用开始事件的载荷。 +type ToolCallStartPayload struct { + Index int `json:"index"` + ID string `json:"id"` + Name string `json:"name"` +} + +// ToolCallDeltaPayload 表示工具调用参数增量事件的载荷。 +type ToolCallDeltaPayload struct { + Index int `json:"index"` + ID string `json:"id"` + ArgumentsDelta string `json:"arguments_delta"` +} + +// MessageDonePayload 表示消息完成事件的载荷。 +type MessageDonePayload struct { + FinishReason string `json:"finish_reason"` + Usage *Usage `json:"usage"` +} + +// --- 构造器 --- + +// NewTextDeltaStreamEvent 创建文本增量流事件。 +func NewTextDeltaStreamEvent(text string) StreamEvent { + return StreamEvent{ + Type: StreamEventTextDelta, + Payload: TextDeltaPayload{Text: text}, + // 兼容层:同步填充弃用字段 + Text: text, + } +} + +// NewToolCallStartStreamEvent 创建工具调用开始流事件。 +func NewToolCallStartStreamEvent(index int, id, name string) StreamEvent { + return StreamEvent{ + Type: StreamEventToolCallStart, + Payload: ToolCallStartPayload{Index: index, ID: id, Name: name}, + // 兼容层:同步填充弃用字段 + ToolCallIndex: index, + ToolCallID: id, + ToolName: name, + } +} + +// NewToolCallDeltaStreamEvent 创建工具调用参数增量流事件。 +func NewToolCallDeltaStreamEvent(index int, id, argumentsDelta string) StreamEvent { + return StreamEvent{ + Type: StreamEventToolCallDelta, + Payload: ToolCallDeltaPayload{Index: index, ID: id, ArgumentsDelta: argumentsDelta}, + // 兼容层:同步填充弃用字段 + ToolCallIndex: index, + ToolCallID: id, + ToolArgumentsDelta: argumentsDelta, + } +} + +// NewMessageDoneStreamEvent 创建消息完成流事件。 +func NewMessageDoneStreamEvent(finishReason string, usage *Usage) StreamEvent { + return StreamEvent{ + Type: StreamEventMessageDone, + Payload: MessageDonePayload{FinishReason: finishReason, Usage: usage}, + // 兼容层:同步填充弃用字段 + FinishReason: finishReason, + Usage: usage, + } +} From bc53f8c147303200b4fc5d6c27f13b2912bab15f Mon Sep 17 00:00:00 2001 From: Yumiue <229866007@qq.com> Date: Sun, 5 Apr 2026 14:24:42 +0800 Subject: [PATCH 02/55] =?UTF-8?q?feat(context):=20=E5=A2=9E=E5=8A=A0?= =?UTF-8?q?=E6=97=A7=E5=B7=A5=E5=85=B7=E7=BB=93=E6=9E=9C=E7=9A=84=E8=AF=BB?= =?UTF-8?q?=E6=97=B6=20Micro=20Compact?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/context/builder.go | 2 +- internal/context/builder_test.go | 54 +++++++ internal/context/microcompact.go | 128 ++++++++++++++++ internal/context/microcompact_test.go | 203 ++++++++++++++++++++++++++ 4 files changed, 386 insertions(+), 1 deletion(-) create mode 100644 internal/context/microcompact.go create mode 100644 internal/context/microcompact_test.go diff --git a/internal/context/builder.go b/internal/context/builder.go index ff1487ea..f1ca2c79 100644 --- a/internal/context/builder.go +++ b/internal/context/builder.go @@ -43,6 +43,6 @@ func (b *DefaultBuilder) Build(ctx context.Context, input BuildInput) (BuildResu return BuildResult{ SystemPrompt: composeSystemPrompt(sections...), - Messages: trimPolicy.Trim(input.Messages), + Messages: microCompactMessages(trimPolicy.Trim(input.Messages)), }, nil } diff --git a/internal/context/builder_test.go b/internal/context/builder_test.go index 42c002ba..9f5098a0 100644 --- a/internal/context/builder_test.go +++ b/internal/context/builder_test.go @@ -150,6 +150,60 @@ func TestDefaultBuilderBuildReturnsPromptSourceError(t *testing.T) { } } +func TestDefaultBuilderBuildAppliesMicroCompactAfterTrim(t *testing.T) { + t.Parallel() + + builder := &DefaultBuilder{ + promptSources: []promptSectionSource{ + stubPromptSectionSource{sections: []promptSection{{title: "Stub", content: "body"}}}, + }, + } + + messages := []provider.Message{ + {Role: provider.RoleUser, Content: "older user"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-1", Content: "old read result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-2", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-3", Name: "webfetch", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, + {Role: provider.RoleUser, Content: "latest explicit instruction"}, + {Role: provider.RoleAssistant, Content: "current reply"}, + } + + got, err := builder.Build(stdcontext.Background(), BuildInput{Messages: messages}) + if err != nil { + t.Fatalf("Build() error = %v", err) + } + if len(got.Messages) != len(messages) { + t.Fatalf("expected builder output to keep message count, got %d want %d", len(got.Messages), len(messages)) + } + if got.Messages[2].Content != microCompactClearedMessage { + t.Fatalf("expected builder output to clear older tool result, got %q", got.Messages[2].Content) + } + if got.Messages[4].Content != "recent bash result" { + t.Fatalf("expected recent tool result to stay visible, got %q", got.Messages[4].Content) + } + if got.Messages[6].Content != "latest webfetch result" { + t.Fatalf("expected latest tool result to stay visible, got %q", got.Messages[6].Content) + } +} + func TestTrimMessagesPreservesToolPairs(t *testing.T) { t.Parallel() diff --git a/internal/context/microcompact.go b/internal/context/microcompact.go new file mode 100644 index 00000000..1eec70c0 --- /dev/null +++ b/internal/context/microcompact.go @@ -0,0 +1,128 @@ +package context + +import ( + "strings" + + "neo-code/internal/context/internalcompact" + "neo-code/internal/provider" +) + +const ( + // microCompactClearedMessage 是旧工具结果被读时微压缩后的占位符文本。 + microCompactClearedMessage = "[Old tool result content cleared]" + // microCompactRetainedToolSpans 定义默认保留原始内容的最近可压缩工具块数量。 + microCompactRetainedToolSpans = 2 +) + +var microCompactableTools = map[string]struct{}{ + "bash": {}, + "webfetch": {}, + "filesystem_read_file": {}, + "filesystem_grep": {}, + "filesystem_glob": {}, + "filesystem_edit": {}, + "filesystem_write_file": {}, +} + +// microCompactMessages 对裁剪后的消息做只读投影式微压缩,仅清理旧工具结果内容。 +func microCompactMessages(messages []provider.Message) []provider.Message { + cloned := cloneContextMessages(messages) + if len(cloned) == 0 { + return cloned + } + + spans := internalcompact.BuildMessageSpans(cloned) + protectedStart, hasProtectedTail := internalcompact.ProtectedTailStart(spans) + retainedCompactableSpans := 0 + + for spanIndex := len(spans) - 1; spanIndex >= 0; spanIndex-- { + span := spans[spanIndex] + if hasProtectedTail && span.Start >= protectedStart { + continue + } + if !isToolCallSpan(cloned, span) { + continue + } + + compactableIDs := compactableToolCallIDs(cloned[span.Start].ToolCalls) + if len(compactableIDs) == 0 { + continue + } + if retainedCompactableSpans < microCompactRetainedToolSpans { + retainedCompactableSpans++ + continue + } + + for messageIndex := span.Start + 1; messageIndex < span.End; messageIndex++ { + if shouldClearToolMessage(cloned[messageIndex], compactableIDs) { + cloned[messageIndex].Content = microCompactClearedMessage + } + } + } + + return cloned +} + +// cloneContextMessages 深拷贝消息切片,避免读时投影污染 runtime 持有的原始会话消息。 +func cloneContextMessages(messages []provider.Message) []provider.Message { + if len(messages) == 0 { + return nil + } + + cloned := make([]provider.Message, 0, len(messages)) + for _, message := range messages { + next := message + next.ToolCalls = append([]provider.ToolCall(nil), message.ToolCalls...) + cloned = append(cloned, next) + } + return cloned +} + +// isToolCallSpan 判断当前 span 是否是由 assistant tool call 起始的原子工具块。 +func isToolCallSpan(messages []provider.Message, span internalcompact.MessageSpan) bool { + if span.Start < 0 || span.Start >= len(messages) { + return false + } + message := messages[span.Start] + return message.Role == provider.RoleAssistant && len(message.ToolCalls) > 0 +} + +// compactableToolCallIDs 返回 assistant tool call 中可参与微压缩的调用 ID 集合。 +func compactableToolCallIDs(calls []provider.ToolCall) map[string]struct{} { + if len(calls) == 0 { + return nil + } + + ids := make(map[string]struct{}, len(calls)) + for _, call := range calls { + toolName := strings.TrimSpace(call.Name) + if _, ok := microCompactableTools[toolName]; !ok { + continue + } + callID := strings.TrimSpace(call.ID) + if callID == "" { + continue + } + ids[callID] = struct{}{} + } + if len(ids) == 0 { + return nil + } + return ids +} + +// shouldClearToolMessage 判断一条 tool 消息是否满足旧结果清理条件。 +func shouldClearToolMessage(message provider.Message, compactableIDs map[string]struct{}) bool { + if message.Role != provider.RoleTool || message.IsError { + return false + } + if compactableIDs == nil { + return false + } + if _, ok := compactableIDs[strings.TrimSpace(message.ToolCallID)]; !ok { + return false + } + + content := strings.TrimSpace(message.Content) + return content != "" && content != microCompactClearedMessage +} diff --git a/internal/context/microcompact_test.go b/internal/context/microcompact_test.go new file mode 100644 index 00000000..b0cce67e --- /dev/null +++ b/internal/context/microcompact_test.go @@ -0,0 +1,203 @@ +package context + +import ( + "testing" + + "neo-code/internal/provider" +) + +func TestMicroCompactMessagesClearsOlderCompactableToolResults(t *testing.T) { + t.Parallel() + + messages := []provider.Message{ + {Role: provider.RoleUser, Content: "older user"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-1", Content: "old read result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-2", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-3", Name: "webfetch", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, + {Role: provider.RoleUser, Content: "latest explicit instruction"}, + {Role: provider.RoleAssistant, Content: "current working reply"}, + } + + got := microCompactMessages(messages) + if len(got) != len(messages) { + t.Fatalf("expected message count to stay unchanged, got %d want %d", len(got), len(messages)) + } + if got[2].Content != microCompactClearedMessage { + t.Fatalf("expected oldest compactable tool result to be cleared, got %q", got[2].Content) + } + if got[4].Content != "recent bash result" { + t.Fatalf("expected recent compactable tool result to be retained, got %q", got[4].Content) + } + if got[6].Content != "latest webfetch result" { + t.Fatalf("expected latest compactable tool result to be retained, got %q", got[6].Content) + } + if messages[2].Content != "old read result" { + t.Fatalf("expected original slice to remain unchanged, got %q", messages[2].Content) + } +} + +func TestMicroCompactMessagesKeepsProtectedTailUntouched(t *testing.T) { + t.Parallel() + + messages := []provider.Message{ + {Role: provider.RoleUser, Content: "older user"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-0", Name: "filesystem_grep", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-0", Content: "old grep result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-1", Content: "recent read result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-2", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + {Role: provider.RoleUser, Content: "latest explicit instruction"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-3", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-3", Content: "tail bash result"}, + } + + got := microCompactMessages(messages) + if got[2].Content != microCompactClearedMessage { + t.Fatalf("expected old tool result before protected tail to be cleared, got %q", got[2].Content) + } + if got[4].Content != "recent read result" { + t.Fatalf("expected recent tool result before protected tail to remain, got %q", got[4].Content) + } + if got[6].Content != "recent bash result" { + t.Fatalf("expected second recent tool result before protected tail to remain, got %q", got[6].Content) + } + if got[9].Content != "tail bash result" { + t.Fatalf("expected protected tail tool result to remain, got %q", got[9].Content) + } +} + +func TestMicroCompactMessagesSkipsNonCompactableErrorsAndOrphans(t *testing.T) { + t.Parallel() + + messages := []provider.Message{ + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-1", Name: "custom_tool", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-1", Content: "custom result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-2", Name: "filesystem_edit", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-2", Content: "edit failed", IsError: true}, + {Role: provider.RoleTool, ToolCallID: "orphan", Content: "orphan result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-3", Name: "filesystem_write_file", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-3", Content: microCompactClearedMessage}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-4", Name: "filesystem_grep", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-4", Content: ""}, + } + + got := microCompactMessages(messages) + if got[1].Content != "custom result" { + t.Fatalf("expected non-compactable tool result to remain, got %q", got[1].Content) + } + if got[3].Content != "edit failed" { + t.Fatalf("expected error tool result to remain, got %q", got[3].Content) + } + if got[4].Content != "orphan result" { + t.Fatalf("expected orphan tool result to remain, got %q", got[4].Content) + } + if got[6].Content != microCompactClearedMessage { + t.Fatalf("expected already cleared content to remain unchanged, got %q", got[6].Content) + } + if got[8].Content != "" { + t.Fatalf("expected empty tool result to remain empty, got %q", got[8].Content) + } +} + +func TestMicroCompactMessagesClearsOnlyCompactableResultsInMixedToolSpan(t *testing.T) { + t.Parallel() + + messages := []provider.Message{ + {Role: provider.RoleUser, Content: "older user"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, + {ID: "call-2", Name: "custom_tool", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-1", Content: "read result"}, + {Role: provider.RoleTool, ToolCallID: "call-2", Content: "custom result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-3", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-3", Content: "recent bash result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-4", Name: "webfetch", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-4", Content: "latest webfetch result"}, + {Role: provider.RoleUser, Content: "latest explicit instruction"}, + {Role: provider.RoleAssistant, Content: "current reply"}, + } + + got := microCompactMessages(messages) + if got[2].Content != microCompactClearedMessage { + t.Fatalf("expected compactable tool result to be cleared, got %q", got[2].Content) + } + if got[3].Content != "custom result" { + t.Fatalf("expected non-compactable tool result in mixed span to remain, got %q", got[3].Content) + } + if len(got[1].ToolCalls) != 2 { + t.Fatalf("expected assistant tool call metadata to remain intact, got %+v", got[1].ToolCalls) + } +} From d12e43dd4e61251d45c3f6ba17e0d04187aa56a3 Mon Sep 17 00:00:00 2001 From: Yumiue <229866007@qq.com> Date: Sun, 5 Apr 2026 14:51:37 +0800 Subject: [PATCH 03/55] fix(context): avoid empty spans consuming micro compact budget --- internal/context/microcompact.go | 13 ++++++ internal/context/microcompact_test.go | 62 +++++++++++++++++++++++++++ 2 files changed, 75 insertions(+) diff --git a/internal/context/microcompact.go b/internal/context/microcompact.go index 1eec70c0..10f5baba 100644 --- a/internal/context/microcompact.go +++ b/internal/context/microcompact.go @@ -48,6 +48,9 @@ func microCompactMessages(messages []provider.Message) []provider.Message { if len(compactableIDs) == 0 { continue } + if !hasCompactableToolContent(cloned, span, compactableIDs) { + continue + } if retainedCompactableSpans < microCompactRetainedToolSpans { retainedCompactableSpans++ continue @@ -111,6 +114,16 @@ func compactableToolCallIDs(calls []provider.ToolCall) map[string]struct{} { return ids } +// hasCompactableToolContent 判断工具块中是否存在会影响保留预算的有效工具结果内容。 +func hasCompactableToolContent(messages []provider.Message, span internalcompact.MessageSpan, compactableIDs map[string]struct{}) bool { + for messageIndex := span.Start + 1; messageIndex < span.End; messageIndex++ { + if shouldClearToolMessage(messages[messageIndex], compactableIDs) { + return true + } + } + return false +} + // shouldClearToolMessage 判断一条 tool 消息是否满足旧结果清理条件。 func shouldClearToolMessage(message provider.Message, compactableIDs map[string]struct{}) bool { if message.Role != provider.RoleTool || message.IsError { diff --git a/internal/context/microcompact_test.go b/internal/context/microcompact_test.go index b0cce67e..77b7b078 100644 --- a/internal/context/microcompact_test.go +++ b/internal/context/microcompact_test.go @@ -201,3 +201,65 @@ func TestMicroCompactMessagesClearsOnlyCompactableResultsInMixedToolSpan(t *test t.Fatalf("expected assistant tool call metadata to remain intact, got %+v", got[1].ToolCalls) } } + +func TestMicroCompactMessagesSkipsEmptyRecentSpansWhenCountingRetainedBudget(t *testing.T) { + t.Parallel() + + messages := []provider.Message{ + {Role: provider.RoleUser, Content: "older user"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-1", Content: "older read result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-2", Name: "filesystem_grep", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-2", Content: "middle grep result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-3", Name: "filesystem_edit", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-3", Content: "near edit result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-4", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-4", Content: "", IsError: true}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-5", Name: "webfetch", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-5", Content: ""}, + {Role: provider.RoleUser, Content: "latest explicit instruction"}, + {Role: provider.RoleAssistant, Content: "current reply"}, + } + + got := microCompactMessages(messages) + if got[2].Content != microCompactClearedMessage { + t.Fatalf("expected oldest valid tool result to be cleared, got %q", got[2].Content) + } + if got[4].Content != "middle grep result" { + t.Fatalf("expected middle valid tool result to remain, got %q", got[4].Content) + } + if got[6].Content != "near edit result" { + t.Fatalf("expected nearer valid tool result to remain, got %q", got[6].Content) + } + if got[8].Content != "" { + t.Fatalf("expected error/empty tool result to remain unchanged, got %q", got[8].Content) + } + if got[10].Content != "" { + t.Fatalf("expected empty recent tool result to remain unchanged, got %q", got[10].Content) + } +} From 87e9752a96a9fc733a4aac31305b6a524b30014a Mon Sep 17 00:00:00 2001 From: phantom5099 <1011668688@qq.com> Date: Sun, 5 Apr 2026 14:53:18 +0800 Subject: [PATCH 04/55] =?UTF-8?q?pref(provider):=E5=8E=BB=E6=8E=89?= =?UTF-8?q?=E4=BA=86provider=E4=B8=BAruntime=E8=BF=AD=E4=BB=A3=E7=95=99?= =?UTF-8?q?=E4=B8=8B=E7=9A=84=E5=85=BC=E5=AE=B9=E4=BB=A3=E7=A0=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/provider/builtin/builtin.go | 4 +- internal/provider/builtin/builtin_test.go | 4 +- internal/provider/catalog/service.go | 4 +- internal/provider/catalog/service_test.go | 6 +- internal/provider/catalog/store.go | 18 +- internal/provider/catalog/store_test.go | 12 +- internal/provider/openai/openai.go | 81 ++--- internal/provider/openai/openai_test.go | 200 ++++++++---- internal/provider/provider.go | 2 +- internal/provider/registry_test.go | 4 +- internal/provider/types.go | 60 +--- internal/runtime/compact_generator.go | 47 ++- internal/runtime/compact_generator_test.go | 53 ++- internal/runtime/runtime.go | 119 ++++++- internal/runtime/runtime_test.go | 354 +++++++-------------- internal/tui/update_test.go | 4 +- 16 files changed, 481 insertions(+), 491 deletions(-) diff --git a/internal/provider/builtin/builtin.go b/internal/provider/builtin/builtin.go index 0c10abe1..33ab5b5e 100644 --- a/internal/provider/builtin/builtin.go +++ b/internal/provider/builtin/builtin.go @@ -9,13 +9,13 @@ import ( func NewRegistry() (*provider.Registry, error) { registry := provider.NewRegistry() - if err := Register(registry); err != nil { + if err := register(registry); err != nil { return nil, err } return registry, nil } -func Register(registry *provider.Registry) error { +func register(registry *provider.Registry) error { if registry == nil { return errors.New("builtin provider registry is nil") } diff --git a/internal/provider/builtin/builtin_test.go b/internal/provider/builtin/builtin_test.go index fb230a0c..a078b3da 100644 --- a/internal/provider/builtin/builtin_test.go +++ b/internal/provider/builtin/builtin_test.go @@ -27,7 +27,7 @@ func TestRegister(t *testing.T) { t.Run("nil registry", func(t *testing.T) { t.Parallel() - err := Register(nil) + err := register(nil) if err == nil { t.Fatal("expected error for nil registry") } @@ -39,7 +39,7 @@ func TestRegister(t *testing.T) { t.Run("valid registry", func(t *testing.T) { t.Parallel() registry := provider.NewRegistry() - err := Register(registry) + err := register(registry) if err != nil { t.Fatalf("Register() error = %v", err) } diff --git a/internal/provider/catalog/service.go b/internal/provider/catalog/service.go index cb229990..6fd9a2ee 100644 --- a/internal/provider/catalog/service.go +++ b/internal/provider/catalog/service.go @@ -29,7 +29,7 @@ type Service struct { func NewService(baseDir string, registry *provider.Registry, store Store) *Service { if store == nil && strings.TrimSpace(baseDir) != "" { - store = NewJSONStore(baseDir) + store = newJSONStore(baseDir) } return &Service{ @@ -174,7 +174,7 @@ func (s *Service) discoverAndPersist(ctx context.Context, providerCfg config.Pro now := s.now() _ = s.store.Save(ctx, ModelCatalog{ - SchemaVersion: SchemaVersion, + SchemaVersion: schemaVersion, Identity: identity, FetchedAt: now, ExpiresAt: now.Add(s.catalogTTL), diff --git a/internal/provider/catalog/service_test.go b/internal/provider/catalog/service_test.go index 7aaa5c34..3d5d0b05 100644 --- a/internal/provider/catalog/service_test.go +++ b/internal/provider/catalog/service_test.go @@ -136,7 +136,7 @@ func TestListProviderModelsReturnsStaleCacheAndRefreshesInBackground(t *testing. store := newMemoryStore() now := time.Date(2026, 4, 2, 12, 0, 0, 0, time.UTC) if err := store.Save(context.Background(), ModelCatalog{ - SchemaVersion: SchemaVersion, + SchemaVersion: schemaVersion, Identity: identity, FetchedAt: now.Add(-48 * time.Hour), ExpiresAt: now.Add(-24 * time.Hour), @@ -234,8 +234,8 @@ func containsModelDescriptorID(models []config.ModelDescriptor, modelID string) type catalogTestProvider struct{} -func (catalogTestProvider) Chat(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) (provider.ChatResponse, error) { - return provider.ChatResponse{}, nil +func (catalogTestProvider) Chat(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { + return nil } type memoryStore struct { diff --git a/internal/provider/catalog/store.go b/internal/provider/catalog/store.go index 522dce7c..96ce6f7b 100644 --- a/internal/provider/catalog/store.go +++ b/internal/provider/catalog/store.go @@ -16,7 +16,7 @@ import ( "neo-code/internal/config" ) -const SchemaVersion = 1 +const schemaVersion = 1 var ErrCatalogNotFound = errors.New("provider: model catalog not found") @@ -39,18 +39,18 @@ type Store interface { Save(ctx context.Context, catalog ModelCatalog) error } -type JSONStore struct { +type jsonStore struct { dir string mu sync.RWMutex } -func NewJSONStore(baseDir string) *JSONStore { - return &JSONStore{ +func newJSONStore(baseDir string) *jsonStore { + return &jsonStore{ dir: filepath.Join(strings.TrimSpace(baseDir), "cache", "models"), } } -func (s *JSONStore) Load(ctx context.Context, identity config.ProviderIdentity) (ModelCatalog, error) { +func (s *jsonStore) Load(ctx context.Context, identity config.ProviderIdentity) (ModelCatalog, error) { if err := ctx.Err(); err != nil { return ModelCatalog{}, err } @@ -80,7 +80,7 @@ func (s *JSONStore) Load(ctx context.Context, identity config.ProviderIdentity) return normalizeCatalog(modelCatalog), nil } -func (s *JSONStore) Save(ctx context.Context, catalog ModelCatalog) error { +func (s *jsonStore) Save(ctx context.Context, catalog ModelCatalog) error { if err := ctx.Err(); err != nil { return err } @@ -111,18 +111,18 @@ func (s *JSONStore) Save(ctx context.Context, catalog ModelCatalog) error { func normalizeCatalog(modelCatalog ModelCatalog) ModelCatalog { if modelCatalog.SchemaVersion == 0 { - modelCatalog.SchemaVersion = SchemaVersion + modelCatalog.SchemaVersion = schemaVersion } modelCatalog.Models = config.MergeModelDescriptors(modelCatalog.Models) return modelCatalog } -func (s *JSONStore) catalogPath(identity config.ProviderIdentity) string { +func (s *jsonStore) catalogPath(identity config.ProviderIdentity) string { sum := sha256.Sum256([]byte(identity.Key())) return filepath.Join(s.dir, hex.EncodeToString(sum[:])+".json") } -func (s *JSONStore) writeCatalogFile(path string, data []byte) error { +func (s *jsonStore) writeCatalogFile(path string, data []byte) error { if err := os.MkdirAll(s.dir, 0o755); err != nil { return fmt.Errorf("provider: create model catalog dir: %w", err) } diff --git a/internal/provider/catalog/store_test.go b/internal/provider/catalog/store_test.go index ed20e435..86c7c098 100644 --- a/internal/provider/catalog/store_test.go +++ b/internal/provider/catalog/store_test.go @@ -14,14 +14,14 @@ import ( func TestJSONStoreRoundTrip(t *testing.T) { t.Parallel() - store := NewJSONStore(t.TempDir()) + store := newJSONStore(t.TempDir()) identity, err := config.NewProviderIdentity("openai", "https://api.openai.com/v1") if err != nil { t.Fatalf("NewProviderIdentity() error = %v", err) } expected := ModelCatalog{ - SchemaVersion: SchemaVersion, + SchemaVersion: schemaVersion, Identity: identity, FetchedAt: time.Date(2026, 4, 2, 10, 0, 0, 0, time.UTC), ExpiresAt: time.Date(2026, 4, 3, 10, 0, 0, 0, time.UTC), @@ -70,7 +70,7 @@ func TestJSONStoreRoundTrip(t *testing.T) { func TestJSONStoreMissingCatalog(t *testing.T) { t.Parallel() - store := NewJSONStore(t.TempDir()) + store := newJSONStore(t.TempDir()) identity, err := config.NewProviderIdentity("openai", "https://api.openai.com/v1") if err != nil { t.Fatalf("NewProviderIdentity() error = %v", err) @@ -86,14 +86,14 @@ func TestJSONStoreSaveReplacesExistingCatalogWithoutTempLeak(t *testing.T) { t.Parallel() baseDir := t.TempDir() - store := NewJSONStore(baseDir) + store := newJSONStore(baseDir) identity, err := config.NewProviderIdentity("openai", "https://api.openai.com/v1") if err != nil { t.Fatalf("NewProviderIdentity() error = %v", err) } first := ModelCatalog{ - SchemaVersion: SchemaVersion, + SchemaVersion: schemaVersion, Identity: identity, FetchedAt: time.Date(2026, 4, 2, 10, 0, 0, 0, time.UTC), ExpiresAt: time.Date(2026, 4, 3, 10, 0, 0, 0, time.UTC), @@ -102,7 +102,7 @@ func TestJSONStoreSaveReplacesExistingCatalogWithoutTempLeak(t *testing.T) { }, } second := ModelCatalog{ - SchemaVersion: SchemaVersion, + SchemaVersion: schemaVersion, Identity: identity, FetchedAt: time.Date(2026, 4, 4, 10, 0, 0, 0, time.UTC), ExpiresAt: time.Date(2026, 4, 5, 10, 0, 0, 0, time.UTC), diff --git a/internal/provider/openai/openai.go b/internal/provider/openai/openai.go index d67ee294..f688943d 100644 --- a/internal/provider/openai/openai.go +++ b/internal/provider/openai/openai.go @@ -10,7 +10,6 @@ import ( "io" "log" "net/http" - "sort" "strings" "time" @@ -29,10 +28,10 @@ type buildOptions struct { transport http.RoundTripper } -type BuildOption func(*buildOptions) +type buildOption func(*buildOptions) -// WithTransport 注入自定义 HTTP Transport(如 RetryTransport)。 -func WithTransport(rt http.RoundTripper) BuildOption { +// withTransport 注入自定义 HTTP Transport(如 RetryTransport)。 +func withTransport(rt http.RoundTripper) buildOption { return func(o *buildOptions) { o.transport = rt } @@ -50,10 +49,10 @@ func Driver() provider.DriverDefinition { return provider.DriverDefinition{ Name: DriverName, Build: func(ctx context.Context, cfg config.ResolvedProviderConfig) (provider.Provider, error) { - return New(cfg, WithTransport(defaultRetryTransport())) + return New(cfg, withTransport(defaultRetryTransport())) }, Discover: func(ctx context.Context, cfg config.ResolvedProviderConfig) ([]config.ModelDescriptor, error) { - p, err := New(cfg, WithTransport(defaultRetryTransport())) + p, err := New(cfg, withTransport(defaultRetryTransport())) if err != nil { return nil, err } @@ -62,7 +61,7 @@ func Driver() provider.DriverDefinition { } } -func New(cfg config.ResolvedProviderConfig, opts ...BuildOption) (*Provider, error) { +func New(cfg config.ResolvedProviderConfig, opts ...buildOption) (*Provider, error) { if err := cfg.Validate(); err != nil { return nil, fmt.Errorf("openai provider: %w", err) } @@ -103,21 +102,21 @@ func (p *Provider) DiscoverModels(ctx context.Context) ([]config.ModelDescriptor return config.MergeModelDescriptors(descriptors), nil } -func (p *Provider) Chat(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) (provider.ChatResponse, error) { +func (p *Provider) Chat(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { payload, err := p.buildRequest(req) if err != nil { - return provider.ChatResponse{}, err + return err } body, err := json.Marshal(payload) if err != nil { - return provider.ChatResponse{}, fmt.Errorf("openai provider: marshal request: %w", err) + return fmt.Errorf("openai provider: marshal request: %w", err) } endpoint := strings.TrimRight(p.cfg.BaseURL, "/") + "/chat/completions" httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(body)) if err != nil { - return provider.ChatResponse{}, fmt.Errorf("openai provider: build request: %w", err) + return fmt.Errorf("openai provider: build request: %w", err) } httpReq.Header.Set("Authorization", "Bearer "+p.cfg.APIKey) httpReq.Header.Set("Content-Type", "application/json") @@ -125,7 +124,7 @@ func (p *Provider) Chat(ctx context.Context, req provider.ChatRequest, events ch resp, err := p.client.Do(httpReq) if err != nil { - return provider.ChatResponse{}, fmt.Errorf("openai provider: send request: %w", err) + return fmt.Errorf("openai provider: send request: %w", err) } defer func(Body io.ReadCloser) { err := Body.Close() @@ -135,7 +134,7 @@ func (p *Provider) Chat(ctx context.Context, req provider.ChatRequest, events ch }(resp.Body) if resp.StatusCode >= http.StatusBadRequest { - return provider.ChatResponse{}, p.parseError(resp) + return p.parseError(resp) } return p.consumeStream(ctx, resp.Body, events) @@ -185,14 +184,13 @@ func (p *Provider) buildRequest(req provider.ChatRequest) (chatCompletionRequest return payload, nil } -func (p *Provider) consumeStream(ctx context.Context, body io.Reader, events chan<- provider.StreamEvent) (provider.ChatResponse, error) { +func (p *Provider) consumeStream(ctx context.Context, body io.Reader, events chan<- provider.StreamEvent) error { reader := bufio.NewReader(body) var ( - contentBuilder strings.Builder - finishReason string - usage provider.Usage - done bool + finishReason string + usage provider.Usage + done bool ) toolCalls := make(map[int]*provider.ToolCall) @@ -222,7 +220,6 @@ func (p *Provider) consumeStream(ctx context.Context, body io.Reader, events cha finishReason = choice.FinishReason } if choice.Delta.Content != "" { - contentBuilder.WriteString(choice.Delta.Content) if err := emitTextDelta(ctx, events, choice.Delta.Content); err != nil { return err } @@ -236,12 +233,12 @@ func (p *Provider) consumeStream(ctx context.Context, body io.Reader, events cha return nil } - // finishStream 统一的流结束处理:发送 message_done 事件并组装最终响应。 - finishStream := func() (provider.ChatResponse, error) { + // finishStream 统一的流结束处理:发送 message_done 事件。 + finishStream := func() error { if err := emitMessageDone(ctx, events, finishReason, &usage); err != nil { - return provider.ChatResponse{}, err + return err } - return finalizeResponse(contentBuilder.String(), toolCalls, finishReason, usage), nil + return nil } flushPendingData := func() error { @@ -254,7 +251,7 @@ func (p *Provider) consumeStream(ctx context.Context, body io.Reader, events cha for { line, err := reader.ReadString('\n') if err != nil && !errors.Is(err, io.EOF) { - return provider.ChatResponse{}, fmt.Errorf("openai provider: read stream: %w", err) + return fmt.Errorf("openai provider: read stream: %w", err) } line = strings.TrimRight(line, "\r\n") @@ -265,7 +262,7 @@ func (p *Provider) consumeStream(ctx context.Context, body io.Reader, events cha dataLines = append(dataLines, strings.TrimSpace(strings.TrimPrefix(trimmed, "data:"))) case trimmed == "": if flushErr := flushPendingData(); flushErr != nil { - return provider.ChatResponse{}, flushErr + return flushErr } if done { return finishStream() @@ -276,7 +273,7 @@ func (p *Provider) consumeStream(ctx context.Context, body io.Reader, events cha if errors.Is(err, io.EOF) { if flushErr := flushPendingData(); flushErr != nil { - return provider.ChatResponse{}, flushErr + return flushErr } return finishStream() } @@ -405,38 +402,6 @@ func mergeToolCallDelta(ctx context.Context, events chan<- provider.StreamEvent, return nil } -// finalizeResponse 将累积的内容、tool calls 和元数据组装为最终 ChatResponse。 -func finalizeResponse(content string, toolCalls map[int]*provider.ToolCall, finishReason string, usage provider.Usage) provider.ChatResponse { - ordered := make([]int, 0, len(toolCalls)) - for index := range toolCalls { - ordered = append(ordered, index) - } - sort.Ints(ordered) - - message := provider.Message{ - Role: provider.RoleAssistant, - Content: content, - } - - for _, index := range ordered { - call := toolCalls[index] - if call == nil { - continue - } - message.ToolCalls = append(message.ToolCalls, *call) - } - - if finishReason == "" && len(message.ToolCalls) > 0 { - finishReason = "tool_calls" - } - - return provider.ChatResponse{ - Message: message, - FinishReason: finishReason, - Usage: usage, - } -} - func emitStreamEvent(ctx context.Context, events chan<- provider.StreamEvent, event provider.StreamEvent) error { if events == nil { return nil diff --git a/internal/provider/openai/openai_test.go b/internal/provider/openai/openai_test.go index cdf62bac..9b0fc108 100644 --- a/internal/provider/openai/openai_test.go +++ b/internal/provider/openai/openai_test.go @@ -33,7 +33,7 @@ func TestWithTransport(t *testing.T) { customTransport := &http.Transport{} cfg := resolvedConfig("", "") - provider, err := New(cfg, WithTransport(customTransport)) + provider, err := New(cfg, withTransport(customTransport)) if err != nil { t.Fatalf("New() error = %v", err) } @@ -117,7 +117,8 @@ func TestEmitToolCallDelta(t *testing.T) { t.Fatalf("emitToolCallDelta() error = %v", err) } got := <-events - if got.Type != domain.StreamEventToolCallDelta || got.ToolCallIndex != 3 || got.ToolArgumentsDelta != `{"path":"main.go"}` || got.ToolCallID != "call_123" { + payload := requireToolCallDeltaPayload(t, got) + if got.Type != domain.StreamEventToolCallDelta || payload.Index != 3 || payload.ArgumentsDelta != `{"path":"main.go"}` || payload.ID != "call_123" { t.Fatalf("unexpected event: %+v", got) } }) @@ -150,7 +151,8 @@ func TestEmitMessageDone(t *testing.T) { t.Fatalf("emitMessageDone() error = %v", err) } got := <-events - if got.Type != domain.StreamEventMessageDone || got.FinishReason != "stop" || got.Usage.TotalTokens != 100 { + payload := requireMessageDonePayload(t, got) + if got.Type != domain.StreamEventMessageDone || payload.FinishReason != "stop" || payload.Usage == nil || payload.Usage.TotalTokens != 100 { t.Fatalf("unexpected event: %+v", got) } }) @@ -290,7 +292,7 @@ func TestProviderChatConsumesSSEAndMergesToolCalls(t *testing.T) { provider.client = server.Client() events := make(chan domain.StreamEvent, 8) - response, err := provider.Chat(context.Background(), domain.ChatRequest{ + err = provider.Chat(context.Background(), domain.ChatRequest{ Model: "gpt-5.4", Messages: []domain.Message{ {Role: "user", Content: "please edit the file"}, @@ -320,38 +322,57 @@ func TestProviderChatConsumesSSEAndMergesToolCalls(t *testing.T) { t.Fatalf("Chat() error = %v", err) } - if response.Message.Content != "Hello world" { - t.Fatalf("expected content %q, got %q", "Hello world", response.Message.Content) + streamEvents := drainStreamEvents(events) + if len(streamEvents) == 0 { + t.Fatal("expected streamed events") } - if response.FinishReason != "tool_calls" { - t.Fatalf("expected finish reason tool_calls, got %q", response.FinishReason) + + var ( + chunks []string + toolCallStartSeen bool + toolCallArgs strings.Builder + messageDone *domain.MessageDonePayload + ) + + for _, event := range streamEvents { + switch event.Type { + case domain.StreamEventTextDelta: + chunks = append(chunks, requireTextDeltaPayload(t, event).Text) + case domain.StreamEventToolCallStart: + payload := requireToolCallStartPayload(t, event) + toolCallStartSeen = true + if payload.Index != 0 || payload.ID != "call_1" || payload.Name != "filesystem_edit" { + t.Fatalf("unexpected tool_call_start payload: %+v", payload) + } + case domain.StreamEventToolCallDelta: + payload := requireToolCallDeltaPayload(t, event) + if payload.Index != 0 || payload.ID != "call_1" { + t.Fatalf("unexpected tool_call_delta payload: %+v", payload) + } + toolCallArgs.WriteString(payload.ArgumentsDelta) + case domain.StreamEventMessageDone: + payload := requireMessageDonePayload(t, event) + messageDone = &payload + } } - if response.Usage.TotalTokens != 15 { - t.Fatalf("expected total tokens 15, got %d", response.Usage.TotalTokens) + + if strings.Join(chunks, "") != "Hello world" { + t.Fatalf("expected streamed chunks to form %q, got %q", "Hello world", strings.Join(chunks, "")) } - if len(response.Message.ToolCalls) != 1 { - t.Fatalf("expected 1 tool call, got %d", len(response.Message.ToolCalls)) + if !toolCallStartSeen { + t.Fatal("expected tool_call_start event") } - - call := response.Message.ToolCalls[0] - if call.ID != "call_1" || call.Name != "filesystem_edit" { - t.Fatalf("unexpected tool call: %+v", call) + if toolCallArgs.String() != `{"path":"main.go","search_string":"old","replace_string":"new"}` { + t.Fatalf("unexpected merged tool arguments: %q", toolCallArgs.String()) } - if call.Arguments != `{"path":"main.go","search_string":"old","replace_string":"new"}` { - t.Fatalf("unexpected merged arguments: %q", call.Arguments) + if messageDone == nil { + t.Fatal("expected message_done event") } - - var chunks []string - for { - select { - case event := <-events: - chunks = append(chunks, event.Text) - default: - if strings.Join(chunks, "") != "Hello world" { - t.Fatalf("expected streamed chunks to form %q, got %q", "Hello world", strings.Join(chunks, "")) - } - return - } + if messageDone.FinishReason != "tool_calls" { + t.Fatalf("expected finish reason %q, got %q", "tool_calls", messageDone.FinishReason) + } + if messageDone.Usage == nil || messageDone.Usage.TotalTokens != 15 { + t.Fatalf("expected total tokens 15, got %+v", messageDone.Usage) } } @@ -395,7 +416,7 @@ func TestProviderChatHTTPErrorResponses(t *testing.T) { } provider.client = server.Client() - _, err = provider.Chat(context.Background(), domain.ChatRequest{ + err = provider.Chat(context.Background(), domain.ChatRequest{ Model: config.OpenAIDefaultModel, }, make(chan domain.StreamEvent, 1)) if err == nil || !strings.Contains(err.Error(), tt.expectErr) { @@ -513,7 +534,7 @@ func TestParseErrorAndEmitTextDelta(t *testing.T) { if err := emitTextDelta(context.Background(), eventCh, "chunk"); err != nil { t.Fatalf("emitTextDelta() error = %v", err) } - if got := <-eventCh; got.Text != "chunk" || got.Type != domain.StreamEventTextDelta { + if got := <-eventCh; got.Type != domain.StreamEventTextDelta || requireTextDeltaPayload(t, got).Text != "chunk" { t.Fatalf("unexpected stream event: %+v", got) } @@ -532,12 +553,63 @@ func TestProviderConsumeStreamRejectsDirtyJSON(t *testing.T) { t.Fatalf("New() error = %v", err) } - _, err = provider.consumeStream(context.Background(), strings.NewReader("data: {not-json}\n\n"), make(chan domain.StreamEvent, 1)) + err = provider.consumeStream(context.Background(), strings.NewReader("data: {not-json}\n\n"), make(chan domain.StreamEvent, 1)) if err == nil || !strings.Contains(err.Error(), "decode stream chunk") { t.Fatalf("expected dirty JSON decode error, got %v", err) } } +func drainStreamEvents(events <-chan domain.StreamEvent) []domain.StreamEvent { + drained := make([]domain.StreamEvent, 0) + for { + select { + case evt, ok := <-events: + if !ok { + return drained + } + drained = append(drained, evt) + default: + return drained + } + } +} + +func requireTextDeltaPayload(t *testing.T, event domain.StreamEvent) domain.TextDeltaPayload { + t.Helper() + payload, ok := event.Payload.(domain.TextDeltaPayload) + if !ok { + t.Fatalf("expected TextDeltaPayload, got %T", event.Payload) + } + return payload +} + +func requireToolCallStartPayload(t *testing.T, event domain.StreamEvent) domain.ToolCallStartPayload { + t.Helper() + payload, ok := event.Payload.(domain.ToolCallStartPayload) + if !ok { + t.Fatalf("expected ToolCallStartPayload, got %T", event.Payload) + } + return payload +} + +func requireToolCallDeltaPayload(t *testing.T, event domain.StreamEvent) domain.ToolCallDeltaPayload { + t.Helper() + payload, ok := event.Payload.(domain.ToolCallDeltaPayload) + if !ok { + t.Fatalf("expected ToolCallDeltaPayload, got %T", event.Payload) + } + return payload +} + +func requireMessageDonePayload(t *testing.T, event domain.StreamEvent) domain.MessageDonePayload { + t.Helper() + payload, ok := event.Payload.(domain.MessageDonePayload) + if !ok { + t.Fatalf("expected MessageDonePayload, got %T", event.Payload) + } + return payload +} + func containsToolRoleMessage(messages []openAIMessage, toolCallID string, content string) bool { for _, message := range messages { if message.Role == "tool" && message.ToolCallID == toolCallID && message.Content == content { @@ -601,7 +673,8 @@ func TestEmitToolCallStartGuards(t *testing.T) { t.Fatalf("emitToolCallStart() error = %v", err) } got := <-events - if got.Type != domain.StreamEventToolCallStart || got.ToolName != "filesystem_edit" || got.ToolCallID != "call-1" || got.ToolCallIndex != 2 { + payload := requireToolCallStartPayload(t, got) + if got.Type != domain.StreamEventToolCallStart || payload.Name != "filesystem_edit" || payload.ID != "call-1" || payload.Index != 2 { t.Fatalf("unexpected event: %+v", got) } @@ -646,16 +719,18 @@ func TestMergeToolCallDeltaEmitsStartWhenNameArrivesLater(t *testing.T) { if start.Type != domain.StreamEventToolCallStart { t.Fatalf("expected tool_call_start event, got %+v", start) } - if start.ToolCallID != "call_late_name" || start.ToolName != "filesystem_edit" { - t.Fatalf("unexpected tool_call_start payload: %+v", start) + startPayload := requireToolCallStartPayload(t, start) + if startPayload.ID != "call_late_name" || startPayload.Name != "filesystem_edit" { + t.Fatalf("unexpected tool_call_start payload: %+v", startPayload) } delta := <-events if delta.Type != domain.StreamEventToolCallDelta { t.Fatalf("expected tool_call_delta event, got %+v", delta) } - if delta.ToolArgumentsDelta != `{"path":"main.go"}` { - t.Fatalf("unexpected tool arguments delta: %+v", delta) + deltaPayload := requireToolCallDeltaPayload(t, delta) + if deltaPayload.ArgumentsDelta != `{"path":"main.go"}` { + t.Fatalf("unexpected tool arguments delta: %+v", deltaPayload) } call := toolCalls[0] @@ -703,7 +778,7 @@ func TestProviderChatEmitsToolCallStartEvent(t *testing.T) { provider.client = server.Client() events := make(chan domain.StreamEvent, 8) - _, err = provider.Chat(context.Background(), domain.ChatRequest{ + err = provider.Chat(context.Background(), domain.ChatRequest{ Model: config.OpenAIDefaultModel, Messages: []domain.Message{{Role: "user", Content: "edit"}}, Tools: []domain.ToolSpec{ @@ -714,17 +789,16 @@ func TestProviderChatEmitsToolCallStartEvent(t *testing.T) { t.Fatalf("Chat() error = %v", err) } - close(events) - var foundToolCallStart bool - for evt := range events { + for _, evt := range drainStreamEvents(events) { if evt.Type == domain.StreamEventToolCallStart { foundToolCallStart = true - if evt.ToolName != "filesystem_edit" { - t.Fatalf("expected ToolName %q, got %q", "filesystem_edit", evt.ToolName) + payload := requireToolCallStartPayload(t, evt) + if payload.Name != "filesystem_edit" { + t.Fatalf("expected ToolName %q, got %q", "filesystem_edit", payload.Name) } - if evt.ToolCallID != "call_tool" { - t.Fatalf("expected ToolCallID %q, got %q", "call_tool", evt.ToolCallID) + if payload.ID != "call_tool" { + t.Fatalf("expected ToolCallID %q, got %q", "call_tool", payload.ID) } } } @@ -819,7 +893,7 @@ func TestProviderChatEmitsFullEventStream(t *testing.T) { provider.client = server.Client() events := make(chan domain.StreamEvent, 16) - _, err = provider.Chat(context.Background(), domain.ChatRequest{ + err = provider.Chat(context.Background(), domain.ChatRequest{ Model: config.OpenAIDefaultModel, Messages: []domain.Message{{Role: "user", Content: "test"}}, }, events) @@ -827,35 +901,35 @@ func TestProviderChatEmitsFullEventStream(t *testing.T) { t.Fatalf("Chat() error = %v", err) } - close(events) - var ( foundTextDelta bool foundToolCallStart bool foundToolCallDelta bool foundMessageDone bool toolCallDeltaContent string - messageDoneEvt *domain.StreamEvent + messageDonePayload *domain.MessageDonePayload ) - for evt := range events { + for _, evt := range drainStreamEvents(events) { switch evt.Type { case domain.StreamEventTextDelta: foundTextDelta = true case domain.StreamEventToolCallStart: foundToolCallStart = true - if evt.ToolName != "filesystem_edit" { - t.Fatalf("expected ToolName %q, got %q", "filesystem_edit", evt.ToolName) + payload := requireToolCallStartPayload(t, evt) + if payload.Name != "filesystem_edit" { + t.Fatalf("expected ToolName %q, got %q", "filesystem_edit", payload.Name) } - if evt.ToolCallIndex != 0 { - t.Fatalf("expected ToolCallIndex %d for tool_call_start, got %d", 0, evt.ToolCallIndex) + if payload.Index != 0 { + t.Fatalf("expected ToolCallIndex %d for tool_call_start, got %d", 0, payload.Index) } case domain.StreamEventToolCallDelta: foundToolCallDelta = true - toolCallDeltaContent += evt.ToolArgumentsDelta + toolCallDeltaContent += requireToolCallDeltaPayload(t, evt).ArgumentsDelta case domain.StreamEventMessageDone: foundMessageDone = true - messageDoneEvt = &evt + payload := requireMessageDonePayload(t, evt) + messageDonePayload = &payload } } @@ -879,16 +953,16 @@ func TestProviderChatEmitsFullEventStream(t *testing.T) { } // 验证 message_done 事件包含正确的字段 - if messageDoneEvt == nil { + if messageDonePayload == nil { t.Fatal("message_done event is nil") } - if messageDoneEvt.FinishReason != "tool_calls" { - t.Fatalf("expected FinishReason %q, got %q", "tool_calls", messageDoneEvt.FinishReason) + if messageDonePayload.FinishReason != "tool_calls" { + t.Fatalf("expected FinishReason %q, got %q", "tool_calls", messageDonePayload.FinishReason) } - if messageDoneEvt.Usage == nil { + if messageDonePayload.Usage == nil { t.Fatal("expected Usage in message_done event") } - if messageDoneEvt.Usage.TotalTokens != 150 { - t.Fatalf("expected TotalTokens %d, got %d", 150, messageDoneEvt.Usage.TotalTokens) + if messageDonePayload.Usage.TotalTokens != 150 { + t.Fatalf("expected TotalTokens %d, got %d", 150, messageDonePayload.Usage.TotalTokens) } } diff --git a/internal/provider/provider.go b/internal/provider/provider.go index 40a838b9..91a99b52 100644 --- a/internal/provider/provider.go +++ b/internal/provider/provider.go @@ -3,5 +3,5 @@ package provider import "context" type Provider interface { - Chat(ctx context.Context, req ChatRequest, events chan<- StreamEvent) (ChatResponse, error) + Chat(ctx context.Context, req ChatRequest, events chan<- StreamEvent) error } diff --git a/internal/provider/registry_test.go b/internal/provider/registry_test.go index f650253d..df57c9da 100644 --- a/internal/provider/registry_test.go +++ b/internal/provider/registry_test.go @@ -12,8 +12,8 @@ import ( type stubProvider struct{} -func (stubProvider) Chat(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) (provider.ChatResponse, error) { - return provider.ChatResponse{}, nil +func (stubProvider) Chat(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { + return nil } func stubDriver(driverType string) provider.DriverDefinition { diff --git a/internal/provider/types.go b/internal/provider/types.go index 294d31b1..278a3a93 100644 --- a/internal/provider/types.go +++ b/internal/provider/types.go @@ -1,11 +1,14 @@ package provider -// Role 常量定义消息角色标识。 const ( - RoleSystem = "system" - RoleUser = "user" + // RoleSystem 标识系统消息。 + RoleSystem = "system" + // RoleUser 标识用户消息。 + RoleUser = "user" + // RoleAssistant 标识助手消息。 RoleAssistant = "assistant" - RoleTool = "tool" + // RoleTool 标识工具结果消息。 + RoleTool = "tool" ) // Message 表示对话中的单条消息。 @@ -39,13 +42,6 @@ type ChatRequest struct { Tools []ToolSpec `json:"tools,omitempty"` } -// ChatResponse 是 provider.Chat() 的返回结果。 -type ChatResponse struct { - Message Message `json:"message"` - FinishReason string `json:"finish_reason"` - Usage Usage `json:"usage"` -} - // Usage 记录本次请求的 token 使用统计。 type Usage struct { InputTokens int `json:"input_tokens"` @@ -59,39 +55,20 @@ type StreamEventType string const ( // StreamEventTextDelta 表示模型输出的文本片段。 StreamEventTextDelta StreamEventType = "text_delta" - // StreamEventToolCallStart 表示模型开始请求工具调用,TUI 可据此展示过渡提示。 + // StreamEventToolCallStart 表示模型开始请求工具调用。 StreamEventToolCallStart StreamEventType = "tool_call_start" // StreamEventToolCallDelta 表示工具调用参数的增量片段。 StreamEventToolCallDelta StreamEventType = "tool_call_delta" - // StreamEventMessageDone 表示本轮消息完成,包含最终统计信息。 + // StreamEventMessageDone 表示本轮消息完成,并携带最终统计信息。 StreamEventMessageDone StreamEventType = "message_done" ) -// StreamEvent 表示 provider 驱动层向 runtime 推送的流式事件。 -// 强制使用 NewXxxStreamEvent 构造器创建实例,禁止直接构造。 +// StreamEvent 表示 provider 向 runtime 推送的流式事件。 type StreamEvent struct { Type StreamEventType `json:"type"` - Payload interface{} `json:"payload,omitempty"` // 强类型载荷,使用类型断言访问 - - // --- 以下字段已弃用,保留仅用于 Phase 1 向后兼容 --- - // Phase 2 将由 Runtime 负责人移除,届时所有消费方应通过 Payload 类型断言访问。 - - // text_delta - Text string `json:"text,omitempty"` // 文本片段 - - // tool_call_start / tool_call_delta - ToolCallIndex int `json:"tool_call_index,omitempty"` // 工具调用索引 - ToolCallID string `json:"tool_call_id,omitempty"` // 工具调用 ID(tool_call_start 时使用) - ToolName string `json:"tool_name,omitempty"` // 工具名称(tool_call_start 时使用) - ToolArgumentsDelta string `json:"tool_arguments_delta,omitempty"` // 参数增量片段(tool_call_delta 时使用) - - // message_done - FinishReason string `json:"finish_reason,omitempty"` // 结束原因(仅 message_done 时有效) - Usage *Usage `json:"usage,omitempty"` // 使用统计(仅 message_done 时有效) + Payload interface{} `json:"payload,omitempty"` } -// --- Payload 强类型定义 --- - // TextDeltaPayload 表示文本增量事件的载荷。 type TextDeltaPayload struct { Text string `json:"text"` @@ -117,15 +94,11 @@ type MessageDonePayload struct { Usage *Usage `json:"usage"` } -// --- 构造器 --- - // NewTextDeltaStreamEvent 创建文本增量流事件。 func NewTextDeltaStreamEvent(text string) StreamEvent { return StreamEvent{ Type: StreamEventTextDelta, Payload: TextDeltaPayload{Text: text}, - // 兼容层:同步填充弃用字段 - Text: text, } } @@ -134,10 +107,6 @@ func NewToolCallStartStreamEvent(index int, id, name string) StreamEvent { return StreamEvent{ Type: StreamEventToolCallStart, Payload: ToolCallStartPayload{Index: index, ID: id, Name: name}, - // 兼容层:同步填充弃用字段 - ToolCallIndex: index, - ToolCallID: id, - ToolName: name, } } @@ -146,10 +115,6 @@ func NewToolCallDeltaStreamEvent(index int, id, argumentsDelta string) StreamEve return StreamEvent{ Type: StreamEventToolCallDelta, Payload: ToolCallDeltaPayload{Index: index, ID: id, ArgumentsDelta: argumentsDelta}, - // 兼容层:同步填充弃用字段 - ToolCallIndex: index, - ToolCallID: id, - ToolArgumentsDelta: argumentsDelta, } } @@ -158,8 +123,5 @@ func NewMessageDoneStreamEvent(finishReason string, usage *Usage) StreamEvent { return StreamEvent{ Type: StreamEventMessageDone, Payload: MessageDonePayload{FinishReason: finishReason, Usage: usage}, - // 兼容层:同步填充弃用字段 - FinishReason: finishReason, - Usage: usage, } } diff --git a/internal/runtime/compact_generator.go b/internal/runtime/compact_generator.go index 87d9d658..49e48755 100644 --- a/internal/runtime/compact_generator.go +++ b/internal/runtime/compact_generator.go @@ -56,22 +56,61 @@ func (g *compactSummaryGenerator) Generate(ctx context.Context, input contextcom if err != nil { return "", err } - resp, err := modelProvider.Chat(ctx, provider.ChatRequest{ + + // 使用流式事件通道收集 compact 摘要响应。 + streamEvents := make(chan provider.StreamEvent, 32) + streamDone := make(chan struct{}) + acc := newStreamAccumulator() + + go func() { + defer close(streamDone) + for { + select { + case event, ok := <-streamEvents: + if !ok { + return + } + switch event.Type { + case provider.StreamEventTextDelta: + if payload, ok := event.Payload.(provider.TextDeltaPayload); ok { + acc.accumulateTextDelta(payload.Text) + } + case provider.StreamEventToolCallStart: + if payload, ok := event.Payload.(provider.ToolCallStartPayload); ok { + acc.accumulateToolCallStart(payload.Index, payload.ID, payload.Name) + } + case provider.StreamEventToolCallDelta: + if payload, ok := event.Payload.(provider.ToolCallDeltaPayload); ok { + acc.accumulateToolCallDelta(payload.Index, payload.ID, payload.ArgumentsDelta) + } + } + case <-ctx.Done(): + return + } + } + }() + + err = modelProvider.Chat(ctx, provider.ChatRequest{ Model: g.model, SystemPrompt: prompt.SystemPrompt, Messages: []provider.Message{{ Role: provider.RoleUser, Content: prompt.UserPrompt, }}, - }, nil) + }, streamEvents) + close(streamEvents) + <-streamDone + if err != nil { return "", err } - if len(resp.Message.ToolCalls) > 0 { + + message := acc.buildMessage() + if len(message.ToolCalls) > 0 { return "", errors.New("runtime: compact summary response must not contain tool calls") } - summary := strings.TrimSpace(resp.Message.Content) + summary := strings.TrimSpace(message.Content) if summary == "" { return "", errors.New("runtime: compact summary response is empty") } diff --git a/internal/runtime/compact_generator_test.go b/internal/runtime/compact_generator_test.go index def57458..e0c348e7 100644 --- a/internal/runtime/compact_generator_test.go +++ b/internal/runtime/compact_generator_test.go @@ -20,28 +20,25 @@ func TestCompactSummaryGeneratorBuildsProviderRequestWithoutTools(t *testing.T) } scripted := &scriptedProvider{ - responses: []provider.ChatResponse{{ - Message: provider.Message{ - Role: provider.RoleAssistant, - Content: strings.Join([]string{ - "[compact_summary]", - "done:", - "- Completed the historical task and kept the final result.", - "", - "in_progress:", - "- Continue from the retained recent window.", - "", - "decisions:", - "- Keep the existing section layout for compatibility.", - "", - "code_changes:", - "- Updated compact summary generation behavior.", - "", - "constraints:", - "- Preserve only the minimum information needed to continue the work.", - }, "\n"), - }, - }}, + streams: [][]provider.StreamEvent{ + {provider.NewTextDeltaStreamEvent(strings.Join([]string{ + "[compact_summary]", + "done:", + "- Completed the historical task and kept the final result.", + "", + "in_progress:", + "- Continue from the retained recent window.", + "", + "decisions:", + "- Keep the existing section layout for compatibility.", + "", + "code_changes:", + "- Updated compact summary generation behavior.", + "", + "constraints:", + "- Preserve only the minimum information needed to continue the work.", + }, "\n"))}, + }, } factory := &scriptedProviderFactory{provider: scripted} generator := newCompactSummaryGenerator(factory, resolvedProvider, "session-model") @@ -116,14 +113,12 @@ func TestCompactSummaryGeneratorRejectsToolCalls(t *testing.T) { } scripted := &scriptedProvider{ - responses: []provider.ChatResponse{{ - Message: provider.Message{ - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ - {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, - }, + streams: [][]provider.StreamEvent{ + { + provider.NewToolCallStartStreamEvent(0, "call-1", "filesystem_read_file"), + provider.NewToolCallDeltaStreamEvent(0, "call-1", "{}"), }, - }}, + }, } generator := newCompactSummaryGenerator(&scriptedProviderFactory{provider: scripted}, resolvedProvider, "session-model") diff --git a/internal/runtime/runtime.go b/internal/runtime/runtime.go index c95bbcd3..d3b3e29f 100644 --- a/internal/runtime/runtime.go +++ b/internal/runtime/runtime.go @@ -8,6 +8,7 @@ import ( "math/rand/v2" "os" "path/filepath" + "sort" "strings" "sync" "time" @@ -29,6 +30,67 @@ const ( providerRetryMaxWait = 5 * time.Second ) +// streamAccumulator 在流式事件处理过程中累积本轮对话需要持久化的助手消息状态, +// 包括文本内容和工具调用列表。 +type streamAccumulator struct { + content strings.Builder + toolCalls map[int]*provider.ToolCall +} + +// newStreamAccumulator 创建并初始化一个空的流式事件累积器。 +func newStreamAccumulator() *streamAccumulator { + return &streamAccumulator{ + toolCalls: make(map[int]*provider.ToolCall), + } +} + +// accumulateTextDelta 累积文本增量片段。 +func (a *streamAccumulator) accumulateTextDelta(text string) { + a.content.WriteString(text) +} + +// accumulateToolCallStart 记录新发现的工具调用(首次出现时创建条目)。 +func (a *streamAccumulator) accumulateToolCallStart(index int, id, name string) { + if _, exists := a.toolCalls[index]; !exists { + a.toolCalls[index] = &provider.ToolCall{ID: id, Name: name} + } +} + +// accumulateToolCallDelta 累积工具调用参数增量。 +func (a *streamAccumulator) accumulateToolCallDelta(index int, id, argumentsDelta string) { + call, exists := a.toolCalls[index] + if !exists { + call = &provider.ToolCall{ID: id} + a.toolCalls[index] = call + } + if name := call.Name; strings.TrimSpace(name) == "" && call.ID != "" { + // 首次出现 delta 时可能还未收到 start 事件,仅记录 ID + } + call.Arguments += argumentsDelta +} + +// buildMessage 从累积状态构建最终的 assistant Message 对象。 +func (a *streamAccumulator) buildMessage() provider.Message { + ordered := make([]int, 0, len(a.toolCalls)) + for index := range a.toolCalls { + ordered = append(ordered, index) + } + sort.Ints(ordered) + + message := provider.Message{ + Role: provider.RoleAssistant, + Content: a.content.String(), + } + for _, index := range ordered { + call := a.toolCalls[index] + if call == nil { + continue + } + message.ToolCalls = append(message.ToolCalls, *call) + } + return message +} + var runtimeSessionWorkdirs = struct { mu sync.RWMutex data map[string]string @@ -169,7 +231,7 @@ func (s *Service) Run(ctx context.Context, input UserInput) error { return s.handleRunError(ctx, input.RunID, session.ID, err) } - resp, err := s.callProviderWithRetry(ctx, input.RunID, session.ID, provider.ChatRequest{ + acc, err := s.callProviderWithRetry(ctx, input.RunID, session.ID, provider.ChatRequest{ Model: cfg.CurrentModel, SystemPrompt: builtContext.SystemPrompt, Messages: builtContext.Messages, @@ -186,7 +248,7 @@ func (s *Service) Run(ctx context.Context, input UserInput) error { session.Provider = cfg.SelectedProvider session.Model = cfg.CurrentModel - assistant := resp.Message + assistant := acc.buildMessage() if strings.TrimSpace(assistant.Role) == "" { assistant.Role = provider.RoleAssistant } @@ -417,9 +479,9 @@ func (s *Service) emit(ctx context.Context, kind EventType, runID string, sessio } } -// forwardProviderEvents 将 provider 流式事件转发为 runtime 事件。 +// forwardProviderEvents 将 provider 流式事件转发为 runtime 事件,同时向 accumulator 累积消息状态。 // 使用 select 同时监听输入通道和 context 取消信号,确保 goroutine 不会因通道阻塞而泄漏。 -func (s *Service) forwardProviderEvents(ctx context.Context, runID string, sessionID string, input <-chan provider.StreamEvent, done chan<- struct{}) { +func (s *Service) forwardProviderEvents(ctx context.Context, runID string, sessionID string, input <-chan provider.StreamEvent, done chan<- struct{}, acc *streamAccumulator) { defer close(done) for { select { @@ -429,9 +491,25 @@ func (s *Service) forwardProviderEvents(ctx context.Context, runID string, sessi } switch event.Type { case provider.StreamEventTextDelta: - s.emit(ctx, EventAgentChunk, runID, sessionID, event.Text) + if payload, ok := event.Payload.(provider.TextDeltaPayload); ok { + s.emit(ctx, EventAgentChunk, runID, sessionID, payload.Text) + if acc != nil { + acc.accumulateTextDelta(payload.Text) + } + } case provider.StreamEventToolCallStart: - s.emit(ctx, EventToolCallThinking, runID, sessionID, event.ToolName) + if payload, ok := event.Payload.(provider.ToolCallStartPayload); ok { + s.emit(ctx, EventToolCallThinking, runID, sessionID, payload.Name) + if acc != nil { + acc.accumulateToolCallStart(payload.Index, payload.ID, payload.Name) + } + } + case provider.StreamEventToolCallDelta: + if payload, ok := event.Payload.(provider.ToolCallDeltaPayload); ok { + if acc != nil { + acc.accumulateToolCallDelta(payload.Index, payload.ID, payload.ArgumentsDelta) + } + } } case <-ctx.Done(): return @@ -489,18 +567,22 @@ func isRetryableProviderError(err error) bool { } // callProviderWithRetry 在可重试的 ProviderError 上自动重试 provider.Chat() 调用。 -// 每次重试都会重新构建 provider 实例和流式事件转发管道。 +// 每次重试都会重新创建 provider 实例、流式事件转发管道和累积器。 // 非可重试错误、context 取消、重试耗尽时直接返回错误。 +// 返回值 acc 保存了本轮流式事件的完整累积状态,供调用方构建 assistant Message 使用。 func (s *Service) callProviderWithRetry( ctx context.Context, runID string, sessionID string, req provider.ChatRequest, -) (provider.ChatResponse, error) { +) (*streamAccumulator, error) { + acc := newStreamAccumulator() var lastErr error for retryAttempt := 0; retryAttempt <= defaultProviderRetryMax; retryAttempt++ { if retryAttempt > 0 { + // 重试时重置累积器,避免混入上轮数据 + acc = newStreamAccumulator() wait := providerRetryBackoff(retryAttempt) s.emit(ctx, EventProviderRetry, runID, sessionID, fmt.Sprintf("retrying provider call (attempt %d/%d, wait=%.1fs)...", @@ -508,45 +590,48 @@ func (s *Service) callProviderWithRetry( select { case <-ctx.Done(): - return provider.ChatResponse{}, ctx.Err() + return nil, ctx.Err() case <-time.After(wait): } } resolvedProvider, err := s.configManager.ResolvedSelectedProvider() if err != nil { - return provider.ChatResponse{}, err + return nil, err } modelProvider, err := s.providerFactory.Build(ctx, resolvedProvider) if err != nil { - return provider.ChatResponse{}, err + return nil, err } streamEvents := make(chan provider.StreamEvent, 32) streamDone := make(chan struct{}) - go s.forwardProviderEvents(ctx, runID, sessionID, streamEvents, streamDone) + go s.forwardProviderEvents(ctx, runID, sessionID, streamEvents, streamDone, acc) - resp, err := modelProvider.Chat(ctx, req, streamEvents) + err = modelProvider.Chat(ctx, req, streamEvents) close(streamEvents) <-streamDone if err == nil { - return resp, nil + return acc, nil } lastErr = err // 非可重试错误或 context 已取消,立即返回。 if !isRetryableProviderError(err) { - return provider.ChatResponse{}, err + return nil, lastErr } if ctx.Err() != nil { - return provider.ChatResponse{}, ctx.Err() + return nil, ctx.Err() } } - return provider.ChatResponse{}, lastErr + if lastErr == nil { + lastErr = errors.New("max retries exceeded") + } + return nil, fmt.Errorf("runtime: max retries exhausted, last error: %w", lastErr) } // providerRetryBackoff 计算指数退避 + 随机抖动的等待时间。 diff --git a/internal/runtime/runtime_test.go b/internal/runtime/runtime_test.go index ff24909e..101a91ad 100644 --- a/internal/runtime/runtime_test.go +++ b/internal/runtime/runtime_test.go @@ -90,14 +90,13 @@ func (s *memoryStore) ListSummaries(ctx context.Context) ([]SessionSummary, erro type scriptedProvider struct { name string - responses []provider.ChatResponse streams [][]provider.StreamEvent requests []provider.ChatRequest callCount int - chatFn func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) (provider.ChatResponse, error) + chatFn func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error } -func (p *scriptedProvider) Chat(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) (provider.ChatResponse, error) { +func (p *scriptedProvider) Chat(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { p.requests = append(p.requests, cloneChatRequest(req)) callIndex := p.callCount @@ -112,15 +111,12 @@ func (p *scriptedProvider) Chat(ctx context.Context, req provider.ChatRequest, e select { case events <- event: case <-ctx.Done(): - return provider.ChatResponse{}, ctx.Err() + return ctx.Err() } } } - if callIndex >= len(p.responses) { - return provider.ChatResponse{}, fmt.Errorf("unexpected provider call %d", callIndex) - } - return p.responses[callIndex], nil + return nil } type scriptedProviderFactory struct { @@ -232,7 +228,6 @@ func TestServiceRun(t *testing.T) { tests := []struct { name string input UserInput - providerResponses []provider.ChatResponse providerStreams [][]provider.StreamEvent registerTool tools.Tool contextBuilder agentcontext.Builder @@ -245,19 +240,10 @@ func TestServiceRun(t *testing.T) { { name: "normal dialogue exits after final assistant reply", input: UserInput{RunID: "run-normal", Content: "hello"}, - providerResponses: []provider.ChatResponse{ - { - Message: provider.Message{ - Role: "assistant", - Content: "plain answer", - }, - FinishReason: "stop", - }, - }, providerStreams: [][]provider.StreamEvent{ { - {Type: provider.StreamEventTextDelta, Text: "plain "}, - {Type: provider.StreamEventTextDelta, Text: "answer"}, + provider.NewTextDeltaStreamEvent("plain "), + provider.NewTextDeltaStreamEvent("answer"), }, }, contextBuilder: &stubContextBuilder{ @@ -293,26 +279,15 @@ func TestServiceRun(t *testing.T) { { name: "tool call triggers execute and follow-up provider round", input: UserInput{RunID: "run-tool", Content: "edit file"}, - providerResponses: []provider.ChatResponse{ + // 第一轮:工具调用事件流(tool_call_start + tool_call_delta) + // 第二轮:普通文本回复 + providerStreams: [][]provider.StreamEvent{ { - Message: provider.Message{ - Role: "assistant", - ToolCalls: []provider.ToolCall{ - { - ID: "call-1", - Name: "filesystem_edit", - Arguments: `{"path":"main.go"}`, - }, - }, - }, - FinishReason: "tool_calls", + provider.NewToolCallStartStreamEvent(0, "call-1", "filesystem_edit"), + provider.NewToolCallDeltaStreamEvent(0, "call-1", `{"path":"main.go"}`), }, { - Message: provider.Message{ - Role: "assistant", - Content: "done", - }, - FinishReason: "stop", + provider.NewTextDeltaStreamEvent("done"), }, }, registerTool: &stubTool{ @@ -371,8 +346,7 @@ func TestServiceRun(t *testing.T) { } scripted := &scriptedProvider{ - responses: tt.providerResponses, - streams: tt.providerStreams, + streams: tt.providerStreams, } factory := &scriptedProviderFactory{provider: scripted} @@ -446,14 +420,8 @@ func TestServiceRunDelegatesToContextBuilder(t *testing.T) { } scripted := &scriptedProvider{ - responses: []provider.ChatResponse{ - { - Message: provider.Message{ - Role: "assistant", - Content: "done", - }, - FinishReason: "stop", - }, + streams: [][]provider.StreamEvent{ + {provider.NewTextDeltaStreamEvent("done")}, }, } @@ -501,10 +469,9 @@ func TestServiceRunPersistsSessionProviderAndModel(t *testing.T) { registry.Register(&stubTool{name: "filesystem_read_file", content: "default"}) scripted := &scriptedProvider{ - responses: []provider.ChatResponse{{ - Message: provider.Message{Role: provider.RoleAssistant, Content: "done"}, - FinishReason: "stop", - }}, + streams: [][]provider.StreamEvent{ + {provider.NewTextDeltaStreamEvent("done")}, + }, } service := NewWithFactory(manager, registry, store, &scriptedProviderFactory{provider: scripted}, nil) @@ -594,23 +561,12 @@ func TestServiceRunUsesToolManager(t *testing.T) { } scripted := &scriptedProvider{ - responses: []provider.ChatResponse{ + streams: [][]provider.StreamEvent{ { - Message: provider.Message{ - Role: "assistant", - ToolCalls: []provider.ToolCall{ - {ID: "call-manager", Name: "filesystem_edit", Arguments: `{"path":"main.go"}`}, - }, - }, - FinishReason: "tool_calls", - }, - { - Message: provider.Message{ - Role: "assistant", - Content: "done", - }, - FinishReason: "stop", + provider.NewToolCallStartStreamEvent(0, "call-manager", "filesystem_edit"), + provider.NewToolCallDeltaStreamEvent(0, "call-manager", `{"path":"main.go"}`), }, + {provider.NewTextDeltaStreamEvent("done")}, }, } @@ -672,20 +628,12 @@ func TestServiceRunEmitsPermissionRequestAndResolvedForAsk(t *testing.T) { } scripted := &scriptedProvider{ - responses: []provider.ChatResponse{ - { - Message: provider.Message{ - Role: "assistant", - ToolCalls: []provider.ToolCall{ - {ID: "call-ask", Name: "webfetch", Arguments: `{"url":"https://example.com/private"}`}, - }, - }, - FinishReason: "tool_calls", - }, + streams: [][]provider.StreamEvent{ { - Message: provider.Message{Role: "assistant", Content: "done"}, - FinishReason: "stop", + provider.NewToolCallStartStreamEvent(0, "call-ask", "webfetch"), + provider.NewToolCallDeltaStreamEvent(0, "call-ask", `{"url":"https://example.com/private"}`), }, + {provider.NewTextDeltaStreamEvent("done")}, }, } @@ -768,20 +716,12 @@ func TestServiceRunEmitsPermissionResolvedForDeny(t *testing.T) { } scripted := &scriptedProvider{ - responses: []provider.ChatResponse{ + streams: [][]provider.StreamEvent{ { - Message: provider.Message{ - Role: "assistant", - ToolCalls: []provider.ToolCall{ - {ID: "call-deny", Name: "bash", Arguments: `{"command":"echo hi"}`}, - }, - }, - FinishReason: "tool_calls", - }, - { - Message: provider.Message{Role: "assistant", Content: "done"}, - FinishReason: "stop", + provider.NewToolCallStartStreamEvent(0, "call-deny", "bash"), + provider.NewToolCallDeltaStreamEvent(0, "call-deny", `{"command":"echo hi"}`), }, + {provider.NewTextDeltaStreamEvent("done")}, }, } @@ -853,14 +793,8 @@ func TestServiceNewWithFactoryDefaultsToolManager(t *testing.T) { store := newMemoryStore() service := NewWithFactory(manager, nil, store, &scriptedProviderFactory{ provider: &scriptedProvider{ - responses: []provider.ChatResponse{ - { - Message: provider.Message{ - Role: provider.RoleAssistant, - Content: "done", - }, - FinishReason: "stop", - }, + streams: [][]provider.StreamEvent{ + {provider.NewTextDeltaStreamEvent("done")}, }, }, }, nil) @@ -902,15 +836,10 @@ func TestServiceRunErrorPaths(t *testing.T) { input: UserInput{RunID: "run-max-loops", Content: "loop"}, maxLoops: 1, provider: &scriptedProvider{ - responses: []provider.ChatResponse{ + streams: [][]provider.StreamEvent{ { - Message: provider.Message{ - Role: "assistant", - ToolCalls: []provider.ToolCall{ - {ID: "loop-call", Name: "filesystem_edit", Arguments: `{"path":"x"}`}, - }, - }, - FinishReason: "tool_calls", + provider.NewToolCallStartStreamEvent(0, "loop-call", "filesystem_edit"), + provider.NewToolCallDeltaStreamEvent(0, "loop-call", `{"path":"x"}`), }, }, }, @@ -946,14 +875,8 @@ func TestServiceRunErrorPaths(t *testing.T) { Content: "continue", }, provider: &scriptedProvider{ - responses: []provider.ChatResponse{ - { - Message: provider.Message{ - Role: "assistant", - Content: "resumed", - }, - FinishReason: "stop", - }, + streams: [][]provider.StreamEvent{ + {provider.NewTextDeltaStreamEvent("resumed")}, }, }, seedSession: &Session{ @@ -984,23 +907,18 @@ func TestServiceRunErrorPaths(t *testing.T) { callIdx := 0 return &scriptedProvider{ name: "retry-then-success", - chatFn: func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) (provider.ChatResponse, error) { + chatFn: func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { callIdx++ if callIdx == 1 { - return provider.ChatResponse{}, &provider.ProviderError{ + return &provider.ProviderError{ StatusCode: 500, Code: provider.ErrorCodeServer, Message: "internal server error", Retryable: true, } } - return provider.ChatResponse{ - Message: provider.Message{ - Role: "assistant", - Content: "recovered", - }, - FinishReason: "stop", - }, nil + events <- provider.NewTextDeltaStreamEvent("recovered") + return nil }, } }(), @@ -1024,8 +942,8 @@ func TestServiceRunErrorPaths(t *testing.T) { input: UserInput{RunID: "run-no-retry", Content: "hello"}, provider: &scriptedProvider{ name: "auth-error-no-retry", - chatFn: func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) (provider.ChatResponse, error) { - return provider.ChatResponse{}, &provider.ProviderError{ + chatFn: func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { + return &provider.ProviderError{ StatusCode: 401, Code: provider.ErrorCodeAuthFailed, Message: "invalid api key", @@ -1047,8 +965,8 @@ func TestServiceRunErrorPaths(t *testing.T) { input: UserInput{RunID: "run-retry-exhausted", Content: "hello"}, provider: &scriptedProvider{ name: "always-500", - chatFn: func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) (provider.ChatResponse, error) { - return provider.ChatResponse{}, &provider.ProviderError{ + chatFn: func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { + return &provider.ProviderError{ StatusCode: 500, Code: provider.ErrorCodeServer, Message: "internal server error", @@ -1131,10 +1049,10 @@ func TestServiceCancelActiveRun(t *testing.T) { started := make(chan struct{}) scripted := &scriptedProvider{ name: "cancel-active-run-provider", - chatFn: func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) (provider.ChatResponse, error) { + chatFn: func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { close(started) <-ctx.Done() - return provider.ChatResponse{}, ctx.Err() + return ctx.Err() }, } @@ -1174,10 +1092,10 @@ func TestServiceRunCanceledByProvider(t *testing.T) { started := make(chan struct{}) scripted := &scriptedProvider{ name: "blocking-provider", - chatFn: func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) (provider.ChatResponse, error) { + chatFn: func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { close(started) <-ctx.Done() - return provider.ChatResponse{}, ctx.Err() + return ctx.Err() }, } @@ -1219,10 +1137,10 @@ func TestServiceRunPreservesProviderErrorAfterCancel(t *testing.T) { providerErr := errors.New("provider failed after cancel") scripted := &scriptedProvider{ name: "provider-error-after-cancel", - chatFn: func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) (provider.ChatResponse, error) { + chatFn: func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { close(started) <-ctx.Done() - return provider.ChatResponse{}, providerErr + return providerErr }, } @@ -1271,15 +1189,10 @@ func TestServiceRunCanceledDuringToolExecution(t *testing.T) { scripted := &scriptedProvider{ name: "tool-cancel-provider", - responses: []provider.ChatResponse{ + streams: [][]provider.StreamEvent{ { - Message: provider.Message{ - Role: "assistant", - ToolCalls: []provider.ToolCall{ - {ID: "cancel-call", Name: "filesystem_edit", Arguments: `{"path":"main.go"}`}, - }, - }, - FinishReason: "tool_calls", + provider.NewToolCallStartStreamEvent(0, "cancel-call", "filesystem_edit"), + provider.NewToolCallDeltaStreamEvent(0, "cancel-call", `{"path":"main.go"}`), }, }, } @@ -1336,15 +1249,10 @@ func TestServiceRunPreservesToolErrorAfterCancel(t *testing.T) { scripted := &scriptedProvider{ name: "tool-error-after-cancel-provider", - responses: []provider.ChatResponse{ + streams: [][]provider.StreamEvent{ { - Message: provider.Message{ - Role: "assistant", - ToolCalls: []provider.ToolCall{ - {ID: "tool-error-call", Name: "filesystem_edit", Arguments: `{"path":"main.go"}`}, - }, - }, - FinishReason: "tool_calls", + provider.NewToolCallStartStreamEvent(0, "tool-error-call", "filesystem_edit"), + provider.NewToolCallDeltaStreamEvent(0, "tool-error-call", `{"path":"main.go"}`), }, }, } @@ -1433,23 +1341,12 @@ func TestServiceRunToolTimeoutIsNotCancellation(t *testing.T) { scripted := &scriptedProvider{ name: "timeout-provider", - responses: []provider.ChatResponse{ + streams: [][]provider.StreamEvent{ { - Message: provider.Message{ - Role: "assistant", - ToolCalls: []provider.ToolCall{ - {ID: "timeout-call", Name: "filesystem_edit", Arguments: `{"path":"main.go"}`}, - }, - }, - FinishReason: "tool_calls", - }, - { - Message: provider.Message{ - Role: "assistant", - Content: "done after timeout", - }, - FinishReason: "stop", + provider.NewToolCallStartStreamEvent(0, "timeout-call", "filesystem_edit"), + provider.NewToolCallDeltaStreamEvent(0, "timeout-call", `{"path":"main.go"}`), }, + {provider.NewTextDeltaStreamEvent("done after timeout")}, }, } @@ -1603,28 +1500,25 @@ func TestServiceCompactUsesSessionProviderAndModelWhenPresent(t *testing.T) { registry.Register(&stubTool{name: "filesystem_read_file", content: "ok"}) scripted := &scriptedProvider{ - responses: []provider.ChatResponse{{ - Message: provider.Message{ - Role: provider.RoleAssistant, - Content: strings.Join([]string{ - "[compact_summary]", - "done:", - "- ok", - "", - "in_progress:", - "- continue", - "", - "decisions:", - "- kept existing provider and model", - "", - "code_changes:", - "- none", - "", - "constraints:", - "- none", - }, "\n"), - }, - }}, + streams: [][]provider.StreamEvent{ + {provider.NewTextDeltaStreamEvent(strings.Join([]string{ + "[compact_summary]", + "done:", + "- ok", + "", + "in_progress:", + "- continue", + "", + "decisions:", + "- kept existing provider and model", + "", + "code_changes:", + "- none", + "", + "constraints:", + "- none", + }, "\n"))}, + }, } factory := &scriptedProviderFactory{provider: scripted} service := NewWithFactory(manager, registry, store, factory, nil) @@ -1678,28 +1572,25 @@ func TestServiceCompactFallsBackToCurrentProviderWhenSessionMetadataMissing(t *t registry.Register(&stubTool{name: "filesystem_read_file", content: "ok"}) scripted := &scriptedProvider{ - responses: []provider.ChatResponse{{ - Message: provider.Message{ - Role: provider.RoleAssistant, - Content: strings.Join([]string{ - "[compact_summary]", - "done:", - "- ok", - "", - "in_progress:", - "- continue", - "", - "decisions:", - "- fallback to current selection", - "", - "code_changes:", - "- none", - "", - "constraints:", - "- none", - }, "\n"), - }, - }}, + streams: [][]provider.StreamEvent{ + {provider.NewTextDeltaStreamEvent(strings.Join([]string{ + "[compact_summary]", + "done:", + "- ok", + "", + "in_progress:", + "- continue", + "", + "decisions:", + "- fallback to current selection", + "", + "code_changes:", + "- none", + "", + "constraints:", + "- none", + }, "\n"))}, + }, } factory := &scriptedProviderFactory{provider: scripted} service := NewWithFactory(manager, registry, store, factory, nil) @@ -1738,20 +1629,12 @@ func TestServiceManualCompactThenRunContinuesToolRound(t *testing.T) { registry.Register(tool) scripted := &scriptedProvider{ - responses: []provider.ChatResponse{ - { - Message: provider.Message{ - Role: "assistant", - ToolCalls: []provider.ToolCall{ - {ID: "call-1", Name: "filesystem_read_file", Arguments: `{"path":"main.go"}`}, - }, - }, - FinishReason: "tool_calls", - }, + streams: [][]provider.StreamEvent{ { - Message: provider.Message{Role: "assistant", Content: "done"}, - FinishReason: "stop", + provider.NewToolCallStartStreamEvent(0, "call-1", "filesystem_read_file"), + provider.NewToolCallDeltaStreamEvent(0, "call-1", `{"path":"main.go"}`), }, + {provider.NewTextDeltaStreamEvent("done")}, }, } @@ -1826,17 +1709,15 @@ func TestServiceSerializesRunAndCompact(t *testing.T) { providerStarted := make(chan struct{}) unblockProvider := make(chan struct{}) scripted := &scriptedProvider{ - chatFn: func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) (provider.ChatResponse, error) { + chatFn: func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { select { case <-providerStarted: default: close(providerStarted) } <-unblockProvider - return provider.ChatResponse{ - Message: provider.Message{Role: provider.RoleAssistant, Content: "done"}, - FinishReason: "stop", - }, nil + events <- provider.NewTextDeltaStreamEvent("done") + return nil }, } @@ -1960,20 +1841,12 @@ func TestServiceRunUsesSessionWorkdirForContextAndTools(t *testing.T) { builder := &stubContextBuilder{} scripted := &scriptedProvider{ - responses: []provider.ChatResponse{ + streams: [][]provider.StreamEvent{ { - Message: provider.Message{ - Role: "assistant", - ToolCalls: []provider.ToolCall{ - {ID: "call-session-workdir", Name: "filesystem_edit", Arguments: `{"path":"main.go"}`}, - }, - }, - FinishReason: "tool_calls", - }, - { - Message: provider.Message{Role: "assistant", Content: "done"}, - FinishReason: "stop", + provider.NewToolCallStartStreamEvent(0, "call-session-workdir", "filesystem_edit"), + provider.NewToolCallDeltaStreamEvent(0, "call-session-workdir", `{"path":"main.go"}`), }, + {provider.NewTextDeltaStreamEvent("done")}, }, } @@ -2010,11 +1883,8 @@ func TestServiceRunUsesInputWorkdirForNewSession(t *testing.T) { registry.Register(&stubTool{name: "filesystem_read_file", content: "default"}) builder := &stubContextBuilder{} scripted := &scriptedProvider{ - responses: []provider.ChatResponse{ - { - Message: provider.Message{Role: "assistant", Content: "done"}, - FinishReason: "stop", - }, + streams: [][]provider.StreamEvent{ + {provider.NewTextDeltaStreamEvent("done")}, }, } diff --git a/internal/tui/update_test.go b/internal/tui/update_test.go index ed307cc0..9e28c76c 100644 --- a/internal/tui/update_test.go +++ b/internal/tui/update_test.go @@ -2603,8 +2603,8 @@ func newTestProviderService(t *testing.T, manager *config.Manager) *config.Selec type tUItestProvider struct{} -func (tUItestProvider) Chat(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) (provider.ChatResponse, error) { - return provider.ChatResponse{}, nil +func (tUItestProvider) Chat(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { + return nil } type tUItestCatalogStore struct { From 7b80117eb58f01b8dffbd87173c33a4d5d34b38d Mon Sep 17 00:00:00 2001 From: phantom5099 <1011668688@qq.com> Date: Sun, 5 Apr 2026 14:57:23 +0800 Subject: [PATCH 05/55] =?UTF-8?q?test(provider):=E8=A1=A5=E9=BD=90?= =?UTF-8?q?=E6=B5=8B=E8=AF=95=E8=A6=86=E7=9B=96=E7=8E=87?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/provider/openai/openai_test.go | 61 +++++++ internal/provider/types_test.go | 229 ++++++++++++++++++++++++ 2 files changed, 290 insertions(+) create mode 100644 internal/provider/types_test.go diff --git a/internal/provider/openai/openai_test.go b/internal/provider/openai/openai_test.go index 9b0fc108..ae0b6b9d 100644 --- a/internal/provider/openai/openai_test.go +++ b/internal/provider/openai/openai_test.go @@ -43,6 +43,67 @@ func TestWithTransport(t *testing.T) { } } +func TestNewValidationErrors(t *testing.T) { + t.Parallel() + + t.Run("empty api key returns error", func(t *testing.T) { + t.Parallel() + cfg := resolvedConfig("", "") + cfg.APIKey = "" + _, err := New(cfg) + if err == nil { + t.Fatal("expected error for empty api key") + } + if !strings.Contains(err.Error(), "api key is empty") { + t.Fatalf("expected api key error, got: %v", err) + } + }) + + t.Run("whitespace-only api key returns error", func(t *testing.T) { + t.Parallel() + cfg := resolvedConfig("", "") + cfg.APIKey = " " + _, err := New(cfg) + if err == nil { + t.Fatal("expected error for whitespace-only api key") + } + }) + + t.Run("invalid config validate fails", func(t *testing.T) { + t.Parallel() + // 空字符串的 BaseURL 和 Model 会导致 Validate 失败(取决于 config 实现) + cfg := config.ResolvedProviderConfig{ + ProviderConfig: config.ProviderConfig{ + Driver: DriverName, + BaseURL: "", + Model: "", + APIKeyEnv: "NONEXISTENT_ENV_VAR_" + t.Name(), + }, + APIKey: "test-key", + } + _, err := New(cfg) + // 验证失败时应该返回错误 + if err != nil { + // 预期行为:config 校验不通过 + return + } + // 如果校验通过了,也接受(取决于具体实现) + }) +} + +func TestNewDefaultTransportWhenNoOption(t *testing.T) { + t.Parallel() + + cfg := resolvedConfig("", "") + provider, err := New(cfg) // 不传任何 buildOption + if err != nil { + t.Fatalf("New() error = %v", err) + } + if provider.client.Transport == nil { + t.Fatal("expected default transport to be set") + } +} + func TestDefaultRetryTransport(t *testing.T) { t.Parallel() diff --git a/internal/provider/types_test.go b/internal/provider/types_test.go new file mode 100644 index 00000000..e5f1c5aa --- /dev/null +++ b/internal/provider/types_test.go @@ -0,0 +1,229 @@ +package provider + +import "testing" + +// --- Role 常量 --- + +func TestRoleConstants(t *testing.T) { + tests := []struct { + name string + got string + expect string + }{ + {"system", RoleSystem, "system"}, + {"user", RoleUser, "user"}, + {"assistant", RoleAssistant, "assistant"}, + {"tool", RoleTool, "tool"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if tt.got != tt.expect { + t.Fatalf("expected %q, got %q", tt.expect, tt.got) + } + }) + } +} + +// --- StreamEventType 常量 --- + +func TestStreamEventConstants(t *testing.T) { + tests := []struct { + name string + got StreamEventType + expect string + }{ + {"text_delta", StreamEventTextDelta, "text_delta"}, + {"tool_call_start", StreamEventToolCallStart, "tool_call_start"}, + {"tool_call_delta", StreamEventToolCallDelta, "tool_call_delta"}, + {"message_done", StreamEventMessageDone, "message_done"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if string(tt.got) != tt.expect { + t.Fatalf("expected %q, got %q", tt.expect, string(tt.got)) + } + }) + } +} + +// --- NewTextDeltaStreamEvent --- + +func TestNewTextDeltaStreamEvent(t *testing.T) { + t.Parallel() + + event := NewTextDeltaStreamEvent("hello") + if event.Type != StreamEventTextDelta { + t.Fatalf("expected type %q, got %q", StreamEventTextDelta, event.Type) + } + + payload, ok := event.Payload.(TextDeltaPayload) + if !ok { + t.Fatal("expected TextDeltaPayload type") + } + if payload.Text != "hello" { + t.Fatalf("expected text %q, got %q", "hello", payload.Text) + } +} + +// --- NewToolCallStartStreamEvent --- + +func TestNewToolCallStartStreamEvent(t *testing.T) { + t.Parallel() + + event := NewToolCallStartStreamEvent(3, "call_1", "edit_file") + if event.Type != StreamEventToolCallStart { + t.Fatalf("expected type %q, got %q", StreamEventToolCallStart, event.Type) + } + + payload, ok := event.Payload.(ToolCallStartPayload) + if !ok { + t.Fatal("expected ToolCallStartPayload type") + } + if payload.Index != 3 || payload.ID != "call_1" || payload.Name != "edit_file" { + t.Fatalf("unexpected payload: %+v", payload) + } +} + +// --- NewToolCallDeltaStreamEvent --- + +func TestNewToolCallDeltaStreamEvent(t *testing.T) { + t.Parallel() + + event := NewToolCallDeltaStreamEvent(1, "call_2", `{"path":"main.go"}`) + if event.Type != StreamEventToolCallDelta { + t.Fatalf("expected type %q, got %q", StreamEventToolCallDelta, event.Type) + } + + payload, ok := event.Payload.(ToolCallDeltaPayload) + if !ok { + t.Fatal("expected ToolCallDeltaPayload type") + } + if payload.Index != 1 || payload.ID != "call_2" || payload.ArgumentsDelta != `{"path":"main.go"}` { + t.Fatalf("unexpected payload: %+v", payload) + } +} + +// --- NewMessageDoneStreamEvent --- + +func TestNewMessageDoneStreamEvent(t *testing.T) { + t.Parallel() + + t.Run("with usage", func(t *testing.T) { + usage := &Usage{TotalTokens: 42} + event := NewMessageDoneStreamEvent("stop", usage) + + if event.Type != StreamEventMessageDone { + t.Fatalf("expected type %q, got %q", StreamEventMessageDone, event.Type) + } + + payload, ok := event.Payload.(MessageDonePayload) + if !ok { + t.Fatal("expected MessageDonePayload type") + } + if payload.FinishReason != "stop" { + t.Fatalf("expected finish reason %q, got %q", "stop", payload.FinishReason) + } + if payload.Usage == nil || payload.Usage.TotalTokens != 42 { + t.Fatalf("unexpected usage: %+v", payload.Usage) + } + }) + + t.Run("nil usage", func(t *testing.T) { + event := NewMessageDoneStreamEvent("tool_calls", nil) + + payload, ok := event.Payload.(MessageDonePayload) + if !ok { + t.Fatal("expected MessageDonePayload type") + } + if payload.FinishReason != "tool_calls" { + t.Fatalf("expected finish reason %q, got %q", "tool_calls", payload.FinishReason) + } + if payload.Usage != nil { + t.Fatal("expected nil usage") + } + }) + + t.Run("empty finish reason", func(t *testing.T) { + event := NewMessageDoneStreamEvent("", nil) + + payload, ok := event.Payload.(MessageDonePayload) + if !ok { + t.Fatal("expected MessageDonePayload type") + } + if payload.FinishReason != "" { + t.Fatalf("expected empty finish reason, got %q", payload.FinishReason) + } + }) +} + +// --- 结构体字段覆盖验证 --- + +func TestMessageStructFields(t *testing.T) { + t.Parallel() + + msg := Message{ + Role: RoleUser, + Content: "hello", + ToolCalls: []ToolCall{{ID: "t1"}}, + ToolCallID: "tc_1", + IsError: true, + } + if msg.Role != RoleUser || msg.Content != "hello" || len(msg.ToolCalls) != 1 || + msg.ToolCallID != "tc_1" || !msg.IsError { + t.Fatalf("message fields not as expected: %+v", msg) + } +} + +func TestToolCallStructFields(t *testing.T) { + t.Parallel() + + tc := ToolCall{ID: "c1", Name: "fn", Arguments: "{}"} + if tc.ID != "c1" || tc.Name != "fn" || tc.Arguments != "{}" { + t.Fatalf("tool call fields not as expected: %+v", tc) + } +} + +func TestToolSpecStructFields(t *testing.T) { + t.Parallel() + + spec := ToolSpec{Name: "read", Description: "read file", Schema: map[string]any{"type": "object"}} + if spec.Name != "read" || spec.Description != "read file" || spec.Schema == nil { + t.Fatalf("tool spec fields not as expected: %+v", spec) + } +} + +func TestChatRequestStructFields(t *testing.T) { + t.Parallel() + + req := ChatRequest{ + Model: "gpt-4", + SystemPrompt: "you are helpful", + Messages: []Message{{Role: RoleUser}}, + Tools: []ToolSpec{{Name: "bash"}}, + } + if req.Model != "gpt-4" || req.SystemPrompt != "you are helpful" || + len(req.Messages) != 1 || len(req.Tools) != 1 { + t.Fatalf("chat request fields not as expected: %+v", req) + } +} + +func TestUsageStructFields(t *testing.T) { + t.Parallel() + + usage := Usage{InputTokens: 10, OutputTokens: 20, TotalTokens: 30} + if usage.InputTokens != 10 || usage.OutputTokens != 20 || usage.TotalTokens != 30 { + t.Fatalf("usage fields not as expected: %+v", usage) + } +} + +func TestStreamEventStructFields(t *testing.T) { + t.Parallel() + + event := StreamEvent{ + Type: StreamEventTextDelta, + Payload: TextDeltaPayload{Text: "hi"}, + } + if event.Type != StreamEventTextDelta { + t.Fatalf("event type not as expected: %s", event.Type) + } +} From 04b273c54fce0c8b6ce736176467616fd314c77e Mon Sep 17 00:00:00 2001 From: Yumiue <229866007@qq.com> Date: Sun, 5 Apr 2026 15:12:20 +0800 Subject: [PATCH 06/55] =?UTF-8?q?fix(context):=20=E5=A2=9E=E5=8A=A0=20micr?= =?UTF-8?q?o=20compact=20=E5=9B=9E=E9=80=80=E5=BC=80=E5=85=B3=E5=B9=B6?= =?UTF-8?q?=E6=94=B6=E6=95=9B=E5=B7=A5=E5=85=B7=E5=90=8D=E5=B8=B8=E9=87=8F?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/guides/configuration.md | 2 ++ internal/config/config_test.go | 10 ++++++ internal/config/loader.go | 3 ++ internal/config/model.go | 1 + internal/context/builder.go | 16 +++++++-- internal/context/builder_test.go | 54 ++++++++++++++++++++++++++++ internal/context/microcompact.go | 15 ++++---- internal/context/types.go | 6 ++++ internal/runtime/runtime.go | 3 ++ internal/runtime/runtime_test.go | 44 +++++++++++++++++++++++ internal/tools/bash/tool.go | 2 +- internal/tools/filesystem/helpers.go | 12 ++++--- internal/tools/names.go | 12 +++++++ internal/tools/webfetch/tool.go | 2 +- 14 files changed, 166 insertions(+), 16 deletions(-) create mode 100644 internal/tools/names.go diff --git a/docs/guides/configuration.md b/docs/guides/configuration.md index fed73849..155e29d9 100644 --- a/docs/guides/configuration.md +++ b/docs/guides/configuration.md @@ -315,6 +315,7 @@ context: manual_strategy: keep_recent manual_keep_recent_messages: 10 max_summary_chars: 1200 + micro_compact_disabled: false ``` ### 字段说明 @@ -324,5 +325,6 @@ context: | `context.compact.manual_strategy` | string | `keep_recent` | 手动 `/compact` 策略,可选 `keep_recent` / `full_replace` | | `context.compact.manual_keep_recent_messages` | int | `10` | `keep_recent` 模式下保留最近 N 条消息;会按 tool call 与 tool result 的原子块整体保留 | | `context.compact.max_summary_chars` | int | `1200` | compact summary 最大字符数 | +| `context.compact.micro_compact_disabled` | bool | `false` | 是否关闭默认启用的读时 micro compact;设为 `true` 可快速回退到仅 trim、不做旧工具结果清理 | 更多行为说明见 [context-compact.md](../context-compact.md)。 diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 4516c5e0..a2504ac2 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -876,10 +876,14 @@ func TestCompactConfigDefaultsAndRoundTrip(t *testing.T) { if compactCfg.MaxSummaryChars != DefaultCompactMaxSummaryChars { t.Fatalf("expected max_summary_chars=%d, got %d", DefaultCompactMaxSummaryChars, compactCfg.MaxSummaryChars) } + if compactCfg.MicroCompactDisabled { + t.Fatalf("expected micro compact to be enabled by default") + } cfg.Context.Compact.ManualStrategy = CompactManualStrategyFullReplace cfg.Context.Compact.ManualKeepRecentMessages = 2 cfg.Context.Compact.MaxSummaryChars = 900 + cfg.Context.Compact.MicroCompactDisabled = true if err := loader.Save(context.Background(), cfg); err != nil { t.Fatalf("Save() error = %v", err) } @@ -894,6 +898,9 @@ func TestCompactConfigDefaultsAndRoundTrip(t *testing.T) { if strings.Contains(text, "manual_keep_recent_spans:") { t.Fatalf("expected persisted config to drop legacy manual_keep_recent_spans key, got:\n%s", text) } + if !strings.Contains(text, "micro_compact_disabled: true") { + t.Fatalf("expected persisted config to include micro_compact_disabled, got:\n%s", text) + } reloaded, err := loader.Load(context.Background()) if err != nil { @@ -908,6 +915,9 @@ func TestCompactConfigDefaultsAndRoundTrip(t *testing.T) { if reloaded.Context.Compact.MaxSummaryChars != 900 { t.Fatalf("expected max_summary_chars=900, got %d", reloaded.Context.Compact.MaxSummaryChars) } + if !reloaded.Context.Compact.MicroCompactDisabled { + t.Fatalf("expected micro_compact_disabled to persist") + } } func TestCompactConfigValidateFailures(t *testing.T) { diff --git a/internal/config/loader.go b/internal/config/loader.go index 6c8ef449..ea3a6c3e 100644 --- a/internal/config/loader.go +++ b/internal/config/loader.go @@ -41,6 +41,7 @@ type persistedCompactConfig struct { ManualStrategy string `yaml:"manual_strategy,omitempty"` ManualKeepRecentMessages int `yaml:"manual_keep_recent_messages,omitempty"` MaxSummaryChars int `yaml:"max_summary_chars,omitempty"` + MicroCompactDisabled bool `yaml:"micro_compact_disabled,omitempty"` } func NewLoader(baseDir string, defaults *Config) *Loader { @@ -217,6 +218,7 @@ func newPersistedContextConfig(cfg ContextConfig) persistedContextConfig { ManualStrategy: cfg.Compact.ManualStrategy, ManualKeepRecentMessages: cfg.Compact.ManualKeepRecentMessages, MaxSummaryChars: cfg.Compact.MaxSummaryChars, + MicroCompactDisabled: cfg.Compact.MicroCompactDisabled, }, } } @@ -228,6 +230,7 @@ func fromPersistedContextConfig(file persistedContextConfig, defaults ContextCon ManualStrategy: strings.TrimSpace(file.Compact.ManualStrategy), ManualKeepRecentMessages: file.Compact.ManualKeepRecentMessages, MaxSummaryChars: file.Compact.MaxSummaryChars, + MicroCompactDisabled: file.Compact.MicroCompactDisabled, }, } out.Compact.ApplyDefaults(defaults.Compact) diff --git a/internal/config/model.go b/internal/config/model.go index e8d1d3c9..75992d1c 100644 --- a/internal/config/model.go +++ b/internal/config/model.go @@ -70,6 +70,7 @@ type CompactConfig struct { ManualStrategy string `yaml:"manual_strategy,omitempty"` ManualKeepRecentMessages int `yaml:"manual_keep_recent_messages,omitempty"` MaxSummaryChars int `yaml:"max_summary_chars,omitempty"` + MicroCompactDisabled bool `yaml:"micro_compact_disabled,omitempty"` } type WebFetchConfig struct { diff --git a/internal/context/builder.go b/internal/context/builder.go index f1ca2c79..eeac9406 100644 --- a/internal/context/builder.go +++ b/internal/context/builder.go @@ -1,6 +1,10 @@ package context -import "context" +import ( + "context" + + "neo-code/internal/provider" +) // DefaultBuilder preserves the current runtime context-building behavior. type DefaultBuilder struct { @@ -43,6 +47,14 @@ func (b *DefaultBuilder) Build(ctx context.Context, input BuildInput) (BuildResu return BuildResult{ SystemPrompt: composeSystemPrompt(sections...), - Messages: microCompactMessages(trimPolicy.Trim(input.Messages)), + Messages: applyReadTimeContextProjection(trimPolicy.Trim(input.Messages), input.Compact), }, nil } + +// applyReadTimeContextProjection 负责在 provider 请求前按开关应用只读上下文投影,避免改写原始会话消息。 +func applyReadTimeContextProjection(messages []provider.Message, options CompactOptions) []provider.Message { + if options.DisableMicroCompact { + return cloneContextMessages(messages) + } + return microCompactMessages(messages) +} diff --git a/internal/context/builder_test.go b/internal/context/builder_test.go index 9f5098a0..8b72951d 100644 --- a/internal/context/builder_test.go +++ b/internal/context/builder_test.go @@ -5,6 +5,7 @@ import ( "fmt" "os" "path/filepath" + "reflect" "strings" "testing" @@ -204,6 +205,59 @@ func TestDefaultBuilderBuildAppliesMicroCompactAfterTrim(t *testing.T) { } } +func TestDefaultBuilderBuildSkipsMicroCompactWhenDisabled(t *testing.T) { + t.Parallel() + + builder := &DefaultBuilder{ + promptSources: []promptSectionSource{ + stubPromptSectionSource{sections: []promptSection{{title: "Stub", content: "body"}}}, + }, + } + + messages := []provider.Message{ + {Role: provider.RoleUser, Content: "older user"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-1", Content: "old read result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-2", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-3", Name: "webfetch", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, + {Role: provider.RoleUser, Content: "latest explicit instruction"}, + {Role: provider.RoleAssistant, Content: "current reply"}, + } + + got, err := builder.Build(stdcontext.Background(), BuildInput{ + Messages: messages, + Compact: CompactOptions{ + DisableMicroCompact: true, + }, + }) + if err != nil { + t.Fatalf("Build() error = %v", err) + } + if !reflect.DeepEqual(got.Messages, messages) { + t.Fatalf("expected messages to remain unchanged when micro compact is disabled, got %+v", got.Messages) + } + if &got.Messages[2] == &messages[2] { + t.Fatalf("expected disabled path to still clone message slice") + } +} + func TestTrimMessagesPreservesToolPairs(t *testing.T) { t.Parallel() diff --git a/internal/context/microcompact.go b/internal/context/microcompact.go index 10f5baba..10a36cde 100644 --- a/internal/context/microcompact.go +++ b/internal/context/microcompact.go @@ -5,6 +5,7 @@ import ( "neo-code/internal/context/internalcompact" "neo-code/internal/provider" + "neo-code/internal/tools" ) const ( @@ -15,13 +16,13 @@ const ( ) var microCompactableTools = map[string]struct{}{ - "bash": {}, - "webfetch": {}, - "filesystem_read_file": {}, - "filesystem_grep": {}, - "filesystem_glob": {}, - "filesystem_edit": {}, - "filesystem_write_file": {}, + tools.ToolNameBash: {}, + tools.ToolNameWebFetch: {}, + tools.ToolNameFilesystemReadFile: {}, + tools.ToolNameFilesystemGrep: {}, + tools.ToolNameFilesystemGlob: {}, + tools.ToolNameFilesystemEdit: {}, + tools.ToolNameFilesystemWriteFile: {}, } // microCompactMessages 对裁剪后的消息做只读投影式微压缩,仅清理旧工具结果内容。 diff --git a/internal/context/types.go b/internal/context/types.go index 6eb516dd..2e406861 100644 --- a/internal/context/types.go +++ b/internal/context/types.go @@ -15,6 +15,7 @@ type Builder interface { type BuildInput struct { Messages []provider.Message Metadata Metadata + Compact CompactOptions } // BuildResult is the provider-facing context produced for a single round. @@ -22,3 +23,8 @@ type BuildResult struct { SystemPrompt string Messages []provider.Message } + +// CompactOptions controls read-time compact behavior inside the context builder. +type CompactOptions struct { + DisableMicroCompact bool +} diff --git a/internal/runtime/runtime.go b/internal/runtime/runtime.go index ae3b380a..2e3b7c29 100644 --- a/internal/runtime/runtime.go +++ b/internal/runtime/runtime.go @@ -157,6 +157,9 @@ func (s *Service) Run(ctx context.Context, input UserInput) error { Provider: cfg.SelectedProvider, Model: cfg.CurrentModel, }, + Compact: agentcontext.CompactOptions{ + DisableMicroCompact: cfg.Context.Compact.MicroCompactDisabled, + }, }) if err != nil { return s.handleRunError(ctx, input.RunID, session.ID, err) diff --git a/internal/runtime/runtime_test.go b/internal/runtime/runtime_test.go index 7f3b5da5..3a4bec17 100644 --- a/internal/runtime/runtime_test.go +++ b/internal/runtime/runtime_test.go @@ -477,6 +477,9 @@ func TestServiceRunDelegatesToContextBuilder(t *testing.T) { if builder.lastInput.Metadata.Model == "" { t.Fatalf("expected model to be forwarded to builder metadata") } + if builder.lastInput.Compact.DisableMicroCompact { + t.Fatalf("expected micro compact to stay enabled by default") + } if len(builder.lastInput.Messages) != 1 || builder.lastInput.Messages[0].Content != "hello" { t.Fatalf("expected persisted session messages to be forwarded, got %+v", builder.lastInput.Messages) } @@ -491,6 +494,47 @@ func TestServiceRunDelegatesToContextBuilder(t *testing.T) { } } +func TestServiceRunCanDisableMicroCompactViaConfig(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + if err := manager.Update(context.Background(), func(cfg *config.Config) error { + cfg.Context.Compact.MicroCompactDisabled = true + return nil + }); err != nil { + t.Fatalf("update config: %v", err) + } + + store := newMemoryStore() + registry := tools.NewRegistry() + registry.Register(&stubTool{name: "filesystem_read_file", content: "default"}) + + builder := &stubContextBuilder{ + buildFn: func(ctx context.Context, input agentcontext.BuildInput) (agentcontext.BuildResult, error) { + return agentcontext.BuildResult{ + SystemPrompt: "delegated prompt", + Messages: append([]provider.Message(nil), input.Messages...), + }, nil + }, + } + + scripted := &scriptedProvider{ + responses: []provider.ChatResponse{{ + Message: provider.Message{Role: provider.RoleAssistant, Content: "done"}, + FinishReason: "stop", + }}, + } + + service := NewWithFactory(manager, registry, store, &scriptedProviderFactory{provider: scripted}, builder) + if err := service.Run(context.Background(), UserInput{RunID: "run-disable-micro-compact", Content: "hello"}); err != nil { + t.Fatalf("Run() error = %v", err) + } + + if !builder.lastInput.Compact.DisableMicroCompact { + t.Fatalf("expected config to disable micro compact in build input") + } +} + func TestServiceRunPersistsSessionProviderAndModel(t *testing.T) { t.Parallel() diff --git a/internal/tools/bash/tool.go b/internal/tools/bash/tool.go index 9facbb88..5c00c2c7 100644 --- a/internal/tools/bash/tool.go +++ b/internal/tools/bash/tool.go @@ -35,7 +35,7 @@ func New(root string, shell string, timeout time.Duration) *Tool { } func (t *Tool) Name() string { - return "bash" + return tools.ToolNameBash } func (t *Tool) Description() string { diff --git a/internal/tools/filesystem/helpers.go b/internal/tools/filesystem/helpers.go index d681ed6e..929d74e1 100644 --- a/internal/tools/filesystem/helpers.go +++ b/internal/tools/filesystem/helpers.go @@ -4,14 +4,16 @@ import ( "os" "path/filepath" "strings" + + "neo-code/internal/tools" ) const ( - readFileToolName = "filesystem_read_file" - writeFileToolName = "filesystem_write_file" - grepToolName = "filesystem_grep" - globToolName = "filesystem_glob" - editToolName = "filesystem_edit" + readFileToolName = tools.ToolNameFilesystemReadFile + writeFileToolName = tools.ToolNameFilesystemWriteFile + grepToolName = tools.ToolNameFilesystemGrep + globToolName = tools.ToolNameFilesystemGlob + editToolName = tools.ToolNameFilesystemEdit ) func effectiveRoot(defaultRoot string, workdir string) string { diff --git a/internal/tools/names.go b/internal/tools/names.go new file mode 100644 index 00000000..833430f7 --- /dev/null +++ b/internal/tools/names.go @@ -0,0 +1,12 @@ +package tools + +// Tool name constants are shared across tool implementations, context policies, and tests. +const ( + ToolNameBash = "bash" + ToolNameWebFetch = "webfetch" + ToolNameFilesystemReadFile = "filesystem_read_file" + ToolNameFilesystemWriteFile = "filesystem_write_file" + ToolNameFilesystemGrep = "filesystem_grep" + ToolNameFilesystemGlob = "filesystem_glob" + ToolNameFilesystemEdit = "filesystem_edit" +) diff --git a/internal/tools/webfetch/tool.go b/internal/tools/webfetch/tool.go index efa5474f..5a77c2b2 100644 --- a/internal/tools/webfetch/tool.go +++ b/internal/tools/webfetch/tool.go @@ -16,7 +16,7 @@ import ( ) const ( - toolName = "webfetch" + toolName = tools.ToolNameWebFetch htmlContentType = "text/html" xhtmlContentType = "application/xhtml+xml" reasonInvalidArguments = "invalid arguments" From 4627929a8ceeba7152eaab1d99f06633cccee490 Mon Sep 17 00:00:00 2001 From: phantom5099 <1011668688@qq.com> Date: Sun, 5 Apr 2026 15:48:15 +0800 Subject: [PATCH 07/55] =?UTF-8?q?fix(provider):=E4=BF=AE=E5=A4=8D=E2=80=9C?= =?UTF-8?q?=E6=98=BE=E5=BC=8F=E6=8A=A5=E9=94=99=E2=80=9D=E6=B5=81=E4=BA=8B?= =?UTF-8?q?=E4=BB=B6=E5=A4=84=E7=90=86=E5=AF=BC=E8=87=B4=E5=8D=A1=E6=AD=BB?= =?UTF-8?q?=E9=97=AE=E9=A2=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/provider/openai/openai_test.go | 24 ++-- internal/provider/types.go | 69 ++++++++-- internal/provider/types_test.go | 80 ++++++++--- internal/runtime/compact_generator.go | 34 ++--- internal/runtime/compact_generator_test.go | 65 +++++++++ internal/runtime/runtime.go | 150 ++++++++++++++++----- internal/runtime/runtime_test.go | 136 +++++++++++++++++++ 7 files changed, 461 insertions(+), 97 deletions(-) diff --git a/internal/provider/openai/openai_test.go b/internal/provider/openai/openai_test.go index ae0b6b9d..29e99654 100644 --- a/internal/provider/openai/openai_test.go +++ b/internal/provider/openai/openai_test.go @@ -637,36 +637,36 @@ func drainStreamEvents(events <-chan domain.StreamEvent) []domain.StreamEvent { func requireTextDeltaPayload(t *testing.T, event domain.StreamEvent) domain.TextDeltaPayload { t.Helper() - payload, ok := event.Payload.(domain.TextDeltaPayload) - if !ok { - t.Fatalf("expected TextDeltaPayload, got %T", event.Payload) + payload, err := event.TextDeltaValue() + if err != nil { + t.Fatalf("TextDeltaValue() error = %v", err) } return payload } func requireToolCallStartPayload(t *testing.T, event domain.StreamEvent) domain.ToolCallStartPayload { t.Helper() - payload, ok := event.Payload.(domain.ToolCallStartPayload) - if !ok { - t.Fatalf("expected ToolCallStartPayload, got %T", event.Payload) + payload, err := event.ToolCallStartValue() + if err != nil { + t.Fatalf("ToolCallStartValue() error = %v", err) } return payload } func requireToolCallDeltaPayload(t *testing.T, event domain.StreamEvent) domain.ToolCallDeltaPayload { t.Helper() - payload, ok := event.Payload.(domain.ToolCallDeltaPayload) - if !ok { - t.Fatalf("expected ToolCallDeltaPayload, got %T", event.Payload) + payload, err := event.ToolCallDeltaValue() + if err != nil { + t.Fatalf("ToolCallDeltaValue() error = %v", err) } return payload } func requireMessageDonePayload(t *testing.T, event domain.StreamEvent) domain.MessageDonePayload { t.Helper() - payload, ok := event.Payload.(domain.MessageDonePayload) - if !ok { - t.Fatalf("expected MessageDonePayload, got %T", event.Payload) + payload, err := event.MessageDoneValue() + if err != nil { + t.Fatalf("MessageDoneValue() error = %v", err) } return payload } diff --git a/internal/provider/types.go b/internal/provider/types.go index 278a3a93..1f689c12 100644 --- a/internal/provider/types.go +++ b/internal/provider/types.go @@ -1,5 +1,7 @@ package provider +import "fmt" + const ( // RoleSystem 标识系统消息。 RoleSystem = "system" @@ -65,8 +67,11 @@ const ( // StreamEvent 表示 provider 向 runtime 推送的流式事件。 type StreamEvent struct { - Type StreamEventType `json:"type"` - Payload interface{} `json:"payload,omitempty"` + Type StreamEventType `json:"type"` + TextDelta *TextDeltaPayload `json:"text_delta,omitempty"` + ToolCallStart *ToolCallStartPayload `json:"tool_call_start,omitempty"` + ToolCallDelta *ToolCallDeltaPayload `json:"tool_call_delta,omitempty"` + MessageDone *MessageDonePayload `json:"message_done,omitempty"` } // TextDeltaPayload 表示文本增量事件的载荷。 @@ -97,31 +102,75 @@ type MessageDonePayload struct { // NewTextDeltaStreamEvent 创建文本增量流事件。 func NewTextDeltaStreamEvent(text string) StreamEvent { return StreamEvent{ - Type: StreamEventTextDelta, - Payload: TextDeltaPayload{Text: text}, + Type: StreamEventTextDelta, + TextDelta: &TextDeltaPayload{Text: text}, } } // NewToolCallStartStreamEvent 创建工具调用开始流事件。 func NewToolCallStartStreamEvent(index int, id, name string) StreamEvent { return StreamEvent{ - Type: StreamEventToolCallStart, - Payload: ToolCallStartPayload{Index: index, ID: id, Name: name}, + Type: StreamEventToolCallStart, + ToolCallStart: &ToolCallStartPayload{Index: index, ID: id, Name: name}, } } // NewToolCallDeltaStreamEvent 创建工具调用参数增量流事件。 func NewToolCallDeltaStreamEvent(index int, id, argumentsDelta string) StreamEvent { return StreamEvent{ - Type: StreamEventToolCallDelta, - Payload: ToolCallDeltaPayload{Index: index, ID: id, ArgumentsDelta: argumentsDelta}, + Type: StreamEventToolCallDelta, + ToolCallDelta: &ToolCallDeltaPayload{Index: index, ID: id, ArgumentsDelta: argumentsDelta}, } } // NewMessageDoneStreamEvent 创建消息完成流事件。 func NewMessageDoneStreamEvent(finishReason string, usage *Usage) StreamEvent { return StreamEvent{ - Type: StreamEventMessageDone, - Payload: MessageDonePayload{FinishReason: finishReason, Usage: usage}, + Type: StreamEventMessageDone, + MessageDone: &MessageDonePayload{FinishReason: finishReason, Usage: usage}, + } +} + +// TextDeltaValue 返回 text_delta 事件的载荷,并在结构缺失时显式报错。 +func (e StreamEvent) TextDeltaValue() (TextDeltaPayload, error) { + if e.Type != StreamEventTextDelta { + return TextDeltaPayload{}, fmt.Errorf("provider: stream event type %q is not text_delta", e.Type) + } + if e.TextDelta == nil { + return TextDeltaPayload{}, fmt.Errorf("provider: text_delta event payload is nil") + } + return *e.TextDelta, nil +} + +// ToolCallStartValue 返回 tool_call_start 事件的载荷,并在结构缺失时显式报错。 +func (e StreamEvent) ToolCallStartValue() (ToolCallStartPayload, error) { + if e.Type != StreamEventToolCallStart { + return ToolCallStartPayload{}, fmt.Errorf("provider: stream event type %q is not tool_call_start", e.Type) + } + if e.ToolCallStart == nil { + return ToolCallStartPayload{}, fmt.Errorf("provider: tool_call_start event payload is nil") + } + return *e.ToolCallStart, nil +} + +// ToolCallDeltaValue 返回 tool_call_delta 事件的载荷,并在结构缺失时显式报错。 +func (e StreamEvent) ToolCallDeltaValue() (ToolCallDeltaPayload, error) { + if e.Type != StreamEventToolCallDelta { + return ToolCallDeltaPayload{}, fmt.Errorf("provider: stream event type %q is not tool_call_delta", e.Type) + } + if e.ToolCallDelta == nil { + return ToolCallDeltaPayload{}, fmt.Errorf("provider: tool_call_delta event payload is nil") + } + return *e.ToolCallDelta, nil +} + +// MessageDoneValue 返回 message_done 事件的载荷,并在结构缺失时显式报错。 +func (e StreamEvent) MessageDoneValue() (MessageDonePayload, error) { + if e.Type != StreamEventMessageDone { + return MessageDonePayload{}, fmt.Errorf("provider: stream event type %q is not message_done", e.Type) + } + if e.MessageDone == nil { + return MessageDonePayload{}, fmt.Errorf("provider: message_done event payload is nil") } + return *e.MessageDone, nil } diff --git a/internal/provider/types_test.go b/internal/provider/types_test.go index e5f1c5aa..5c822ee8 100644 --- a/internal/provider/types_test.go +++ b/internal/provider/types_test.go @@ -1,6 +1,9 @@ package provider -import "testing" +import ( + "encoding/json" + "testing" +) // --- Role 常量 --- @@ -56,9 +59,9 @@ func TestNewTextDeltaStreamEvent(t *testing.T) { t.Fatalf("expected type %q, got %q", StreamEventTextDelta, event.Type) } - payload, ok := event.Payload.(TextDeltaPayload) - if !ok { - t.Fatal("expected TextDeltaPayload type") + payload, err := event.TextDeltaValue() + if err != nil { + t.Fatalf("TextDeltaValue() error = %v", err) } if payload.Text != "hello" { t.Fatalf("expected text %q, got %q", "hello", payload.Text) @@ -75,9 +78,9 @@ func TestNewToolCallStartStreamEvent(t *testing.T) { t.Fatalf("expected type %q, got %q", StreamEventToolCallStart, event.Type) } - payload, ok := event.Payload.(ToolCallStartPayload) - if !ok { - t.Fatal("expected ToolCallStartPayload type") + payload, err := event.ToolCallStartValue() + if err != nil { + t.Fatalf("ToolCallStartValue() error = %v", err) } if payload.Index != 3 || payload.ID != "call_1" || payload.Name != "edit_file" { t.Fatalf("unexpected payload: %+v", payload) @@ -94,9 +97,9 @@ func TestNewToolCallDeltaStreamEvent(t *testing.T) { t.Fatalf("expected type %q, got %q", StreamEventToolCallDelta, event.Type) } - payload, ok := event.Payload.(ToolCallDeltaPayload) - if !ok { - t.Fatal("expected ToolCallDeltaPayload type") + payload, err := event.ToolCallDeltaValue() + if err != nil { + t.Fatalf("ToolCallDeltaValue() error = %v", err) } if payload.Index != 1 || payload.ID != "call_2" || payload.ArgumentsDelta != `{"path":"main.go"}` { t.Fatalf("unexpected payload: %+v", payload) @@ -116,9 +119,9 @@ func TestNewMessageDoneStreamEvent(t *testing.T) { t.Fatalf("expected type %q, got %q", StreamEventMessageDone, event.Type) } - payload, ok := event.Payload.(MessageDonePayload) - if !ok { - t.Fatal("expected MessageDonePayload type") + payload, err := event.MessageDoneValue() + if err != nil { + t.Fatalf("MessageDoneValue() error = %v", err) } if payload.FinishReason != "stop" { t.Fatalf("expected finish reason %q, got %q", "stop", payload.FinishReason) @@ -131,9 +134,9 @@ func TestNewMessageDoneStreamEvent(t *testing.T) { t.Run("nil usage", func(t *testing.T) { event := NewMessageDoneStreamEvent("tool_calls", nil) - payload, ok := event.Payload.(MessageDonePayload) - if !ok { - t.Fatal("expected MessageDonePayload type") + payload, err := event.MessageDoneValue() + if err != nil { + t.Fatalf("MessageDoneValue() error = %v", err) } if payload.FinishReason != "tool_calls" { t.Fatalf("expected finish reason %q, got %q", "tool_calls", payload.FinishReason) @@ -146,9 +149,9 @@ func TestNewMessageDoneStreamEvent(t *testing.T) { t.Run("empty finish reason", func(t *testing.T) { event := NewMessageDoneStreamEvent("", nil) - payload, ok := event.Payload.(MessageDonePayload) - if !ok { - t.Fatal("expected MessageDonePayload type") + payload, err := event.MessageDoneValue() + if err != nil { + t.Fatalf("MessageDoneValue() error = %v", err) } if payload.FinishReason != "" { t.Fatalf("expected empty finish reason, got %q", payload.FinishReason) @@ -220,10 +223,45 @@ func TestStreamEventStructFields(t *testing.T) { t.Parallel() event := StreamEvent{ - Type: StreamEventTextDelta, - Payload: TextDeltaPayload{Text: "hi"}, + Type: StreamEventTextDelta, + TextDelta: &TextDeltaPayload{Text: "hi"}, } if event.Type != StreamEventTextDelta { t.Fatalf("event type not as expected: %s", event.Type) } + if event.TextDelta == nil || event.TextDelta.Text != "hi" { + t.Fatalf("event text_delta not as expected: %+v", event.TextDelta) + } +} + +func TestStreamEventJSONRoundTrip(t *testing.T) { + t.Parallel() + + original := NewToolCallDeltaStreamEvent(2, "call-7", `{"path":"main.go"}`) + data, err := json.Marshal(original) + if err != nil { + t.Fatalf("json.Marshal() error = %v", err) + } + + var decoded StreamEvent + if err := json.Unmarshal(data, &decoded); err != nil { + t.Fatalf("json.Unmarshal() error = %v", err) + } + + payload, err := decoded.ToolCallDeltaValue() + if err != nil { + t.Fatalf("ToolCallDeltaValue() error = %v", err) + } + if payload.Index != 2 || payload.ID != "call-7" || payload.ArgumentsDelta != `{"path":"main.go"}` { + t.Fatalf("unexpected round-trip payload: %+v", payload) + } +} + +func TestStreamEventValueAccessorsRejectMissingPayload(t *testing.T) { + t.Parallel() + + event := StreamEvent{Type: StreamEventTextDelta} + if _, err := event.TextDeltaValue(); err == nil { + t.Fatal("expected TextDeltaValue() to reject missing payload") + } } diff --git a/internal/runtime/compact_generator.go b/internal/runtime/compact_generator.go index 49e48755..20be11a1 100644 --- a/internal/runtime/compact_generator.go +++ b/internal/runtime/compact_generator.go @@ -59,30 +59,24 @@ func (g *compactSummaryGenerator) Generate(ctx context.Context, input contextcom // 使用流式事件通道收集 compact 摘要响应。 streamEvents := make(chan provider.StreamEvent, 32) - streamDone := make(chan struct{}) + streamDone := make(chan error, 1) acc := newStreamAccumulator() go func() { - defer close(streamDone) + var streamErr error + defer func() { + streamDone <- streamErr + }() + for { select { case event, ok := <-streamEvents: if !ok { return } - switch event.Type { - case provider.StreamEventTextDelta: - if payload, ok := event.Payload.(provider.TextDeltaPayload); ok { - acc.accumulateTextDelta(payload.Text) - } - case provider.StreamEventToolCallStart: - if payload, ok := event.Payload.(provider.ToolCallStartPayload); ok { - acc.accumulateToolCallStart(payload.Index, payload.ID, payload.Name) - } - case provider.StreamEventToolCallDelta: - if payload, ok := event.Payload.(provider.ToolCallDeltaPayload); ok { - acc.accumulateToolCallDelta(payload.Index, payload.ID, payload.ArgumentsDelta) - } + if err := handleProviderStreamEvent(event, acc, nil, nil); err != nil && streamErr == nil { + // 记录首个协议错误后继续排空事件通道,避免 provider 在后续发送时阻塞。 + streamErr = err } case <-ctx.Done(): return @@ -99,13 +93,19 @@ func (g *compactSummaryGenerator) Generate(ctx context.Context, input contextcom }}, }, streamEvents) close(streamEvents) - <-streamDone + streamErr := <-streamDone if err != nil { return "", err } + if streamErr != nil { + return "", streamErr + } - message := acc.buildMessage() + message, err := acc.buildMessage() + if err != nil { + return "", err + } if len(message.ToolCalls) > 0 { return "", errors.New("runtime: compact summary response must not contain tool calls") } diff --git a/internal/runtime/compact_generator_test.go b/internal/runtime/compact_generator_test.go index e0c348e7..e82c625f 100644 --- a/internal/runtime/compact_generator_test.go +++ b/internal/runtime/compact_generator_test.go @@ -4,6 +4,7 @@ import ( "context" "strings" "testing" + "time" "neo-code/internal/config" contextcompact "neo-code/internal/context/compact" @@ -133,3 +134,67 @@ func TestCompactSummaryGeneratorRejectsToolCalls(t *testing.T) { t.Fatalf("expected tool call rejection, got %v", err) } } + +func TestCompactSummaryGeneratorRejectsMalformedStreamEvent(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + resolvedProvider, err := resolvedProviderForTests(manager.Get(), config.OpenAIName) + if err != nil { + t.Fatalf("resolve provider: %v", err) + } + + scripted := &scriptedProvider{ + streams: [][]provider.StreamEvent{ + { + {Type: provider.StreamEventTextDelta}, + }, + }, + } + generator := newCompactSummaryGenerator(&scriptedProviderFactory{provider: scripted}, resolvedProvider, "session-model") + + _, err = generator.Generate(context.Background(), contextcompact.SummaryInput{ + Mode: contextcompact.ModeManual, + Config: manager.Get().Context.Compact, + }) + if err == nil || !strings.Contains(err.Error(), "text_delta event payload is nil") { + t.Fatalf("expected malformed stream event rejection, got %v", err) + } +} + +func TestCompactSummaryGeneratorMalformedStreamEventDoesNotDeadlock(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + resolvedProvider, err := resolvedProviderForTests(manager.Get(), config.OpenAIName) + if err != nil { + t.Fatalf("resolve provider: %v", err) + } + + stream := []provider.StreamEvent{{Type: provider.StreamEventTextDelta}} + for i := 0; i < 40; i++ { + stream = append(stream, provider.NewTextDeltaStreamEvent("ignored")) + } + scripted := &scriptedProvider{ + streams: [][]provider.StreamEvent{stream}, + } + generator := newCompactSummaryGenerator(&scriptedProviderFactory{provider: scripted}, resolvedProvider, "session-model") + + errCh := make(chan error, 1) + go func() { + _, genErr := generator.Generate(context.Background(), contextcompact.SummaryInput{ + Mode: contextcompact.ModeManual, + Config: manager.Get().Context.Compact, + }) + errCh <- genErr + }() + + select { + case genErr := <-errCh: + if genErr == nil || !strings.Contains(genErr.Error(), "text_delta event payload is nil") { + t.Fatalf("expected malformed stream event rejection, got %v", genErr) + } + case <-time.After(2 * time.Second): + t.Fatal("expected compact generation to fail instead of deadlocking on malformed stream event") + } +} diff --git a/internal/runtime/runtime.go b/internal/runtime/runtime.go index d3b3e29f..03d6d556 100644 --- a/internal/runtime/runtime.go +++ b/internal/runtime/runtime.go @@ -49,28 +49,38 @@ func (a *streamAccumulator) accumulateTextDelta(text string) { a.content.WriteString(text) } +// ensureToolCall 返回指定索引的工具调用条目,不存在时会先创建占位对象。 +func (a *streamAccumulator) ensureToolCall(index int) *provider.ToolCall { + call, exists := a.toolCalls[index] + if !exists { + call = &provider.ToolCall{} + a.toolCalls[index] = call + } + return call +} + // accumulateToolCallStart 记录新发现的工具调用(首次出现时创建条目)。 func (a *streamAccumulator) accumulateToolCallStart(index int, id, name string) { - if _, exists := a.toolCalls[index]; !exists { - a.toolCalls[index] = &provider.ToolCall{ID: id, Name: name} + call := a.ensureToolCall(index) + if strings.TrimSpace(id) != "" { + call.ID = id + } + if strings.TrimSpace(name) != "" { + call.Name = name } } // accumulateToolCallDelta 累积工具调用参数增量。 func (a *streamAccumulator) accumulateToolCallDelta(index int, id, argumentsDelta string) { - call, exists := a.toolCalls[index] - if !exists { - call = &provider.ToolCall{ID: id} - a.toolCalls[index] = call - } - if name := call.Name; strings.TrimSpace(name) == "" && call.ID != "" { - // 首次出现 delta 时可能还未收到 start 事件,仅记录 ID + call := a.ensureToolCall(index) + if strings.TrimSpace(id) != "" { + call.ID = id } call.Arguments += argumentsDelta } -// buildMessage 从累积状态构建最终的 assistant Message 对象。 -func (a *streamAccumulator) buildMessage() provider.Message { +// buildMessage 从累积状态构建最终的 assistant Message 对象,并校验工具调用元数据是否完整。 +func (a *streamAccumulator) buildMessage() (provider.Message, error) { ordered := make([]int, 0, len(a.toolCalls)) for index := range a.toolCalls { ordered = append(ordered, index) @@ -86,9 +96,15 @@ func (a *streamAccumulator) buildMessage() provider.Message { if call == nil { continue } + if strings.TrimSpace(call.ID) == "" { + return provider.Message{}, fmt.Errorf("runtime: provider emitted tool call %d without id", index) + } + if strings.TrimSpace(call.Name) == "" { + return provider.Message{}, fmt.Errorf("runtime: provider emitted tool call %d without name", index) + } message.ToolCalls = append(message.ToolCalls, *call) } - return message + return message, nil } var runtimeSessionWorkdirs = struct { @@ -248,7 +264,10 @@ func (s *Service) Run(ctx context.Context, input UserInput) error { session.Provider = cfg.SelectedProvider session.Model = cfg.CurrentModel - assistant := acc.buildMessage() + assistant, err := acc.buildMessage() + if err != nil { + return s.handleRunError(ctx, input.RunID, session.ID, err) + } if strings.TrimSpace(assistant.Role) == "" { assistant.Role = provider.RoleAssistant } @@ -479,37 +498,88 @@ func (s *Service) emit(ctx context.Context, kind EventType, runID string, sessio } } +// handleProviderStreamEvent 解析并应用单条 provider 流式事件,缺失载荷或未知类型时返回错误。 +func handleProviderStreamEvent( + event provider.StreamEvent, + acc *streamAccumulator, + onTextDelta func(string), + onToolCallStart func(provider.ToolCallStartPayload), +) error { + switch event.Type { + case provider.StreamEventTextDelta: + payload, err := event.TextDeltaValue() + if err != nil { + return err + } + if onTextDelta != nil { + onTextDelta(payload.Text) + } + if acc != nil { + acc.accumulateTextDelta(payload.Text) + } + case provider.StreamEventToolCallStart: + payload, err := event.ToolCallStartValue() + if err != nil { + return err + } + if onToolCallStart != nil { + onToolCallStart(payload) + } + if acc != nil { + acc.accumulateToolCallStart(payload.Index, payload.ID, payload.Name) + } + case provider.StreamEventToolCallDelta: + payload, err := event.ToolCallDeltaValue() + if err != nil { + return err + } + if acc != nil { + acc.accumulateToolCallDelta(payload.Index, payload.ID, payload.ArgumentsDelta) + } + case provider.StreamEventMessageDone: + if _, err := event.MessageDoneValue(); err != nil { + return err + } + default: + return fmt.Errorf("runtime: unsupported provider stream event type %q", event.Type) + } + return nil +} + // forwardProviderEvents 将 provider 流式事件转发为 runtime 事件,同时向 accumulator 累积消息状态。 // 使用 select 同时监听输入通道和 context 取消信号,确保 goroutine 不会因通道阻塞而泄漏。 -func (s *Service) forwardProviderEvents(ctx context.Context, runID string, sessionID string, input <-chan provider.StreamEvent, done chan<- struct{}, acc *streamAccumulator) { - defer close(done) +func (s *Service) forwardProviderEvents( + ctx context.Context, + runID string, + sessionID string, + input <-chan provider.StreamEvent, + done chan<- error, + acc *streamAccumulator, +) { + var forwardErr error + defer func() { + done <- forwardErr + }() + for { select { case event, ok := <-input: if !ok { return } - switch event.Type { - case provider.StreamEventTextDelta: - if payload, ok := event.Payload.(provider.TextDeltaPayload); ok { - s.emit(ctx, EventAgentChunk, runID, sessionID, payload.Text) - if acc != nil { - acc.accumulateTextDelta(payload.Text) - } - } - case provider.StreamEventToolCallStart: - if payload, ok := event.Payload.(provider.ToolCallStartPayload); ok { + err := handleProviderStreamEvent( + event, + acc, + func(text string) { + s.emit(ctx, EventAgentChunk, runID, sessionID, text) + }, + func(payload provider.ToolCallStartPayload) { s.emit(ctx, EventToolCallThinking, runID, sessionID, payload.Name) - if acc != nil { - acc.accumulateToolCallStart(payload.Index, payload.ID, payload.Name) - } - } - case provider.StreamEventToolCallDelta: - if payload, ok := event.Payload.(provider.ToolCallDeltaPayload); ok { - if acc != nil { - acc.accumulateToolCallDelta(payload.Index, payload.ID, payload.ArgumentsDelta) - } - } + }, + ) + if err != nil && forwardErr == nil { + // 记录首个协议错误后继续排空事件通道,避免 provider 在后续发送时阻塞。 + forwardErr = err } case <-ctx.Done(): return @@ -606,12 +676,18 @@ func (s *Service) callProviderWithRetry( } streamEvents := make(chan provider.StreamEvent, 32) - streamDone := make(chan struct{}) + streamDone := make(chan error, 1) go s.forwardProviderEvents(ctx, runID, sessionID, streamEvents, streamDone, acc) err = modelProvider.Chat(ctx, req, streamEvents) close(streamEvents) - <-streamDone + forwardErr := <-streamDone + if forwardErr != nil { + if err != nil { + return nil, fmt.Errorf("runtime: provider stream handling failed after provider error: %v: %w", err, forwardErr) + } + return nil, forwardErr + } if err == nil { return acc, nil diff --git a/internal/runtime/runtime_test.go b/internal/runtime/runtime_test.go index 101a91ad..aac914a8 100644 --- a/internal/runtime/runtime_test.go +++ b/internal/runtime/runtime_test.go @@ -383,6 +383,142 @@ func TestServiceRun(t *testing.T) { } } +func TestServiceRunMergesLateToolCallMetadata(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + store := newMemoryStore() + tool := &stubTool{name: "filesystem_edit", content: "tool output"} + registry := tools.NewRegistry() + registry.Register(tool) + + scripted := &scriptedProvider{ + streams: [][]provider.StreamEvent{ + { + provider.NewToolCallDeltaStreamEvent(0, "", `{"path":"main.go"`), + provider.NewToolCallStartStreamEvent(0, "call-late", "filesystem_edit"), + provider.NewToolCallDeltaStreamEvent(0, "call-late", `}`), + }, + {provider.NewTextDeltaStreamEvent("done")}, + }, + } + + service := NewWithFactory(manager, registry, store, &scriptedProviderFactory{provider: scripted}, nil) + if err := service.Run(context.Background(), UserInput{RunID: "run-late-tool-metadata", Content: "edit"}); err != nil { + t.Fatalf("Run() error = %v", err) + } + + if tool.callCount != 1 { + t.Fatalf("expected tool to execute once, got %d", tool.callCount) + } + if tool.lastInput.ID != "call-late" { + t.Fatalf("expected merged tool call id %q, got %q", "call-late", tool.lastInput.ID) + } + if tool.lastInput.Name != "filesystem_edit" { + t.Fatalf("expected merged tool name %q, got %q", "filesystem_edit", tool.lastInput.Name) + } + if got := string(tool.lastInput.Arguments); got != `{"path":"main.go"}` { + t.Fatalf("expected merged tool arguments %q, got %q", `{"path":"main.go"}`, got) + } + + session := onlySession(t, store) + if len(session.Messages) < 3 { + t.Fatalf("expected assistant/tool follow-up messages, got %+v", session.Messages) + } + if len(session.Messages[1].ToolCalls) != 1 { + t.Fatalf("expected persisted assistant tool call, got %+v", session.Messages[1]) + } + if session.Messages[1].ToolCalls[0].ID != "call-late" || session.Messages[1].ToolCalls[0].Name != "filesystem_edit" { + t.Fatalf("expected merged assistant tool call metadata, got %+v", session.Messages[1].ToolCalls[0]) + } + if session.Messages[2].ToolCallID != "call-late" { + t.Fatalf("expected tool result to reference merged tool call id, got %+v", session.Messages[2]) + } +} + +func TestServiceRunRejectsToolCallWithoutID(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + store := newMemoryStore() + tool := &stubTool{name: "filesystem_edit", content: "tool output"} + registry := tools.NewRegistry() + registry.Register(tool) + + scripted := &scriptedProvider{ + streams: [][]provider.StreamEvent{ + { + provider.NewToolCallStartStreamEvent(0, "", "filesystem_edit"), + provider.NewToolCallDeltaStreamEvent(0, "", `{}`), + }, + }, + } + + service := NewWithFactory(manager, registry, store, &scriptedProviderFactory{provider: scripted}, nil) + err := service.Run(context.Background(), UserInput{RunID: "run-missing-tool-id", Content: "edit"}) + if err == nil || !containsError(err, "without id") { + t.Fatalf("expected missing tool id error, got %v", err) + } + if tool.callCount != 0 { + t.Fatalf("expected tool execution to be blocked, got %d calls", tool.callCount) + } +} + +func TestServiceRunRejectsMalformedProviderStreamEvent(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + store := newMemoryStore() + registry := tools.NewRegistry() + registry.Register(&stubTool{name: "filesystem_read_file", content: "default"}) + + scripted := &scriptedProvider{ + streams: [][]provider.StreamEvent{ + { + {Type: provider.StreamEventTextDelta}, + }, + }, + } + + service := NewWithFactory(manager, registry, store, &scriptedProviderFactory{provider: scripted}, nil) + err := service.Run(context.Background(), UserInput{RunID: "run-malformed-stream-event", Content: "hello"}) + if err == nil || !containsError(err, "text_delta event payload is nil") { + t.Fatalf("expected malformed stream event error, got %v", err) + } +} + +func TestServiceRunMalformedProviderStreamEventDoesNotDeadlock(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + store := newMemoryStore() + registry := tools.NewRegistry() + registry.Register(&stubTool{name: "filesystem_read_file", content: "default"}) + + stream := []provider.StreamEvent{{Type: provider.StreamEventTextDelta}} + for i := 0; i < 40; i++ { + stream = append(stream, provider.NewTextDeltaStreamEvent("ignored")) + } + scripted := &scriptedProvider{ + streams: [][]provider.StreamEvent{stream}, + } + + service := NewWithFactory(manager, registry, store, &scriptedProviderFactory{provider: scripted}, nil) + errCh := make(chan error, 1) + go func() { + errCh <- service.Run(context.Background(), UserInput{RunID: "run-malformed-stream-no-deadlock", Content: "hello"}) + }() + + select { + case err := <-errCh: + if err == nil || !containsError(err, "text_delta event payload is nil") { + t.Fatalf("expected malformed stream event error, got %v", err) + } + case <-time.After(2 * time.Second): + t.Fatal("expected run to fail instead of deadlocking on malformed stream event") + } +} + type stubCompactRunner struct { runFn func(ctx context.Context, input contextcompact.Input) (contextcompact.Result, error) calls []contextcompact.Input From b3b0f88b02fb2a4883e5bc504734225a9e593e7d Mon Sep 17 00:00:00 2001 From: Yumiue <229866007@qq.com> Date: Sun, 5 Apr 2026 14:24:42 +0800 Subject: [PATCH 08/55] =?UTF-8?q?feat(context):=20=E5=A2=9E=E5=8A=A0?= =?UTF-8?q?=E6=97=A7=E5=B7=A5=E5=85=B7=E7=BB=93=E6=9E=9C=E7=9A=84=E8=AF=BB?= =?UTF-8?q?=E6=97=B6=20Micro=20Compact?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/context/builder.go | 2 +- internal/context/builder_test.go | 54 +++++++ internal/context/microcompact.go | 128 ++++++++++++++++ internal/context/microcompact_test.go | 203 ++++++++++++++++++++++++++ 4 files changed, 386 insertions(+), 1 deletion(-) create mode 100644 internal/context/microcompact.go create mode 100644 internal/context/microcompact_test.go diff --git a/internal/context/builder.go b/internal/context/builder.go index ff1487ea..f1ca2c79 100644 --- a/internal/context/builder.go +++ b/internal/context/builder.go @@ -43,6 +43,6 @@ func (b *DefaultBuilder) Build(ctx context.Context, input BuildInput) (BuildResu return BuildResult{ SystemPrompt: composeSystemPrompt(sections...), - Messages: trimPolicy.Trim(input.Messages), + Messages: microCompactMessages(trimPolicy.Trim(input.Messages)), }, nil } diff --git a/internal/context/builder_test.go b/internal/context/builder_test.go index 42c002ba..9f5098a0 100644 --- a/internal/context/builder_test.go +++ b/internal/context/builder_test.go @@ -150,6 +150,60 @@ func TestDefaultBuilderBuildReturnsPromptSourceError(t *testing.T) { } } +func TestDefaultBuilderBuildAppliesMicroCompactAfterTrim(t *testing.T) { + t.Parallel() + + builder := &DefaultBuilder{ + promptSources: []promptSectionSource{ + stubPromptSectionSource{sections: []promptSection{{title: "Stub", content: "body"}}}, + }, + } + + messages := []provider.Message{ + {Role: provider.RoleUser, Content: "older user"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-1", Content: "old read result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-2", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-3", Name: "webfetch", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, + {Role: provider.RoleUser, Content: "latest explicit instruction"}, + {Role: provider.RoleAssistant, Content: "current reply"}, + } + + got, err := builder.Build(stdcontext.Background(), BuildInput{Messages: messages}) + if err != nil { + t.Fatalf("Build() error = %v", err) + } + if len(got.Messages) != len(messages) { + t.Fatalf("expected builder output to keep message count, got %d want %d", len(got.Messages), len(messages)) + } + if got.Messages[2].Content != microCompactClearedMessage { + t.Fatalf("expected builder output to clear older tool result, got %q", got.Messages[2].Content) + } + if got.Messages[4].Content != "recent bash result" { + t.Fatalf("expected recent tool result to stay visible, got %q", got.Messages[4].Content) + } + if got.Messages[6].Content != "latest webfetch result" { + t.Fatalf("expected latest tool result to stay visible, got %q", got.Messages[6].Content) + } +} + func TestTrimMessagesPreservesToolPairs(t *testing.T) { t.Parallel() diff --git a/internal/context/microcompact.go b/internal/context/microcompact.go new file mode 100644 index 00000000..1eec70c0 --- /dev/null +++ b/internal/context/microcompact.go @@ -0,0 +1,128 @@ +package context + +import ( + "strings" + + "neo-code/internal/context/internalcompact" + "neo-code/internal/provider" +) + +const ( + // microCompactClearedMessage 是旧工具结果被读时微压缩后的占位符文本。 + microCompactClearedMessage = "[Old tool result content cleared]" + // microCompactRetainedToolSpans 定义默认保留原始内容的最近可压缩工具块数量。 + microCompactRetainedToolSpans = 2 +) + +var microCompactableTools = map[string]struct{}{ + "bash": {}, + "webfetch": {}, + "filesystem_read_file": {}, + "filesystem_grep": {}, + "filesystem_glob": {}, + "filesystem_edit": {}, + "filesystem_write_file": {}, +} + +// microCompactMessages 对裁剪后的消息做只读投影式微压缩,仅清理旧工具结果内容。 +func microCompactMessages(messages []provider.Message) []provider.Message { + cloned := cloneContextMessages(messages) + if len(cloned) == 0 { + return cloned + } + + spans := internalcompact.BuildMessageSpans(cloned) + protectedStart, hasProtectedTail := internalcompact.ProtectedTailStart(spans) + retainedCompactableSpans := 0 + + for spanIndex := len(spans) - 1; spanIndex >= 0; spanIndex-- { + span := spans[spanIndex] + if hasProtectedTail && span.Start >= protectedStart { + continue + } + if !isToolCallSpan(cloned, span) { + continue + } + + compactableIDs := compactableToolCallIDs(cloned[span.Start].ToolCalls) + if len(compactableIDs) == 0 { + continue + } + if retainedCompactableSpans < microCompactRetainedToolSpans { + retainedCompactableSpans++ + continue + } + + for messageIndex := span.Start + 1; messageIndex < span.End; messageIndex++ { + if shouldClearToolMessage(cloned[messageIndex], compactableIDs) { + cloned[messageIndex].Content = microCompactClearedMessage + } + } + } + + return cloned +} + +// cloneContextMessages 深拷贝消息切片,避免读时投影污染 runtime 持有的原始会话消息。 +func cloneContextMessages(messages []provider.Message) []provider.Message { + if len(messages) == 0 { + return nil + } + + cloned := make([]provider.Message, 0, len(messages)) + for _, message := range messages { + next := message + next.ToolCalls = append([]provider.ToolCall(nil), message.ToolCalls...) + cloned = append(cloned, next) + } + return cloned +} + +// isToolCallSpan 判断当前 span 是否是由 assistant tool call 起始的原子工具块。 +func isToolCallSpan(messages []provider.Message, span internalcompact.MessageSpan) bool { + if span.Start < 0 || span.Start >= len(messages) { + return false + } + message := messages[span.Start] + return message.Role == provider.RoleAssistant && len(message.ToolCalls) > 0 +} + +// compactableToolCallIDs 返回 assistant tool call 中可参与微压缩的调用 ID 集合。 +func compactableToolCallIDs(calls []provider.ToolCall) map[string]struct{} { + if len(calls) == 0 { + return nil + } + + ids := make(map[string]struct{}, len(calls)) + for _, call := range calls { + toolName := strings.TrimSpace(call.Name) + if _, ok := microCompactableTools[toolName]; !ok { + continue + } + callID := strings.TrimSpace(call.ID) + if callID == "" { + continue + } + ids[callID] = struct{}{} + } + if len(ids) == 0 { + return nil + } + return ids +} + +// shouldClearToolMessage 判断一条 tool 消息是否满足旧结果清理条件。 +func shouldClearToolMessage(message provider.Message, compactableIDs map[string]struct{}) bool { + if message.Role != provider.RoleTool || message.IsError { + return false + } + if compactableIDs == nil { + return false + } + if _, ok := compactableIDs[strings.TrimSpace(message.ToolCallID)]; !ok { + return false + } + + content := strings.TrimSpace(message.Content) + return content != "" && content != microCompactClearedMessage +} diff --git a/internal/context/microcompact_test.go b/internal/context/microcompact_test.go new file mode 100644 index 00000000..b0cce67e --- /dev/null +++ b/internal/context/microcompact_test.go @@ -0,0 +1,203 @@ +package context + +import ( + "testing" + + "neo-code/internal/provider" +) + +func TestMicroCompactMessagesClearsOlderCompactableToolResults(t *testing.T) { + t.Parallel() + + messages := []provider.Message{ + {Role: provider.RoleUser, Content: "older user"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-1", Content: "old read result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-2", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-3", Name: "webfetch", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, + {Role: provider.RoleUser, Content: "latest explicit instruction"}, + {Role: provider.RoleAssistant, Content: "current working reply"}, + } + + got := microCompactMessages(messages) + if len(got) != len(messages) { + t.Fatalf("expected message count to stay unchanged, got %d want %d", len(got), len(messages)) + } + if got[2].Content != microCompactClearedMessage { + t.Fatalf("expected oldest compactable tool result to be cleared, got %q", got[2].Content) + } + if got[4].Content != "recent bash result" { + t.Fatalf("expected recent compactable tool result to be retained, got %q", got[4].Content) + } + if got[6].Content != "latest webfetch result" { + t.Fatalf("expected latest compactable tool result to be retained, got %q", got[6].Content) + } + if messages[2].Content != "old read result" { + t.Fatalf("expected original slice to remain unchanged, got %q", messages[2].Content) + } +} + +func TestMicroCompactMessagesKeepsProtectedTailUntouched(t *testing.T) { + t.Parallel() + + messages := []provider.Message{ + {Role: provider.RoleUser, Content: "older user"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-0", Name: "filesystem_grep", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-0", Content: "old grep result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-1", Content: "recent read result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-2", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + {Role: provider.RoleUser, Content: "latest explicit instruction"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-3", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-3", Content: "tail bash result"}, + } + + got := microCompactMessages(messages) + if got[2].Content != microCompactClearedMessage { + t.Fatalf("expected old tool result before protected tail to be cleared, got %q", got[2].Content) + } + if got[4].Content != "recent read result" { + t.Fatalf("expected recent tool result before protected tail to remain, got %q", got[4].Content) + } + if got[6].Content != "recent bash result" { + t.Fatalf("expected second recent tool result before protected tail to remain, got %q", got[6].Content) + } + if got[9].Content != "tail bash result" { + t.Fatalf("expected protected tail tool result to remain, got %q", got[9].Content) + } +} + +func TestMicroCompactMessagesSkipsNonCompactableErrorsAndOrphans(t *testing.T) { + t.Parallel() + + messages := []provider.Message{ + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-1", Name: "custom_tool", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-1", Content: "custom result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-2", Name: "filesystem_edit", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-2", Content: "edit failed", IsError: true}, + {Role: provider.RoleTool, ToolCallID: "orphan", Content: "orphan result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-3", Name: "filesystem_write_file", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-3", Content: microCompactClearedMessage}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-4", Name: "filesystem_grep", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-4", Content: ""}, + } + + got := microCompactMessages(messages) + if got[1].Content != "custom result" { + t.Fatalf("expected non-compactable tool result to remain, got %q", got[1].Content) + } + if got[3].Content != "edit failed" { + t.Fatalf("expected error tool result to remain, got %q", got[3].Content) + } + if got[4].Content != "orphan result" { + t.Fatalf("expected orphan tool result to remain, got %q", got[4].Content) + } + if got[6].Content != microCompactClearedMessage { + t.Fatalf("expected already cleared content to remain unchanged, got %q", got[6].Content) + } + if got[8].Content != "" { + t.Fatalf("expected empty tool result to remain empty, got %q", got[8].Content) + } +} + +func TestMicroCompactMessagesClearsOnlyCompactableResultsInMixedToolSpan(t *testing.T) { + t.Parallel() + + messages := []provider.Message{ + {Role: provider.RoleUser, Content: "older user"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, + {ID: "call-2", Name: "custom_tool", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-1", Content: "read result"}, + {Role: provider.RoleTool, ToolCallID: "call-2", Content: "custom result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-3", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-3", Content: "recent bash result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-4", Name: "webfetch", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-4", Content: "latest webfetch result"}, + {Role: provider.RoleUser, Content: "latest explicit instruction"}, + {Role: provider.RoleAssistant, Content: "current reply"}, + } + + got := microCompactMessages(messages) + if got[2].Content != microCompactClearedMessage { + t.Fatalf("expected compactable tool result to be cleared, got %q", got[2].Content) + } + if got[3].Content != "custom result" { + t.Fatalf("expected non-compactable tool result in mixed span to remain, got %q", got[3].Content) + } + if len(got[1].ToolCalls) != 2 { + t.Fatalf("expected assistant tool call metadata to remain intact, got %+v", got[1].ToolCalls) + } +} From dea11ddd04ce0b76f05d599490484efcfb3276b2 Mon Sep 17 00:00:00 2001 From: Yumiue <229866007@qq.com> Date: Sun, 5 Apr 2026 14:51:37 +0800 Subject: [PATCH 09/55] fix(context): avoid empty spans consuming micro compact budget --- internal/context/microcompact.go | 13 ++++++ internal/context/microcompact_test.go | 62 +++++++++++++++++++++++++++ 2 files changed, 75 insertions(+) diff --git a/internal/context/microcompact.go b/internal/context/microcompact.go index 1eec70c0..10f5baba 100644 --- a/internal/context/microcompact.go +++ b/internal/context/microcompact.go @@ -48,6 +48,9 @@ func microCompactMessages(messages []provider.Message) []provider.Message { if len(compactableIDs) == 0 { continue } + if !hasCompactableToolContent(cloned, span, compactableIDs) { + continue + } if retainedCompactableSpans < microCompactRetainedToolSpans { retainedCompactableSpans++ continue @@ -111,6 +114,16 @@ func compactableToolCallIDs(calls []provider.ToolCall) map[string]struct{} { return ids } +// hasCompactableToolContent 判断工具块中是否存在会影响保留预算的有效工具结果内容。 +func hasCompactableToolContent(messages []provider.Message, span internalcompact.MessageSpan, compactableIDs map[string]struct{}) bool { + for messageIndex := span.Start + 1; messageIndex < span.End; messageIndex++ { + if shouldClearToolMessage(messages[messageIndex], compactableIDs) { + return true + } + } + return false +} + // shouldClearToolMessage 判断一条 tool 消息是否满足旧结果清理条件。 func shouldClearToolMessage(message provider.Message, compactableIDs map[string]struct{}) bool { if message.Role != provider.RoleTool || message.IsError { diff --git a/internal/context/microcompact_test.go b/internal/context/microcompact_test.go index b0cce67e..77b7b078 100644 --- a/internal/context/microcompact_test.go +++ b/internal/context/microcompact_test.go @@ -201,3 +201,65 @@ func TestMicroCompactMessagesClearsOnlyCompactableResultsInMixedToolSpan(t *test t.Fatalf("expected assistant tool call metadata to remain intact, got %+v", got[1].ToolCalls) } } + +func TestMicroCompactMessagesSkipsEmptyRecentSpansWhenCountingRetainedBudget(t *testing.T) { + t.Parallel() + + messages := []provider.Message{ + {Role: provider.RoleUser, Content: "older user"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-1", Content: "older read result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-2", Name: "filesystem_grep", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-2", Content: "middle grep result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-3", Name: "filesystem_edit", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-3", Content: "near edit result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-4", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-4", Content: "", IsError: true}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-5", Name: "webfetch", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-5", Content: ""}, + {Role: provider.RoleUser, Content: "latest explicit instruction"}, + {Role: provider.RoleAssistant, Content: "current reply"}, + } + + got := microCompactMessages(messages) + if got[2].Content != microCompactClearedMessage { + t.Fatalf("expected oldest valid tool result to be cleared, got %q", got[2].Content) + } + if got[4].Content != "middle grep result" { + t.Fatalf("expected middle valid tool result to remain, got %q", got[4].Content) + } + if got[6].Content != "near edit result" { + t.Fatalf("expected nearer valid tool result to remain, got %q", got[6].Content) + } + if got[8].Content != "" { + t.Fatalf("expected error/empty tool result to remain unchanged, got %q", got[8].Content) + } + if got[10].Content != "" { + t.Fatalf("expected empty recent tool result to remain unchanged, got %q", got[10].Content) + } +} From d1aea902101531278259689a8216211116eb2dce Mon Sep 17 00:00:00 2001 From: Yumiue <229866007@qq.com> Date: Sun, 5 Apr 2026 15:12:20 +0800 Subject: [PATCH 10/55] =?UTF-8?q?fix(context):=20=E5=A2=9E=E5=8A=A0=20micr?= =?UTF-8?q?o=20compact=20=E5=9B=9E=E9=80=80=E5=BC=80=E5=85=B3=E5=B9=B6?= =?UTF-8?q?=E6=94=B6=E6=95=9B=E5=B7=A5=E5=85=B7=E5=90=8D=E5=B8=B8=E9=87=8F?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/guides/configuration.md | 2 ++ internal/config/config_test.go | 10 ++++++ internal/config/loader.go | 3 ++ internal/config/model.go | 1 + internal/context/builder.go | 16 +++++++-- internal/context/builder_test.go | 54 ++++++++++++++++++++++++++++ internal/context/microcompact.go | 15 ++++---- internal/context/types.go | 6 ++++ internal/runtime/runtime.go | 3 ++ internal/runtime/runtime_test.go | 44 +++++++++++++++++++++++ internal/tools/bash/tool.go | 2 +- internal/tools/filesystem/helpers.go | 12 ++++--- internal/tools/names.go | 12 +++++++ internal/tools/webfetch/tool.go | 2 +- 14 files changed, 166 insertions(+), 16 deletions(-) create mode 100644 internal/tools/names.go diff --git a/docs/guides/configuration.md b/docs/guides/configuration.md index fed73849..155e29d9 100644 --- a/docs/guides/configuration.md +++ b/docs/guides/configuration.md @@ -315,6 +315,7 @@ context: manual_strategy: keep_recent manual_keep_recent_messages: 10 max_summary_chars: 1200 + micro_compact_disabled: false ``` ### 字段说明 @@ -324,5 +325,6 @@ context: | `context.compact.manual_strategy` | string | `keep_recent` | 手动 `/compact` 策略,可选 `keep_recent` / `full_replace` | | `context.compact.manual_keep_recent_messages` | int | `10` | `keep_recent` 模式下保留最近 N 条消息;会按 tool call 与 tool result 的原子块整体保留 | | `context.compact.max_summary_chars` | int | `1200` | compact summary 最大字符数 | +| `context.compact.micro_compact_disabled` | bool | `false` | 是否关闭默认启用的读时 micro compact;设为 `true` 可快速回退到仅 trim、不做旧工具结果清理 | 更多行为说明见 [context-compact.md](../context-compact.md)。 diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 4516c5e0..a2504ac2 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -876,10 +876,14 @@ func TestCompactConfigDefaultsAndRoundTrip(t *testing.T) { if compactCfg.MaxSummaryChars != DefaultCompactMaxSummaryChars { t.Fatalf("expected max_summary_chars=%d, got %d", DefaultCompactMaxSummaryChars, compactCfg.MaxSummaryChars) } + if compactCfg.MicroCompactDisabled { + t.Fatalf("expected micro compact to be enabled by default") + } cfg.Context.Compact.ManualStrategy = CompactManualStrategyFullReplace cfg.Context.Compact.ManualKeepRecentMessages = 2 cfg.Context.Compact.MaxSummaryChars = 900 + cfg.Context.Compact.MicroCompactDisabled = true if err := loader.Save(context.Background(), cfg); err != nil { t.Fatalf("Save() error = %v", err) } @@ -894,6 +898,9 @@ func TestCompactConfigDefaultsAndRoundTrip(t *testing.T) { if strings.Contains(text, "manual_keep_recent_spans:") { t.Fatalf("expected persisted config to drop legacy manual_keep_recent_spans key, got:\n%s", text) } + if !strings.Contains(text, "micro_compact_disabled: true") { + t.Fatalf("expected persisted config to include micro_compact_disabled, got:\n%s", text) + } reloaded, err := loader.Load(context.Background()) if err != nil { @@ -908,6 +915,9 @@ func TestCompactConfigDefaultsAndRoundTrip(t *testing.T) { if reloaded.Context.Compact.MaxSummaryChars != 900 { t.Fatalf("expected max_summary_chars=900, got %d", reloaded.Context.Compact.MaxSummaryChars) } + if !reloaded.Context.Compact.MicroCompactDisabled { + t.Fatalf("expected micro_compact_disabled to persist") + } } func TestCompactConfigValidateFailures(t *testing.T) { diff --git a/internal/config/loader.go b/internal/config/loader.go index 6c8ef449..ea3a6c3e 100644 --- a/internal/config/loader.go +++ b/internal/config/loader.go @@ -41,6 +41,7 @@ type persistedCompactConfig struct { ManualStrategy string `yaml:"manual_strategy,omitempty"` ManualKeepRecentMessages int `yaml:"manual_keep_recent_messages,omitempty"` MaxSummaryChars int `yaml:"max_summary_chars,omitempty"` + MicroCompactDisabled bool `yaml:"micro_compact_disabled,omitempty"` } func NewLoader(baseDir string, defaults *Config) *Loader { @@ -217,6 +218,7 @@ func newPersistedContextConfig(cfg ContextConfig) persistedContextConfig { ManualStrategy: cfg.Compact.ManualStrategy, ManualKeepRecentMessages: cfg.Compact.ManualKeepRecentMessages, MaxSummaryChars: cfg.Compact.MaxSummaryChars, + MicroCompactDisabled: cfg.Compact.MicroCompactDisabled, }, } } @@ -228,6 +230,7 @@ func fromPersistedContextConfig(file persistedContextConfig, defaults ContextCon ManualStrategy: strings.TrimSpace(file.Compact.ManualStrategy), ManualKeepRecentMessages: file.Compact.ManualKeepRecentMessages, MaxSummaryChars: file.Compact.MaxSummaryChars, + MicroCompactDisabled: file.Compact.MicroCompactDisabled, }, } out.Compact.ApplyDefaults(defaults.Compact) diff --git a/internal/config/model.go b/internal/config/model.go index e8d1d3c9..75992d1c 100644 --- a/internal/config/model.go +++ b/internal/config/model.go @@ -70,6 +70,7 @@ type CompactConfig struct { ManualStrategy string `yaml:"manual_strategy,omitempty"` ManualKeepRecentMessages int `yaml:"manual_keep_recent_messages,omitempty"` MaxSummaryChars int `yaml:"max_summary_chars,omitempty"` + MicroCompactDisabled bool `yaml:"micro_compact_disabled,omitempty"` } type WebFetchConfig struct { diff --git a/internal/context/builder.go b/internal/context/builder.go index f1ca2c79..eeac9406 100644 --- a/internal/context/builder.go +++ b/internal/context/builder.go @@ -1,6 +1,10 @@ package context -import "context" +import ( + "context" + + "neo-code/internal/provider" +) // DefaultBuilder preserves the current runtime context-building behavior. type DefaultBuilder struct { @@ -43,6 +47,14 @@ func (b *DefaultBuilder) Build(ctx context.Context, input BuildInput) (BuildResu return BuildResult{ SystemPrompt: composeSystemPrompt(sections...), - Messages: microCompactMessages(trimPolicy.Trim(input.Messages)), + Messages: applyReadTimeContextProjection(trimPolicy.Trim(input.Messages), input.Compact), }, nil } + +// applyReadTimeContextProjection 负责在 provider 请求前按开关应用只读上下文投影,避免改写原始会话消息。 +func applyReadTimeContextProjection(messages []provider.Message, options CompactOptions) []provider.Message { + if options.DisableMicroCompact { + return cloneContextMessages(messages) + } + return microCompactMessages(messages) +} diff --git a/internal/context/builder_test.go b/internal/context/builder_test.go index 9f5098a0..8b72951d 100644 --- a/internal/context/builder_test.go +++ b/internal/context/builder_test.go @@ -5,6 +5,7 @@ import ( "fmt" "os" "path/filepath" + "reflect" "strings" "testing" @@ -204,6 +205,59 @@ func TestDefaultBuilderBuildAppliesMicroCompactAfterTrim(t *testing.T) { } } +func TestDefaultBuilderBuildSkipsMicroCompactWhenDisabled(t *testing.T) { + t.Parallel() + + builder := &DefaultBuilder{ + promptSources: []promptSectionSource{ + stubPromptSectionSource{sections: []promptSection{{title: "Stub", content: "body"}}}, + }, + } + + messages := []provider.Message{ + {Role: provider.RoleUser, Content: "older user"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-1", Content: "old read result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-2", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-3", Name: "webfetch", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, + {Role: provider.RoleUser, Content: "latest explicit instruction"}, + {Role: provider.RoleAssistant, Content: "current reply"}, + } + + got, err := builder.Build(stdcontext.Background(), BuildInput{ + Messages: messages, + Compact: CompactOptions{ + DisableMicroCompact: true, + }, + }) + if err != nil { + t.Fatalf("Build() error = %v", err) + } + if !reflect.DeepEqual(got.Messages, messages) { + t.Fatalf("expected messages to remain unchanged when micro compact is disabled, got %+v", got.Messages) + } + if &got.Messages[2] == &messages[2] { + t.Fatalf("expected disabled path to still clone message slice") + } +} + func TestTrimMessagesPreservesToolPairs(t *testing.T) { t.Parallel() diff --git a/internal/context/microcompact.go b/internal/context/microcompact.go index 10f5baba..10a36cde 100644 --- a/internal/context/microcompact.go +++ b/internal/context/microcompact.go @@ -5,6 +5,7 @@ import ( "neo-code/internal/context/internalcompact" "neo-code/internal/provider" + "neo-code/internal/tools" ) const ( @@ -15,13 +16,13 @@ const ( ) var microCompactableTools = map[string]struct{}{ - "bash": {}, - "webfetch": {}, - "filesystem_read_file": {}, - "filesystem_grep": {}, - "filesystem_glob": {}, - "filesystem_edit": {}, - "filesystem_write_file": {}, + tools.ToolNameBash: {}, + tools.ToolNameWebFetch: {}, + tools.ToolNameFilesystemReadFile: {}, + tools.ToolNameFilesystemGrep: {}, + tools.ToolNameFilesystemGlob: {}, + tools.ToolNameFilesystemEdit: {}, + tools.ToolNameFilesystemWriteFile: {}, } // microCompactMessages 对裁剪后的消息做只读投影式微压缩,仅清理旧工具结果内容。 diff --git a/internal/context/types.go b/internal/context/types.go index 6eb516dd..2e406861 100644 --- a/internal/context/types.go +++ b/internal/context/types.go @@ -15,6 +15,7 @@ type Builder interface { type BuildInput struct { Messages []provider.Message Metadata Metadata + Compact CompactOptions } // BuildResult is the provider-facing context produced for a single round. @@ -22,3 +23,8 @@ type BuildResult struct { SystemPrompt string Messages []provider.Message } + +// CompactOptions controls read-time compact behavior inside the context builder. +type CompactOptions struct { + DisableMicroCompact bool +} diff --git a/internal/runtime/runtime.go b/internal/runtime/runtime.go index c95bbcd3..1ab3ec9b 100644 --- a/internal/runtime/runtime.go +++ b/internal/runtime/runtime.go @@ -157,6 +157,9 @@ func (s *Service) Run(ctx context.Context, input UserInput) error { Provider: cfg.SelectedProvider, Model: cfg.CurrentModel, }, + Compact: agentcontext.CompactOptions{ + DisableMicroCompact: cfg.Context.Compact.MicroCompactDisabled, + }, }) if err != nil { return s.handleRunError(ctx, input.RunID, session.ID, err) diff --git a/internal/runtime/runtime_test.go b/internal/runtime/runtime_test.go index ff24909e..9cd8ed79 100644 --- a/internal/runtime/runtime_test.go +++ b/internal/runtime/runtime_test.go @@ -478,6 +478,9 @@ func TestServiceRunDelegatesToContextBuilder(t *testing.T) { if builder.lastInput.Metadata.Model == "" { t.Fatalf("expected model to be forwarded to builder metadata") } + if builder.lastInput.Compact.DisableMicroCompact { + t.Fatalf("expected micro compact to stay enabled by default") + } if len(builder.lastInput.Messages) != 1 || builder.lastInput.Messages[0].Content != "hello" { t.Fatalf("expected persisted session messages to be forwarded, got %+v", builder.lastInput.Messages) } @@ -492,6 +495,47 @@ func TestServiceRunDelegatesToContextBuilder(t *testing.T) { } } +func TestServiceRunCanDisableMicroCompactViaConfig(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + if err := manager.Update(context.Background(), func(cfg *config.Config) error { + cfg.Context.Compact.MicroCompactDisabled = true + return nil + }); err != nil { + t.Fatalf("update config: %v", err) + } + + store := newMemoryStore() + registry := tools.NewRegistry() + registry.Register(&stubTool{name: "filesystem_read_file", content: "default"}) + + builder := &stubContextBuilder{ + buildFn: func(ctx context.Context, input agentcontext.BuildInput) (agentcontext.BuildResult, error) { + return agentcontext.BuildResult{ + SystemPrompt: "delegated prompt", + Messages: append([]provider.Message(nil), input.Messages...), + }, nil + }, + } + + scripted := &scriptedProvider{ + responses: []provider.ChatResponse{{ + Message: provider.Message{Role: provider.RoleAssistant, Content: "done"}, + FinishReason: "stop", + }}, + } + + service := NewWithFactory(manager, registry, store, &scriptedProviderFactory{provider: scripted}, builder) + if err := service.Run(context.Background(), UserInput{RunID: "run-disable-micro-compact", Content: "hello"}); err != nil { + t.Fatalf("Run() error = %v", err) + } + + if !builder.lastInput.Compact.DisableMicroCompact { + t.Fatalf("expected config to disable micro compact in build input") + } +} + func TestServiceRunPersistsSessionProviderAndModel(t *testing.T) { t.Parallel() diff --git a/internal/tools/bash/tool.go b/internal/tools/bash/tool.go index f2559d97..5c67c719 100644 --- a/internal/tools/bash/tool.go +++ b/internal/tools/bash/tool.go @@ -45,7 +45,7 @@ func NewWithExecutor(root string, shell string, timeout time.Duration, executor } func (t *Tool) Name() string { - return "bash" + return tools.ToolNameBash } func (t *Tool) Description() string { diff --git a/internal/tools/filesystem/helpers.go b/internal/tools/filesystem/helpers.go index d681ed6e..929d74e1 100644 --- a/internal/tools/filesystem/helpers.go +++ b/internal/tools/filesystem/helpers.go @@ -4,14 +4,16 @@ import ( "os" "path/filepath" "strings" + + "neo-code/internal/tools" ) const ( - readFileToolName = "filesystem_read_file" - writeFileToolName = "filesystem_write_file" - grepToolName = "filesystem_grep" - globToolName = "filesystem_glob" - editToolName = "filesystem_edit" + readFileToolName = tools.ToolNameFilesystemReadFile + writeFileToolName = tools.ToolNameFilesystemWriteFile + grepToolName = tools.ToolNameFilesystemGrep + globToolName = tools.ToolNameFilesystemGlob + editToolName = tools.ToolNameFilesystemEdit ) func effectiveRoot(defaultRoot string, workdir string) string { diff --git a/internal/tools/names.go b/internal/tools/names.go new file mode 100644 index 00000000..833430f7 --- /dev/null +++ b/internal/tools/names.go @@ -0,0 +1,12 @@ +package tools + +// Tool name constants are shared across tool implementations, context policies, and tests. +const ( + ToolNameBash = "bash" + ToolNameWebFetch = "webfetch" + ToolNameFilesystemReadFile = "filesystem_read_file" + ToolNameFilesystemWriteFile = "filesystem_write_file" + ToolNameFilesystemGrep = "filesystem_grep" + ToolNameFilesystemGlob = "filesystem_glob" + ToolNameFilesystemEdit = "filesystem_edit" +) diff --git a/internal/tools/webfetch/tool.go b/internal/tools/webfetch/tool.go index efa5474f..5a77c2b2 100644 --- a/internal/tools/webfetch/tool.go +++ b/internal/tools/webfetch/tool.go @@ -16,7 +16,7 @@ import ( ) const ( - toolName = "webfetch" + toolName = tools.ToolNameWebFetch htmlContentType = "text/html" xhtmlContentType = "application/xhtml+xml" reasonInvalidArguments = "invalid arguments" From 104ac2ca67ac91171efd596d013bea796b65d462 Mon Sep 17 00:00:00 2001 From: Cai_Tang <106404101+Cai-Tang-www@users.noreply.github.com> Date: Sun, 5 Apr 2026 16:01:03 +0800 Subject: [PATCH 11/55] feat(security/runtime): add session remember once/always/reject --- internal/runtime/runtime.go | 6 +- internal/runtime/runtime_test.go | 97 +++++++++++++++ internal/tools/manager.go | 83 ++++++++++++- internal/tools/manager_test.go | 203 ++++++++++++++++++++++++++++++- internal/tools/session_memory.go | 142 +++++++++++++++++++++ 5 files changed, 522 insertions(+), 9 deletions(-) create mode 100644 internal/tools/session_memory.go diff --git a/internal/runtime/runtime.go b/internal/runtime/runtime.go index c95bbcd3..ff00ad59 100644 --- a/internal/runtime/runtime.go +++ b/internal/runtime/runtime.go @@ -578,6 +578,7 @@ type permissionEventView struct { decision string reason string ruleID string + scope string resolvedAs string } @@ -609,6 +610,7 @@ func permissionEventFromError(err error) (permissionEventView, bool) { decision: decision, reason: reason, ruleID: strings.TrimSpace(permissionErr.RuleID()), + scope: strings.TrimSpace(permissionErr.RememberScope()), resolvedAs: resolvedAs, }, true } @@ -624,7 +626,7 @@ func (v permissionEventView) toRequestPayload() PermissionRequestPayload { Decision: v.decision, Reason: v.reason, RuleID: v.ruleID, - RememberScope: "", + RememberScope: v.scope, } } @@ -639,7 +641,7 @@ func (v permissionEventView) toResolvedPayload() PermissionResolvedPayload { Decision: v.decision, Reason: v.reason, RuleID: v.ruleID, - RememberScope: "", + RememberScope: v.scope, ResolvedAs: v.resolvedAs, } } diff --git a/internal/runtime/runtime_test.go b/internal/runtime/runtime_test.go index ff24909e..e3b1770e 100644 --- a/internal/runtime/runtime_test.go +++ b/internal/runtime/runtime_test.go @@ -431,6 +431,9 @@ func TestServiceRunDelegatesToContextBuilder(t *testing.T) { manager := newRuntimeConfigManager(t) store := newMemoryStore() + session := newSession("memory reject") + session.ID = "session-memory-reject" + store.sessions[session.ID] = cloneSession(session) registry := tools.NewRegistry() registry.Register(&stubTool{name: "filesystem_read_file", content: "default"}) @@ -650,6 +653,9 @@ func TestServiceRunEmitsPermissionRequestAndResolvedForAsk(t *testing.T) { manager := newRuntimeConfigManager(t) store := newMemoryStore() + session := newSession("memory reject") + session.ID = "session-memory-reject" + store.sessions[session.ID] = cloneSession(session) registry := tools.NewRegistry() tool := &stubTool{name: "webfetch", content: "should-not-run"} registry.Register(tool) @@ -821,6 +827,97 @@ func TestServiceRunEmitsPermissionResolvedForDeny(t *testing.T) { t.Fatalf("expected permission resolved event payload") } +func TestServiceRunEmitsRememberScopeWhenSessionRejectMemoryHits(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + store := newMemoryStore() + session := newSession("memory reject") + session.ID = "session-memory-reject" + store.sessions[session.ID] = cloneSession(session) + registry := tools.NewRegistry() + tool := &stubTool{name: "webfetch", content: "should-not-run"} + registry.Register(tool) + + engine, err := security.NewStaticGateway(security.DecisionAllow, []security.Rule{ + { + ID: "ask-webfetch", + Type: security.ActionTypeRead, + Resource: "webfetch", + Decision: security.DecisionAsk, + Reason: "requires approval", + }, + }) + if err != nil { + t.Fatalf("new static gateway: %v", err) + } + toolManager, err := tools.NewManager(registry, engine, nil) + if err != nil { + t.Fatalf("new tool manager: %v", err) + } + if err := toolManager.RememberSessionDecision("session-memory-reject", security.Action{ + Type: security.ActionTypeRead, + Payload: security.ActionPayload{ + ToolName: "webfetch", + Resource: "webfetch", + Operation: "fetch", + TargetType: security.TargetTypeURL, + Target: "https://example.com/private", + }, + }, tools.SessionPermissionScopeReject); err != nil { + t.Fatalf("remember session reject: %v", err) + } + + scripted := &scriptedProvider{ + responses: []provider.ChatResponse{ + { + Message: provider.Message{ + Role: "assistant", + ToolCalls: []provider.ToolCall{ + {ID: "call-memory-reject", Name: "webfetch", Arguments: `{"url":"https://example.com/private"}`}, + }, + }, + FinishReason: "tool_calls", + }, + { + Message: provider.Message{Role: "assistant", Content: "done"}, + FinishReason: "stop", + }, + }, + } + + service := NewWithFactory(manager, toolManager, store, &scriptedProviderFactory{provider: scripted}, nil) + if err := service.Run(context.Background(), UserInput{ + SessionID: "session-memory-reject", + RunID: "run-memory-reject", + Content: "fetch private", + }); err != nil { + t.Fatalf("Run() error = %v", err) + } + if tool.callCount != 0 { + t.Fatalf("expected remembered reject to skip tool execution, got %d", tool.callCount) + } + + events := collectRuntimeEvents(service.Events()) + assertEventSequence(t, events, []EventType{EventPermissionResolved, EventToolResult, EventAgentDone}) + assertNoEventType(t, events, EventPermissionRequest) + + for _, event := range events { + if event.Type != EventPermissionResolved { + continue + } + payload, ok := event.Payload.(PermissionResolvedPayload) + if !ok { + t.Fatalf("expected PermissionResolvedPayload, got %#v", event.Payload) + } + if payload.RememberScope != string(tools.SessionPermissionScopeReject) { + t.Fatalf("expected remember_scope reject, got %+v", payload) + } + return + } + t.Fatalf("expected permission resolved event payload") +} + func TestServiceRunHandlesToolManagerSpecError(t *testing.T) { manager := newRuntimeConfigManager(t) store := newMemoryStore() diff --git a/internal/tools/manager.go b/internal/tools/manager.go index 4c099c4c..33735ca4 100644 --- a/internal/tools/manager.go +++ b/internal/tools/manager.go @@ -50,6 +50,7 @@ type PermissionDecisionError struct { action security.Action reason string ruleID string + scope SessionPermissionScope } // Error returns a stable error message for the blocked tool call. @@ -112,12 +113,21 @@ func (e *PermissionDecisionError) RuleID() string { return strings.TrimSpace(e.ruleID) } +// RememberScope 返回触发该权限结果时命中的会话记忆范围。 +func (e *PermissionDecisionError) RememberScope() string { + if e == nil { + return "" + } + return strings.TrimSpace(string(e.scope)) +} + // DefaultManager routes tool calls through the permission engine, workspace // sandbox, and executor. type DefaultManager struct { - executor Executor - engine security.PermissionEngine - sandbox WorkspaceSandbox + executor Executor + engine security.PermissionEngine + sandbox WorkspaceSandbox + sessionDecisions *sessionPermissionMemory } // NewManager creates a manager that wraps an executor with security checks. @@ -137,9 +147,10 @@ func NewManager(executor Executor, engine security.PermissionEngine, sandbox Wor } return &DefaultManager{ - executor: executor, - engine: engine, - sandbox: sandbox, + executor: executor, + engine: engine, + sandbox: sandbox, + sessionDecisions: newSessionPermissionMemory(), }, nil } @@ -175,6 +186,22 @@ func (m *DefaultManager) Execute(ctx context.Context, input ToolCallInput) (Tool result.ToolCallID = input.ID return result, err } + if m.sessionDecisions != nil { + if rememberedDecision, rememberedScope, ok := m.sessionDecisions.resolve(input.SessionID, action); ok { + decision = security.CheckResult{ + Decision: rememberedDecision, + Action: action, + Reason: sessionDecisionReason(rememberedScope), + } + if rememberedScope != "" { + decision.Rule = &security.Rule{ + ID: "session-memory:" + string(rememberedScope), + Decision: rememberedDecision, + Reason: decision.Reason, + } + } + } + } if decision.Decision != security.DecisionAllow { result := blockedToolResult(input, decision) return result, permissionErrorFromDecision(decision) @@ -194,6 +221,17 @@ func (m *DefaultManager) Execute(ctx context.Context, input ToolCallInput) (Tool return m.executor.Execute(ctx, input) } +// RememberSessionDecision 记录会话内权限记忆,用于后续同类 action 快速决策。 +func (m *DefaultManager) RememberSessionDecision(sessionID string, action security.Action, scope SessionPermissionScope) error { + if m == nil { + return errors.New("tools: manager is nil") + } + if m.sessionDecisions == nil { + m.sessionDecisions = newSessionPermissionMemory() + } + return m.sessionDecisions.remember(sessionID, action, scope) +} + func blockedToolResult(input ToolCallInput, decision security.CheckResult) ToolResult { reason := "permission denied" if decision.Decision == security.DecisionAsk { @@ -219,6 +257,39 @@ func permissionErrorFromDecision(decision security.CheckResult) error { action: decision.Action, reason: decision.Reason, ruleID: ruleID, + scope: extractRememberScope(decision), + } +} + +// extractRememberScope 从决策规则中提取会话记忆范围。 +func extractRememberScope(decision security.CheckResult) SessionPermissionScope { + if decision.Rule == nil { + return "" + } + ruleID := strings.TrimSpace(decision.Rule.ID) + switch ruleID { + case "session-memory:" + string(SessionPermissionScopeOnce): + return SessionPermissionScopeOnce + case "session-memory:" + string(SessionPermissionScopeAlways): + return SessionPermissionScopeAlways + case "session-memory:" + string(SessionPermissionScopeReject): + return SessionPermissionScopeReject + default: + return "" + } +} + +// sessionDecisionReason 生成会话记忆命中的统一原因文本。 +func sessionDecisionReason(scope SessionPermissionScope) string { + switch scope { + case SessionPermissionScopeOnce: + return "session permission remembered: once" + case SessionPermissionScopeAlways: + return "session permission remembered: always(session)" + case SessionPermissionScopeReject: + return "session permission remembered: reject" + default: + return "session permission remembered" } } diff --git a/internal/tools/manager_test.go b/internal/tools/manager_test.go index a4a2fe98..477c7481 100644 --- a/internal/tools/manager_test.go +++ b/internal/tools/manager_test.go @@ -377,6 +377,9 @@ func TestPermissionDecisionError(t *testing.T) { if err.Action().Type != security.ActionTypeRead { t.Fatalf("expected action type read, got %q", err.Action().Type) } + if err.RememberScope() != "" { + t.Fatalf("expected empty remember scope, got %q", err.RememberScope()) + } if errors.Is(err, context.Canceled) { t.Fatalf("permission error should not match unrelated errors") } @@ -391,13 +394,211 @@ func TestPermissionDecisionError(t *testing.T) { if denyErr.ToolName() != "" { t.Fatalf("expected empty tool name, got %q", denyErr.ToolName()) } + if denyErr.RememberScope() != "" { + t.Fatalf("expected empty remember scope, got %q", denyErr.RememberScope()) + } var nilErr *PermissionDecisionError - if nilErr.Error() != "" || nilErr.Decision() != "" || nilErr.ToolName() != "" { + if nilErr.Error() != "" || nilErr.Decision() != "" || nilErr.ToolName() != "" || nilErr.RememberScope() != "" { t.Fatalf("expected nil permission error helpers to be empty") } } +func TestDefaultManagerSessionPermissionMemory(t *testing.T) { + t.Parallel() + + newAskManager := func(t *testing.T) (*DefaultManager, *managerStubTool) { + t.Helper() + registry := NewRegistry() + webTool := &managerStubTool{name: "webfetch", content: "ok"} + registry.Register(webTool) + engine, err := security.NewStaticGateway(security.DecisionAllow, []security.Rule{ + { + ID: "ask-webfetch", + Type: security.ActionTypeRead, + Resource: "webfetch", + Decision: security.DecisionAsk, + Reason: "requires approval", + }, + }) + if err != nil { + t.Fatalf("new engine: %v", err) + } + manager, err := NewManager(registry, engine, nil) + if err != nil { + t.Fatalf("new manager: %v", err) + } + return manager, webTool + } + + t.Run("once allows only first follow-up", func(t *testing.T) { + t.Parallel() + manager, webTool := newAskManager(t) + input := ToolCallInput{ + ID: "call-once", + Name: "webfetch", + Arguments: []byte(`{"url":"https://example.com/once"}`), + SessionID: "session-once", + } + + _, err := manager.Execute(context.Background(), input) + var permissionErr *PermissionDecisionError + if !errors.As(err, &permissionErr) || permissionErr.Decision() != "ask" { + t.Fatalf("expected initial ask decision, got %v", err) + } + if rememberErr := manager.RememberSessionDecision(input.SessionID, permissionErr.Action(), SessionPermissionScopeOnce); rememberErr != nil { + t.Fatalf("remember once: %v", rememberErr) + } + + result, err := manager.Execute(context.Background(), input) + if err != nil { + t.Fatalf("expected remembered once allow, got %v", err) + } + if result.IsError { + t.Fatalf("expected non-error result, got %+v", result) + } + if webTool.callCount != 1 { + t.Fatalf("expected tool call count 1 after once allow, got %d", webTool.callCount) + } + + _, err = manager.Execute(context.Background(), input) + if !errors.As(err, &permissionErr) || permissionErr.Decision() != "ask" { + t.Fatalf("expected ask after once consumed, got %v", err) + } + }) + + t.Run("always(session) keeps allowing in same session", func(t *testing.T) { + t.Parallel() + manager, webTool := newAskManager(t) + input := ToolCallInput{ + ID: "call-always", + Name: "webfetch", + Arguments: []byte(`{"url":"https://example.com/always"}`), + SessionID: "session-always", + } + + _, err := manager.Execute(context.Background(), input) + var permissionErr *PermissionDecisionError + if !errors.As(err, &permissionErr) || permissionErr.Decision() != "ask" { + t.Fatalf("expected initial ask decision, got %v", err) + } + if rememberErr := manager.RememberSessionDecision(input.SessionID, permissionErr.Action(), SessionPermissionScopeAlways); rememberErr != nil { + t.Fatalf("remember always: %v", rememberErr) + } + + for i := 0; i < 2; i++ { + if _, err := manager.Execute(context.Background(), input); err != nil { + t.Fatalf("expected always allow on iteration %d, got %v", i, err) + } + } + if webTool.callCount != 2 { + t.Fatalf("expected tool to execute twice, got %d", webTool.callCount) + } + }) + + t.Run("reject denies in same session and keeps scope metadata", func(t *testing.T) { + t.Parallel() + manager, webTool := newAskManager(t) + input := ToolCallInput{ + ID: "call-reject", + Name: "webfetch", + Arguments: []byte(`{"url":"https://example.com/reject"}`), + SessionID: "session-reject", + } + + _, err := manager.Execute(context.Background(), input) + var permissionErr *PermissionDecisionError + if !errors.As(err, &permissionErr) || permissionErr.Decision() != "ask" { + t.Fatalf("expected initial ask decision, got %v", err) + } + if rememberErr := manager.RememberSessionDecision(input.SessionID, permissionErr.Action(), SessionPermissionScopeReject); rememberErr != nil { + t.Fatalf("remember reject: %v", rememberErr) + } + + _, err = manager.Execute(context.Background(), input) + if !errors.As(err, &permissionErr) { + t.Fatalf("expected permission error, got %v", err) + } + if permissionErr.Decision() != "deny" { + t.Fatalf("expected deny from remembered reject, got %q", permissionErr.Decision()) + } + if permissionErr.RememberScope() != string(SessionPermissionScopeReject) { + t.Fatalf("expected reject remember scope, got %q", permissionErr.RememberScope()) + } + if webTool.callCount != 0 { + t.Fatalf("expected rejected call to skip tool execution, got %d", webTool.callCount) + } + }) + + t.Run("session memory does not leak across sessions", func(t *testing.T) { + t.Parallel() + manager, _ := newAskManager(t) + inputA := ToolCallInput{ + ID: "call-session-a", + Name: "webfetch", + Arguments: []byte(`{"url":"https://example.com/session-a"}`), + SessionID: "session-a", + } + inputB := ToolCallInput{ + ID: "call-session-b", + Name: "webfetch", + Arguments: []byte(`{"url":"https://example.com/session-a"}`), + SessionID: "session-b", + } + + _, err := manager.Execute(context.Background(), inputA) + var permissionErr *PermissionDecisionError + if !errors.As(err, &permissionErr) { + t.Fatalf("expected permission ask on session A, got %v", err) + } + if rememberErr := manager.RememberSessionDecision(inputA.SessionID, permissionErr.Action(), SessionPermissionScopeAlways); rememberErr != nil { + t.Fatalf("remember session A always: %v", rememberErr) + } + if _, err := manager.Execute(context.Background(), inputA); err != nil { + t.Fatalf("expected session A to be allowed, got %v", err) + } + + _, err = manager.Execute(context.Background(), inputB) + if !errors.As(err, &permissionErr) || permissionErr.Decision() != "ask" { + t.Fatalf("expected session B remain ask, got %v", err) + } + }) + + t.Run("action matching remains exact by structured target", func(t *testing.T) { + t.Parallel() + manager, _ := newAskManager(t) + inputA := ToolCallInput{ + ID: "call-target-a", + Name: "webfetch", + Arguments: []byte(`{"url":"https://example.com/a"}`), + SessionID: "session-target", + } + inputB := ToolCallInput{ + ID: "call-target-b", + Name: "webfetch", + Arguments: []byte(`{"url":"https://example.com/b"}`), + SessionID: "session-target", + } + + _, err := manager.Execute(context.Background(), inputA) + var permissionErr *PermissionDecisionError + if !errors.As(err, &permissionErr) { + t.Fatalf("expected permission ask on target A, got %v", err) + } + if rememberErr := manager.RememberSessionDecision(inputA.SessionID, permissionErr.Action(), SessionPermissionScopeAlways); rememberErr != nil { + t.Fatalf("remember target A: %v", rememberErr) + } + if _, err := manager.Execute(context.Background(), inputA); err != nil { + t.Fatalf("expected target A to be allowed, got %v", err) + } + + _, err = manager.Execute(context.Background(), inputB) + if !errors.As(err, &permissionErr) || permissionErr.Decision() != "ask" { + t.Fatalf("expected target B to stay ask, got %v", err) + } + }) +} + func TestBuildPermissionAction(t *testing.T) { t.Parallel() diff --git a/internal/tools/session_memory.go b/internal/tools/session_memory.go new file mode 100644 index 00000000..be8fb5ba --- /dev/null +++ b/internal/tools/session_memory.go @@ -0,0 +1,142 @@ +package tools + +import ( + "errors" + "fmt" + "strings" + "sync" + + "neo-code/internal/security" +) + +// SessionPermissionScope 表示 session 级权限记忆的作用范围。 +type SessionPermissionScope string + +const ( + // SessionPermissionScopeOnce 表示仅当前一次请求放行。 + SessionPermissionScopeOnce SessionPermissionScope = "once" + // SessionPermissionScopeAlways 表示当前会话内同类请求持续放行。 + SessionPermissionScopeAlways SessionPermissionScope = "always_session" + // SessionPermissionScopeReject 表示当前会话内同类请求持续拒绝。 + SessionPermissionScopeReject SessionPermissionScope = "reject" +) + +type sessionPermissionEntry struct { + decision security.Decision + scope SessionPermissionScope + remaining int +} + +// sessionPermissionMemory 管理按 session/action 维度的审批记忆。 +type sessionPermissionMemory struct { + mu sync.Mutex + entries map[string]map[string]sessionPermissionEntry +} + +// newSessionPermissionMemory 创建 session 级权限记忆存储。 +func newSessionPermissionMemory() *sessionPermissionMemory { + return &sessionPermissionMemory{ + entries: make(map[string]map[string]sessionPermissionEntry), + } +} + +// remember 记录一条 session 级权限决策。 +func (m *sessionPermissionMemory) remember(sessionID string, action security.Action, scope SessionPermissionScope) error { + trimmedSessionID := strings.TrimSpace(sessionID) + if trimmedSessionID == "" { + return errors.New("tools: session id is empty") + } + if err := action.Validate(); err != nil { + return err + } + + var entry sessionPermissionEntry + switch scope { + case SessionPermissionScopeOnce: + entry = sessionPermissionEntry{ + decision: security.DecisionAllow, + scope: scope, + remaining: 1, + } + case SessionPermissionScopeAlways: + entry = sessionPermissionEntry{ + decision: security.DecisionAllow, + scope: scope, + remaining: -1, + } + case SessionPermissionScopeReject: + entry = sessionPermissionEntry{ + decision: security.DecisionDeny, + scope: scope, + remaining: -1, + } + default: + return fmt.Errorf("tools: unsupported session permission scope %q", scope) + } + + actionKey := sessionPermissionActionKey(action) + m.mu.Lock() + defer m.mu.Unlock() + sessionEntries, ok := m.entries[trimmedSessionID] + if !ok { + sessionEntries = make(map[string]sessionPermissionEntry) + m.entries[trimmedSessionID] = sessionEntries + } + sessionEntries[actionKey] = entry + return nil +} + +// resolve 查询并按 scope 规则消费 session 级权限记忆。 +func (m *sessionPermissionMemory) resolve(sessionID string, action security.Action) (security.Decision, SessionPermissionScope, bool) { + trimmedSessionID := strings.TrimSpace(sessionID) + if trimmedSessionID == "" { + return "", "", false + } + actionKey := sessionPermissionActionKey(action) + + m.mu.Lock() + defer m.mu.Unlock() + + sessionEntries, ok := m.entries[trimmedSessionID] + if !ok { + return "", "", false + } + entry, ok := sessionEntries[actionKey] + if !ok { + return "", "", false + } + + if entry.scope == SessionPermissionScopeOnce && entry.remaining > 0 { + entry.remaining-- + if entry.remaining <= 0 { + delete(sessionEntries, actionKey) + } else { + sessionEntries[actionKey] = entry + } + } + + if len(sessionEntries) == 0 { + delete(m.entries, trimmedSessionID) + } + + return entry.decision, entry.scope, true +} + +// sessionPermissionActionKey 基于结构化 action 生成稳定匹配键。 +func sessionPermissionActionKey(action security.Action) string { + normalizedTool := strings.ToLower(strings.TrimSpace(action.Payload.ToolName)) + normalizedResource := strings.ToLower(strings.TrimSpace(action.Payload.Resource)) + normalizedOperation := strings.ToLower(strings.TrimSpace(action.Payload.Operation)) + normalizedTargetType := strings.ToLower(strings.TrimSpace(string(action.Payload.TargetType))) + normalizedTarget := strings.TrimSpace(action.Payload.Target) + normalizedTarget = strings.ReplaceAll(normalizedTarget, "\r\n", "\n") + normalizedTarget = strings.ReplaceAll(normalizedTarget, "\r", "\n") + return strings.Join([]string{ + string(action.Type), + normalizedTool, + normalizedResource, + normalizedOperation, + normalizedTargetType, + normalizedTarget, + }, "|") +} From b60d0f29c6b4383131c5648236c97b22fc2ad16c Mon Sep 17 00:00:00 2001 From: Cai_Tang <106404101+Cai-Tang-www@users.noreply.github.com> Date: Sun, 5 Apr 2026 17:10:36 +0800 Subject: [PATCH 12/55] feat(security): complete permission ask loop and policy engine --- internal/app/bootstrap.go | 2 +- internal/app/bootstrap_test.go | 32 +- internal/runtime/events.go | 6 + internal/runtime/permission.go | 279 ++++++++++++++++++ internal/runtime/runtime.go | 21 +- internal/runtime/runtime_test.go | 94 ++++-- internal/security/policy.go | 490 +++++++++++++++++++++++++++++++ internal/security/policy_test.go | 174 +++++++++++ internal/tools/manager.go | 1 + internal/tools/manager_test.go | 84 +++++- internal/tools/registry.go | 6 + internal/tools/session_memory.go | 46 ++- internal/tui/app.go | 75 ++--- internal/tui/state.go | 8 + internal/tui/update.go | 113 +++++++ internal/tui/update_test.go | 7 + 16 files changed, 1340 insertions(+), 98 deletions(-) create mode 100644 internal/runtime/permission.go create mode 100644 internal/security/policy.go create mode 100644 internal/security/policy_test.go diff --git a/internal/app/bootstrap.go b/internal/app/bootstrap.go index d87241ea..bd4ea261 100644 --- a/internal/app/bootstrap.go +++ b/internal/app/bootstrap.go @@ -99,7 +99,7 @@ func buildToolRegistry(cfg config.Config) *tools.Registry { } func buildToolManager(registry *tools.Registry) (tools.Manager, error) { - engine, err := security.NewStaticGateway(security.DecisionAllow, nil) + engine, err := security.NewRecommendedPolicyEngine() if err != nil { return nil, err } diff --git a/internal/app/bootstrap_test.go b/internal/app/bootstrap_test.go index bb204cbd..beefb244 100644 --- a/internal/app/bootstrap_test.go +++ b/internal/app/bootstrap_test.go @@ -132,16 +132,13 @@ func TestBuildToolManagerWrapsRegistry(t *testing.T) { t.Fatalf("expected 1 spec, got %+v", specs) } - result, execErr := manager.Execute(context.Background(), tools.ToolCallInput{ + _, execErr := manager.Execute(context.Background(), tools.ToolCallInput{ Name: "bash", Arguments: []byte(`{"command":"echo hi"}`), Workdir: workdir, }) - if execErr != nil { - t.Fatalf("Execute() error = %v", execErr) - } - if result.Content != "ok" { - t.Fatalf("expected ok result, got %+v", result) + if execErr == nil { + t.Fatalf("expected bash to require approval by default policy") } _, execErr = manager.Execute(context.Background(), tools.ToolCallInput{ @@ -154,6 +151,29 @@ func TestBuildToolManagerWrapsRegistry(t *testing.T) { } } +func TestBuildToolManagerAllowsWebfetchWhitelist(t *testing.T) { + t.Parallel() + + registry := tools.NewRegistry() + registry.Register(stubToolForBootstrap{name: "webfetch", content: "ok"}) + manager, err := buildToolManager(registry) + if err != nil { + t.Fatalf("buildToolManager() error = %v", err) + } + + result, execErr := manager.Execute(context.Background(), tools.ToolCallInput{ + Name: "webfetch", + Arguments: []byte(`{"url":"https://github.com/1024XEngineer/neo-code"}`), + Workdir: t.TempDir(), + }) + if execErr != nil { + t.Fatalf("expected whitelist webfetch allow, got %v", execErr) + } + if result.Content != "ok" { + t.Fatalf("expected ok result, got %+v", result) + } +} + func TestEnsureConsoleUTF8SetsOutputThenInput(t *testing.T) { originalOutput := setConsoleOutputCodePage originalInput := setConsoleInputCodePage diff --git a/internal/runtime/events.go b/internal/runtime/events.go index 594e8f5e..704ff14a 100644 --- a/internal/runtime/events.go +++ b/internal/runtime/events.go @@ -15,7 +15,10 @@ type RuntimeEvent struct { // PermissionRequestPayload 描述一次需要审批的权限请求上下文。 type PermissionRequestPayload struct { + RequestID string + ToolCallID string ToolName string + ToolCategory string ActionType string Operation string TargetType string @@ -28,7 +31,10 @@ type PermissionRequestPayload struct { // PermissionResolvedPayload 描述权限请求被运行时处理后的最终状态。 type PermissionResolvedPayload struct { + RequestID string + ToolCallID string ToolName string + ToolCategory string ActionType string Operation string TargetType string diff --git a/internal/runtime/permission.go b/internal/runtime/permission.go new file mode 100644 index 00000000..52f59908 --- /dev/null +++ b/internal/runtime/permission.go @@ -0,0 +1,279 @@ +package runtime + +import ( + "context" + "errors" + "fmt" + "strings" + "sync" + "time" + + "neo-code/internal/provider" + "neo-code/internal/security" + "neo-code/internal/tools" +) + +// PermissionResolutionInput 描述一次来自界面的权限审批决定。 +type PermissionResolutionInput struct { + RequestID string + Decision PermissionResolutionDecision +} + +// PermissionResolutionDecision 表示用户在权限提示中的最终选择。 +type PermissionResolutionDecision string + +const ( + PermissionResolutionAllowOnce PermissionResolutionDecision = "allow_once" + PermissionResolutionAllowSession PermissionResolutionDecision = "allow_session" + PermissionResolutionReject PermissionResolutionDecision = "reject" +) + +type permissionExecutionInput struct { + RunID string + SessionID string + Call provider.ToolCall + Workdir string + ToolTimeout time.Duration +} + +type pendingPermissionRequest struct { + RequestID string + RunID string + SessionID string + Call provider.ToolCall + Action security.Action + ResultCh chan PermissionResolutionDecision +} + +var runtimePendingPermissions = struct { + mu sync.Mutex + nextID uint64 + byRun map[*Service]*pendingPermissionRequest +}{ + byRun: make(map[*Service]*pendingPermissionRequest), +} + +// ResolvePermission 接收 UI 的审批决定,并唤醒 runtime 中等待的工具调用。 +func (s *Service) ResolvePermission(ctx context.Context, input PermissionResolutionInput) error { + requestID := strings.TrimSpace(input.RequestID) + if requestID == "" { + return errors.New("runtime: permission request id is empty") + } + decision := normalizePermissionResolutionDecision(input.Decision) + if decision == "" { + return fmt.Errorf("runtime: unsupported permission decision %q", input.Decision) + } + + runtimePendingPermissions.mu.Lock() + pending := runtimePendingPermissions.byRun[s] + runtimePendingPermissions.mu.Unlock() + if pending == nil || pending.RequestID != requestID { + return fmt.Errorf("runtime: permission request %q not found", requestID) + } + + select { + case pending.ResultCh <- decision: + return nil + case <-ctx.Done(): + return ctx.Err() + } +} + +// executeToolCallWithPermission 执行工具调用并处理 ask 决策的显式审批闭环。 +func (s *Service) executeToolCallWithPermission(ctx context.Context, input permissionExecutionInput) (tools.ToolResult, error) { + callInput := tools.ToolCallInput{ + ID: input.Call.ID, + Name: input.Call.Name, + Arguments: []byte(input.Call.Arguments), + Workdir: input.Workdir, + SessionID: input.SessionID, + EmitChunk: func(chunk []byte) { + s.emit(ctx, EventToolChunk, input.RunID, input.SessionID, string(chunk)) + }, + } + + runCtx, cancel := context.WithTimeout(ctx, input.ToolTimeout) + result, execErr := s.toolManager.Execute(runCtx, callInput) + cancel() + if execErr == nil { + return result, nil + } + + var permissionErr *tools.PermissionDecisionError + if !errors.As(execErr, &permissionErr) { + return result, execErr + } + if !strings.EqualFold(permissionErr.Decision(), string(security.DecisionAsk)) { + return result, execErr + } + + decision, scope, requestID, err := s.awaitPermissionDecision(ctx, input, permissionErr) + if err != nil { + return result, err + } + + if decision == PermissionResolutionReject { + if scope == tools.SessionPermissionScopeReject { + if rememberErr := s.toolManager.RememberSessionDecision(input.SessionID, permissionErr.Action(), scope); rememberErr != nil { + return tools.ToolResult{}, rememberErr + } + } + s.emit(ctx, EventPermissionResolved, input.RunID, input.SessionID, PermissionResolvedPayload{ + RequestID: requestID, + ToolCallID: input.Call.ID, + ToolName: input.Call.Name, + ToolCategory: permissionToolCategory(permissionErr.Action()), + ActionType: string(permissionErr.Action().Type), + Operation: permissionErr.Action().Payload.Operation, + TargetType: string(permissionErr.Action().Payload.TargetType), + Target: permissionErr.Action().Payload.Target, + Decision: "deny", + Reason: "permission rejected by user", + RuleID: permissionErr.RuleID(), + RememberScope: string(scope), + ResolvedAs: "rejected", + }) + return tools.ToolResult{ + ToolCallID: input.Call.ID, + Name: input.Call.Name, + Content: "tool error\ntool: " + input.Call.Name + "\nreason: permission rejected by user", + IsError: true, + }, errors.New("tools: permission rejected by user") + } + + if rememberErr := s.toolManager.RememberSessionDecision(input.SessionID, permissionErr.Action(), scope); rememberErr != nil { + return tools.ToolResult{}, rememberErr + } + s.emit(ctx, EventPermissionResolved, input.RunID, input.SessionID, PermissionResolvedPayload{ + RequestID: requestID, + ToolCallID: input.Call.ID, + ToolName: input.Call.Name, + ToolCategory: permissionToolCategory(permissionErr.Action()), + ActionType: string(permissionErr.Action().Type), + Operation: permissionErr.Action().Payload.Operation, + TargetType: string(permissionErr.Action().Payload.TargetType), + Target: permissionErr.Action().Payload.Target, + Decision: "allow", + Reason: "permission approved by user", + RuleID: permissionErr.RuleID(), + RememberScope: string(scope), + ResolvedAs: "approved", + }) + + retryCtx, retryCancel := context.WithTimeout(ctx, input.ToolTimeout) + retryResult, retryErr := s.toolManager.Execute(retryCtx, callInput) + retryCancel() + return retryResult, retryErr +} + +// awaitPermissionDecision 发送 permission_request 事件,并等待 UI 回传审批结果。 +func (s *Service) awaitPermissionDecision( + ctx context.Context, + input permissionExecutionInput, + permissionErr *tools.PermissionDecisionError, +) (PermissionResolutionDecision, tools.SessionPermissionScope, string, error) { + request := registerPendingPermission(s, input, permissionErr.Action()) + defer clearPendingPermission(s, request.RequestID) + + s.emit(ctx, EventPermissionRequest, input.RunID, input.SessionID, PermissionRequestPayload{ + RequestID: request.RequestID, + ToolCallID: input.Call.ID, + ToolName: input.Call.Name, + ToolCategory: permissionToolCategory(permissionErr.Action()), + ActionType: string(permissionErr.Action().Type), + Operation: permissionErr.Action().Payload.Operation, + TargetType: string(permissionErr.Action().Payload.TargetType), + Target: permissionErr.Action().Payload.Target, + Decision: permissionErr.Decision(), + Reason: permissionErr.Reason(), + RuleID: permissionErr.RuleID(), + RememberScope: string(tools.SessionPermissionScopeAlways), + }) + + select { + case <-ctx.Done(): + return "", "", request.RequestID, ctx.Err() + case decision := <-request.ResultCh: + scope, err := rememberScopeFromDecision(decision) + if err != nil { + return "", "", request.RequestID, err + } + return decision, scope, request.RequestID, nil + } +} + +// registerPendingPermission 注册待审批请求并返回 request id。 +func registerPendingPermission(s *Service, input permissionExecutionInput, action security.Action) pendingPermissionRequest { + runtimePendingPermissions.mu.Lock() + defer runtimePendingPermissions.mu.Unlock() + + runtimePendingPermissions.nextID++ + request := pendingPermissionRequest{ + RequestID: fmt.Sprintf("perm-%d", runtimePendingPermissions.nextID), + RunID: input.RunID, + SessionID: input.SessionID, + Call: input.Call, + Action: action, + ResultCh: make(chan PermissionResolutionDecision, 1), + } + runtimePendingPermissions.byRun[s] = &request + return request +} + +// clearPendingPermission 清理待审批请求,避免跨请求误用。 +func clearPendingPermission(s *Service, requestID string) { + runtimePendingPermissions.mu.Lock() + defer runtimePendingPermissions.mu.Unlock() + + current := runtimePendingPermissions.byRun[s] + if current != nil && current.RequestID == requestID { + delete(runtimePendingPermissions.byRun, s) + } +} + +// normalizePermissionResolutionDecision 将审批决定归一化为受支持枚举值。 +func normalizePermissionResolutionDecision(decision PermissionResolutionDecision) PermissionResolutionDecision { + switch strings.ToLower(strings.TrimSpace(string(decision))) { + case "allow_once", "once", "y", "yes": + return PermissionResolutionAllowOnce + case "allow_session", "always", "always_session", "a": + return PermissionResolutionAllowSession + case "reject", "deny", "n", "no", "r": + return PermissionResolutionReject + default: + return "" + } +} + +// rememberScopeFromDecision 将审批决定映射到工具层权限记忆范围。 +func rememberScopeFromDecision(decision PermissionResolutionDecision) (tools.SessionPermissionScope, error) { + switch decision { + case PermissionResolutionAllowOnce: + return tools.SessionPermissionScopeOnce, nil + case PermissionResolutionAllowSession: + return tools.SessionPermissionScopeAlways, nil + case PermissionResolutionReject: + return tools.SessionPermissionScopeReject, nil + default: + return "", fmt.Errorf("runtime: unsupported permission decision %q", decision) + } +} + +// permissionToolCategory 将 action 归一化为工具类别标签,供审批展示和记忆使用。 +func permissionToolCategory(action security.Action) string { + resource := strings.ToLower(strings.TrimSpace(action.Payload.Resource)) + switch action.Type { + case security.ActionTypeRead: + if strings.HasPrefix(resource, "filesystem_") { + return "filesystem_read" + } + case security.ActionTypeWrite: + if strings.HasPrefix(resource, "filesystem_") { + return "filesystem_write" + } + } + if resource != "" { + return resource + } + return strings.ToLower(strings.TrimSpace(action.Payload.ToolName)) +} diff --git a/internal/runtime/runtime.go b/internal/runtime/runtime.go index ff00ad59..216ce178 100644 --- a/internal/runtime/runtime.go +++ b/internal/runtime/runtime.go @@ -39,6 +39,7 @@ var runtimeSessionWorkdirs = struct { type Runtime interface { Run(ctx context.Context, input UserInput) error Compact(ctx context.Context, input CompactInput) (CompactResult, error) + ResolvePermission(ctx context.Context, input PermissionResolutionInput) error CancelActiveRun() bool Events() <-chan RuntimeEvent ListSessions(ctx context.Context) ([]SessionSummary, error) @@ -218,18 +219,13 @@ func (s *Service) Run(ctx context.Context, input UserInput) error { } s.emit(ctx, EventToolStart, input.RunID, session.ID, call) - runCtx, cancel := context.WithTimeout(ctx, time.Duration(cfg.ToolTimeoutSec)*time.Second) - result, execErr := s.toolManager.Execute(runCtx, tools.ToolCallInput{ - ID: call.ID, - Name: call.Name, - Arguments: []byte(call.Arguments), - Workdir: activeWorkdir, - SessionID: session.ID, - EmitChunk: func(chunk []byte) { - s.emit(ctx, EventToolChunk, input.RunID, session.ID, string(chunk)) - }, + result, execErr := s.executeToolCallWithPermission(ctx, permissionExecutionInput{ + RunID: input.RunID, + SessionID: session.ID, + Call: call, + Workdir: activeWorkdir, + ToolTimeout: time.Duration(cfg.ToolTimeoutSec) * time.Second, }) - cancel() if s.isRunCanceled(execErr) { return s.handleRunError(ctx, input.RunID, session.ID, execErr) } @@ -243,9 +239,6 @@ func (s *Service) Run(ctx context.Context, input UserInput) error { result.Content = execErr.Error() } if permissionEvent, ok := permissionEventFromError(execErr); ok { - if permissionEvent.decision == "ask" { - s.emit(ctx, EventPermissionRequest, input.RunID, session.ID, permissionEvent.toRequestPayload()) - } s.emit(ctx, EventPermissionResolved, input.RunID, session.ID, permissionEvent.toResolvedPayload()) } diff --git a/internal/runtime/runtime_test.go b/internal/runtime/runtime_test.go index e3b1770e..50123578 100644 --- a/internal/runtime/runtime_test.go +++ b/internal/runtime/runtime_test.go @@ -205,6 +205,12 @@ type stubToolManager struct { listCalls int executeCalls int lastInput tools.ToolCallInput + rememberErr error + remembered []struct { + sessionID string + action security.Action + scope tools.SessionPermissionScope + } } func (m *stubToolManager) ListAvailableSpecs(ctx context.Context, input tools.SpecListInput) ([]provider.ToolSpec, error) { @@ -228,6 +234,19 @@ func (m *stubToolManager) Execute(ctx context.Context, input tools.ToolCallInput return result, m.err } +func (m *stubToolManager) RememberSessionDecision(sessionID string, action security.Action, scope tools.SessionPermissionScope) error { + m.remembered = append(m.remembered, struct { + sessionID string + action security.Action + scope tools.SessionPermissionScope + }{ + sessionID: sessionID, + action: action, + scope: scope, + }) + return m.rememberErr +} + func TestServiceRun(t *testing.T) { tests := []struct { name string @@ -648,7 +667,7 @@ func TestServiceRunUsesToolManager(t *testing.T) { } } -func TestServiceRunEmitsPermissionRequestAndResolvedForAsk(t *testing.T) { +func TestServiceRunWaitsForPermissionResolutionAndContinues(t *testing.T) { t.Parallel() manager := newRuntimeConfigManager(t) @@ -657,7 +676,7 @@ func TestServiceRunEmitsPermissionRequestAndResolvedForAsk(t *testing.T) { session.ID = "session-memory-reject" store.sessions[session.ID] = cloneSession(session) registry := tools.NewRegistry() - tool := &stubTool{name: "webfetch", content: "should-not-run"} + tool := &stubTool{name: "webfetch", content: "fetched"} registry.Register(tool) engine, err := security.NewStaticGateway(security.DecisionAllow, []security.Rule{ @@ -696,34 +715,62 @@ func TestServiceRunEmitsPermissionRequestAndResolvedForAsk(t *testing.T) { } service := NewWithFactory(manager, toolManager, store, &scriptedProviderFactory{provider: scripted}, nil) - if err := service.Run(context.Background(), UserInput{RunID: "run-permission-ask", Content: "fetch private"}); err != nil { + runErrCh := make(chan error, 1) + go func() { + runErrCh <- service.Run(context.Background(), UserInput{RunID: "run-permission-ask", Content: "fetch private"}) + }() + + var requestPayload PermissionRequestPayload + deadline := time.After(3 * time.Second) +waitRequest: + for { + select { + case <-deadline: + t.Fatalf("timed out waiting permission request event") + case event := <-service.Events(): + if event.Type != EventPermissionRequest { + continue + } + payload, ok := event.Payload.(PermissionRequestPayload) + if !ok { + t.Fatalf("expected PermissionRequestPayload, got %#v", event.Payload) + } + requestPayload = payload + break waitRequest + } + } + + if strings.TrimSpace(requestPayload.RequestID) == "" { + t.Fatalf("expected non-empty permission request id") + } + if requestPayload.ToolName != "webfetch" || requestPayload.Decision != "ask" { + t.Fatalf("unexpected permission request payload: %+v", requestPayload) + } + + if err := service.ResolvePermission(context.Background(), PermissionResolutionInput{ + RequestID: requestPayload.RequestID, + Decision: PermissionResolutionAllowSession, + }); err != nil { + t.Fatalf("ResolvePermission() error = %v", err) + } + if err := <-runErrCh; err != nil { t.Fatalf("Run() error = %v", err) } - if tool.callCount != 0 { - t.Fatalf("expected blocked tool not to execute, got %d", tool.callCount) + if tool.callCount != 1 { + t.Fatalf("expected allowed tool to execute once, got %d", tool.callCount) } events := collectRuntimeEvents(service.Events()) assertEventSequence(t, events, []EventType{ - EventPermissionRequest, EventPermissionResolved, EventToolResult, EventAgentDone, }) assertNoEventType(t, events, EventError) - var ( - requestPayload PermissionRequestPayload - resolvedPayload PermissionResolvedPayload - ) + var resolvedPayload PermissionResolvedPayload for _, event := range events { switch event.Type { - case EventPermissionRequest: - payload, ok := event.Payload.(PermissionRequestPayload) - if !ok { - t.Fatalf("expected PermissionRequestPayload, got %#v", event.Payload) - } - requestPayload = payload case EventPermissionResolved: payload, ok := event.Payload.(PermissionResolvedPayload) if !ok { @@ -733,17 +780,14 @@ func TestServiceRunEmitsPermissionRequestAndResolvedForAsk(t *testing.T) { } } - if requestPayload.ToolName != "webfetch" || requestPayload.Decision != "ask" { - t.Fatalf("unexpected permission request payload: %+v", requestPayload) - } - if requestPayload.RuleID != "ask-webfetch" { - t.Fatalf("expected rule id ask-webfetch, got %+v", requestPayload) - } - if resolvedPayload.ToolName != "webfetch" || resolvedPayload.Decision != "ask" { + if resolvedPayload.ToolName != "webfetch" || resolvedPayload.Decision != "allow" { t.Fatalf("unexpected permission resolved payload: %+v", resolvedPayload) } - if resolvedPayload.ResolvedAs != "rejected" { - t.Fatalf("expected resolved_as rejected, got %+v", resolvedPayload) + if resolvedPayload.ResolvedAs != "approved" { + t.Fatalf("expected resolved_as approved, got %+v", resolvedPayload) + } + if resolvedPayload.RememberScope != string(tools.SessionPermissionScopeAlways) { + t.Fatalf("expected remember scope always_session, got %+v", resolvedPayload) } } diff --git a/internal/security/policy.go b/internal/security/policy.go new file mode 100644 index 00000000..0b7060cc --- /dev/null +++ b/internal/security/policy.go @@ -0,0 +1,490 @@ +package security + +import ( + "context" + "fmt" + "net/url" + "path/filepath" + "sort" + "strings" +) + +// PolicyRule 描述一条可组合的权限策略规则。 +// 规则按 Priority 从高到低匹配,同优先级保持声明顺序。 +type PolicyRule struct { + ID string + Priority int + Decision Decision + Reason string + + ActionTypes []ActionType + ResourcePatterns []string + ToolCategories []string + TargetTypes []TargetType + + PathPatterns []string + PathSegmentKeywords []string + PathBasenamePatterns []string + RequireSensitivePath bool + + HostPatterns []string + RequireHostMatch bool + RequireHostMissing bool +} + +// PolicyEngine 基于结构化命中条件执行权限决策。 +type PolicyEngine struct { + defaultDecision Decision + rules []PolicyRule +} + +type compiledRule struct { + rule PolicyRule + order int +} + +type actionView struct { + action Action + resource string + toolCategory string + targetType TargetType + target string + targetPath string + host string + sensitive bool +} + +// NewPolicyEngine 创建支持优先级与多条件匹配的权限引擎。 +func NewPolicyEngine(defaultDecision Decision, rules []PolicyRule) (*PolicyEngine, error) { + if defaultDecision == "" { + defaultDecision = DecisionAllow + } + if err := defaultDecision.Validate(); err != nil { + return nil, err + } + + compiled := make([]compiledRule, 0, len(rules)) + for idx := range rules { + rule := rules[idx] + if strings.TrimSpace(rule.ID) == "" { + return nil, fmt.Errorf("security: policy rule id is empty at index %d", idx) + } + if err := rule.Decision.Validate(); err != nil { + return nil, fmt.Errorf("security: policy rule %q: %w", rule.ID, err) + } + for _, actionType := range rule.ActionTypes { + if actionType == "" { + continue + } + if err := actionType.Validate(); err != nil { + return nil, fmt.Errorf("security: policy rule %q: %w", rule.ID, err) + } + } + compiled = append(compiled, compiledRule{ + rule: normalizePolicyRule(rule), + order: idx, + }) + } + + sort.SliceStable(compiled, func(i, j int) bool { + if compiled[i].rule.Priority == compiled[j].rule.Priority { + return compiled[i].order < compiled[j].order + } + return compiled[i].rule.Priority > compiled[j].rule.Priority + }) + + sortedRules := make([]PolicyRule, 0, len(compiled)) + for _, item := range compiled { + sortedRules = append(sortedRules, item.rule) + } + + return &PolicyEngine{ + defaultDecision: defaultDecision, + rules: sortedRules, + }, nil +} + +// Check 返回首条命中规则;若无命中则返回默认决策。 +func (e *PolicyEngine) Check(ctx context.Context, action Action) (CheckResult, error) { + if err := ctx.Err(); err != nil { + return CheckResult{}, err + } + if err := action.Validate(); err != nil { + return CheckResult{}, err + } + + view := newActionView(action) + for _, rule := range e.rules { + if !matchesPolicyRule(rule, view) { + continue + } + matchedRule := Rule{ + ID: rule.ID, + Type: action.Type, + Resource: action.Payload.Resource, + Decision: rule.Decision, + Reason: rule.Reason, + } + return CheckResult{ + Decision: rule.Decision, + Action: action, + Rule: &matchedRule, + Reason: strings.TrimSpace(rule.Reason), + }, nil + } + + return CheckResult{ + Decision: e.defaultDecision, + Action: action, + }, nil +} + +// NewRecommendedPolicyEngine 返回推荐安全策略: +// bash=ask, filesystem write=ask, filesystem read敏感路径=ask/deny, webfetch白名单allow其余ask。 +func NewRecommendedPolicyEngine() (*PolicyEngine, error) { + const ( + reasonAskBash = "bash command requires approval" + reasonAskFilesystemWrite = "filesystem write requires approval" + reasonDenyPrivateKeyRead = "reading private key material is blocked" + reasonAskSensitiveRead = "reading sensitive path requires approval" + reasonAllowWebfetchDomain = "approved web domain" + reasonAskWebfetchDomain = "external web domain requires approval" + ) + + rules := []PolicyRule{ + { + ID: "deny-sensitive-private-keys", + Priority: 1000, + Decision: DecisionDeny, + Reason: reasonDenyPrivateKeyRead, + ActionTypes: []ActionType{ActionTypeRead}, + ToolCategories: []string{"filesystem_read"}, + PathBasenamePatterns: []string{"id_rsa", "id_dsa", "id_ecdsa", "id_ed25519", "*.pem", "*.p12", "*.pfx", "*.key"}, + }, + { + ID: "ask-sensitive-filesystem-read", + Priority: 900, + Decision: DecisionAsk, + Reason: reasonAskSensitiveRead, + ActionTypes: []ActionType{ActionTypeRead}, + ToolCategories: []string{"filesystem_read"}, + RequireSensitivePath: true, + PathSegmentKeywords: []string{"secrets", ".ssh", ".gnupg", ".aws", ".config"}, + PathBasenamePatterns: []string{".env", ".env.*", "*.env", "*.secret", "*.secrets", "*.token"}, + ResourcePatterns: []string{"filesystem_read_*", "filesystem_grep", "filesystem_glob"}, + TargetTypes: []TargetType{TargetTypePath, TargetTypeDirectory}, + RequireHostMissing: false, + RequireHostMatch: false, + }, + { + ID: "ask-all-bash", + Priority: 800, + Decision: DecisionAsk, + Reason: reasonAskBash, + ActionTypes: []ActionType{ActionTypeBash}, + ResourcePatterns: []string{"bash"}, + }, + { + ID: "ask-filesystem-write", + Priority: 780, + Decision: DecisionAsk, + Reason: reasonAskFilesystemWrite, + ActionTypes: []ActionType{ActionTypeWrite}, + ResourcePatterns: []string{"filesystem_write_*", "filesystem_edit"}, + }, + { + ID: "allow-webfetch-whitelist", + Priority: 760, + Decision: DecisionAllow, + Reason: reasonAllowWebfetchDomain, + ActionTypes: []ActionType{ActionTypeRead}, + ResourcePatterns: []string{"webfetch"}, + HostPatterns: []string{"github.com", "*.github.com", "docs.*"}, + RequireHostMatch: true, + }, + { + ID: "ask-webfetch-non-whitelist", + Priority: 740, + Decision: DecisionAsk, + Reason: reasonAskWebfetchDomain, + ActionTypes: []ActionType{ActionTypeRead}, + ResourcePatterns: []string{"webfetch"}, + HostPatterns: []string{"github.com", "*.github.com", "docs.*"}, + RequireHostMissing: true, + }, + } + + return NewPolicyEngine(DecisionAllow, rules) +} + +func normalizePolicyRule(rule PolicyRule) PolicyRule { + rule.ID = strings.TrimSpace(rule.ID) + rule.Reason = strings.TrimSpace(rule.Reason) + rule.ActionTypes = normalizeActionTypes(rule.ActionTypes) + rule.ResourcePatterns = normalizeLowerList(rule.ResourcePatterns) + rule.ToolCategories = normalizeLowerList(rule.ToolCategories) + rule.TargetTypes = normalizeTargetTypes(rule.TargetTypes) + rule.PathPatterns = normalizePathPatterns(rule.PathPatterns) + rule.PathSegmentKeywords = normalizeLowerList(rule.PathSegmentKeywords) + rule.PathBasenamePatterns = normalizeLowerList(rule.PathBasenamePatterns) + rule.HostPatterns = normalizeHostPatterns(rule.HostPatterns) + return rule +} + +func normalizeActionTypes(values []ActionType) []ActionType { + out := make([]ActionType, 0, len(values)) + for _, value := range values { + if strings.TrimSpace(string(value)) == "" { + continue + } + out = append(out, ActionType(strings.TrimSpace(string(value)))) + } + return out +} + +func normalizeTargetTypes(values []TargetType) []TargetType { + out := make([]TargetType, 0, len(values)) + for _, value := range values { + if strings.TrimSpace(string(value)) == "" { + continue + } + out = append(out, TargetType(strings.TrimSpace(string(value)))) + } + return out +} + +func normalizeLowerList(values []string) []string { + out := make([]string, 0, len(values)) + for _, value := range values { + trimmed := strings.ToLower(strings.TrimSpace(value)) + if trimmed == "" { + continue + } + out = append(out, trimmed) + } + return out +} + +func normalizePathPatterns(values []string) []string { + out := make([]string, 0, len(values)) + for _, value := range values { + trimmed := filepath.ToSlash(strings.ToLower(strings.TrimSpace(value))) + if trimmed == "" { + continue + } + out = append(out, trimmed) + } + return out +} + +func normalizeHostPatterns(values []string) []string { + out := make([]string, 0, len(values)) + for _, value := range values { + trimmed := strings.ToLower(strings.TrimSpace(value)) + trimmed = strings.TrimPrefix(trimmed, ".") + if trimmed == "" { + continue + } + out = append(out, trimmed) + } + return out +} + +func newActionView(action Action) actionView { + resource := strings.ToLower(strings.TrimSpace(action.Payload.Resource)) + target := strings.TrimSpace(action.Payload.Target) + host := "" + if action.Payload.TargetType == TargetTypeURL || resource == "webfetch" { + host = parseURLHost(target) + } + targetPath := filepath.ToSlash(strings.ToLower(strings.TrimSpace(target))) + category := deriveToolCategory(action) + sensitive := classifySensitivePath(targetPath) + + return actionView{ + action: action, + resource: resource, + toolCategory: category, + targetType: action.Payload.TargetType, + target: target, + targetPath: targetPath, + host: host, + sensitive: sensitive, + } +} + +func matchesPolicyRule(rule PolicyRule, view actionView) bool { + if len(rule.ActionTypes) > 0 { + matched := false + for _, actionType := range rule.ActionTypes { + if view.action.Type == actionType { + matched = true + break + } + } + if !matched { + return false + } + } + + if len(rule.ResourcePatterns) > 0 && !matchesAnyPattern(view.resource, rule.ResourcePatterns) { + return false + } + if len(rule.ToolCategories) > 0 && !containsString(rule.ToolCategories, view.toolCategory) { + return false + } + + if len(rule.TargetTypes) > 0 { + matched := false + for _, targetType := range rule.TargetTypes { + if view.targetType == targetType { + matched = true + break + } + } + if !matched { + return false + } + } + + if rule.RequireSensitivePath && !view.sensitive { + return false + } + pathMatcherCount := len(rule.PathPatterns) + len(rule.PathSegmentKeywords) + len(rule.PathBasenamePatterns) + if pathMatcherCount > 0 { + pathMatched := false + if len(rule.PathPatterns) > 0 && matchesAnyPattern(view.targetPath, rule.PathPatterns) { + pathMatched = true + } + if len(rule.PathSegmentKeywords) > 0 && matchesPathSegmentKeyword(view.targetPath, rule.PathSegmentKeywords) { + pathMatched = true + } + if len(rule.PathBasenamePatterns) > 0 && matchesPathBasenamePattern(view.targetPath, rule.PathBasenamePatterns) { + pathMatched = true + } + if !pathMatched { + return false + } + } + + hostMatched := len(rule.HostPatterns) == 0 || matchesAnyPattern(view.host, rule.HostPatterns) + if rule.RequireHostMatch && !hostMatched { + return false + } + if rule.RequireHostMissing && hostMatched { + return false + } + + return true +} + +func deriveToolCategory(action Action) string { + resource := strings.ToLower(strings.TrimSpace(action.Payload.Resource)) + switch action.Type { + case ActionTypeRead: + if strings.HasPrefix(resource, "filesystem_") { + return "filesystem_read" + } + case ActionTypeWrite: + if strings.HasPrefix(resource, "filesystem_") { + return "filesystem_write" + } + case ActionTypeBash: + return "bash" + } + if resource != "" { + return resource + } + return strings.ToLower(strings.TrimSpace(action.Payload.ToolName)) +} + +func classifySensitivePath(normalizedTargetPath string) bool { + if normalizedTargetPath == "" { + return false + } + return matchesPathSegmentKeyword(normalizedTargetPath, []string{"secrets", ".ssh", ".gnupg", ".aws", ".config"}) || + matchesPathBasenamePattern(normalizedTargetPath, []string{".env", ".env.*", "*.env", "*.secret", "*.secrets", "*.token", "*.key", "*.pem", "id_rsa", "id_ed25519"}) +} + +func matchesPathSegmentKeyword(normalizedTargetPath string, keywords []string) bool { + if normalizedTargetPath == "" || len(keywords) == 0 { + return false + } + segments := strings.Split(normalizedTargetPath, "/") + for _, segment := range segments { + token := strings.ToLower(strings.TrimSpace(segment)) + if token == "" { + continue + } + for _, keyword := range keywords { + if token == keyword || strings.Contains(token, keyword) { + return true + } + } + } + return false +} + +func matchesPathBasenamePattern(normalizedTargetPath string, patterns []string) bool { + if normalizedTargetPath == "" || len(patterns) == 0 { + return false + } + base := strings.ToLower(filepath.Base(normalizedTargetPath)) + for _, pattern := range patterns { + matched, err := filepath.Match(pattern, base) + if err != nil { + continue + } + if matched { + return true + } + } + return false +} + +func matchesAnyPattern(value string, patterns []string) bool { + if len(patterns) == 0 { + return true + } + normalized := strings.ToLower(strings.TrimSpace(value)) + for _, pattern := range patterns { + p := strings.ToLower(strings.TrimSpace(pattern)) + if p == "" { + continue + } + matched, err := filepath.Match(p, normalized) + if err == nil && matched { + return true + } + if p == normalized { + return true + } + if strings.HasPrefix(p, "*.") && strings.HasSuffix(normalized, p[1:]) { + return true + } + } + return false +} + +func parseURLHost(raw string) string { + trimmed := strings.TrimSpace(raw) + if trimmed == "" { + return "" + } + parsed, err := url.Parse(trimmed) + if err != nil || parsed == nil { + return "" + } + host := strings.ToLower(strings.TrimSpace(parsed.Hostname())) + return strings.TrimPrefix(host, ".") +} + +func containsString(values []string, target string) bool { + target = strings.ToLower(strings.TrimSpace(target)) + for _, value := range values { + if strings.EqualFold(strings.TrimSpace(value), target) { + return true + } + } + return false +} diff --git a/internal/security/policy_test.go b/internal/security/policy_test.go new file mode 100644 index 00000000..227b6766 --- /dev/null +++ b/internal/security/policy_test.go @@ -0,0 +1,174 @@ +package security + +import ( + "context" + "testing" +) + +func TestPolicyEngineRecommendedRules(t *testing.T) { + t.Parallel() + + engine, err := NewRecommendedPolicyEngine() + if err != nil { + t.Fatalf("new recommended engine: %v", err) + } + + tests := []struct { + name string + action Action + wantDecision Decision + wantRuleID string + }{ + { + name: "bash always ask", + action: Action{ + Type: ActionTypeBash, + Payload: ActionPayload{ + ToolName: "bash", + Resource: "bash", + Operation: "command", + TargetType: TargetTypeCommand, + Target: "ls -la", + }, + }, + wantDecision: DecisionAsk, + wantRuleID: "ask-all-bash", + }, + { + name: "filesystem write ask", + action: Action{ + Type: ActionTypeWrite, + Payload: ActionPayload{ + ToolName: "filesystem_write_file", + Resource: "filesystem_write_file", + Operation: "write_file", + TargetType: TargetTypePath, + Target: "README.md", + }, + }, + wantDecision: DecisionAsk, + wantRuleID: "ask-filesystem-write", + }, + { + name: "filesystem read sensitive path ask", + action: Action{ + Type: ActionTypeRead, + Payload: ActionPayload{ + ToolName: "filesystem_read_file", + Resource: "filesystem_read_file", + Operation: "read_file", + TargetType: TargetTypePath, + Target: ".env.production", + }, + }, + wantDecision: DecisionAsk, + wantRuleID: "ask-sensitive-filesystem-read", + }, + { + name: "filesystem read private key deny", + action: Action{ + Type: ActionTypeRead, + Payload: ActionPayload{ + ToolName: "filesystem_read_file", + Resource: "filesystem_read_file", + Operation: "read_file", + TargetType: TargetTypePath, + Target: "C:/Users/test/.ssh/id_rsa", + }, + }, + wantDecision: DecisionDeny, + wantRuleID: "deny-sensitive-private-keys", + }, + { + name: "filesystem read normal source allow", + action: Action{ + Type: ActionTypeRead, + Payload: ActionPayload{ + ToolName: "filesystem_read_file", + Resource: "filesystem_read_file", + Operation: "read_file", + TargetType: TargetTypePath, + Target: "internal/runtime/runtime.go", + }, + }, + wantDecision: DecisionAllow, + wantRuleID: "", + }, + { + name: "webfetch whitelist allow", + action: Action{ + Type: ActionTypeRead, + Payload: ActionPayload{ + ToolName: "webfetch", + Resource: "webfetch", + Operation: "fetch", + TargetType: TargetTypeURL, + Target: "https://github.com/1024XEngineer/neo-code", + }, + }, + wantDecision: DecisionAllow, + wantRuleID: "allow-webfetch-whitelist", + }, + { + name: "webfetch non-whitelist ask", + action: Action{ + Type: ActionTypeRead, + Payload: ActionPayload{ + ToolName: "webfetch", + Resource: "webfetch", + Operation: "fetch", + TargetType: TargetTypeURL, + Target: "https://example.com", + }, + }, + wantDecision: DecisionAsk, + wantRuleID: "ask-webfetch-non-whitelist", + }, + } + + for _, tt := range tests { + tt := tt + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + result, checkErr := engine.Check(context.Background(), tt.action) + if checkErr != nil { + t.Fatalf("Check() error = %v", checkErr) + } + if result.Decision != tt.wantDecision { + t.Fatalf("expected decision %q, got %q", tt.wantDecision, result.Decision) + } + if tt.wantRuleID == "" { + if result.Rule != nil { + t.Fatalf("expected no matched rule, got %+v", result.Rule) + } + return + } + if result.Rule == nil || result.Rule.ID != tt.wantRuleID { + t.Fatalf("expected rule id %q, got %+v", tt.wantRuleID, result.Rule) + } + }) + } +} + +func TestNewPolicyEngineValidation(t *testing.T) { + t.Parallel() + + _, err := NewPolicyEngine(Decision("invalid"), nil) + if err == nil { + t.Fatalf("expected invalid default decision error") + } + + _, err = NewPolicyEngine(DecisionAllow, []PolicyRule{ + {ID: "", Decision: DecisionAsk}, + }) + if err == nil { + t.Fatalf("expected missing rule id error") + } + + _, err = NewPolicyEngine(DecisionAllow, []PolicyRule{ + {ID: "r1", Decision: Decision("invalid")}, + }) + if err == nil { + t.Fatalf("expected invalid rule decision error") + } +} diff --git a/internal/tools/manager.go b/internal/tools/manager.go index 33735ca4..2829506a 100644 --- a/internal/tools/manager.go +++ b/internal/tools/manager.go @@ -20,6 +20,7 @@ type SpecListInput struct { type Manager interface { ListAvailableSpecs(ctx context.Context, input SpecListInput) ([]provider.ToolSpec, error) Execute(ctx context.Context, input ToolCallInput) (ToolResult, error) + RememberSessionDecision(sessionID string, action security.Action, scope SessionPermissionScope) error } // Executor is the concrete tool execution layer under the manager. diff --git a/internal/tools/manager_test.go b/internal/tools/manager_test.go index 477c7481..bbecc902 100644 --- a/internal/tools/manager_test.go +++ b/internal/tools/manager_test.go @@ -564,7 +564,7 @@ func TestDefaultManagerSessionPermissionMemory(t *testing.T) { } }) - t.Run("action matching remains exact by structured target", func(t *testing.T) { + t.Run("category matching shares decision across same tool category", func(t *testing.T) { t.Parallel() manager, _ := newAskManager(t) inputA := ToolCallInput{ @@ -592,9 +592,87 @@ func TestDefaultManagerSessionPermissionMemory(t *testing.T) { t.Fatalf("expected target A to be allowed, got %v", err) } - _, err = manager.Execute(context.Background(), inputB) + if _, err := manager.Execute(context.Background(), inputB); err != nil { + t.Fatalf("expected target B to inherit same-category allow, got %v", err) + } + }) + + t.Run("filesystem read category applies across file/grep/glob", func(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + readTool := &managerStubTool{name: "filesystem_read_file", content: "ok"} + grepTool := &managerStubTool{name: "filesystem_grep", content: "ok"} + globTool := &managerStubTool{name: "filesystem_glob", content: "ok"} + registry.Register(readTool) + registry.Register(grepTool) + registry.Register(globTool) + + engine, err := security.NewStaticGateway(security.DecisionAllow, []security.Rule{ + { + ID: "ask-filesystem-read", + Type: security.ActionTypeRead, + Resource: "filesystem_read_file", + Decision: security.DecisionAsk, + Reason: "requires approval", + }, + { + ID: "ask-filesystem-grep", + Type: security.ActionTypeRead, + Resource: "filesystem_grep", + Decision: security.DecisionAsk, + Reason: "requires approval", + }, + { + ID: "ask-filesystem-glob", + Type: security.ActionTypeRead, + Resource: "filesystem_glob", + Decision: security.DecisionAsk, + Reason: "requires approval", + }, + }) + if err != nil { + t.Fatalf("new engine: %v", err) + } + manager, err := NewManager(registry, engine, nil) + if err != nil { + t.Fatalf("new manager: %v", err) + } + + sessionID := "session-fs-read" + readInput := ToolCallInput{ + ID: "call-read", + Name: "filesystem_read_file", + Arguments: []byte(`{"path":"README.md"}`), + SessionID: sessionID, + } + grepInput := ToolCallInput{ + ID: "call-grep", + Name: "filesystem_grep", + Arguments: []byte(`{"dir":"internal","pattern":"TODO"}`), + SessionID: sessionID, + } + globInput := ToolCallInput{ + ID: "call-glob", + Name: "filesystem_glob", + Arguments: []byte(`{"dir":"internal","pattern":"*.go"}`), + SessionID: sessionID, + } + + _, err = manager.Execute(context.Background(), readInput) + var permissionErr *PermissionDecisionError if !errors.As(err, &permissionErr) || permissionErr.Decision() != "ask" { - t.Fatalf("expected target B to stay ask, got %v", err) + t.Fatalf("expected initial read ask, got %v", err) + } + if rememberErr := manager.RememberSessionDecision(sessionID, permissionErr.Action(), SessionPermissionScopeAlways); rememberErr != nil { + t.Fatalf("remember filesystem read category: %v", rememberErr) + } + + if _, err := manager.Execute(context.Background(), grepInput); err != nil { + t.Fatalf("expected grep allow via filesystem_read category, got %v", err) + } + if _, err := manager.Execute(context.Background(), globInput); err != nil { + t.Fatalf("expected glob allow via filesystem_read category, got %v", err) } }) } diff --git a/internal/tools/registry.go b/internal/tools/registry.go index c057ee86..b512c323 100644 --- a/internal/tools/registry.go +++ b/internal/tools/registry.go @@ -7,6 +7,7 @@ import ( "strings" "neo-code/internal/provider" + "neo-code/internal/security" ) type Registry struct { @@ -94,3 +95,8 @@ func (r *Registry) Execute(ctx context.Context, input ToolCallInput) (ToolResult } return result, nil } + +// RememberSessionDecision 对纯 Registry 管理器不生效,保留接口以满足 runtime 依赖。 +func (r *Registry) RememberSessionDecision(sessionID string, action security.Action, scope SessionPermissionScope) error { + return errors.New("tools: session permission memory is unsupported by registry manager") +} diff --git a/internal/tools/session_memory.go b/internal/tools/session_memory.go index be8fb5ba..9400acf9 100644 --- a/internal/tools/session_memory.go +++ b/internal/tools/session_memory.go @@ -124,19 +124,41 @@ func (m *sessionPermissionMemory) resolve(sessionID string, action security.Acti // sessionPermissionActionKey 基于结构化 action 生成稳定匹配键。 func sessionPermissionActionKey(action security.Action) string { - normalizedTool := strings.ToLower(strings.TrimSpace(action.Payload.ToolName)) - normalizedResource := strings.ToLower(strings.TrimSpace(action.Payload.Resource)) - normalizedOperation := strings.ToLower(strings.TrimSpace(action.Payload.Operation)) - normalizedTargetType := strings.ToLower(strings.TrimSpace(string(action.Payload.TargetType))) - normalizedTarget := strings.TrimSpace(action.Payload.Target) - normalizedTarget = strings.ReplaceAll(normalizedTarget, "\r\n", "\n") - normalizedTarget = strings.ReplaceAll(normalizedTarget, "\r", "\n") return strings.Join([]string{ string(action.Type), - normalizedTool, - normalizedResource, - normalizedOperation, - normalizedTargetType, - normalizedTarget, + sessionPermissionCategory(action), }, "|") } + +// sessionPermissionCategory 将安全动作归一为稳定的工具类别。 +// 类别用于 once/always/reject 的 session 级记忆,不再按具体 target 区分。 +func sessionPermissionCategory(action security.Action) string { + resource := strings.ToLower(strings.TrimSpace(action.Payload.Resource)) + switch action.Type { + case security.ActionTypeRead: + if strings.HasPrefix(resource, "filesystem_") { + return "filesystem_read" + } + if resource == "webfetch" { + return "webfetch" + } + case security.ActionTypeWrite: + if strings.HasPrefix(resource, "filesystem_") { + return "filesystem_write" + } + case security.ActionTypeBash: + return "bash" + case security.ActionTypeMCP: + target := strings.ToLower(strings.TrimSpace(action.Payload.Target)) + if target != "" { + return "mcp:" + target + } + return "mcp" + } + + toolName := strings.ToLower(strings.TrimSpace(action.Payload.ToolName)) + if toolName != "" { + return toolName + } + return resource +} diff --git a/internal/tui/app.go b/internal/tui/app.go index 2fdfaf81..eae67f5d 100644 --- a/internal/tui/app.go +++ b/internal/tui/app.go @@ -20,43 +20,44 @@ import ( ) type App struct { - state UIState - configManager *config.Manager - providerSvc ProviderController - runtime agentruntime.Runtime - keys keyMap - help help.Model - spinner spinner.Model - sessions list.Model - commandMenu list.Model - commandMenuMeta commandMenuMeta - providerPicker list.Model - modelPicker list.Model - fileBrowser filepicker.Model - progress progress.Model - transcript viewport.Model - activity viewport.Model - input textarea.Model - markdownRenderer markdownContentRenderer - codeCopyBlocks map[int]string - pendingCopyID int - nowFn func() time.Time - lastInputEditAt time.Time - lastPasteLikeAt time.Time - inputBurstStart time.Time - inputBurstCount int - pasteMode bool - activeMessages []provider.Message - activities []activityEntry - fileCandidates []string - modelRefreshID string - focus panel - runProgressValue float64 - runProgressKnown bool - runProgressLabel string - width int - height int - styles styles + state UIState + configManager *config.Manager + providerSvc ProviderController + runtime agentruntime.Runtime + keys keyMap + help help.Model + spinner spinner.Model + sessions list.Model + commandMenu list.Model + commandMenuMeta commandMenuMeta + providerPicker list.Model + modelPicker list.Model + fileBrowser filepicker.Model + progress progress.Model + transcript viewport.Model + activity viewport.Model + input textarea.Model + markdownRenderer markdownContentRenderer + codeCopyBlocks map[int]string + pendingCopyID int + nowFn func() time.Time + lastInputEditAt time.Time + lastPasteLikeAt time.Time + inputBurstStart time.Time + inputBurstCount int + pasteMode bool + pendingPermission *pendingPermissionPrompt + activeMessages []provider.Message + activities []activityEntry + fileCandidates []string + modelRefreshID string + focus panel + runProgressValue float64 + runProgressKnown bool + runProgressLabel string + width int + height int + styles styles } func New(cfg *config.Config, configManager *config.Manager, runtime agentruntime.Runtime, providerSvc ProviderController) (App, error) { diff --git a/internal/tui/state.go b/internal/tui/state.go index a4200e1c..adcafbf7 100644 --- a/internal/tui/state.go +++ b/internal/tui/state.go @@ -58,6 +58,14 @@ type activityEntry struct { IsError bool } +type pendingPermissionPrompt struct { + RequestID string + ToolCallID string + ToolName string + ToolCategory string + Target string +} + type commandMenuMeta struct { Title string } diff --git a/internal/tui/update.go b/internal/tui/update.go index db804480..0fdd6c52 100644 --- a/internal/tui/update.go +++ b/internal/tui/update.go @@ -24,6 +24,11 @@ import ( type RuntimeMsg struct{ Event agentruntime.RuntimeEvent } type RuntimeClosedMsg struct{} type runFinishedMsg struct{ err error } +type permissionResolveResultMsg struct { + requestID string + decision agentruntime.PermissionResolutionDecision + err error +} type modelCatalogRefreshMsg struct { providerID string models []config.ModelDescriptor @@ -91,6 +96,7 @@ func (a App) Update(msg tea.Msg) (tea.Model, tea.Cmd) { case runFinishedMsg: if typed.err != nil { a.state.IsAgentRunning = false + a.pendingPermission = nil a.clearRunProgress() a.state.StreamingReply = false a.state.CurrentTool = "" @@ -108,6 +114,17 @@ func (a App) Update(msg tea.Msg) (tea.Model, tea.Cmd) { _ = a.refreshSessions() a.syncActiveSessionTitle() return a, tea.Batch(cmds...) + case permissionResolveResultMsg: + if typed.err != nil { + a.state.ExecutionError = typed.err.Error() + a.state.StatusText = typed.err.Error() + a.appendActivity("permission", "Submit permission failed", typed.err.Error(), true) + return a, tea.Batch(cmds...) + } + a.state.ExecutionError = "" + a.state.StatusText = "Permission decision submitted" + a.appendActivity("permission", "Submitted permission decision", string(typed.decision), false) + return a, tea.Batch(cmds...) case modelCatalogRefreshMsg: if strings.EqualFold(a.modelRefreshID, typed.providerID) { a.modelRefreshID = "" @@ -236,6 +253,14 @@ func (a App) Update(msg tea.Msg) (tea.Model, tea.Cmd) { a.applyComponentLayout(true) return a, tea.Batch(cmds...) } + if a.pendingPermission != nil { + if permissionCmd, handled := a.handlePermissionDecisionKey(typed); handled { + if permissionCmd != nil { + cmds = append(cmds, permissionCmd) + } + return a, tea.Batch(cmds...) + } + } if a.state.IsAgentRunning && key.Matches(typed, a.keys.CancelAgent) { if a.runtime.CancelActiveRun() { a.state.StatusText = statusCanceling @@ -725,6 +750,7 @@ func (a *App) handleRuntimeEvent(event agentruntime.RuntimeEvent) bool { a.state.IsAgentRunning = false a.state.StreamingReply = false a.state.CurrentTool = "" + a.pendingPermission = nil a.clearRunProgress() if strings.TrimSpace(a.state.ExecutionError) == "" { a.state.StatusText = statusReady @@ -737,6 +763,7 @@ func (a *App) handleRuntimeEvent(event agentruntime.RuntimeEvent) bool { a.state.IsAgentRunning = false a.state.StreamingReply = false a.state.CurrentTool = "" + a.pendingPermission = nil a.state.ExecutionError = "" a.state.StatusText = statusCanceled a.clearRunProgress() @@ -746,6 +773,7 @@ func (a *App) handleRuntimeEvent(event agentruntime.RuntimeEvent) bool { a.state.IsAgentRunning = false a.state.StreamingReply = false a.state.CurrentTool = "" + a.pendingPermission = nil a.clearRunProgress() if payload, ok := event.Payload.(string); ok { a.state.ExecutionError = payload @@ -758,6 +786,47 @@ func (a *App) handleRuntimeEvent(event agentruntime.RuntimeEvent) bool { a.runProgressKnown = false a.appendActivity("provider", "Retrying provider call", payload, false) } + case agentruntime.EventPermissionRequest: + payload, ok := event.Payload.(agentruntime.PermissionRequestPayload) + if !ok { + return transcriptDirty + } + a.pendingPermission = &pendingPermissionPrompt{ + RequestID: strings.TrimSpace(payload.RequestID), + ToolCallID: strings.TrimSpace(payload.ToolCallID), + ToolName: strings.TrimSpace(payload.ToolName), + ToolCategory: strings.TrimSpace(payload.ToolCategory), + Target: strings.TrimSpace(payload.Target), + } + prompt := fmt.Sprintf( + "[Permission] Tool=%s Category=%s Target=%s | y=once, a=always(session), n=reject(session)", + fallback(payload.ToolName, "-"), + fallback(payload.ToolCategory, "-"), + fallback(payload.Target, "-"), + ) + a.state.StatusText = "Permission required (y/a/n)" + a.appendActivity("permission", "Awaiting permission decision", prompt, false) + a.appendInlineMessage(roleSystem, prompt) + transcriptDirty = true + case agentruntime.EventPermissionResolved: + payload, ok := event.Payload.(agentruntime.PermissionResolvedPayload) + if !ok { + return transcriptDirty + } + if a.pendingPermission != nil && strings.TrimSpace(a.pendingPermission.RequestID) != "" && + strings.EqualFold(a.pendingPermission.RequestID, strings.TrimSpace(payload.RequestID)) { + a.pendingPermission = nil + } + resolved := fmt.Sprintf( + "[Permission] %s %s (%s)", + fallback(payload.ToolName, "-"), + fallback(payload.ResolvedAs, "-"), + fallback(payload.RememberScope, "-"), + ) + a.appendActivity("permission", "Permission resolved", resolved, strings.EqualFold(payload.ResolvedAs, "rejected")) + if strings.EqualFold(payload.ResolvedAs, "approved") { + a.state.StatusText = "Permission approved, executing tool..." + } case agentruntime.EventCompactDone: payload, ok := event.Payload.(agentruntime.CompactDonePayload) if !ok { @@ -1396,6 +1465,31 @@ func ListenForRuntimeEvent(sub <-chan agentruntime.RuntimeEvent) tea.Cmd { } } +// handlePermissionDecisionKey 处理待审批状态下的快捷授权键位。 +func (a *App) handlePermissionDecisionKey(msg tea.KeyMsg) (tea.Cmd, bool) { + if a.pendingPermission == nil { + return nil, false + } + + keyText := strings.ToLower(strings.TrimSpace(msg.String())) + var decision agentruntime.PermissionResolutionDecision + switch keyText { + case "y": + decision = agentruntime.PermissionResolutionAllowOnce + case "a": + decision = agentruntime.PermissionResolutionAllowSession + case "n": + decision = agentruntime.PermissionResolutionReject + default: + return nil, false + } + + requestID := strings.TrimSpace(a.pendingPermission.RequestID) + a.state.StatusText = "Submitting permission decision..." + a.state.ExecutionError = "" + return runResolvePermission(a.runtime, requestID, decision), true +} + func runAgent(runtime agentruntime.Runtime, sessionID string, workdir string, content string) tea.Cmd { return func() tea.Msg { err := runtime.Run(context.Background(), agentruntime.UserInput{ @@ -1407,6 +1501,25 @@ func runAgent(runtime agentruntime.Runtime, sessionID string, workdir string, co } } +// runResolvePermission 在独立命令中提交权限审批决定,避免阻塞 UI 事件循环。 +func runResolvePermission( + runtime agentruntime.Runtime, + requestID string, + decision agentruntime.PermissionResolutionDecision, +) tea.Cmd { + return func() tea.Msg { + err := runtime.ResolvePermission(context.Background(), agentruntime.PermissionResolutionInput{ + RequestID: requestID, + Decision: decision, + }) + return permissionResolveResultMsg{ + requestID: requestID, + decision: decision, + err: err, + } + } +} + func runSessionWorkdirCommand( runtime agentruntime.Runtime, sessionID string, diff --git a/internal/tui/update_test.go b/internal/tui/update_test.go index ed307cc0..38993eea 100644 --- a/internal/tui/update_test.go +++ b/internal/tui/update_test.go @@ -38,6 +38,8 @@ type stubRuntime struct { setWorkdirErr error setResult *agentruntime.Session setCalls int + resolveInputs []agentruntime.PermissionResolutionInput + resolveErr error cancelCalls int cancelResult bool } @@ -78,6 +80,11 @@ func (r *stubRuntime) Compact(ctx context.Context, input agentruntime.CompactInp return r.compactResult, r.compactErr } +func (r *stubRuntime) ResolvePermission(ctx context.Context, input agentruntime.PermissionResolutionInput) error { + r.resolveInputs = append(r.resolveInputs, input) + return r.resolveErr +} + func (r *stubRuntime) Events() <-chan agentruntime.RuntimeEvent { return r.events } From bcfcb63183a6d82467c831c71701cdbac0bc3e58 Mon Sep 17 00:00:00 2001 From: Cai_Tang <106404101+Cai-Tang-www@users.noreply.github.com> Date: Sun, 5 Apr 2026 17:22:48 +0800 Subject: [PATCH 13/55] test(runtime,tui): improve permission loop coverage --- internal/runtime/permission_test.go | 285 ++++++++++++++++++++++++++++ internal/tui/update_test.go | 83 ++++++++ 2 files changed, 368 insertions(+) create mode 100644 internal/runtime/permission_test.go diff --git a/internal/runtime/permission_test.go b/internal/runtime/permission_test.go new file mode 100644 index 00000000..403a8732 --- /dev/null +++ b/internal/runtime/permission_test.go @@ -0,0 +1,285 @@ +package runtime + +import ( + "context" + "errors" + "testing" + "time" + + "neo-code/internal/provider" + "neo-code/internal/security" + "neo-code/internal/tools" +) + +func TestResolvePermissionValidation(t *testing.T) { + t.Parallel() + + service := NewWithFactory( + newRuntimeConfigManager(t), + &stubToolManager{}, + newMemoryStore(), + &scriptedProviderFactory{provider: &scriptedProvider{}}, + nil, + ) + + if err := service.ResolvePermission(context.Background(), PermissionResolutionInput{}); err == nil { + t.Fatalf("expected empty request id error") + } + if err := service.ResolvePermission(context.Background(), PermissionResolutionInput{ + RequestID: "perm-1", + Decision: PermissionResolutionDecision("invalid"), + }); err == nil { + t.Fatalf("expected invalid decision error") + } + if err := service.ResolvePermission(context.Background(), PermissionResolutionInput{ + RequestID: "perm-not-found", + Decision: PermissionResolutionAllowOnce, + }); err == nil { + t.Fatalf("expected request not found error") + } +} + +func TestResolvePermissionSuccess(t *testing.T) { + t.Parallel() + + service := NewWithFactory( + newRuntimeConfigManager(t), + &stubToolManager{}, + newMemoryStore(), + &scriptedProviderFactory{provider: &scriptedProvider{}}, + nil, + ) + + request := registerPendingPermission(service, permissionExecutionInput{ + RunID: "run-permission", + SessionID: "session-permission", + Call: provider.ToolCall{ + ID: "call-1", + Name: "webfetch", + }, + }, security.Action{ + Type: security.ActionTypeRead, + Payload: security.ActionPayload{ + ToolName: "webfetch", + Resource: "webfetch", + Operation: "fetch", + TargetType: security.TargetTypeURL, + Target: "https://example.com", + }, + }) + defer clearPendingPermission(service, request.RequestID) + + errCh := make(chan error, 1) + go func() { + errCh <- service.ResolvePermission(context.Background(), PermissionResolutionInput{ + RequestID: request.RequestID, + Decision: PermissionResolutionAllowSession, + }) + }() + + select { + case resolved := <-request.ResultCh: + if resolved != PermissionResolutionAllowSession { + t.Fatalf("expected allow session decision, got %q", resolved) + } + case <-time.After(2 * time.Second): + t.Fatalf("timed out waiting permission resolution") + } + + if err := <-errCh; err != nil { + t.Fatalf("ResolvePermission() error = %v", err) + } +} + +func TestServiceRunPermissionRejectFlow(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + store := newMemoryStore() + registry := tools.NewRegistry() + tool := &stubTool{name: "webfetch", content: "should-not-run"} + registry.Register(tool) + + engine, err := security.NewStaticGateway(security.DecisionAllow, []security.Rule{ + { + ID: "ask-webfetch", + Type: security.ActionTypeRead, + Resource: "webfetch", + Decision: security.DecisionAsk, + Reason: "requires approval", + }, + }) + if err != nil { + t.Fatalf("new static gateway: %v", err) + } + toolManager, err := tools.NewManager(registry, engine, nil) + if err != nil { + t.Fatalf("new tool manager: %v", err) + } + + scripted := &scriptedProvider{ + responses: []provider.ChatResponse{ + { + Message: provider.Message{ + Role: "assistant", + ToolCalls: []provider.ToolCall{ + {ID: "call-ask-reject", Name: "webfetch", Arguments: `{"url":"https://example.com/private"}`}, + }, + }, + FinishReason: "tool_calls", + }, + { + Message: provider.Message{Role: "assistant", Content: "done"}, + FinishReason: "stop", + }, + }, + } + + service := NewWithFactory(manager, toolManager, store, &scriptedProviderFactory{provider: scripted}, nil) + runErrCh := make(chan error, 1) + go func() { + runErrCh <- service.Run(context.Background(), UserInput{RunID: "run-permission-reject", Content: "fetch private"}) + }() + + var requestPayload PermissionRequestPayload +waitRequest: + for { + select { + case <-time.After(3 * time.Second): + t.Fatalf("timed out waiting permission request") + case event := <-service.Events(): + if event.Type != EventPermissionRequest { + continue + } + payload, ok := event.Payload.(PermissionRequestPayload) + if !ok { + t.Fatalf("expected permission request payload, got %#v", event.Payload) + } + requestPayload = payload + break waitRequest + } + } + + if err := service.ResolvePermission(context.Background(), PermissionResolutionInput{ + RequestID: requestPayload.RequestID, + Decision: PermissionResolutionReject, + }); err != nil { + t.Fatalf("ResolvePermission() error = %v", err) + } + if err := <-runErrCh; err != nil { + t.Fatalf("Run() error = %v", err) + } + + if tool.callCount != 0 { + t.Fatalf("expected tool not executed after reject, got %d", tool.callCount) + } + + events := collectRuntimeEvents(service.Events()) + assertEventSequence(t, events, []EventType{EventPermissionResolved, EventToolResult, EventAgentDone}) + + found := false + for _, event := range events { + if event.Type != EventPermissionResolved { + continue + } + payload, ok := event.Payload.(PermissionResolvedPayload) + if !ok { + t.Fatalf("expected permission resolved payload, got %#v", event.Payload) + } + if payload.Decision == "deny" && payload.ResolvedAs == "rejected" { + found = true + } + } + if !found { + t.Fatalf("expected user reject resolved payload") + } +} + +func TestPermissionHelpers(t *testing.T) { + t.Parallel() + + if got := normalizePermissionResolutionDecision(PermissionResolutionDecision("Y")); got != PermissionResolutionAllowOnce { + t.Fatalf("expected Y => allow_once, got %q", got) + } + if got := normalizePermissionResolutionDecision(PermissionResolutionDecision("a")); got != PermissionResolutionAllowSession { + t.Fatalf("expected a => allow_session, got %q", got) + } + if got := normalizePermissionResolutionDecision(PermissionResolutionDecision("n")); got != PermissionResolutionReject { + t.Fatalf("expected n => reject, got %q", got) + } + if got := normalizePermissionResolutionDecision(PermissionResolutionDecision("???")); got != "" { + t.Fatalf("expected unknown => empty, got %q", got) + } + + if scope, err := rememberScopeFromDecision(PermissionResolutionAllowOnce); err != nil || scope != tools.SessionPermissionScopeOnce { + t.Fatalf("expected once scope, got %q / %v", scope, err) + } + if scope, err := rememberScopeFromDecision(PermissionResolutionAllowSession); err != nil || scope != tools.SessionPermissionScopeAlways { + t.Fatalf("expected always scope, got %q / %v", scope, err) + } + if scope, err := rememberScopeFromDecision(PermissionResolutionReject); err != nil || scope != tools.SessionPermissionScopeReject { + t.Fatalf("expected reject scope, got %q / %v", scope, err) + } + if _, err := rememberScopeFromDecision(PermissionResolutionDecision("invalid")); err == nil { + t.Fatalf("expected invalid decision error") + } + + category := permissionToolCategory(security.Action{ + Type: security.ActionTypeRead, + Payload: security.ActionPayload{ + ToolName: "filesystem_grep", + Resource: "filesystem_grep", + }, + }) + if category != "filesystem_read" { + t.Fatalf("expected filesystem_read category, got %q", category) + } + + category = permissionToolCategory(security.Action{ + Type: security.ActionTypeRead, + Payload: security.ActionPayload{ + ToolName: "webfetch", + Resource: "webfetch", + }, + }) + if category != "webfetch" { + t.Fatalf("expected webfetch category, got %q", category) + } +} + +func TestResolvePermissionCanceledContext(t *testing.T) { + t.Parallel() + + service := NewWithFactory( + newRuntimeConfigManager(t), + &stubToolManager{}, + newMemoryStore(), + &scriptedProviderFactory{provider: &scriptedProvider{}}, + nil, + ) + request := registerPendingPermission(service, permissionExecutionInput{ + RunID: "run-canceled", + SessionID: "session-canceled", + Call: provider.ToolCall{ + ID: "call-canceled", + Name: "webfetch", + }, + }, security.Action{ + Type: security.ActionTypeRead, + Payload: security.ActionPayload{ + ToolName: "webfetch", + Resource: "webfetch", + }, + }) + defer clearPendingPermission(service, request.RequestID) + request.ResultCh <- PermissionResolutionAllowOnce + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + if err := service.ResolvePermission(ctx, PermissionResolutionInput{ + RequestID: request.RequestID, + Decision: PermissionResolutionAllowOnce, + }); !errors.Is(err, context.Canceled) { + t.Fatalf("expected context canceled, got %v", err) + } +} diff --git a/internal/tui/update_test.go b/internal/tui/update_test.go index 38993eea..b88b81b9 100644 --- a/internal/tui/update_test.go +++ b/internal/tui/update_test.go @@ -496,6 +496,89 @@ func TestRunAgentWorkdirForwarding(t *testing.T) { }) } +func TestHandlePermissionDecisionKey(t *testing.T) { + t.Parallel() + + manager := newTestConfigManager(t) + runtime := newStubRuntime() + app, err := New(nil, manager, runtime, newTestProviderService(t, manager)) + if err != nil { + t.Fatalf("New() error = %v", err) + } + app.pendingPermission = &pendingPermissionPrompt{ + RequestID: "perm-1", + ToolName: "webfetch", + } + + tests := []struct { + name string + key tea.KeyMsg + wantSent agentruntime.PermissionResolutionDecision + handled bool + }{ + { + name: "y maps to allow once", + key: tea.KeyMsg{Type: tea.KeyRunes, Runes: []rune{'y'}}, + wantSent: agentruntime.PermissionResolutionAllowOnce, + handled: true, + }, + { + name: "a maps to allow session", + key: tea.KeyMsg{Type: tea.KeyRunes, Runes: []rune{'a'}}, + wantSent: agentruntime.PermissionResolutionAllowSession, + handled: true, + }, + { + name: "n maps to reject", + key: tea.KeyMsg{Type: tea.KeyRunes, Runes: []rune{'n'}}, + wantSent: agentruntime.PermissionResolutionReject, + handled: true, + }, + { + name: "other key ignored", + key: tea.KeyMsg{Type: tea.KeyRunes, Runes: []rune{'x'}}, + handled: false, + }, + } + + for _, tt := range tests { + tt := tt + t.Run(tt.name, func(t *testing.T) { + cmd, handled := app.handlePermissionDecisionKey(tt.key) + if handled != tt.handled { + t.Fatalf("expected handled=%v, got %v", tt.handled, handled) + } + if !tt.handled { + if cmd != nil { + t.Fatalf("expected nil cmd for unhandled key") + } + return + } + if cmd == nil { + t.Fatalf("expected resolve cmd") + } + msg := cmd() + result, ok := msg.(permissionResolveResultMsg) + if !ok { + t.Fatalf("expected permissionResolveResultMsg, got %T", msg) + } + if result.err != nil { + t.Fatalf("expected nil resolve error, got %v", result.err) + } + if len(runtime.resolveInputs) == 0 { + t.Fatalf("expected runtime resolve inputs") + } + last := runtime.resolveInputs[len(runtime.resolveInputs)-1] + if last.Decision != tt.wantSent { + t.Fatalf("expected decision %q, got %q", tt.wantSent, last.Decision) + } + if strings.TrimSpace(last.RequestID) != "perm-1" { + t.Fatalf("expected request id perm-1, got %q", last.RequestID) + } + }) + } +} + func TestAppUpdateModelPickerAndRuntimeMessages(t *testing.T) { tests := []struct { name string From d17cd209b86c9bf42ebdd09237d9ab94ac8b254e Mon Sep 17 00:00:00 2001 From: Yumiue <229866007@qq.com> Date: Mon, 6 Apr 2026 11:43:16 +0800 Subject: [PATCH 14/55] =?UTF-8?q?fix:=20=E6=94=B6=E6=95=9B=20micro=20compa?= =?UTF-8?q?ct=20=E5=B7=A5=E5=85=B7=E7=AD=96=E7=95=A5=E6=9D=A5=E6=BA=90?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/context-compact.md | 6 ++ docs/guides/configuration.md | 2 + internal/app/bootstrap.go | 2 +- internal/app/bootstrap_test.go | 3 + internal/context/builder.go | 19 +++++-- internal/context/builder_test.go | 48 ++++++++++++++++ internal/context/microcompact.go | 29 +++++----- internal/context/microcompact_test.go | 63 ++++++++++++++++++--- internal/context/types.go | 6 ++ internal/runtime/runtime.go | 2 +- internal/runtime/runtime_test.go | 75 +++++++++++++++++++++++++ internal/tools/bash/tool.go | 5 ++ internal/tools/filesystem/edit.go | 5 ++ internal/tools/filesystem/glob.go | 5 ++ internal/tools/filesystem/grep.go | 5 ++ internal/tools/filesystem/read_file.go | 5 ++ internal/tools/filesystem/write_file.go | 5 ++ internal/tools/manager.go | 16 ++++++ internal/tools/manager_test.go | 2 + internal/tools/micro_compact_policy.go | 11 ++++ internal/tools/registry.go | 30 +++++++++- internal/tools/registry_test.go | 26 +++++++++ internal/tools/types.go | 1 + internal/tools/webfetch/tool.go | 5 ++ 24 files changed, 345 insertions(+), 31 deletions(-) create mode 100644 internal/tools/micro_compact_policy.go diff --git a/docs/context-compact.md b/docs/context-compact.md index b55c082a..43e43fa1 100644 --- a/docs/context-compact.md +++ b/docs/context-compact.md @@ -19,6 +19,7 @@ context: manual_strategy: keep_recent manual_keep_recent_messages: 10 max_summary_chars: 1200 + micro_compact_disabled: false ``` - `manual_strategy` @@ -27,6 +28,11 @@ context: 在 `keep_recent` 模式下保留最近消息数量,并按 tool call 与 tool result 的原子块整体保留。 - `max_summary_chars` 控制 compact summary 的最大字符数。 +- `micro_compact_disabled` + 控制是否关闭默认启用的读时 micro compact;设为 `true` 时会回退为仅 trim、不清理旧 tool result。 + +新增工具时,micro compact 策略不再由 `context` 层静态白名单维护,而是由 `internal/tools` 中的工具实现声明。 +默认情况下,已注册工具都会参与 micro compact;只有显式声明保留历史结果的工具才会跳过旧结果清理。 ## 执行链路 diff --git a/docs/guides/configuration.md b/docs/guides/configuration.md index 155e29d9..26a5800e 100644 --- a/docs/guides/configuration.md +++ b/docs/guides/configuration.md @@ -327,4 +327,6 @@ context: | `context.compact.max_summary_chars` | int | `1200` | compact summary 最大字符数 | | `context.compact.micro_compact_disabled` | bool | `false` | 是否关闭默认启用的读时 micro compact;设为 `true` 可快速回退到仅 trim、不做旧工具结果清理 | +新增工具默认会参与 micro compact;如果某个工具的历史结果必须保留,需要在 `internal/tools` 的工具实现中显式声明保留策略。 + 更多行为说明见 [context-compact.md](../context-compact.md)。 diff --git a/internal/app/bootstrap.go b/internal/app/bootstrap.go index d87241ea..fa23fed3 100644 --- a/internal/app/bootstrap.go +++ b/internal/app/bootstrap.go @@ -68,7 +68,7 @@ func NewProgram(ctx context.Context) (*tea.Program, error) { toolManager, sessionStore, providerRegistry, - agentcontext.NewBuilder(), + agentcontext.NewBuilderWithToolPolicies(toolRegistry), ) tuiApp, err := tui.New(&cfg, manager, runtimeSvc, providerSelection) diff --git a/internal/app/bootstrap_test.go b/internal/app/bootstrap_test.go index bb204cbd..d3d8e340 100644 --- a/internal/app/bootstrap_test.go +++ b/internal/app/bootstrap_test.go @@ -218,6 +218,9 @@ type stubToolForBootstrap struct { func (s stubToolForBootstrap) Name() string { return s.name } func (s stubToolForBootstrap) Description() string { return "stub" } func (s stubToolForBootstrap) Schema() map[string]any { return map[string]any{"type": "object"} } +func (s stubToolForBootstrap) MicroCompactPolicy() tools.MicroCompactPolicy { + return tools.MicroCompactPolicyCompact +} func (s stubToolForBootstrap) Execute(ctx context.Context, call tools.ToolCallInput) (tools.ToolResult, error) { return tools.ToolResult{Name: s.name, Content: s.content}, nil } diff --git a/internal/context/builder.go b/internal/context/builder.go index eeac9406..db2e77c1 100644 --- a/internal/context/builder.go +++ b/internal/context/builder.go @@ -8,12 +8,18 @@ import ( // DefaultBuilder preserves the current runtime context-building behavior. type DefaultBuilder struct { - promptSources []promptSectionSource - trimPolicy messageTrimPolicy + promptSources []promptSectionSource + trimPolicy messageTrimPolicy + microCompactPolicies MicroCompactPolicySource } // NewBuilder returns the default context builder implementation. func NewBuilder() Builder { + return NewBuilderWithToolPolicies(nil) +} + +// NewBuilderWithToolPolicies 返回带工具 micro compact 策略源的默认上下文构建器。 +func NewBuilderWithToolPolicies(policies MicroCompactPolicySource) Builder { systemSource := &systemStateSource{gitRunner: runGitCommand} return &DefaultBuilder{ promptSources: []promptSectionSource{ @@ -21,7 +27,8 @@ func NewBuilder() Builder { &projectRulesSource{}, systemSource, }, - trimPolicy: spanMessageTrimPolicy{}, + trimPolicy: spanMessageTrimPolicy{}, + microCompactPolicies: policies, } } @@ -47,14 +54,14 @@ func (b *DefaultBuilder) Build(ctx context.Context, input BuildInput) (BuildResu return BuildResult{ SystemPrompt: composeSystemPrompt(sections...), - Messages: applyReadTimeContextProjection(trimPolicy.Trim(input.Messages), input.Compact), + Messages: applyReadTimeContextProjection(trimPolicy.Trim(input.Messages), input.Compact, b.microCompactPolicies), }, nil } // applyReadTimeContextProjection 负责在 provider 请求前按开关应用只读上下文投影,避免改写原始会话消息。 -func applyReadTimeContextProjection(messages []provider.Message, options CompactOptions) []provider.Message { +func applyReadTimeContextProjection(messages []provider.Message, options CompactOptions, policies MicroCompactPolicySource) []provider.Message { if options.DisableMicroCompact { return cloneContextMessages(messages) } - return microCompactMessages(messages) + return microCompactMessagesWithPolicies(messages, policies) } diff --git a/internal/context/builder_test.go b/internal/context/builder_test.go index 8b72951d..8254a25f 100644 --- a/internal/context/builder_test.go +++ b/internal/context/builder_test.go @@ -11,6 +11,7 @@ import ( "neo-code/internal/context/internalcompact" "neo-code/internal/provider" + "neo-code/internal/tools" ) type stubPromptSectionSource struct { @@ -258,6 +259,53 @@ func TestDefaultBuilderBuildSkipsMicroCompactWhenDisabled(t *testing.T) { } } +func TestDefaultBuilderBuildHonorsToolMicroCompactPolicies(t *testing.T) { + t.Parallel() + + builder := &DefaultBuilder{ + promptSources: []promptSectionSource{ + stubPromptSectionSource{sections: []promptSection{{title: "Stub", content: "body"}}}, + }, + microCompactPolicies: stubMicroCompactPolicySource{ + "custom_tool": tools.MicroCompactPolicyPreserveHistory, + }, + } + + messages := []provider.Message{ + {Role: provider.RoleUser, Content: "older user"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-1", Name: "custom_tool", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-1", Content: "old custom result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-2", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-3", Name: "webfetch", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, + {Role: provider.RoleUser, Content: "latest explicit instruction"}, + } + + got, err := builder.Build(stdcontext.Background(), BuildInput{Messages: messages}) + if err != nil { + t.Fatalf("Build() error = %v", err) + } + if got.Messages[2].Content != "old custom result" { + t.Fatalf("expected preserved tool result to remain, got %q", got.Messages[2].Content) + } +} + func TestTrimMessagesPreservesToolPairs(t *testing.T) { t.Parallel() diff --git a/internal/context/microcompact.go b/internal/context/microcompact.go index 10a36cde..e90a5d24 100644 --- a/internal/context/microcompact.go +++ b/internal/context/microcompact.go @@ -15,18 +15,13 @@ const ( microCompactRetainedToolSpans = 2 ) -var microCompactableTools = map[string]struct{}{ - tools.ToolNameBash: {}, - tools.ToolNameWebFetch: {}, - tools.ToolNameFilesystemReadFile: {}, - tools.ToolNameFilesystemGrep: {}, - tools.ToolNameFilesystemGlob: {}, - tools.ToolNameFilesystemEdit: {}, - tools.ToolNameFilesystemWriteFile: {}, -} - // microCompactMessages 对裁剪后的消息做只读投影式微压缩,仅清理旧工具结果内容。 func microCompactMessages(messages []provider.Message) []provider.Message { + return microCompactMessagesWithPolicies(messages, nil) +} + +// microCompactMessagesWithPolicies 按工具策略对裁剪后的消息做只读投影式微压缩。 +func microCompactMessagesWithPolicies(messages []provider.Message, policies MicroCompactPolicySource) []provider.Message { cloned := cloneContextMessages(messages) if len(cloned) == 0 { return cloned @@ -45,7 +40,7 @@ func microCompactMessages(messages []provider.Message) []provider.Message { continue } - compactableIDs := compactableToolCallIDs(cloned[span.Start].ToolCalls) + compactableIDs := compactableToolCallIDs(cloned[span.Start].ToolCalls, policies) if len(compactableIDs) == 0 { continue } @@ -92,7 +87,7 @@ func isToolCallSpan(messages []provider.Message, span internalcompact.MessageSpa } // compactableToolCallIDs 返回 assistant tool call 中可参与微压缩的调用 ID 集合。 -func compactableToolCallIDs(calls []provider.ToolCall) map[string]struct{} { +func compactableToolCallIDs(calls []provider.ToolCall, policies MicroCompactPolicySource) map[string]struct{} { if len(calls) == 0 { return nil } @@ -100,7 +95,7 @@ func compactableToolCallIDs(calls []provider.ToolCall) map[string]struct{} { ids := make(map[string]struct{}, len(calls)) for _, call := range calls { toolName := strings.TrimSpace(call.Name) - if _, ok := microCompactableTools[toolName]; !ok { + if !toolParticipatesInMicroCompact(toolName, policies) { continue } callID := strings.TrimSpace(call.ID) @@ -115,6 +110,14 @@ func compactableToolCallIDs(calls []provider.ToolCall) map[string]struct{} { return ids } +// toolParticipatesInMicroCompact 判断工具是否应参与 micro compact;未知工具默认视为可压缩。 +func toolParticipatesInMicroCompact(toolName string, policies MicroCompactPolicySource) bool { + if policies == nil { + return true + } + return policies.MicroCompactPolicy(toolName) != tools.MicroCompactPolicyPreserveHistory +} + // hasCompactableToolContent 判断工具块中是否存在会影响保留预算的有效工具结果内容。 func hasCompactableToolContent(messages []provider.Message, span internalcompact.MessageSpan, compactableIDs map[string]struct{}) bool { for messageIndex := span.Start + 1; messageIndex < span.End; messageIndex++ { diff --git a/internal/context/microcompact_test.go b/internal/context/microcompact_test.go index 77b7b078..22a95fa4 100644 --- a/internal/context/microcompact_test.go +++ b/internal/context/microcompact_test.go @@ -4,8 +4,18 @@ import ( "testing" "neo-code/internal/provider" + "neo-code/internal/tools" ) +type stubMicroCompactPolicySource map[string]tools.MicroCompactPolicy + +func (s stubMicroCompactPolicySource) MicroCompactPolicy(name string) tools.MicroCompactPolicy { + if policy, ok := s[name]; ok { + return policy + } + return tools.MicroCompactPolicyCompact +} + func TestMicroCompactMessagesClearsOlderCompactableToolResults(t *testing.T) { t.Parallel() @@ -105,7 +115,7 @@ func TestMicroCompactMessagesKeepsProtectedTailUntouched(t *testing.T) { } } -func TestMicroCompactMessagesSkipsNonCompactableErrorsAndOrphans(t *testing.T) { +func TestMicroCompactMessagesKeepsPreservedToolsErrorsAndOrphans(t *testing.T) { t.Parallel() messages := []provider.Message{ @@ -140,9 +150,11 @@ func TestMicroCompactMessagesSkipsNonCompactableErrorsAndOrphans(t *testing.T) { {Role: provider.RoleTool, ToolCallID: "call-4", Content: ""}, } - got := microCompactMessages(messages) + got := microCompactMessagesWithPolicies(messages, stubMicroCompactPolicySource{ + "custom_tool": tools.MicroCompactPolicyPreserveHistory, + }) if got[1].Content != "custom result" { - t.Fatalf("expected non-compactable tool result to remain, got %q", got[1].Content) + t.Fatalf("expected preserved tool result to remain, got %q", got[1].Content) } if got[3].Content != "edit failed" { t.Fatalf("expected error tool result to remain, got %q", got[3].Content) @@ -158,7 +170,7 @@ func TestMicroCompactMessagesSkipsNonCompactableErrorsAndOrphans(t *testing.T) { } } -func TestMicroCompactMessagesClearsOnlyCompactableResultsInMixedToolSpan(t *testing.T) { +func TestMicroCompactMessagesClearsOnlyNonPreservedResultsInMixedToolSpan(t *testing.T) { t.Parallel() messages := []provider.Message{ @@ -190,18 +202,55 @@ func TestMicroCompactMessagesClearsOnlyCompactableResultsInMixedToolSpan(t *test {Role: provider.RoleAssistant, Content: "current reply"}, } - got := microCompactMessages(messages) + got := microCompactMessagesWithPolicies(messages, stubMicroCompactPolicySource{ + "custom_tool": tools.MicroCompactPolicyPreserveHistory, + }) if got[2].Content != microCompactClearedMessage { - t.Fatalf("expected compactable tool result to be cleared, got %q", got[2].Content) + t.Fatalf("expected default compactable tool result to be cleared, got %q", got[2].Content) } if got[3].Content != "custom result" { - t.Fatalf("expected non-compactable tool result in mixed span to remain, got %q", got[3].Content) + t.Fatalf("expected preserved tool result in mixed span to remain, got %q", got[3].Content) } if len(got[1].ToolCalls) != 2 { t.Fatalf("expected assistant tool call metadata to remain intact, got %+v", got[1].ToolCalls) } } +func TestMicroCompactMessagesTreatsNewToolsAsCompactableByDefault(t *testing.T) { + t.Parallel() + + messages := []provider.Message{ + {Role: provider.RoleUser, Content: "older user"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-1", Name: "repo_search", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-1", Content: "old repo search result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-2", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-3", Name: "webfetch", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, + {Role: provider.RoleUser, Content: "latest explicit instruction"}, + } + + got := microCompactMessagesWithPolicies(messages, stubMicroCompactPolicySource{}) + if got[2].Content != microCompactClearedMessage { + t.Fatalf("expected new tool result to be compacted by default, got %q", got[2].Content) + } +} + func TestMicroCompactMessagesSkipsEmptyRecentSpansWhenCountingRetainedBudget(t *testing.T) { t.Parallel() diff --git a/internal/context/types.go b/internal/context/types.go index 2e406861..e1bb0c93 100644 --- a/internal/context/types.go +++ b/internal/context/types.go @@ -4,6 +4,7 @@ import ( "context" "neo-code/internal/provider" + "neo-code/internal/tools" ) // Builder builds the provider-facing context for a single model round. @@ -24,6 +25,11 @@ type BuildResult struct { Messages []provider.Message } +// MicroCompactPolicySource 定义 context 读取工具 micro compact 策略的最小依赖。 +type MicroCompactPolicySource interface { + MicroCompactPolicy(name string) tools.MicroCompactPolicy +} + // CompactOptions controls read-time compact behavior inside the context builder. type CompactOptions struct { DisableMicroCompact bool diff --git a/internal/runtime/runtime.go b/internal/runtime/runtime.go index 1ab3ec9b..55402042 100644 --- a/internal/runtime/runtime.go +++ b/internal/runtime/runtime.go @@ -86,7 +86,7 @@ func NewWithFactory( toolManager = tools.NewRegistry() } if contextBuilder == nil { - contextBuilder = agentcontext.NewBuilder() + contextBuilder = agentcontext.NewBuilderWithToolPolicies(toolManager) } return &Service{ diff --git a/internal/runtime/runtime_test.go b/internal/runtime/runtime_test.go index 9cd8ed79..7b22a68d 100644 --- a/internal/runtime/runtime_test.go +++ b/internal/runtime/runtime_test.go @@ -144,6 +144,7 @@ type stubTool struct { content string isError bool err error + policy tools.MicroCompactPolicy callCount int lastInput tools.ToolCallInput executeFn func(ctx context.Context, input tools.ToolCallInput) (tools.ToolResult, error) @@ -161,6 +162,10 @@ func (t *stubTool) Schema() map[string]any { return map[string]any{"type": "object"} } +func (t *stubTool) MicroCompactPolicy() tools.MicroCompactPolicy { + return t.policy +} + func (t *stubTool) Execute(ctx context.Context, input tools.ToolCallInput) (tools.ToolResult, error) { t.callCount++ t.lastInput = input @@ -202,6 +207,7 @@ type stubToolManager struct { result tools.ToolResult err error listErr error + policies map[string]tools.MicroCompactPolicy listCalls int executeCalls int lastInput tools.ToolCallInput @@ -218,6 +224,13 @@ func (m *stubToolManager) ListAvailableSpecs(ctx context.Context, input tools.Sp return append([]provider.ToolSpec(nil), m.specs...), nil } +func (m *stubToolManager) MicroCompactPolicy(name string) tools.MicroCompactPolicy { + if policy, ok := m.policies[name]; ok { + return policy + } + return tools.MicroCompactPolicyCompact +} + func (m *stubToolManager) Execute(ctx context.Context, input tools.ToolCallInput) (tools.ToolResult, error) { m.executeCalls++ m.lastInput = input @@ -566,6 +579,68 @@ func TestServiceRunPersistsSessionProviderAndModel(t *testing.T) { } } +func TestServiceRunDefaultBuilderUsesToolManagerMicroCompactPolicies(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + store := newMemoryStore() + registry := tools.NewRegistry() + registry.Register(&stubTool{name: "preserve_tool", content: "default", policy: tools.MicroCompactPolicyPreserveHistory}) + registry.Register(&stubTool{name: "bash", content: "default"}) + registry.Register(&stubTool{name: "webfetch", content: "default"}) + + session := newSession("preserve history") + session.ID = "session-preserve-history" + session.Messages = []provider.Message{ + {Role: provider.RoleUser, Content: "older user"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-1", Name: "preserve_tool", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-1", Content: "preserved result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-2", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-3", Name: "webfetch", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, + } + store.sessions[session.ID] = cloneSession(session) + + scripted := &scriptedProvider{ + responses: []provider.ChatResponse{{ + Message: provider.Message{Role: provider.RoleAssistant, Content: "done"}, + FinishReason: "stop", + }}, + } + + service := NewWithFactory(manager, registry, store, &scriptedProviderFactory{provider: scripted}, nil) + if err := service.Run(context.Background(), UserInput{ + SessionID: session.ID, + RunID: "run-preserve-history-policy", + Content: "latest explicit instruction", + }); err != nil { + t.Fatalf("Run() error = %v", err) + } + + if len(scripted.requests) != 1 { + t.Fatalf("expected 1 provider request, got %d", len(scripted.requests)) + } + if got := scripted.requests[0].Messages[2].Content; got != "preserved result" { + t.Fatalf("expected preserved tool result to remain visible, got %q", got) + } +} + func TestServiceRunFailurePreservesExistingSessionProviderAndModel(t *testing.T) { t.Parallel() diff --git a/internal/tools/bash/tool.go b/internal/tools/bash/tool.go index 5c67c719..e02bce21 100644 --- a/internal/tools/bash/tool.go +++ b/internal/tools/bash/tool.go @@ -69,6 +69,11 @@ func (t *Tool) Schema() map[string]any { } } +// MicroCompactPolicy 声明 bash 工具的历史结果默认参与 micro compact 清理。 +func (t *Tool) MicroCompactPolicy() tools.MicroCompactPolicy { + return tools.MicroCompactPolicyCompact +} + func (t *Tool) Execute(ctx context.Context, call tools.ToolCallInput) (tools.ToolResult, error) { var in input if err := json.Unmarshal(call.Arguments, &in); err != nil { diff --git a/internal/tools/filesystem/edit.go b/internal/tools/filesystem/edit.go index 069830f7..631e1082 100644 --- a/internal/tools/filesystem/edit.go +++ b/internal/tools/filesystem/edit.go @@ -55,6 +55,11 @@ func (t *EditTool) Schema() map[string]any { } } +// MicroCompactPolicy 声明编辑工具的历史结果默认参与 micro compact 清理。 +func (t *EditTool) MicroCompactPolicy() tools.MicroCompactPolicy { + return tools.MicroCompactPolicyCompact +} + func (t *EditTool) Execute(ctx context.Context, input tools.ToolCallInput) (tools.ToolResult, error) { var args editInput if err := json.Unmarshal(input.Arguments, &args); err != nil { diff --git a/internal/tools/filesystem/glob.go b/internal/tools/filesystem/glob.go index fb5c1c74..d2858f3d 100644 --- a/internal/tools/filesystem/glob.go +++ b/internal/tools/filesystem/glob.go @@ -51,6 +51,11 @@ func (t *GlobTool) Schema() map[string]any { } } +// MicroCompactPolicy 声明 glob 工具的历史结果默认参与 micro compact 清理。 +func (t *GlobTool) MicroCompactPolicy() tools.MicroCompactPolicy { + return tools.MicroCompactPolicyCompact +} + func (t *GlobTool) Execute(ctx context.Context, input tools.ToolCallInput) (tools.ToolResult, error) { var args globInput if err := json.Unmarshal(input.Arguments, &args); err != nil { diff --git a/internal/tools/filesystem/grep.go b/internal/tools/filesystem/grep.go index 3c993016..5ef5c8fc 100644 --- a/internal/tools/filesystem/grep.go +++ b/internal/tools/filesystem/grep.go @@ -60,6 +60,11 @@ func (t *GrepTool) Schema() map[string]any { } } +// MicroCompactPolicy 声明 grep 工具的历史结果默认参与 micro compact 清理。 +func (t *GrepTool) MicroCompactPolicy() tools.MicroCompactPolicy { + return tools.MicroCompactPolicyCompact +} + func (t *GrepTool) Execute(ctx context.Context, input tools.ToolCallInput) (tools.ToolResult, error) { var args grepInput if err := json.Unmarshal(input.Arguments, &args); err != nil { diff --git a/internal/tools/filesystem/read_file.go b/internal/tools/filesystem/read_file.go index 030b0b4b..7919fdeb 100644 --- a/internal/tools/filesystem/read_file.go +++ b/internal/tools/filesystem/read_file.go @@ -47,6 +47,11 @@ func (t *ReadFileTool) Schema() map[string]any { } } +// MicroCompactPolicy 声明读文件工具的历史结果默认参与 micro compact 清理。 +func (t *ReadFileTool) MicroCompactPolicy() tools.MicroCompactPolicy { + return tools.MicroCompactPolicyCompact +} + func (t *ReadFileTool) Execute(ctx context.Context, input tools.ToolCallInput) (tools.ToolResult, error) { var args readFileInput if err := json.Unmarshal(input.Arguments, &args); err != nil { diff --git a/internal/tools/filesystem/write_file.go b/internal/tools/filesystem/write_file.go index 1b5069f6..f68c6548 100644 --- a/internal/tools/filesystem/write_file.go +++ b/internal/tools/filesystem/write_file.go @@ -50,6 +50,11 @@ func (t *WriteFileTool) Schema() map[string]any { } } +// MicroCompactPolicy 声明写文件工具的历史结果默认参与 micro compact 清理。 +func (t *WriteFileTool) MicroCompactPolicy() tools.MicroCompactPolicy { + return tools.MicroCompactPolicyCompact +} + func (t *WriteFileTool) Execute(ctx context.Context, input tools.ToolCallInput) (tools.ToolResult, error) { var args writeFileInput if err := json.Unmarshal(input.Arguments, &args); err != nil { diff --git a/internal/tools/manager.go b/internal/tools/manager.go index 4c099c4c..8dd7870a 100644 --- a/internal/tools/manager.go +++ b/internal/tools/manager.go @@ -19,6 +19,7 @@ type SpecListInput struct { // Manager is the runtime-facing tool execution and schema exposure boundary. type Manager interface { ListAvailableSpecs(ctx context.Context, input SpecListInput) ([]provider.ToolSpec, error) + MicroCompactPolicy(name string) MicroCompactPolicy Execute(ctx context.Context, input ToolCallInput) (ToolResult, error) } @@ -29,6 +30,10 @@ type Executor interface { Supports(name string) bool } +type microCompactPolicyExecutor interface { + MicroCompactPolicy(name string) MicroCompactPolicy +} + // WorkspaceSandbox enforces workspace-oriented constraints before execution. type WorkspaceSandbox interface { Check(ctx context.Context, action security.Action) (*security.WorkspaceExecutionPlan, error) @@ -151,6 +156,17 @@ func (m *DefaultManager) ListAvailableSpecs(ctx context.Context, input SpecListI return m.executor.ListAvailableSpecs(ctx, input) } +// MicroCompactPolicy 返回工具的 micro compact 策略;无法判断时按默认可压缩处理。 +func (m *DefaultManager) MicroCompactPolicy(name string) MicroCompactPolicy { + if m == nil || m.executor == nil { + return MicroCompactPolicyCompact + } + if source, ok := m.executor.(microCompactPolicyExecutor); ok { + return source.MicroCompactPolicy(name) + } + return MicroCompactPolicyCompact +} + // Execute runs the tool if the permission engine allows it and the sandbox // check passes. func (m *DefaultManager) Execute(ctx context.Context, input ToolCallInput) (ToolResult, error) { diff --git a/internal/tools/manager_test.go b/internal/tools/manager_test.go index a4a2fe98..e2efd02f 100644 --- a/internal/tools/manager_test.go +++ b/internal/tools/manager_test.go @@ -25,6 +25,8 @@ func (t *managerStubTool) Description() string { return "stub tool" } func (t *managerStubTool) Schema() map[string]any { return map[string]any{"type": "object"} } +func (t *managerStubTool) MicroCompactPolicy() MicroCompactPolicy { return MicroCompactPolicyCompact } + func (t *managerStubTool) Execute(ctx context.Context, call ToolCallInput) (ToolResult, error) { t.callCount++ t.lastCall = call diff --git a/internal/tools/micro_compact_policy.go b/internal/tools/micro_compact_policy.go new file mode 100644 index 00000000..9225e9bc --- /dev/null +++ b/internal/tools/micro_compact_policy.go @@ -0,0 +1,11 @@ +package tools + +// MicroCompactPolicy 描述工具历史结果参与 read-time micro compact 的策略。 +type MicroCompactPolicy string + +const ( + // MicroCompactPolicyCompact 表示工具历史结果默认参与 micro compact 清理。 + MicroCompactPolicyCompact MicroCompactPolicy = "" + // MicroCompactPolicyPreserveHistory 表示工具历史结果应显式保留,不参与 micro compact 清理。 + MicroCompactPolicyPreserveHistory MicroCompactPolicy = "preserve_history" +) diff --git a/internal/tools/registry.go b/internal/tools/registry.go index c057ee86..66a9fc33 100644 --- a/internal/tools/registry.go +++ b/internal/tools/registry.go @@ -10,12 +10,14 @@ import ( ) type Registry struct { - tools map[string]Tool + tools map[string]Tool + microCompactPolicies map[string]MicroCompactPolicy } func NewRegistry() *Registry { return &Registry{ - tools: map[string]Tool{}, + tools: map[string]Tool{}, + microCompactPolicies: map[string]MicroCompactPolicy{}, } } @@ -23,7 +25,14 @@ func (r *Registry) Register(tool Tool) { if tool == nil { return } - r.tools[strings.ToLower(tool.Name())] = tool + name := strings.ToLower(tool.Name()) + r.tools[name] = tool + switch tool.MicroCompactPolicy() { + case MicroCompactPolicyPreserveHistory: + r.microCompactPolicies[name] = MicroCompactPolicyPreserveHistory + default: + r.microCompactPolicies[name] = MicroCompactPolicyCompact + } } func (r *Registry) Get(name string) (Tool, error) { @@ -40,6 +49,21 @@ func (r *Registry) Supports(name string) bool { return err == nil } +// MicroCompactPolicy 返回指定工具名的 micro compact 策略;未知工具按默认可压缩处理。 +func (r *Registry) MicroCompactPolicy(name string) MicroCompactPolicy { + if r == nil { + return MicroCompactPolicyCompact + } + policy, ok := r.microCompactPolicies[strings.ToLower(strings.TrimSpace(name))] + if !ok { + return MicroCompactPolicyCompact + } + if policy == MicroCompactPolicyPreserveHistory { + return MicroCompactPolicyPreserveHistory + } + return MicroCompactPolicyCompact +} + func (r *Registry) GetSpecs() []provider.ToolSpec { names := make([]string, 0, len(r.tools)) for name := range r.tools { diff --git a/internal/tools/registry_test.go b/internal/tools/registry_test.go index e8d2db35..ec09852a 100644 --- a/internal/tools/registry_test.go +++ b/internal/tools/registry_test.go @@ -11,6 +11,7 @@ type stubTool struct { name string description string schema map[string]any + policy MicroCompactPolicy result ToolResult err error } @@ -20,6 +21,9 @@ func (s stubTool) Description() string { return s.description } func (s stubTool) Schema() map[string]any { return s.schema } +func (s stubTool) MicroCompactPolicy() MicroCompactPolicy { + return s.policy +} func (s stubTool) Execute(ctx context.Context, call ToolCallInput) (ToolResult, error) { return s.result, s.err } @@ -159,6 +163,12 @@ func TestRegistryHelpers(t *testing.T) { if registry.Supports("missing") { t.Fatalf("did not expect registry to support missing tool") } + if registry.MicroCompactPolicy("a_tool") != MicroCompactPolicyCompact { + t.Fatalf("expected compact policy default for a_tool") + } + if registry.MicroCompactPolicy("missing") != MicroCompactPolicyCompact { + t.Fatalf("expected compact policy default for unknown tool") + } schemas := registry.ListSchemas() if len(schemas) != 1 || schemas[0].Name != "a_tool" { @@ -180,3 +190,19 @@ func TestRegistryHelpers(t *testing.T) { t.Fatalf("expected context canceled, got %v", err) } } + +func TestRegistryMicroCompactPolicyPreserveHistory(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + registry.Register(stubTool{ + name: "custom_tool", + description: "preserve history", + schema: map[string]any{"type": "object"}, + policy: MicroCompactPolicyPreserveHistory, + }) + + if got := registry.MicroCompactPolicy("custom_tool"); got != MicroCompactPolicyPreserveHistory { + t.Fatalf("expected preserve history policy, got %q", got) + } +} diff --git a/internal/tools/types.go b/internal/tools/types.go index 39147ee9..4a0a0a32 100644 --- a/internal/tools/types.go +++ b/internal/tools/types.go @@ -11,6 +11,7 @@ type Tool interface { Name() string Description() string Schema() map[string]any + MicroCompactPolicy() MicroCompactPolicy Execute(ctx context.Context, call ToolCallInput) (ToolResult, error) } diff --git a/internal/tools/webfetch/tool.go b/internal/tools/webfetch/tool.go index 5a77c2b2..83daa3f5 100644 --- a/internal/tools/webfetch/tool.go +++ b/internal/tools/webfetch/tool.go @@ -87,6 +87,11 @@ func (t *Tool) Schema() map[string]any { } } +// MicroCompactPolicy 声明 webfetch 工具的历史结果默认参与 micro compact 清理。 +func (t *Tool) MicroCompactPolicy() tools.MicroCompactPolicy { + return tools.MicroCompactPolicyCompact +} + func (t *Tool) Execute(ctx context.Context, call tools.ToolCallInput) (tools.ToolResult, error) { in, err := decodeInput(call.Arguments) if err != nil { From 19e3b5a2df5ce1f74ea9005a5d7d4882bdf92a2b Mon Sep 17 00:00:00 2001 From: Cai_Tang <106404101+Cai-Tang-www@users.noreply.github.com> Date: Mon, 6 Apr 2026 12:56:03 +0800 Subject: [PATCH 15/55] fix(security): enforce deny precedence and tighten webfetch whitelist --- internal/runtime/runtime_test.go | 44 ++++++ internal/security/policy.go | 4 +- internal/security/policy_test.go | 15 ++ internal/tools/manager.go | 8 +- internal/tools/manager_test.go | 60 ++++++++ internal/tools/registry_test.go | 15 ++ internal/tools/session_memory_test.go | 211 ++++++++++++++++++++++++++ 7 files changed, 354 insertions(+), 3 deletions(-) create mode 100644 internal/tools/session_memory_test.go diff --git a/internal/runtime/runtime_test.go b/internal/runtime/runtime_test.go index 50123578..2f65bb2a 100644 --- a/internal/runtime/runtime_test.go +++ b/internal/runtime/runtime_test.go @@ -2559,3 +2559,47 @@ func TestProviderRetryBackoff(t *testing.T) { }) } } + +func TestPermissionEventViewPayloadMapping(t *testing.T) { + t.Parallel() + + view := permissionEventView{ + toolName: "webfetch", + actionType: string(security.ActionTypeRead), + operation: "fetch", + targetType: string(security.TargetTypeURL), + target: "https://example.com", + decision: "ask", + reason: "need approval", + ruleID: "rule-1", + scope: string(tools.SessionPermissionScopeAlways), + resolvedAs: "rejected", + } + + requestPayload := view.toRequestPayload() + if requestPayload.ToolName != view.toolName || + requestPayload.ActionType != view.actionType || + requestPayload.Operation != view.operation || + requestPayload.TargetType != view.targetType || + requestPayload.Target != view.target || + requestPayload.Decision != view.decision || + requestPayload.Reason != view.reason || + requestPayload.RuleID != view.ruleID || + requestPayload.RememberScope != view.scope { + t.Fatalf("unexpected request payload: %+v", requestPayload) + } + + resolvedPayload := view.toResolvedPayload() + if resolvedPayload.ToolName != view.toolName || + resolvedPayload.ActionType != view.actionType || + resolvedPayload.Operation != view.operation || + resolvedPayload.TargetType != view.targetType || + resolvedPayload.Target != view.target || + resolvedPayload.Decision != view.decision || + resolvedPayload.Reason != view.reason || + resolvedPayload.RuleID != view.ruleID || + resolvedPayload.RememberScope != view.scope || + resolvedPayload.ResolvedAs != view.resolvedAs { + t.Fatalf("unexpected resolved payload: %+v", resolvedPayload) + } +} diff --git a/internal/security/policy.go b/internal/security/policy.go index 0b7060cc..f414655c 100644 --- a/internal/security/policy.go +++ b/internal/security/policy.go @@ -199,7 +199,7 @@ func NewRecommendedPolicyEngine() (*PolicyEngine, error) { Reason: reasonAllowWebfetchDomain, ActionTypes: []ActionType{ActionTypeRead}, ResourcePatterns: []string{"webfetch"}, - HostPatterns: []string{"github.com", "*.github.com", "docs.*"}, + HostPatterns: []string{"github.com", "*.github.com"}, RequireHostMatch: true, }, { @@ -209,7 +209,7 @@ func NewRecommendedPolicyEngine() (*PolicyEngine, error) { Reason: reasonAskWebfetchDomain, ActionTypes: []ActionType{ActionTypeRead}, ResourcePatterns: []string{"webfetch"}, - HostPatterns: []string{"github.com", "*.github.com", "docs.*"}, + HostPatterns: []string{"github.com", "*.github.com"}, RequireHostMissing: true, }, } diff --git a/internal/security/policy_test.go b/internal/security/policy_test.go index 227b6766..ab12d2ee 100644 --- a/internal/security/policy_test.go +++ b/internal/security/policy_test.go @@ -124,6 +124,21 @@ func TestPolicyEngineRecommendedRules(t *testing.T) { wantDecision: DecisionAsk, wantRuleID: "ask-webfetch-non-whitelist", }, + { + name: "webfetch docs wildcard host is not implicitly trusted", + action: Action{ + Type: ActionTypeRead, + Payload: ActionPayload{ + ToolName: "webfetch", + Resource: "webfetch", + Operation: "fetch", + TargetType: TargetTypeURL, + Target: "https://docs.attacker.com", + }, + }, + wantDecision: DecisionAsk, + wantRuleID: "ask-webfetch-non-whitelist", + }, } for _, tt := range tests { diff --git a/internal/tools/manager.go b/internal/tools/manager.go index 2829506a..64ecb01a 100644 --- a/internal/tools/manager.go +++ b/internal/tools/manager.go @@ -187,7 +187,13 @@ func (m *DefaultManager) Execute(ctx context.Context, input ToolCallInput) (Tool result.ToolCallID = input.ID return result, err } - if m.sessionDecisions != nil { + // deny 规则始终优先,避免 session 记忆覆盖硬性安全策略。 + if decision.Decision == security.DecisionDeny { + result := blockedToolResult(input, decision) + return result, permissionErrorFromDecision(decision) + } + // session 记忆仅用于自动处理 ask,不提升原本已 allow 的策略结果。 + if decision.Decision == security.DecisionAsk && m.sessionDecisions != nil { if rememberedDecision, rememberedScope, ok := m.sessionDecisions.resolve(input.SessionID, action); ok { decision = security.CheckResult{ Decision: rememberedDecision, diff --git a/internal/tools/manager_test.go b/internal/tools/manager_test.go index bbecc902..064af6d8 100644 --- a/internal/tools/manager_test.go +++ b/internal/tools/manager_test.go @@ -675,6 +675,66 @@ func TestDefaultManagerSessionPermissionMemory(t *testing.T) { t.Fatalf("expected glob allow via filesystem_read category, got %v", err) } }) + + t.Run("remembered allow does not override hard deny", func(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + readTool := &managerStubTool{name: "filesystem_read_file", content: "ok"} + registry.Register(readTool) + + engine, err := security.NewStaticGateway(security.DecisionAllow, []security.Rule{ + { + ID: "deny-private-key", + Type: security.ActionTypeRead, + Resource: "filesystem_read_file", + Decision: security.DecisionDeny, + Reason: "private key blocked", + }, + }) + if err != nil { + t.Fatalf("new engine: %v", err) + } + manager, err := NewManager(registry, engine, nil) + if err != nil { + t.Fatalf("new manager: %v", err) + } + + sessionID := "session-deny-priority" + action := security.Action{ + Type: security.ActionTypeRead, + Payload: security.ActionPayload{ + ToolName: "filesystem_read_file", + Resource: "filesystem_read_file", + Operation: "read_file", + TargetType: security.TargetTypePath, + Target: "README.md", + }, + } + if err := manager.RememberSessionDecision(sessionID, action, SessionPermissionScopeAlways); err != nil { + t.Fatalf("remember allow: %v", err) + } + + _, execErr := manager.Execute(context.Background(), ToolCallInput{ + ID: "call-deny-priority", + Name: "filesystem_read_file", + Arguments: []byte(`{"path":"C:/Users/test/.ssh/id_rsa"}`), + SessionID: sessionID, + }) + var permissionErr *PermissionDecisionError + if !errors.As(execErr, &permissionErr) { + t.Fatalf("expected permission error, got %v", execErr) + } + if permissionErr.Decision() != "deny" { + t.Fatalf("expected hard deny to win, got %q", permissionErr.Decision()) + } + if permissionErr.RuleID() != "deny-private-key" { + t.Fatalf("expected deny rule id, got %q", permissionErr.RuleID()) + } + if readTool.callCount != 0 { + t.Fatalf("expected blocked call not to execute tool, got %d", readTool.callCount) + } + }) } func TestBuildPermissionAction(t *testing.T) { diff --git a/internal/tools/registry_test.go b/internal/tools/registry_test.go index e8d2db35..c6cd0ad1 100644 --- a/internal/tools/registry_test.go +++ b/internal/tools/registry_test.go @@ -5,6 +5,8 @@ import ( "errors" "strings" "testing" + + "neo-code/internal/security" ) type stubTool struct { @@ -180,3 +182,16 @@ func TestRegistryHelpers(t *testing.T) { t.Fatalf("expected context canceled, got %v", err) } } + +func TestRegistryRememberSessionDecisionUnsupported(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + err := registry.RememberSessionDecision("session-1", security.Action{}, SessionPermissionScopeAlways) + if err == nil { + t.Fatalf("expected unsupported error") + } + if !strings.Contains(err.Error(), "unsupported") { + t.Fatalf("expected unsupported error, got %v", err) + } +} diff --git a/internal/tools/session_memory_test.go b/internal/tools/session_memory_test.go new file mode 100644 index 00000000..9ca0fa0a --- /dev/null +++ b/internal/tools/session_memory_test.go @@ -0,0 +1,211 @@ +package tools + +import ( + "strings" + "testing" + + "neo-code/internal/security" +) + +func TestSessionPermissionMemoryRememberAndResolve(t *testing.T) { + t.Parallel() + + action := security.Action{ + Type: security.ActionTypeRead, + Payload: security.ActionPayload{ + ToolName: "webfetch", + Resource: "webfetch", + }, + } + + t.Run("once decision is consumed after first resolve", func(t *testing.T) { + t.Parallel() + + memory := newSessionPermissionMemory() + if err := memory.remember("session-1", action, SessionPermissionScopeOnce); err != nil { + t.Fatalf("remember() error = %v", err) + } + + decision, scope, ok := memory.resolve("session-1", action) + if !ok || decision != security.DecisionAllow || scope != SessionPermissionScopeOnce { + t.Fatalf("expected once allow decision, got decision=%q scope=%q ok=%v", decision, scope, ok) + } + + _, _, ok = memory.resolve("session-1", action) + if ok { + t.Fatalf("expected once decision to be consumed") + } + }) + + t.Run("always decision keeps applying", func(t *testing.T) { + t.Parallel() + + memory := newSessionPermissionMemory() + if err := memory.remember("session-2", action, SessionPermissionScopeAlways); err != nil { + t.Fatalf("remember() error = %v", err) + } + + for i := 0; i < 2; i++ { + decision, scope, ok := memory.resolve("session-2", action) + if !ok || decision != security.DecisionAllow || scope != SessionPermissionScopeAlways { + t.Fatalf("expected always allow decision, got decision=%q scope=%q ok=%v", decision, scope, ok) + } + } + }) + + t.Run("reject decision keeps applying", func(t *testing.T) { + t.Parallel() + + memory := newSessionPermissionMemory() + if err := memory.remember("session-3", action, SessionPermissionScopeReject); err != nil { + t.Fatalf("remember() error = %v", err) + } + + decision, scope, ok := memory.resolve("session-3", action) + if !ok || decision != security.DecisionDeny || scope != SessionPermissionScopeReject { + t.Fatalf("expected reject deny decision, got decision=%q scope=%q ok=%v", decision, scope, ok) + } + }) +} + +func TestSessionPermissionMemoryValidationAndMisses(t *testing.T) { + t.Parallel() + + memory := newSessionPermissionMemory() + validAction := security.Action{ + Type: security.ActionTypeRead, + Payload: security.ActionPayload{ + ToolName: "webfetch", + Resource: "webfetch", + }, + } + + if err := memory.remember(" ", validAction, SessionPermissionScopeAlways); err == nil { + t.Fatalf("expected empty session id error") + } + + invalidAction := security.Action{ + Type: security.ActionTypeRead, + Payload: security.ActionPayload{ + ToolName: "", + Resource: "webfetch", + }, + } + if err := memory.remember("session", invalidAction, SessionPermissionScopeAlways); err == nil { + t.Fatalf("expected invalid action error") + } + + if err := memory.remember("session", validAction, SessionPermissionScope("bad")); err == nil { + t.Fatalf("expected unsupported scope error") + } + + if _, _, ok := memory.resolve(" ", validAction); ok { + t.Fatalf("expected empty session resolve miss") + } + if _, _, ok := memory.resolve("missing", validAction); ok { + t.Fatalf("expected missing session resolve miss") + } +} + +func TestSessionPermissionCategoryAndActionKey(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + action security.Action + expected string + }{ + { + name: "filesystem read category", + action: security.Action{ + Type: security.ActionTypeRead, + Payload: security.ActionPayload{ + Resource: "filesystem_grep", + }, + }, + expected: "filesystem_read", + }, + { + name: "filesystem write category", + action: security.Action{ + Type: security.ActionTypeWrite, + Payload: security.ActionPayload{ + Resource: "filesystem_edit", + }, + }, + expected: "filesystem_write", + }, + { + name: "webfetch category", + action: security.Action{ + Type: security.ActionTypeRead, + Payload: security.ActionPayload{ + Resource: "webfetch", + }, + }, + expected: "webfetch", + }, + { + name: "bash category", + action: security.Action{ + Type: security.ActionTypeBash, + }, + expected: "bash", + }, + { + name: "mcp with target", + action: security.Action{ + Type: security.ActionTypeMCP, + Payload: security.ActionPayload{ + Target: "Server-A", + }, + }, + expected: "mcp:server-a", + }, + { + name: "mcp without target", + action: security.Action{ + Type: security.ActionTypeMCP, + }, + expected: "mcp", + }, + { + name: "fallback to tool name", + action: security.Action{ + Type: security.ActionTypeRead, + Payload: security.ActionPayload{ + ToolName: "CustomTool", + Resource: "other_resource", + }, + }, + expected: "customtool", + }, + { + name: "fallback to resource", + action: security.Action{ + Type: security.ActionTypeRead, + Payload: security.ActionPayload{ + Resource: "custom_resource", + }, + }, + expected: "custom_resource", + }, + } + + for _, tt := range tests { + tt := tt + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + + got := sessionPermissionCategory(tt.action) + if got != tt.expected { + t.Fatalf("sessionPermissionCategory() = %q, want %q", got, tt.expected) + } + + key := sessionPermissionActionKey(tt.action) + if !strings.HasSuffix(key, "|"+tt.expected) { + t.Fatalf("sessionPermissionActionKey() = %q, expected suffix %q", key, "|"+tt.expected) + } + }) + } +} From 098a26338f0f70328da90aaf886aa6f89c45fc83 Mon Sep 17 00:00:00 2001 From: Yumiue <229866007@qq.com> Date: Mon, 6 Apr 2026 13:45:05 +0800 Subject: [PATCH 16/55] =?UTF-8?q?test:=20=E8=A1=A5=E5=85=85=20micro=20comp?= =?UTF-8?q?act=20=E7=AD=96=E7=95=A5=E5=9B=9E=E5=BD=92=E6=B5=8B=E8=AF=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/context/builder_test.go | 42 +++++++++++++++++++++ internal/runtime/runtime_test.go | 63 ++++++++++++++++++++++++++++++++ internal/tools/manager_test.go | 59 +++++++++++++++++++++++++++++- 3 files changed, 163 insertions(+), 1 deletion(-) diff --git a/internal/context/builder_test.go b/internal/context/builder_test.go index 8254a25f..e76d4389 100644 --- a/internal/context/builder_test.go +++ b/internal/context/builder_test.go @@ -306,6 +306,48 @@ func TestDefaultBuilderBuildHonorsToolMicroCompactPolicies(t *testing.T) { } } +func TestNewBuilderWithToolPoliciesUsesProvidedPolicySource(t *testing.T) { + t.Parallel() + + builder := NewBuilderWithToolPolicies(stubMicroCompactPolicySource{ + "custom_tool": tools.MicroCompactPolicyPreserveHistory, + }) + + messages := []provider.Message{ + {Role: provider.RoleUser, Content: "older user"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-1", Name: "custom_tool", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-1", Content: "old custom result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-2", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-3", Name: "webfetch", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, + {Role: provider.RoleUser, Content: "latest explicit instruction"}, + } + + got, err := builder.Build(stdcontext.Background(), BuildInput{Messages: messages}) + if err != nil { + t.Fatalf("Build() error = %v", err) + } + if got.Messages[2].Content != "old custom result" { + t.Fatalf("expected preserved tool result to remain, got %q", got.Messages[2].Content) + } +} + func TestTrimMessagesPreservesToolPairs(t *testing.T) { t.Parallel() diff --git a/internal/runtime/runtime_test.go b/internal/runtime/runtime_test.go index 7b22a68d..35d782e7 100644 --- a/internal/runtime/runtime_test.go +++ b/internal/runtime/runtime_test.go @@ -641,6 +641,69 @@ func TestServiceRunDefaultBuilderUsesToolManagerMicroCompactPolicies(t *testing. } } +func TestServiceRunDefaultBuilderUsesGenericToolManagerMicroCompactPolicies(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + store := newMemoryStore() + toolManager := &stubToolManager{ + policies: map[string]tools.MicroCompactPolicy{ + "preserve_tool": tools.MicroCompactPolicyPreserveHistory, + }, + } + + session := newSession("preserve history by manager") + session.ID = "session-preserve-history-manager" + session.Messages = []provider.Message{ + {Role: provider.RoleUser, Content: "older user"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-1", Name: "preserve_tool", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-1", Content: "preserved result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-2", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "call-3", Name: "webfetch", Arguments: "{}"}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, + } + store.sessions[session.ID] = cloneSession(session) + + scripted := &scriptedProvider{ + responses: []provider.ChatResponse{{ + Message: provider.Message{Role: provider.RoleAssistant, Content: "done"}, + FinishReason: "stop", + }}, + } + + service := NewWithFactory(manager, toolManager, store, &scriptedProviderFactory{provider: scripted}, nil) + if err := service.Run(context.Background(), UserInput{ + SessionID: session.ID, + RunID: "run-preserve-history-generic-manager", + Content: "latest explicit instruction", + }); err != nil { + t.Fatalf("Run() error = %v", err) + } + + if len(scripted.requests) != 1 { + t.Fatalf("expected 1 provider request, got %d", len(scripted.requests)) + } + if got := scripted.requests[0].Messages[2].Content; got != "preserved result" { + t.Fatalf("expected preserved tool result to remain visible, got %q", got) + } +} + func TestServiceRunFailurePreservesExistingSessionProviderAndModel(t *testing.T) { t.Parallel() diff --git a/internal/tools/manager_test.go b/internal/tools/manager_test.go index e2efd02f..9a1cac9e 100644 --- a/internal/tools/manager_test.go +++ b/internal/tools/manager_test.go @@ -8,6 +8,7 @@ import ( "strings" "testing" + "neo-code/internal/provider" "neo-code/internal/security" ) @@ -15,6 +16,7 @@ type managerStubTool struct { name string content string err error + policy MicroCompactPolicy callCount int lastCall ToolCallInput } @@ -25,7 +27,7 @@ func (t *managerStubTool) Description() string { return "stub tool" } func (t *managerStubTool) Schema() map[string]any { return map[string]any{"type": "object"} } -func (t *managerStubTool) MicroCompactPolicy() MicroCompactPolicy { return MicroCompactPolicyCompact } +func (t *managerStubTool) MicroCompactPolicy() MicroCompactPolicy { return t.policy } func (t *managerStubTool) Execute(ctx context.Context, call ToolCallInput) (ToolResult, error) { t.callCount++ @@ -42,6 +44,21 @@ type stubSandbox struct { lastAction security.Action } +type executorWithoutMicroCompactPolicy struct{} + +func (executorWithoutMicroCompactPolicy) ListAvailableSpecs(ctx context.Context, input SpecListInput) ([]provider.ToolSpec, error) { + if err := ctx.Err(); err != nil { + return nil, err + } + return nil, nil +} + +func (executorWithoutMicroCompactPolicy) Execute(ctx context.Context, call ToolCallInput) (ToolResult, error) { + return ToolResult{}, ctx.Err() +} + +func (executorWithoutMicroCompactPolicy) Supports(name string) bool { return false } + func (s *stubSandbox) Check(ctx context.Context, action security.Action) (*security.WorkspaceExecutionPlan, error) { s.callCount++ s.lastAction = action @@ -70,6 +87,46 @@ func TestDefaultManagerListAvailableSpecs(t *testing.T) { } } +func TestDefaultManagerMicroCompactPolicy(t *testing.T) { + t.Parallel() + + t.Run("nil manager defaults to compact", func(t *testing.T) { + t.Parallel() + + var manager *DefaultManager + if got := manager.MicroCompactPolicy("custom_tool"); got != MicroCompactPolicyCompact { + t.Fatalf("expected compact default, got %q", got) + } + }) + + t.Run("executor without policy support defaults to compact", func(t *testing.T) { + t.Parallel() + + manager, err := NewManager(executorWithoutMicroCompactPolicy{}, nil, nil) + if err != nil { + t.Fatalf("new manager: %v", err) + } + if got := manager.MicroCompactPolicy("custom_tool"); got != MicroCompactPolicyCompact { + t.Fatalf("expected compact default, got %q", got) + } + }) + + t.Run("executor policy is forwarded", func(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + registry.Register(&managerStubTool{name: "preserve_tool", policy: MicroCompactPolicyPreserveHistory}) + + manager, err := NewManager(registry, nil, nil) + if err != nil { + t.Fatalf("new manager: %v", err) + } + if got := manager.MicroCompactPolicy("preserve_tool"); got != MicroCompactPolicyPreserveHistory { + t.Fatalf("expected preserve history, got %q", got) + } + }) +} + func TestDefaultManagerListAvailableSpecsBoundaries(t *testing.T) { t.Parallel() From c25604438189fed694fdaddc4e3976b20447627b Mon Sep 17 00:00:00 2001 From: Yumiue <229866007@qq.com> Date: Mon, 6 Apr 2026 13:58:09 +0800 Subject: [PATCH 17/55] =?UTF-8?q?test:=20=E8=A1=A5=E5=85=85=20micro=20comp?= =?UTF-8?q?act=20=E8=BE=B9=E7=95=8C=E6=B5=8B=E8=AF=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/context/microcompact_test.go | 34 ++++++++++++++++ internal/tools/manager_test.go | 57 ++++++++++++++++++++++++++- internal/tools/registry_test.go | 21 ++++++++++ 3 files changed, 111 insertions(+), 1 deletion(-) diff --git a/internal/context/microcompact_test.go b/internal/context/microcompact_test.go index 22a95fa4..a690b4c6 100644 --- a/internal/context/microcompact_test.go +++ b/internal/context/microcompact_test.go @@ -64,6 +64,27 @@ func TestMicroCompactMessagesClearsOlderCompactableToolResults(t *testing.T) { } } +func TestMicroCompactMessagesHandlesEmptyAndInvalidSpanInputs(t *testing.T) { + t.Parallel() + + if got := microCompactMessages(nil); got != nil { + t.Fatalf("expected nil input to remain nil, got %+v", got) + } + + assistantOnly := []provider.Message{ + { + Role: provider.RoleAssistant, + ToolCalls: []provider.ToolCall{ + {ID: "", Name: "bash", Arguments: "{}"}, + }, + }, + } + got := microCompactMessagesWithPolicies(assistantOnly, stubMicroCompactPolicySource{}) + if len(got) != 1 || len(got[0].ToolCalls) != 1 { + t.Fatalf("expected invalid tool call id path to keep message untouched, got %+v", got) + } +} + func TestMicroCompactMessagesKeepsProtectedTailUntouched(t *testing.T) { t.Parallel() @@ -312,3 +333,16 @@ func TestMicroCompactMessagesSkipsEmptyRecentSpansWhenCountingRetainedBudget(t * t.Fatalf("expected empty recent tool result to remain unchanged, got %q", got[10].Content) } } + +func TestMicroCompactMessagesSkipsToolMessagesWhenCompactableIDsMissing(t *testing.T) { + t.Parallel() + + messages := []provider.Message{ + {Role: provider.RoleTool, ToolCallID: "orphan", Content: "orphan result"}, + } + + got := microCompactMessagesWithPolicies(messages, stubMicroCompactPolicySource{}) + if got[0].Content != "orphan result" { + t.Fatalf("expected orphan tool result to remain, got %q", got[0].Content) + } +} diff --git a/internal/tools/manager_test.go b/internal/tools/manager_test.go index 9a1cac9e..7c51a7c1 100644 --- a/internal/tools/manager_test.go +++ b/internal/tools/manager_test.go @@ -40,6 +40,7 @@ func (t *managerStubTool) Execute(ctx context.Context, call ToolCallInput) (Tool type stubSandbox struct { err error + plan *security.WorkspaceExecutionPlan callCount int lastAction security.Action } @@ -65,7 +66,7 @@ func (s *stubSandbox) Check(ctx context.Context, action security.Action) (*secur if err := ctx.Err(); err != nil { return nil, err } - return nil, s.err + return s.plan, s.err } func TestDefaultManagerListAvailableSpecs(t *testing.T) { @@ -402,6 +403,43 @@ func TestDefaultManagerExecuteWithWorkspaceSandbox(t *testing.T) { } } +func TestDefaultManagerExecuteForwardsWorkspacePlanToTool(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + tool := &managerStubTool{name: "filesystem_write_file", content: "ok"} + registry.Register(tool) + + engine, err := security.NewStaticGateway(security.DecisionAllow, nil) + if err != nil { + t.Fatalf("new engine: %v", err) + } + plan := &security.WorkspaceExecutionPlan{ + Root: "workspace-root", + Target: "workspace-root/notes.txt", + RequestedTarget: "notes.txt", + } + manager, err := NewManager(registry, engine, &stubSandbox{plan: plan}) + if err != nil { + t.Fatalf("new manager: %v", err) + } + + result, execErr := manager.Execute(context.Background(), ToolCallInput{ + Name: "filesystem_write_file", + Arguments: []byte(`{"path":"notes.txt","content":"hello"}`), + Workdir: t.TempDir(), + }) + if execErr != nil { + t.Fatalf("unexpected error: %v", execErr) + } + if result.Content != "ok" { + t.Fatalf("expected ok result, got %+v", result) + } + if tool.lastCall.WorkspacePlan == nil || tool.lastCall.WorkspacePlan.Target != plan.Target { + t.Fatalf("expected workspace plan to be forwarded, got %+v", tool.lastCall.WorkspacePlan) + } +} + func TestPermissionDecisionError(t *testing.T) { t.Parallel() @@ -455,6 +493,23 @@ func TestPermissionDecisionError(t *testing.T) { if nilErr.Error() != "" || nilErr.Decision() != "" || nilErr.ToolName() != "" { t.Fatalf("expected nil permission error helpers to be empty") } + if nilErr.Reason() != "" || nilErr.RuleID() != "" || nilErr.Action() != (security.Action{}) { + t.Fatalf("expected nil permission error extended helpers to be empty") + } + + defaultAsk := &PermissionDecisionError{decision: security.DecisionAsk} + if !strings.Contains(defaultAsk.Error(), "permission approval required") { + t.Fatalf("expected default ask message, got %q", defaultAsk.Error()) + } +} + +func TestNewManagerRejectsNilExecutor(t *testing.T) { + t.Parallel() + + manager, err := NewManager(nil, nil, nil) + if err == nil || !strings.Contains(err.Error(), "executor is nil") { + t.Fatalf("expected nil executor error, got manager=%v err=%v", manager, err) + } } func TestBuildPermissionAction(t *testing.T) { diff --git a/internal/tools/registry_test.go b/internal/tools/registry_test.go index ec09852a..acf66f05 100644 --- a/internal/tools/registry_test.go +++ b/internal/tools/registry_test.go @@ -206,3 +206,24 @@ func TestRegistryMicroCompactPolicyPreserveHistory(t *testing.T) { t.Fatalf("expected preserve history policy, got %q", got) } } + +func TestRegistryMicroCompactPolicyNormalizesNameAndNilRegistry(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + registry.Register(stubTool{ + name: "Custom_Tool", + description: "preserve history", + schema: map[string]any{"type": "object"}, + policy: MicroCompactPolicyPreserveHistory, + }) + + if got := registry.MicroCompactPolicy(" custom_tool "); got != MicroCompactPolicyPreserveHistory { + t.Fatalf("expected normalized preserve history policy, got %q", got) + } + + var nilRegistry *Registry + if got := nilRegistry.MicroCompactPolicy("whatever"); got != MicroCompactPolicyCompact { + t.Fatalf("expected nil registry default compact policy, got %q", got) + } +} From ca06d7889c1aa2487362c6996b45d2cea900e350 Mon Sep 17 00:00:00 2001 From: Cai_Tang <106404101+Cai-Tang-www@users.noreply.github.com> Date: Mon, 6 Apr 2026 17:48:23 +0800 Subject: [PATCH 18/55] fix(permission): harden approval flow and narrow session memory scope --- internal/runtime/permission.go | 44 ++++++++++------ internal/runtime/permission_test.go | 52 ++++++++++++++++++ internal/runtime/runtime_test.go | 3 ++ internal/tools/manager_test.go | 2 +- internal/tools/session_memory.go | 47 ++++++++++++++++- internal/tools/session_memory_test.go | 76 ++++++++++++++++++++++++++- internal/tui/state.go | 1 + internal/tui/update.go | 8 +++ internal/tui/update_test.go | 42 +++++++++++++++ 9 files changed, 255 insertions(+), 20 deletions(-) diff --git a/internal/runtime/permission.go b/internal/runtime/permission.go index 52f59908..6782e9be 100644 --- a/internal/runtime/permission.go +++ b/internal/runtime/permission.go @@ -43,6 +43,7 @@ type pendingPermissionRequest struct { Call provider.ToolCall Action security.Action ResultCh chan PermissionResolutionDecision + Submitted bool } var runtimePendingPermissions = struct { @@ -63,19 +64,31 @@ func (s *Service) ResolvePermission(ctx context.Context, input PermissionResolut if decision == "" { return fmt.Errorf("runtime: unsupported permission decision %q", input.Decision) } + if err := ctx.Err(); err != nil { + return err + } runtimePendingPermissions.mu.Lock() pending := runtimePendingPermissions.byRun[s] - runtimePendingPermissions.mu.Unlock() if pending == nil || pending.RequestID != requestID { + runtimePendingPermissions.mu.Unlock() return fmt.Errorf("runtime: permission request %q not found", requestID) } + // Submitted 标记用于避免重复提交同一个 request 导致阻塞。 + if pending.Submitted { + runtimePendingPermissions.mu.Unlock() + return nil + } + pending.Submitted = true + resultCh := pending.ResultCh + runtimePendingPermissions.mu.Unlock() + // 非阻塞提交,避免 UI 重复触发导致写满 channel 后长时间卡住。 select { - case pending.ResultCh <- decision: + case resultCh <- decision: + return nil + default: return nil - case <-ctx.Done(): - return ctx.Err() } } @@ -176,18 +189,17 @@ func (s *Service) awaitPermissionDecision( defer clearPendingPermission(s, request.RequestID) s.emit(ctx, EventPermissionRequest, input.RunID, input.SessionID, PermissionRequestPayload{ - RequestID: request.RequestID, - ToolCallID: input.Call.ID, - ToolName: input.Call.Name, - ToolCategory: permissionToolCategory(permissionErr.Action()), - ActionType: string(permissionErr.Action().Type), - Operation: permissionErr.Action().Payload.Operation, - TargetType: string(permissionErr.Action().Payload.TargetType), - Target: permissionErr.Action().Payload.Target, - Decision: permissionErr.Decision(), - Reason: permissionErr.Reason(), - RuleID: permissionErr.RuleID(), - RememberScope: string(tools.SessionPermissionScopeAlways), + RequestID: request.RequestID, + ToolCallID: input.Call.ID, + ToolName: input.Call.Name, + ToolCategory: permissionToolCategory(permissionErr.Action()), + ActionType: string(permissionErr.Action().Type), + Operation: permissionErr.Action().Payload.Operation, + TargetType: string(permissionErr.Action().Payload.TargetType), + Target: permissionErr.Action().Payload.Target, + Decision: permissionErr.Decision(), + Reason: permissionErr.Reason(), + RuleID: permissionErr.RuleID(), }) select { diff --git a/internal/runtime/permission_test.go b/internal/runtime/permission_test.go index 403a8732..a32ce0fd 100644 --- a/internal/runtime/permission_test.go +++ b/internal/runtime/permission_test.go @@ -91,6 +91,58 @@ func TestResolvePermissionSuccess(t *testing.T) { } } +func TestResolvePermissionDuplicateSubmissionIsNonBlocking(t *testing.T) { + t.Parallel() + + service := NewWithFactory( + newRuntimeConfigManager(t), + &stubToolManager{}, + newMemoryStore(), + &scriptedProviderFactory{provider: &scriptedProvider{}}, + nil, + ) + + request := registerPendingPermission(service, permissionExecutionInput{ + RunID: "run-permission-dup", + SessionID: "session-permission-dup", + Call: provider.ToolCall{ + ID: "call-dup", + Name: "webfetch", + }, + }, security.Action{ + Type: security.ActionTypeRead, + Payload: security.ActionPayload{ + ToolName: "webfetch", + Resource: "webfetch", + }, + }) + defer clearPendingPermission(service, request.RequestID) + + if err := service.ResolvePermission(context.Background(), PermissionResolutionInput{ + RequestID: request.RequestID, + Decision: PermissionResolutionAllowOnce, + }); err != nil { + t.Fatalf("first ResolvePermission() error = %v", err) + } + + secondDone := make(chan error, 1) + go func() { + secondDone <- service.ResolvePermission(context.Background(), PermissionResolutionInput{ + RequestID: request.RequestID, + Decision: PermissionResolutionAllowSession, + }) + }() + + select { + case err := <-secondDone: + if err != nil { + t.Fatalf("second ResolvePermission() error = %v", err) + } + case <-time.After(500 * time.Millisecond): + t.Fatalf("second ResolvePermission() should not block") + } +} + func TestServiceRunPermissionRejectFlow(t *testing.T) { t.Parallel() diff --git a/internal/runtime/runtime_test.go b/internal/runtime/runtime_test.go index 2f65bb2a..33e3590b 100644 --- a/internal/runtime/runtime_test.go +++ b/internal/runtime/runtime_test.go @@ -743,6 +743,9 @@ waitRequest: if strings.TrimSpace(requestPayload.RequestID) == "" { t.Fatalf("expected non-empty permission request id") } + if strings.TrimSpace(requestPayload.RememberScope) != "" { + t.Fatalf("expected empty remember scope for permission_request, got %q", requestPayload.RememberScope) + } if requestPayload.ToolName != "webfetch" || requestPayload.Decision != "ask" { t.Fatalf("unexpected permission request payload: %+v", requestPayload) } diff --git a/internal/tools/manager_test.go b/internal/tools/manager_test.go index 064af6d8..1292de68 100644 --- a/internal/tools/manager_test.go +++ b/internal/tools/manager_test.go @@ -643,7 +643,7 @@ func TestDefaultManagerSessionPermissionMemory(t *testing.T) { readInput := ToolCallInput{ ID: "call-read", Name: "filesystem_read_file", - Arguments: []byte(`{"path":"README.md"}`), + Arguments: []byte(`{"path":"internal/README.md"}`), SessionID: sessionID, } grepInput := ToolCallInput{ diff --git a/internal/tools/session_memory.go b/internal/tools/session_memory.go index 9400acf9..91006a14 100644 --- a/internal/tools/session_memory.go +++ b/internal/tools/session_memory.go @@ -3,6 +3,8 @@ package tools import ( "errors" "fmt" + "net/url" + "path/filepath" "strings" "sync" @@ -127,11 +129,12 @@ func sessionPermissionActionKey(action security.Action) string { return strings.Join([]string{ string(action.Type), sessionPermissionCategory(action), + sessionPermissionTargetScope(action), }, "|") } // sessionPermissionCategory 将安全动作归一为稳定的工具类别。 -// 类别用于 once/always/reject 的 session 级记忆,不再按具体 target 区分。 +// 类别用于聚合同类工具,再配合 target scope 控制最小授权范围。 func sessionPermissionCategory(action security.Action) string { resource := strings.ToLower(strings.TrimSpace(action.Payload.Resource)) switch action.Type { @@ -162,3 +165,45 @@ func sessionPermissionCategory(action security.Action) string { } return resource } + +// sessionPermissionTargetScope 基于 action 的 target 生成最小授权范围键。 +func sessionPermissionTargetScope(action security.Action) string { + target := strings.TrimSpace(action.Payload.Target) + if target == "" { + return "*" + } + + switch action.Payload.TargetType { + case security.TargetTypeURL: + return normalizePermissionURLTarget(target) + case security.TargetTypePath: + return normalizePermissionPathTarget(filepath.Dir(target)) + case security.TargetTypeDirectory: + return normalizePermissionPathTarget(target) + default: + return strings.ToLower(target) + } +} + +// normalizePermissionURLTarget 将 URL 归一到 host[:port] 维度。 +func normalizePermissionURLTarget(raw string) string { + parsed, err := url.Parse(strings.TrimSpace(raw)) + if err != nil || strings.TrimSpace(parsed.Host) == "" { + return strings.ToLower(strings.TrimSpace(raw)) + } + + host := strings.ToLower(strings.TrimSpace(parsed.Hostname())) + if port := strings.TrimSpace(parsed.Port()); port != "" { + host += ":" + port + } + return host +} + +// normalizePermissionPathTarget 统一路径分隔并按平台无关形式生成匹配键。 +func normalizePermissionPathTarget(raw string) string { + cleaned := filepath.Clean(strings.TrimSpace(raw)) + if cleaned == "." || cleaned == "" { + return "." + } + return strings.ToLower(filepath.ToSlash(cleaned)) +} diff --git a/internal/tools/session_memory_test.go b/internal/tools/session_memory_test.go index 9ca0fa0a..194b316e 100644 --- a/internal/tools/session_memory_test.go +++ b/internal/tools/session_memory_test.go @@ -203,9 +203,81 @@ func TestSessionPermissionCategoryAndActionKey(t *testing.T) { } key := sessionPermissionActionKey(tt.action) - if !strings.HasSuffix(key, "|"+tt.expected) { - t.Fatalf("sessionPermissionActionKey() = %q, expected suffix %q", key, "|"+tt.expected) + if !strings.Contains(key, "|"+tt.expected+"|") { + t.Fatalf("sessionPermissionActionKey() = %q, expected category token %q", key, "|"+tt.expected+"|") } }) } } + +func TestSessionPermissionMemoryResolveRequiresTargetScopeMatch(t *testing.T) { + t.Parallel() + + memory := newSessionPermissionMemory() + sessionID := "session-target-scope" + + webAction := security.Action{ + Type: security.ActionTypeRead, + Payload: security.ActionPayload{ + ToolName: "webfetch", + Resource: "webfetch", + TargetType: security.TargetTypeURL, + Target: "https://docs.github.com/en/rest", + }, + } + if err := memory.remember(sessionID, webAction, SessionPermissionScopeAlways); err != nil { + t.Fatalf("remember web action: %v", err) + } + + sameHost := security.Action{ + Type: security.ActionTypeRead, + Payload: security.ActionPayload{ + ToolName: "webfetch", + Resource: "webfetch", + TargetType: security.TargetTypeURL, + Target: "https://docs.github.com/en/actions", + }, + } + if _, _, ok := memory.resolve(sessionID, sameHost); !ok { + t.Fatalf("expected same host/path scope web action to hit memory") + } + + differentHost := security.Action{ + Type: security.ActionTypeRead, + Payload: security.ActionPayload{ + ToolName: "webfetch", + Resource: "webfetch", + TargetType: security.TargetTypeURL, + Target: "https://example.com/en/actions", + }, + } + if _, _, ok := memory.resolve(sessionID, differentHost); ok { + t.Fatalf("expected different host web action to miss memory") + } + + fileAction := security.Action{ + Type: security.ActionTypeRead, + Payload: security.ActionPayload{ + ToolName: "filesystem_read_file", + Resource: "filesystem_read_file", + TargetType: security.TargetTypePath, + Target: "src/main.go", + }, + } + if err := memory.remember(sessionID, fileAction, SessionPermissionScopeAlways); err != nil { + t.Fatalf("remember file action: %v", err) + } + + otherFile := security.Action{ + Type: security.ActionTypeRead, + Payload: security.ActionPayload{ + ToolName: "filesystem_read_file", + Resource: "filesystem_read_file", + TargetType: security.TargetTypePath, + Target: "secrets/secret.key", + }, + } + if _, _, ok := memory.resolve(sessionID, otherFile); ok { + t.Fatalf("expected different path file action to miss memory") + } +} diff --git a/internal/tui/state.go b/internal/tui/state.go index adcafbf7..1f939dad 100644 --- a/internal/tui/state.go +++ b/internal/tui/state.go @@ -64,6 +64,7 @@ type pendingPermissionPrompt struct { ToolName string ToolCategory string Target string + Submitted bool } type commandMenuMeta struct { diff --git a/internal/tui/update.go b/internal/tui/update.go index 0fdd6c52..4836fd28 100644 --- a/internal/tui/update.go +++ b/internal/tui/update.go @@ -116,6 +116,9 @@ func (a App) Update(msg tea.Msg) (tea.Model, tea.Cmd) { return a, tea.Batch(cmds...) case permissionResolveResultMsg: if typed.err != nil { + if a.pendingPermission != nil && strings.EqualFold(strings.TrimSpace(a.pendingPermission.RequestID), strings.TrimSpace(typed.requestID)) { + a.pendingPermission.Submitted = false + } a.state.ExecutionError = typed.err.Error() a.state.StatusText = typed.err.Error() a.appendActivity("permission", "Submit permission failed", typed.err.Error(), true) @@ -1483,8 +1486,13 @@ func (a *App) handlePermissionDecisionKey(msg tea.KeyMsg) (tea.Cmd, bool) { default: return nil, false } + if a.pendingPermission.Submitted { + a.state.StatusText = "Permission decision already submitted, waiting runtime..." + return nil, true + } requestID := strings.TrimSpace(a.pendingPermission.RequestID) + a.pendingPermission.Submitted = true a.state.StatusText = "Submitting permission decision..." a.state.ExecutionError = "" return runResolvePermission(a.runtime, requestID, decision), true diff --git a/internal/tui/update_test.go b/internal/tui/update_test.go index b88b81b9..85f6b1be 100644 --- a/internal/tui/update_test.go +++ b/internal/tui/update_test.go @@ -544,6 +544,9 @@ func TestHandlePermissionDecisionKey(t *testing.T) { for _, tt := range tests { tt := tt t.Run(tt.name, func(t *testing.T) { + app.pendingPermission.Submitted = false + runtime.resolveInputs = nil + cmd, handled := app.handlePermissionDecisionKey(tt.key) if handled != tt.handled { t.Fatalf("expected handled=%v, got %v", tt.handled, handled) @@ -579,6 +582,45 @@ func TestHandlePermissionDecisionKey(t *testing.T) { } } +func TestHandlePermissionDecisionKeyIgnoresRepeatedSubmission(t *testing.T) { + t.Parallel() + + manager := newTestConfigManager(t) + runtime := newStubRuntime() + app, err := New(nil, manager, runtime, newTestProviderService(t, manager)) + if err != nil { + t.Fatalf("New() error = %v", err) + } + app.pendingPermission = &pendingPermissionPrompt{ + RequestID: "perm-repeat", + ToolName: "webfetch", + } + + cmd, handled := app.handlePermissionDecisionKey(tea.KeyMsg{Type: tea.KeyRunes, Runes: []rune{'y'}}) + if !handled || cmd == nil { + t.Fatalf("expected first permission key to be handled with cmd") + } + msg := cmd() + result, ok := msg.(permissionResolveResultMsg) + if !ok || result.err != nil { + t.Fatalf("expected successful permissionResolveResultMsg, got %#v", msg) + } + if len(runtime.resolveInputs) != 1 { + t.Fatalf("expected one resolve call after first submission, got %d", len(runtime.resolveInputs)) + } + + cmd, handled = app.handlePermissionDecisionKey(tea.KeyMsg{Type: tea.KeyRunes, Runes: []rune{'a'}}) + if !handled { + t.Fatalf("expected repeated permission key to be consumed") + } + if cmd != nil { + t.Fatalf("expected repeated submission to skip runtime command") + } + if len(runtime.resolveInputs) != 1 { + t.Fatalf("expected resolve call count unchanged after repeat, got %d", len(runtime.resolveInputs)) + } +} + func TestAppUpdateModelPickerAndRuntimeMessages(t *testing.T) { tests := []struct { name string From ab510d7d6ff0bc3e1e541e1e567aea962a12f5fc Mon Sep 17 00:00:00 2001 From: phantom5099 <1011668688@qq.com> Date: Mon, 6 Apr 2026 18:56:37 +0800 Subject: [PATCH 19/55] =?UTF-8?q?feat(provider):=E5=A2=9E=E5=8A=A0?= =?UTF-8?q?=E9=80=8F=E6=98=8E=E9=87=8D=E4=BC=A0=E4=BB=A5=E5=8F=8A=E8=A7=A3?= =?UTF-8?q?=E5=86=B3=E7=BC=93=E5=86=B2=E5=8C=BA=E6=BA=A2=E5=87=BA=E9=97=AE?= =?UTF-8?q?=E9=A2=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/provider/errors.go | 47 +++++++++ internal/provider/openai/openai.go | 129 ++++++++++++++++++++---- internal/provider/openai/openai_test.go | 2 +- internal/provider/openai/sse_reader.go | 68 +++++++++++++ 4 files changed, 226 insertions(+), 20 deletions(-) create mode 100644 internal/provider/openai/sse_reader.go diff --git a/internal/provider/errors.go b/internal/provider/errors.go index 49936b9e..29697cca 100644 --- a/internal/provider/errors.go +++ b/internal/provider/errors.go @@ -1,8 +1,10 @@ package provider import ( + "context" "errors" "fmt" + "net" "net/http" ) @@ -10,6 +12,11 @@ import ( var ( ErrDriverNotFound = errors.New("provider driver not found") ErrDriverAlreadyRegistered = errors.New("provider: driver already registered") + + // 流级哨兵错误,用于区分可恢复/不可恢复的流中断原因。 + ErrStreamInterrupted = errors.New("provider: stream interrupted") + ErrLineTooLong = errors.New("provider: SSE line exceeds max length") + ErrStreamTooLarge = errors.New("provider: stream total size exceeds limit") ) type ProviderErrorCode string @@ -100,3 +107,43 @@ func NewTimeoutProviderError(message string) *ProviderError { Retryable: true, // 超时默认可重试 } } + +// IsRecoverableStreamError 判断流读取错误是否可通过透明重连恢复。 +// +// 不可恢复的情况: +// - context 取消/超时(调用方主动终止) +// - 缓冲区溢出(重连只会再次溢出) +// - 认证失败等业务错误(重连无意义) +// +// 可恢复的情况: +// - ProviderError 且 Retryable=true(5xx、429 等) +// - 网络层临时错误(*net.OpError) +// - ErrStreamInterrupted(通用流中断标记) +func IsRecoverableStreamError(err error) bool { + if err == nil { + return false + } + // context 取消 → 不可恢复 + if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + return false + } + // 缓冲区溢出 → 不可恢复(重连同样会溢出) + if errors.Is(err, ErrLineTooLong) || errors.Is(err, ErrStreamTooLarge) { + return false + } + // 流中断标记 → 可恢复 + if errors.Is(err, ErrStreamInterrupted) { + return true + } + // ProviderError → 依据 Retryable 字段 + var pErr *ProviderError + if errors.As(err, &pErr) { + return pErr.Retryable + } + // 网络层临时故障(连接重置、超时等)→ 可恢复 + var netErr *net.OpError + if errors.As(err, &netErr) { + return true + } + return false +} diff --git a/internal/provider/openai/openai.go b/internal/provider/openai/openai.go index f688943d..a7b1c7ac 100644 --- a/internal/provider/openai/openai.go +++ b/internal/provider/openai/openai.go @@ -1,7 +1,6 @@ package openai import ( - "bufio" "bytes" "context" "encoding/json" @@ -102,7 +101,54 @@ func (p *Provider) DiscoverModels(ctx context.Context) ([]config.ModelDescriptor return config.MergeModelDescriptors(descriptors), nil } +// Chat 发起 SSE 流式对话请求,支持透明重连。 +// +// 流中途断连时,将已累积的 assistant 消息(文本 + tool call)注入请求上下文, +// 利用 OpenAI 多轮对话语义实现断点续传,对上层调用方透明。 +// 最多重连 maxReconnects 次;不可恢复错误直接返回。 func (p *Provider) Chat(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { + const maxReconnects = 3 + + // 跨重连周期持久化的累积状态:已收到的文本和 tool call + var ( + accumText strings.Builder + accumCalls map[int]*provider.ToolCall + ) + + for attempt := 0; attempt <= maxReconnects; attempt++ { + if attempt > 0 { + // 将已累积内容作为 assistant 消息注入,使新请求能从断点继续 + req.Messages = append(req.Messages, + p.buildAssistantMsg(&accumText, accumCalls)) + // 指数退避等待 + backoff := time.Duration(1< 0 { + calls := make([]provider.ToolCall, 0, len(accumCalls)) + for _, c := range accumCalls { + calls = append(calls, *c) + } + msg.ToolCalls = calls + } + return msg +} + +// mergeToolCallDeltaWithAccum 在 mergeToolCallDelta 的基础上, +// 同步将 tool call 累积状态写入跨周期的 accumCalls(*map[int]*ToolCall)。 +func mergeToolCallDeltaWithAccum( + ctx context.Context, + events chan<- provider.StreamEvent, + accumCalls *map[int]*provider.ToolCall, + delta toolCallDelta, +) error { + if *accumCalls == nil { + *accumCalls = make(map[int]*provider.ToolCall) + } + + // 先确保 accumCalls 中有对应条目 + call, exists := (*accumCalls)[delta.Index] + if !exists { + call = &provider.ToolCall{} + (*accumCalls)[delta.Index] = call + } + + // 复用原有逻辑处理事件发送和局部累积 + return mergeToolCallDelta(ctx, events, *accumCalls, delta) +} + func emitStreamEvent(ctx context.Context, events chan<- provider.StreamEvent, event provider.StreamEvent) error { if events == nil { return nil diff --git a/internal/provider/openai/openai_test.go b/internal/provider/openai/openai_test.go index 29e99654..60c10ee6 100644 --- a/internal/provider/openai/openai_test.go +++ b/internal/provider/openai/openai_test.go @@ -614,7 +614,7 @@ func TestProviderConsumeStreamRejectsDirtyJSON(t *testing.T) { t.Fatalf("New() error = %v", err) } - err = provider.consumeStream(context.Background(), strings.NewReader("data: {not-json}\n\n"), make(chan domain.StreamEvent, 1)) + err = provider.consumeStream(context.Background(), strings.NewReader("data: {not-json}\n\n"), make(chan domain.StreamEvent, 1), &strings.Builder{}, new(map[int]*domain.ToolCall)) if err == nil || !strings.Contains(err.Error(), "decode stream chunk") { t.Fatalf("expected dirty JSON decode error, got %v", err) } diff --git a/internal/provider/openai/sse_reader.go b/internal/provider/openai/sse_reader.go new file mode 100644 index 00000000..e1c8330a --- /dev/null +++ b/internal/provider/openai/sse_reader.go @@ -0,0 +1,68 @@ +package openai + +import ( + "bufio" + "errors" + "io" + + "neo-code/internal/provider" +) + +// 单行与总量上限,防止恶意或异常数据导致内存无限增长。 +const ( + maxSSELineSize = 256 * 1024 // L1: 单行 256KB + maxStreamTotalSize = 10 << 20 // L3: 总量 10MB +) + +// boundedSSEReader 对 bufio.Reader 包装两级有界检查: +// - L1: 每次读取的行不超过 maxSSELineSize +// - L3: 累计读取字节数不超过 maxStreamTotalSize +// +// 纯同步设计,无 goroutine/channel,适用于 SSE 顺序消费场景。 +type boundedSSEReader struct { + reader *bufio.Reader + totalRead int64 +} + +// newBoundedSSEReader 创建有界 SSE 行读取器。 +func newBoundedSSEReader(r io.Reader) *boundedSSEReader { + return &boundedSSEReader{ + reader: bufio.NewReader(r), + } +} + +// ReadLine 读取一行(以 \n 分隔),同时执行 L1 和 L3 检查。 +// 返回去除尾部 \r\n 的行内容;遇到 io.EOF 时返回空字符串和 nil。 +func (r *boundedSSEReader) ReadLine() (string, error) { + line, err := r.reader.ReadString('\n') + + // L3: 总量检查(在 L1 之后、返回前统一判断) + r.totalRead += int64(len(line)) + if r.totalRead > maxStreamTotalSize { + return "", provider.ErrStreamTooLarge + } + + if err != nil && !errors.Is(err, io.EOF) { + return "", err + } + + // L1: 单行长度检查(不含末尾 \n) + rawLen := len(line) + if rawLen > 0 && line[rawLen-1] == '\n' { + rawLen-- + } + if rawLen > maxSSELineSize { + return "", provider.ErrLineTooLong + } + + // 去除尾部 \r\n + return trimLineEnding(line), err +} + +// trimLineEnding 移除行尾的 \r\n 或 \n。 +func trimLineEnding(line string) string { + for len(line) > 0 && (line[len(line)-1] == '\n' || line[len(line)-1] == '\r') { + line = line[:len(line)-1] + } + return line +} From 3612e9ff9e1959d4a00855e769e939973e69e2d0 Mon Sep 17 00:00:00 2001 From: phantom5099 <1011668688@qq.com> Date: Mon, 6 Apr 2026 19:23:38 +0800 Subject: [PATCH 20/55] =?UTF-8?q?test(provider):=E8=A1=A5=E5=85=A8?= =?UTF-8?q?=E6=B5=8B=E8=AF=95=E7=8E=87?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/provider/errors_test.go | 131 +++++++ internal/provider/openai/openai.go | 6 +- internal/provider/openai/openai_test.go | 375 ++++++++++++++++++++ internal/provider/openai/sse_reader_test.go | 193 ++++++++++ 4 files changed, 704 insertions(+), 1 deletion(-) create mode 100644 internal/provider/openai/sse_reader_test.go diff --git a/internal/provider/errors_test.go b/internal/provider/errors_test.go index e6e9ac4d..38e6b8c6 100644 --- a/internal/provider/errors_test.go +++ b/internal/provider/errors_test.go @@ -1,8 +1,10 @@ package provider import ( + "context" "errors" "fmt" + "io" "net/http" "strings" "testing" @@ -170,3 +172,132 @@ func TestProviderError_As(t *testing.T) { t.Fatalf("expected retryable") } } + +// --- IsRecoverableStreamError 全分支覆盖 --- + +func TestIsRecoverableStreamError_Nil(t *testing.T) { + t.Parallel() + if IsRecoverableStreamError(nil) { + t.Fatal("nil error should not be recoverable") + } +} + +func TestIsRecoverableStreamError_ContextErrors_NotRecoverable(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + err error + }{ + {"context.Canceled", context.Canceled}, + {"context.DeadlineExceeded", context.DeadlineExceeded}, + {"wrapped Canceled", fmt.Errorf("wrap: %w", context.Canceled)}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + if IsRecoverableStreamError(tt.err) { + t.Fatalf("%v should not be recoverable", tt.err) + } + }) + } +} + +func TestIsRecoverableStreamError_BufferOverflow_NotRecoverable(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + sentinel error + }{ + {"ErrLineTooLong", ErrLineTooLong}, + {"ErrStreamTooLarge", ErrStreamTooLarge}, + {"wrapped ErrLineTooLong", fmt.Errorf("read: %w", ErrLineTooLong)}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + if IsRecoverableStreamError(tt.sentinel) { + t.Fatalf("%v should not be recoverable", tt.sentinel) + } + }) + } +} + +func TestIsRecoverableStreamError_StreamInterrupted_Recoverable(t *testing.T) { + t.Parallel() + if !IsRecoverableStreamError(ErrStreamInterrupted) { + t.Fatal("ErrStreamInterrupted should be recoverable") + } + wrapped := fmt.Errorf("stream broken: %w", ErrStreamInterrupted) + if !IsRecoverableStreamError(wrapped) { + t.Fatal("wrapped ErrStreamInterrupted should be recoverable") + } +} + +func TestIsRecoverableStreamError_ProviderError_ByRetryableField(t *testing.T) { + t.Parallel() + + retryable := NewProviderErrorFromStatus(http.StatusTooManyRequests, "rate limit") + if !IsRecoverableStreamError(retryable) { + t.Fatal("429 ProviderError should be recoverable") + } + + serverErr := NewProviderErrorFromStatus(http.StatusInternalServerError, "internal") + if !IsRecoverableStreamError(serverErr) { + t.Fatal("5xx ProviderError should be recoverable") + } + + authErr := NewProviderErrorFromStatus(http.StatusUnauthorized, "bad key") + if IsRecoverableStreamError(authErr) { + t.Fatal("401 ProviderError should NOT be recoverable") + } + + clientErr := NewProviderErrorFromStatus(http.StatusBadRequest, "bad request") + if IsRecoverableStreamError(clientErr) { + t.Fatal("400 ProviderError should NOT be recoverable") + } + + wrappedRetryable := fmt.Errorf("layer1: %w", retryable) + if !IsRecoverableStreamError(wrappedRetryable) { + t.Fatal("wrapped retryable ProviderError should be recoverable") + } +} + +func TestIsRecoverableStreamError_NetOpError_Recoverable(t *testing.T) { + t.Parallel() + + // 模拟网络错误:使用一个包含 "connection reset" 的通用 error + // net.OpError 需要真实网络操作才能产生,这里用包装方式模拟 + genericNetErr := fmt.Errorf("net error: connection reset by peer") + // 注意:真实的 *net.OpError 需要 errors.As 匹配 + // 此处验证非上述已知不可恢复类型时默认返回 false + if IsRecoverableStreamError(genericNetErr) { + // 通用 error(非 OpError/ProviderError/哨兵)默认不恢复 + t.Fatal("generic non-net error should not be recoverable") + } +} + +func TestIsRecoverableStreamError_UnknownError_NotRecoverable(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + err error + }{ + {"generic error", errors.New("something went wrong")}, + {"io.EOF", io.EOF}, + {"io.ErrClosedPipe", io.ErrClosedPipe}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + if IsRecoverableStreamError(tt.err) { + t.Fatalf("%v should not be recoverable", tt.err) + } + }) + } +} diff --git a/internal/provider/openai/openai.go b/internal/provider/openai/openai.go index a7b1c7ac..2479de5f 100644 --- a/internal/provider/openai/openai.go +++ b/internal/provider/openai/openai.go @@ -302,7 +302,11 @@ func (p *Provider) consumeStream( line, err := reader.ReadLine() if err != nil && !errors.Is(err, io.EOF) { - // 非 EOF 的读取错误统一包装为流中断,交由 Chat() 判断是否可重连 + // 非 EOF 的读取错误:先刷新缓冲的 data 行,再包装为流中断, + // 避免中断前最后一段数据丢失。 + if flushErr := flushPendingData(); flushErr != nil { + return flushErr + } return fmt.Errorf("%w: %v", provider.ErrStreamInterrupted, err) } diff --git a/internal/provider/openai/openai_test.go b/internal/provider/openai/openai_test.go index 60c10ee6..a803fb64 100644 --- a/internal/provider/openai/openai_test.go +++ b/internal/provider/openai/openai_test.go @@ -3,6 +3,8 @@ package openai import ( "context" "encoding/json" + "errors" + "io" "net/http" "net/http/httptest" "strings" @@ -1027,3 +1029,376 @@ func TestProviderChatEmitsFullEventStream(t *testing.T) { t.Fatalf("expected TotalTokens %d, got %d", 150, messageDonePayload.Usage.TotalTokens) } } + +// --- 透明重连测试 --- + +func TestProviderChatReconnect_OnRecoverableError(t *testing.T) { + t.Setenv(config.OpenAIDefaultAPIKeyEnv, "test-key") + + attempt := 0 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + attempt++ + w.Header().Set("Content-Type", "text/event-stream") + + if attempt == 1 { + // 第一次请求:返回 5xx(可恢复) + w.WriteHeader(http.StatusInternalServerError) + _, _ = w.Write([]byte(`{"error":{"message":"temporarily unavailable"}}`)) + return + } + + // 第二次请求:正常返回 + writeSSEChunk(t, w, map[string]any{ + "choices": []map[string]any{ + {"index": 0, "delta": map[string]any{"content": "recovered"}}, + }, + }) + _, _ = w.Write([]byte("data: [DONE]\n\n")) + })) + defer server.Close() + + p, err := New(resolvedConfig(server.URL, config.OpenAIDefaultModel)) + if err != nil { + t.Fatalf("New() error = %v", err) + } + p.client = server.Client() + + events := make(chan domain.StreamEvent, 8) + err = p.Chat(context.Background(), domain.ChatRequest{ + Model: config.OpenAIDefaultModel, + Messages: []domain.Message{{Role: "user", Content: "hello"}}, + }, events) + if err != nil { + t.Fatalf("Chat() should succeed after reconnect, got: %v", err) + } + + drained := drainStreamEvents(events) + var foundText bool + for _, evt := range drained { + if evt.Type == domain.StreamEventTextDelta { + foundText = true + } + } + if !foundText { + t.Fatal("expected text_delta event after reconnect") + } + if attempt < 2 { + t.Fatalf("expected at least 2 attempts, got %d", attempt) + } +} + +func TestProviderChatReconnect_NonRecoverableError_StopsImmediately(t *testing.T) { + t.Setenv(config.OpenAIDefaultAPIKeyEnv, "test-key") + + attempt := 0 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + attempt++ + // 返回 401(不可恢复)→ 应立即停止,不重试 + w.WriteHeader(http.StatusUnauthorized) + _, _ = w.Write([]byte(`{"error":{"message":"invalid key"}}`)) + })) + defer server.Close() + + p, err := New(resolvedConfig(server.URL, config.OpenAIDefaultModel)) + if err != nil { + t.Fatalf("New() error = %v", err) + } + p.client = server.Client() + + err = p.Chat(context.Background(), domain.ChatRequest{ + Model: config.OpenAIDefaultModel, + Messages: []domain.Message{{Role: "user", Content: "hello"}}, + }, make(chan domain.StreamEvent, 1)) + if err == nil { + t.Fatal("expected error for 401") + } + if !strings.Contains(err.Error(), "invalid key") { + t.Fatalf("expected auth error, got: %v", err) + } + if attempt > 1 { + t.Fatalf("non-recoverable error should stop immediately, but got %d attempts", attempt) + } +} + +func TestProviderChatReconnect_MaxRetriesExhausted(t *testing.T) { + t.Setenv(config.OpenAIDefaultAPIKeyEnv, "test-key") + + attempt := 0 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + attempt++ + w.WriteHeader(http.StatusBadGateway) + _, _ = w.Write([]byte(`bad gateway`)) + })) + defer server.Close() + + p, err := New(resolvedConfig(server.URL, config.OpenAIDefaultModel)) + if err != nil { + t.Fatalf("New() error = %v", err) + } + p.client = server.Client() + + err = p.Chat(context.Background(), domain.ChatRequest{ + Model: config.OpenAIDefaultModel, + Messages: []domain.Message{{Role: "user", Content: "hello"}}, + }, make(chan domain.StreamEvent, 1)) + if err == nil { + t.Fatal("expected error after exhausting retries") + } + // 初始1次 + 最大3次重连 = 最多4次尝试 + if attempt > 4 { + t.Fatalf("too many attempts: %d (max should be 4)", attempt) + } +} + +func TestProviderChatReconnect_InjectsAccumulatedContext(t *testing.T) { + t.Setenv(config.OpenAIDefaultAPIKeyEnv, "test-key") + + p, err := New(resolvedConfig(config.OpenAIDefaultBaseURL, config.OpenAIDefaultModel)) + if err != nil { + t.Fatalf("New() error = %v", err) + } + + // 构造一个先返回有效 SSE 数据再中断的 reader,验证 consumeStream + // 在中断前正确累积 accumText 和 accumCalls,且错误为 ErrStreamInterrupted。 + sseData := "data: {\"choices\":[{\"index\":0,\"delta\":{\"content\":\"partial \"}}]}\n\n" + + "data: {\"choices\":[{\"index\":0,\"delta\":{\"tool_calls\":[{\"index\":0,\"id\":\"c1\",\"type\":\"function\",\"function\":{\"name\":\"bash\",\"arguments\":\"run\"}}]}}]}\n" + + reader := io.MultiReader(strings.NewReader(sseData), &errReader{err: io.ErrClosedPipe}) + events := make(chan domain.StreamEvent, 8) + accumText := &strings.Builder{} + accumCalls := make(map[int]*domain.ToolCall) + + err = p.consumeStream(context.Background(), reader, events, accumText, &accumCalls) + if err == nil { + t.Fatal("expected error from interrupted stream") + } + if !errors.Is(err, domain.ErrStreamInterrupted) { + t.Fatalf("expected ErrStreamInterrupted, got: %v", err) + } + + // 验证累积状态:文本和 tool call 都应已保留 + if accumText.String() != "partial " { + t.Fatalf("expected accumText %q, got %q", "partial ", accumText.String()) + } + call, ok := accumCalls[0] + if !ok { + t.Fatal("expected accumCalls[0] to exist") + } + if call.ID != "c1" || call.Name != "bash" { + t.Fatalf("expected tool call c1/bash, got %+v", call) + } + if call.Arguments != "run" { + t.Fatalf("expected arguments %q, got %q", "run", call.Arguments) + } + + // 验证累积状态可用于 buildAssistantMsg + msg := p.buildAssistantMsg(accumText, accumCalls) + if msg.Role != domain.RoleAssistant { + t.Fatalf("expected role assistant, got %q", msg.Role) + } + if !strings.Contains(msg.Content, "partial ") { + t.Fatalf("expected assistant content to contain 'partial', got %q", msg.Content) + } + if len(msg.ToolCalls) != 1 || msg.ToolCalls[0].ID != "c1" { + t.Fatalf("expected assistant tool calls to contain c1, got %+v", msg.ToolCalls) + } +} + +// --- 辅助方法测试 --- + +func TestBuildAssistantMsg_TextOnly(t *testing.T) { + t.Parallel() + + p, _ := New(resolvedConfig("", "")) + var accumText strings.Builder + accumText.WriteString("hello world") + + msg := p.buildAssistantMsg(&accumText, nil) + if msg.Role != domain.RoleAssistant { + t.Fatalf("expected role assistant, got %q", msg.Role) + } + if msg.Content != "hello world" { + t.Fatalf("expected content %q, got %q", "hello world", msg.Content) + } + if len(msg.ToolCalls) != 0 { + t.Fatalf("expected no tool calls, got %+v", msg.ToolCalls) + } +} + +func TestBuildAssistantMsg_WithToolCalls(t *testing.T) { + t.Parallel() + + p, _ := New(resolvedConfig("", "")) + var accumText strings.Builder + accumText.WriteString("done") + accumCalls := map[int]*domain.ToolCall{ + 0: {ID: "call_1", Name: "edit", Arguments: `{"path":"f.go"}`}, + 1: {ID: "call_2", Name: "read", Arguments: `{"path":"f.go"}`}, + } + + msg := p.buildAssistantMsg(&accumText, accumCalls) + if msg.Content != "done" { + t.Fatalf("content mismatch") + } + if len(msg.ToolCalls) != 2 { + t.Fatalf("expected 2 tool calls, got %d", len(msg.ToolCalls)) + } + if msg.ToolCalls[0].Name != "edit" || msg.ToolCalls[1].Name != "read" { + t.Fatalf("unexpected tool calls: %+v", msg.ToolCalls) + } +} + +func TestBuildAssistantMsg_EmptyAccum(t *testing.T) { + t.Parallel() + + p, _ := New(resolvedConfig("", "")) + var accumText strings.Builder + + msg := p.buildAssistantMsg(&accumText, nil) + if msg.Content != "" { + t.Fatalf("expected empty content, got %q", msg.Content) + } + if msg.ToolCalls != nil { + t.Fatal("expected nil ToolCalls when accum is nil") + } +} + +func TestMergeToolCallDeltaWithAccum_SyncsExternalState(t *testing.T) { + t.Parallel() + + events := make(chan domain.StreamEvent, 4) + accumCalls := make(map[int]*domain.ToolCall) + + delta1 := toolCallDelta{ + Index: 0, + ID: "call_acc", + Function: openAIFunctionCall{ + Name: "bash", + Arguments: `{"cmd":"ls"`, + }, + } + if err := mergeToolCallDeltaWithAccum(context.Background(), events, &accumCalls, delta1); err != nil { + t.Fatalf("first delta error = %v", err) + } + + delta2 := toolCallDelta{ + Index: 0, + Function: openAIFunctionCall{ + Arguments: `"}`, + }, + } + if err := mergeToolCallDeltaWithAccum(context.Background(), events, &accumCalls, delta2); err != nil { + t.Fatalf("second delta error = %v", err) + } + + // 验证外部 accumCalls 状态已同步 + call, ok := accumCalls[0] + if !ok { + t.Fatal("expected accumCalls[0] to exist") + } + if call.ID != "call_acc" || call.Name != "bash" { + t.Fatalf("unexpected call state: %+v", call) + } + if call.Arguments != `{"cmd":"ls""}` { + t.Fatalf("expected arguments %q, got %q", `{"cmd":"ls""}`, call.Arguments) + } +} + +func TestMergeToolCallDeltaWithAccum_NilMapInitializes(t *testing.T) { + t.Parallel() + + var accumCalls map[int]*domain.ToolCall // nil map(非指针) + + delta := toolCallDelta{ + Index: 2, + ID: "call_nil", + Function: openAIFunctionCall{ + Name: "read", + }, + } + events := make(chan domain.StreamEvent, 2) + if err := mergeToolCallDeltaWithAccum(context.Background(), events, &accumCalls, delta); err != nil { + t.Fatalf("error = %v", err) + } + + if accumCalls == nil { + t.Fatal("expected accumCalls to be initialized from nil") + } + if accumCalls[2] == nil || accumCalls[2].Name != "read" { + t.Fatalf("unexpected accumCalls[2]: %+v", accumCalls[2]) + } +} + +// --- consumeStream 错误包装测试 --- + +func TestConsumeStream_WrapsNonEOFAsInterrupted(t *testing.T) { + t.Setenv(config.OpenAIDefaultAPIKeyEnv, "test-key") + + p, err := New(resolvedConfig("", "")) + if err != nil { + t.Fatalf("New() error = %v", err) + } + + // 使用一个会触发读取错误的 source(模拟网络断开) + errReader := &errReader{err: io.ErrClosedPipe} + accumText := &strings.Builder{} + accums := make(map[int]*domain.ToolCall) + + err = p.consumeStream(context.Background(), errReader, make(chan domain.StreamEvent, 1), accumText, &accums) + if err == nil { + t.Fatal("expected error for broken reader") + } + if !errors.Is(err, domain.ErrStreamInterrupted) { + t.Fatalf("expected ErrStreamInterrupted wrapping, got: %v", err) + } +} + +// TestConsumeStream_FlushesPendingDataOnNonEOFError 验证非 EOF 读取错误发生前, +// 已缓冲但尚未刷新的 data: 行仍会被处理(不会因中断而丢失)。 +func TestConsumeStream_FlushesPendingDataOnNonEOFError(t *testing.T) { + t.Setenv(config.OpenAIDefaultAPIKeyEnv, "test-key") + + p, err := New(resolvedConfig("", "")) + if err != nil { + t.Fatalf("New() error = %v", err) + } + + // 构造一段 SSE 数据:包含一个有效 data 行,但紧跟一个错误而非空行。 + // 这模拟了流中断前最后一帧数据还没来得及被空行触发刷新的场景。 + sseData := `data: {"id":"a","object":"chat.completion.chunk","choices":[{"delta":{"content":"hello"},"finish_reason":""}]} +` // 注意:这里有换行符,但由于紧接着是 error,不会被空行刷新 + body := io.MultiReader(strings.NewReader(sseData), &errReader{err: io.ErrClosedPipe}) + + events := make(chan domain.StreamEvent, 10) + accumText := &strings.Builder{} + accums := make(map[int]*domain.ToolCall) + + err = p.consumeStream(context.Background(), body, events, accumText, &accums) + if err == nil { + t.Fatal("expected error for broken reader") + } + if !errors.Is(err, domain.ErrStreamInterrupted) { + t.Fatalf("expected ErrStreamInterrupted, got: %v", err) + } + + // 关键断言:中断前的 data 行必须已被刷新处理,文本累积不为空。 + if accumText.String() != "hello" { + t.Fatalf("expected accumText 'hello', got %q", accumText.String()) + } +} + +// errReader 是一个每次 ReadLine 都返回指定错误的测试辅助类型。 +type errReader struct { + err error +} + +func (e *errReader) Read(p []byte) (int, error) { + return 0, e.err +} + +// roundTripperFunc 将函数适配为 http.RoundTripper 接口,用于测试中 mock HTTP 行为。 +type roundTripperFunc func(*http.Request) (*http.Response, error) + +func (f roundTripperFunc) RoundTrip(req *http.Request) (*http.Response, error) { + return f(req) +} diff --git a/internal/provider/openai/sse_reader_test.go b/internal/provider/openai/sse_reader_test.go new file mode 100644 index 00000000..f0d8093b --- /dev/null +++ b/internal/provider/openai/sse_reader_test.go @@ -0,0 +1,193 @@ +package openai + +import ( + "errors" + "io" + "strings" + "testing" + + "neo-code/internal/provider" +) + +// --- boundedSSEReader 单元测试 --- + +func TestBoundedSSEReader_ReadLine_Normal(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + input string + want string + isEOF bool + }{ + { + name: "single line with newline", + input: "data: hello\n", + want: "data: hello", + isEOF: false, + }, + { + name: "line with CRLF", + input: "data: world\r\n", + want: "data: world", + isEOF: false, + }, + { + name: "empty line", + input: "\n", + want: "", + isEOF: false, + }, + { + name: "SSE comment line", + input: ": heartbeat\n", + want: ": heartbeat", + isEOF: false, + }, + { + name: "EOF without trailing newline (io.EOF)", + input: "data: partial", + want: "data: partial", + isEOF: true, + }, + } + + for _, tt := range tests { + tt := tt + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + r := newBoundedSSEReader(strings.NewReader(tt.input)) + got, err := r.ReadLine() + if got != tt.want { + t.Fatalf("ReadLine() = %q, want %q", got, tt.want) + } + if tt.isEOF && !errors.Is(err, io.EOF) { + t.Fatalf("expected io.EOF, got %v", err) + } + if !tt.isEOF && err != nil { + t.Fatalf("unexpected error: %v", err) + } + }) + } +} + +func TestBoundedSSEReader_ReadLine_MultipleLines(t *testing.T) { + t.Parallel() + + r := newBoundedSSEReader(strings.NewReader("line1\nline2\n\nline4\n")) + + line1, err := r.ReadLine() + if err != nil || line1 != "line1" { + t.Fatalf("first line: got %q, err = %v", line1, err) + } + + line2, err := r.ReadLine() + if err != nil || line2 != "line2" { + t.Fatalf("second line: got %q, err = %v", line2, err) + } + + // 空行 + empty, err := r.ReadLine() + if err != nil || empty != "" { + t.Fatalf("empty line: got %q, err = %v", empty, err) + } + + line4, err := r.ReadLine() + if err != nil || line4 != "line4" { + t.Fatalf("fourth line: got %q, err = %v", line4, err) + } + + // EOF + _, err = r.ReadLine() + if !errors.Is(err, io.EOF) { + t.Fatalf("expected EOF after all lines, got %v", err) + } +} + +func TestBoundedSSEReader_L1_LineTooLong(t *testing.T) { + t.Parallel() + + longLine := strings.Repeat("x", maxSSELineSize+1) + "\n" + r := newBoundedSSEReader(strings.NewReader(longLine)) + + _, err := r.ReadLine() + if err == nil { + t.Fatal("expected ErrLineTooLong for oversized line") + } + if !errors.Is(err, provider.ErrLineTooLong) { + t.Fatalf("expected ErrLineTooLong, got %v", err) + } +} + +func TestBoundedSSEReader_L1_BoundaryExactLimit(t *testing.T) { + t.Parallel() + + // 恰好等于上限的行应该正常通过(不含 \n) + exactLine := strings.Repeat("a", maxSSELineSize) + "\n" + r := newBoundedSSEReader(strings.NewReader(exactLine)) + + got, err := r.ReadLine() + if err != nil { + t.Fatalf("unexpected error at exact limit: %v", err) + } + if len(got) != maxSSELineSize { + t.Fatalf("expected line length %d, got %d", maxSSELineSize, len(got)) + } +} + +func TestBoundedSSEReader_L3_StreamTooLarge(t *testing.T) { + t.Parallel() + + // 构造输入:多行小内容但总量超过 maxStreamTotalSize + var sb strings.Builder + for i := 0; i < 100; i++ { + sb.WriteString("data: chunk\n") + } + // 填充到超过总量上限 + remaining := maxStreamTotalSize - int64(sb.Len()) + 1 + sb.WriteString(strings.Repeat("x", int(remaining)) + "\n") + + r := newBoundedSSEReader(strings.NewReader(sb.String())) + + // 前面的行应能正常读取 + for i := 0; i < 100; i++ { + _, err := r.ReadLine() + if err != nil { + t.Fatalf("unexpected error on normal line %d: %v", i, err) + } + } + + // 超限的行应返回 ErrStreamTooLarge + _, err := r.ReadLine() + if err == nil { + t.Fatal("expected ErrStreamTooLarge") + } + if !errors.Is(err, provider.ErrStreamTooLarge) { + t.Fatalf("expected ErrStreamTooLarge, got %v", err) + } +} + +func TestTrimLineEnding(t *testing.T) { + t.Parallel() + + tests := []struct { + input string + want string + }{ + {"hello\n", "hello"}, + {"hello\r\n", "hello"}, + {"hello\r\n\n", "hello"}, // 连续换行符全部去除 + {"hello", "hello"}, + {"\n", ""}, + {"\r\n", ""}, + {"\r", ""}, // 孤立 \r 也去除 + {"", ""}, + } + + for _, tt := range tests { + got := trimLineEnding(tt.input) + if got != tt.want { + t.Fatalf("trimLineEnding(%q) = %q, want %q", tt.input, got, tt.want) + } + } +} From aefa4fab7c74975ef832babd239d5c0f7ff98f12 Mon Sep 17 00:00:00 2001 From: phantom5099 <1011668688@qq.com> Date: Mon, 6 Apr 2026 19:44:46 +0800 Subject: [PATCH 21/55] =?UTF-8?q?fix(provider):=E4=BF=AE=E5=A4=8DSSE=20Rea?= =?UTF-8?q?der=20=E7=BC=93=E5=86=B2=E5=8C=BA=E6=BA=A2=E5=87=BA=E4=BB=A5?= =?UTF-8?q?=E5=8F=8A=E9=87=8D=E8=BF=9E=E6=B6=88=E6=81=AF=E6=B1=A1=E6=9F=93?= =?UTF-8?q?=E9=A3=8E=E9=99=A9?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/provider/openai/openai.go | 17 +- internal/provider/openai/openai_test.go | 222 +++++++++++++++++++- internal/provider/openai/sse_reader.go | 30 ++- internal/provider/openai/sse_reader_test.go | 91 +++++++- 4 files changed, 337 insertions(+), 23 deletions(-) diff --git a/internal/provider/openai/openai.go b/internal/provider/openai/openai.go index 2479de5f..c71e1cca 100644 --- a/internal/provider/openai/openai.go +++ b/internal/provider/openai/openai.go @@ -109,6 +109,10 @@ func (p *Provider) DiscoverModels(ctx context.Context) ([]config.ModelDescriptor func (p *Provider) Chat(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { const maxReconnects = 3 + // 保存原始消息列表的副本,避免重连时反复 append 到同一个切片导致上下文污染 + originalMessages := make([]provider.Message, len(req.Messages)) + copy(originalMessages, req.Messages) + // 跨重连周期持久化的累积状态:已收到的文本和 tool call var ( accumText strings.Builder @@ -117,9 +121,16 @@ func (p *Provider) Chat(ctx context.Context, req provider.ChatRequest, events ch for attempt := 0; attempt <= maxReconnects; attempt++ { if attempt > 0 { - // 将已累积内容作为 assistant 消息注入,使新请求能从断点继续 - req.Messages = append(req.Messages, - p.buildAssistantMsg(&accumText, accumCalls)) + // 从原始消息出发构造本次请求的完整消息列表 + req.Messages = make([]provider.Message, len(originalMessages), len(originalMessages)+1) + copy(req.Messages, originalMessages) + + // 仅在有实际累积内容时注入 assistant 快照,避免插入空消息 + if accumText.Len() > 0 || len(accumCalls) > 0 { + req.Messages = append(req.Messages, + p.buildAssistantMsg(&accumText, accumCalls)) + } + // 指数退避等待 backoff := time.Duration(1< maxStreamTotalSize { - return "", provider.ErrStreamTooLarge + // L1: 缓冲区溢出 → 单行超过 maxSSELineSize(触发在读取过程中,而非读完后) + if errors.Is(err, bufio.ErrBufferFull) { + return "", provider.ErrLineTooLong } if err != nil && !errors.Is(err, io.EOF) { return "", err } - // L1: 单行长度检查(不含末尾 \n) + // L1 兜底:行内容长度检查(不含末尾 \n) rawLen := len(line) if rawLen > 0 && line[rawLen-1] == '\n' { rawLen-- @@ -55,8 +61,14 @@ func (r *boundedSSEReader) ReadLine() (string, error) { return "", provider.ErrLineTooLong } - // 去除尾部 \r\n - return trimLineEnding(line), err + // L3: 总量检查 + r.totalRead += int64(len(line)) + if r.totalRead > maxStreamTotalSize { + return "", provider.ErrStreamTooLarge + } + + // 将 []byte 转为 string(ReadSlice 返回的底层数据在下次读取时会被覆盖) + return trimLineEnding(string(line)), err } // trimLineEnding 移除行尾的 \r\n 或 \n。 diff --git a/internal/provider/openai/sse_reader_test.go b/internal/provider/openai/sse_reader_test.go index f0d8093b..eb7a1803 100644 --- a/internal/provider/openai/sse_reader_test.go +++ b/internal/provider/openai/sse_reader_test.go @@ -138,26 +138,28 @@ func TestBoundedSSEReader_L1_BoundaryExactLimit(t *testing.T) { func TestBoundedSSEReader_L3_StreamTooLarge(t *testing.T) { t.Parallel() - // 构造输入:多行小内容但总量超过 maxStreamTotalSize + // 构造输入:每行 1KB(远小于 maxSSELineSize),行数足够多使总量超过 maxStreamTotalSize + line := strings.Repeat("x", 1024) + "\n" // 1025 bytes per line + lineSize := int64(len(line)) + var sb strings.Builder - for i := 0; i < 100; i++ { - sb.WriteString("data: chunk\n") + linesToWrite := int(maxStreamTotalSize/lineSize) + 1 + for range linesToWrite { + sb.WriteString(line) } - // 填充到超过总量上限 - remaining := maxStreamTotalSize - int64(sb.Len()) + 1 - sb.WriteString(strings.Repeat("x", int(remaining)) + "\n") r := newBoundedSSEReader(strings.NewReader(sb.String())) // 前面的行应能正常读取 - for i := 0; i < 100; i++ { + expectedNormal := int(maxStreamTotalSize / lineSize) + for range expectedNormal { _, err := r.ReadLine() if err != nil { - t.Fatalf("unexpected error on normal line %d: %v", i, err) + t.Fatalf("unexpected error on normal line: %v", err) } } - // 超限的行应返回 ErrStreamTooLarge + // 超限的行应返回 ErrStreamTooLarge(而非 ErrLineTooLong) _, err := r.ReadLine() if err == nil { t.Fatal("expected ErrStreamTooLarge") @@ -191,3 +193,74 @@ func TestTrimLineEnding(t *testing.T) { } } } + +// TestBoundedSSEReader_L1_NoNewlineEOFAtLimit 验证恰好等于缓冲区大小 +// 但以 EOF 结尾(无 \n)的行能被正常读取并返回 io.EOF。 +func TestBoundedSSEReader_L1_NoNewlineEOFAtLimit(t *testing.T) { + t.Parallel() + + // 不含 \n 的行,长度恰好等于 maxSSELineSize,以 EOF 结尾 + exactLine := strings.Repeat("b", maxSSELineSize) + r := newBoundedSSEReader(strings.NewReader(exactLine)) + + got, err := r.ReadLine() + if err == nil || !errors.Is(err, io.EOF) { + t.Fatalf("expected io.EOF for line without trailing newline, got err=%v", err) + } + if len(got) != maxSSELineSize { + t.Fatalf("expected line length %d, got %d", maxSSELineSize, len(got)) + } +} + +// TestBoundedSSEReader_L1_NoNewlineEOFExceedsLimit 验证超过缓冲区大小 +// 且以 EOF 结尾(无 \n)的行返回 ErrLineTooLong。 +func TestBoundedSSEReader_L1_NoNewlineEOFExceedsLimit(t *testing.T) { + t.Parallel() + + // 不含 \n,长度超过 maxSSELineSize + longLine := strings.Repeat("c", maxSSELineSize+1) + r := newBoundedSSEReader(strings.NewReader(longLine)) + + _, err := r.ReadLine() + if err == nil { + t.Fatal("expected ErrLineTooLong for oversized line without newline") + } + if !errors.Is(err, provider.ErrLineTooLong) { + t.Fatalf("expected ErrLineTooLong, got %v", err) + } +} + +// TestBoundedSSEReader_UnderlyingErrorPropagation 验证底层 reader 的非 EOF 错误 +// 会被正确传播,而不是被 L1/L3 吞掉。 +func TestBoundedSSEReader_UnderlyingErrorPropagation(t *testing.T) { + t.Parallel() + + r := newBoundedSSEReader(&errReader{err: io.ErrClosedPipe}) + _, err := r.ReadLine() + if err == nil { + t.Fatal("expected error from broken reader") + } + if !errors.Is(err, io.ErrClosedPipe) { + t.Fatalf("expected io.ErrClosedPipe, got %v", err) + } +} + +// TestBoundedSSEReader_L1_ThenNormalRead 验证 L1 触发后 reader 状态仍然一致, +// 后续正常行仍可继续读取(如果调用方选择恢复)。 +func TestBoundedSSEReader_L1_ThenNormalRead(t *testing.T) { + t.Parallel() + + // 第一行超长触发 L1,第二行正常 + input := strings.Repeat("x", maxSSELineSize+10) + "\nnormal line\n" + r := newBoundedSSEReader(strings.NewReader(input)) + + // 第一行应返回 ErrLineTooLong + _, err := r.ReadLine() + if !errors.Is(err, provider.ErrLineTooLong) { + t.Fatalf("first line: expected ErrLineTooLong, got %v", err) + } + + // 注意:L1 触发后 bufio.Reader 内部可能已消耗了部分后续数据(缓冲区残留), + // 因此此测试仅验证 L1 错误被正确返回,不要求后续行一定可读。 + // 这符合实际使用场景——L1 触发后调用方会终止流消费。 +} From 8aadd974fa38ccdc51b0e325f452c8ee3219cde4 Mon Sep 17 00:00:00 2001 From: phantom5099 <1011668688@qq.com> Date: Mon, 6 Apr 2026 20:07:31 +0800 Subject: [PATCH 22/55] =?UTF-8?q?fix(provider):=E4=BF=AE=E5=A4=8D=E9=87=8D?= =?UTF-8?q?=E8=AF=95=E6=AC=A1=E6=95=B0=E8=BF=87=E5=A4=9A=E4=BB=A5=E5=8F=8A?= =?UTF-8?q?=E9=94=99=E8=AF=AF=E4=B8=A2=E5=A4=B1=E9=97=AE=E9=A2=98?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/provider/errors.go | 20 ++++++ internal/provider/errors_test.go | 86 +++++++++++++++++++++++++ internal/provider/openai/openai.go | 11 ++-- internal/provider/openai/openai_test.go | 8 +++ 4 files changed, 120 insertions(+), 5 deletions(-) diff --git a/internal/provider/errors.go b/internal/provider/errors.go index 29697cca..a2fee38f 100644 --- a/internal/provider/errors.go +++ b/internal/provider/errors.go @@ -108,6 +108,26 @@ func NewTimeoutProviderError(message string) *ProviderError { } } +// MarkNonRetryable 将错误标记为不可重试,用于防止上层重试叠加放大。 +// +// 若错误链中包含 *ProviderError,返回其 Retryable=false 的副本; +// 否则将原始错误包装为 *ProviderError{Code: ErrorCodeUnknown, Retryable: false}。 +// 原始错误通过 Unwrap 保留,不影响 errors.Is/As 对原始哨兵的匹配。 +func MarkNonRetryable(err error) error { + var pErr *ProviderError + if errors.As(err, &pErr) { + clone := *pErr + clone.Retryable = false + return &clone + } + return &ProviderError{ + StatusCode: 0, + Code: ErrorCodeUnknown, + Message: err.Error(), + Retryable: false, + } +} + // IsRecoverableStreamError 判断流读取错误是否可通过透明重连恢复。 // // 不可恢复的情况: diff --git a/internal/provider/errors_test.go b/internal/provider/errors_test.go index 38e6b8c6..eed23978 100644 --- a/internal/provider/errors_test.go +++ b/internal/provider/errors_test.go @@ -301,3 +301,89 @@ func TestIsRecoverableStreamError_UnknownError_NotRecoverable(t *testing.T) { }) } } + +// --- MarkNonRetryable 测试 --- + +func TestMarkNonRetryable_ProviderError(t *testing.T) { + t.Parallel() + + // Retryable=true 的 ProviderError → Retryable=false + retryable := NewProviderErrorFromStatus(http.StatusInternalServerError, "internal") + if !retryable.Retryable { + t.Fatal("setup: expected retryable") + } + + marked := MarkNonRetryable(retryable) + var pErr *ProviderError + if !errors.As(marked, &pErr) { + t.Fatal("marked error should be *ProviderError") + } + if pErr.Retryable { + t.Fatal("MarkNonRetryable should set Retryable=false") + } + if pErr.StatusCode != 500 || pErr.Code != ErrorCodeServer { + t.Fatalf("MarkNonRetryable should preserve StatusCode and Code, got status=%d code=%s", pErr.StatusCode, pErr.Code) + } + + // 原始对象不受影响 + if !retryable.Retryable { + t.Fatal("original ProviderError should not be mutated") + } +} + +func TestMarkNonRetryable_WrappedProviderError(t *testing.T) { + t.Parallel() + + inner := NewProviderErrorFromStatus(http.StatusTooManyRequests, "rate limited") + wrapped := fmt.Errorf("layer: %w", inner) + + marked := MarkNonRetryable(wrapped) + var pErr *ProviderError + if !errors.As(marked, &pErr) { + t.Fatal("marked error should contain *ProviderError") + } + if pErr.Retryable { + t.Fatal("MarkNonRetryable on wrapped should set Retryable=false") + } +} + +func TestMarkNonRetryable_NonProviderError(t *testing.T) { + t.Parallel() + + // 非 ProviderError → 包装为 ProviderError{Retryable: false} + generic := errors.New("some error") + marked := MarkNonRetryable(generic) + + var pErr *ProviderError + if !errors.As(marked, &pErr) { + t.Fatal("marked error should be *ProviderError") + } + if pErr.Retryable { + t.Fatal("should be non-retryable") + } + if pErr.Code != ErrorCodeUnknown { + t.Fatalf("expected code %s, got %s", ErrorCodeUnknown, pErr.Code) + } + if !strings.Contains(pErr.Message, "some error") { + t.Fatalf("message should contain original error text, got: %s", pErr.Message) + } +} + +func TestMarkNonRetryable_AlreadyNonRetryable(t *testing.T) { + t.Parallel() + + // 已经 Retryable=false 的 ProviderError → 保持 false + nonRetryable := NewProviderErrorFromStatus(http.StatusUnauthorized, "bad key") + if nonRetryable.Retryable { + t.Fatal("setup: 401 should be non-retryable") + } + + marked := MarkNonRetryable(nonRetryable) + var pErr *ProviderError + if !errors.As(marked, &pErr) { + t.Fatal("marked error should be *ProviderError") + } + if pErr.Retryable { + t.Fatal("should still be non-retryable") + } +} diff --git a/internal/provider/openai/openai.go b/internal/provider/openai/openai.go index c71e1cca..6e3a59c7 100644 --- a/internal/provider/openai/openai.go +++ b/internal/provider/openai/openai.go @@ -144,9 +144,13 @@ func (p *Provider) Chat(ctx context.Context, req provider.ChatRequest, events ch if err == nil { return nil } - if !provider.IsRecoverableStreamError(err) || attempt == maxReconnects { + if !provider.IsRecoverableStreamError(err) { return err } + // 可恢复但重连次数已耗尽 → 标记为不可重试,防止上层 runtime 重试叠加放大。 + if attempt == maxReconnects { + return provider.MarkNonRetryable(err) + } } return nil // unreachable,但满足编译器 } @@ -318,7 +322,7 @@ func (p *Provider) consumeStream( if flushErr := flushPendingData(); flushErr != nil { return flushErr } - return fmt.Errorf("%w: %v", provider.ErrStreamInterrupted, err) + return fmt.Errorf("%w: %w", provider.ErrStreamInterrupted, err) } trimmed := line @@ -415,9 +419,6 @@ func emitToolCallDelta(ctx context.Context, events chan<- provider.StreamEvent, // emitMessageDone 发送消息完成事件。 func emitMessageDone(ctx context.Context, events chan<- provider.StreamEvent, finishReason string, usage *provider.Usage) error { - if events == nil { - return nil - } return emitStreamEvent(ctx, events, provider.NewMessageDoneStreamEvent(finishReason, usage)) } diff --git a/internal/provider/openai/openai_test.go b/internal/provider/openai/openai_test.go index db61452d..9085e54b 100644 --- a/internal/provider/openai/openai_test.go +++ b/internal/provider/openai/openai_test.go @@ -1146,6 +1146,14 @@ func TestProviderChatReconnect_MaxRetriesExhausted(t *testing.T) { if err == nil { t.Fatal("expected error after exhausting retries") } + // 重连耗尽后,错误应被标记为不可重试,防止上层 runtime 再次重试叠加放大。 + var pErr *domain.ProviderError + if !errors.As(err, &pErr) { + t.Fatalf("expected *ProviderError after retry exhaustion, got: %T: %v", err, err) + } + if pErr.Retryable { + t.Fatal("error should be non-retryable after reconnect exhaustion") + } // 初始1次 + 最大3次重连 = 最多4次尝试 if attempt > 4 { t.Fatalf("too many attempts: %d (max should be 4)", attempt) From c9a0ec18d98b4fc6fb428ed5aff0bf26764d85da Mon Sep 17 00:00:00 2001 From: Yumiue <229866007@qq.com> Date: Mon, 6 Apr 2026 22:05:44 +0800 Subject: [PATCH 23/55] =?UTF-8?q?docs:=E5=88=A0=E9=99=A4=E8=BF=87=E6=9C=9F?= =?UTF-8?q?mvp=20=E6=96=87=E4=BB=B6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/neocode-coding-agent-mvp-architecture.md | 521 ------------------ 1 file changed, 521 deletions(-) delete mode 100644 docs/neocode-coding-agent-mvp-architecture.md diff --git a/docs/neocode-coding-agent-mvp-architecture.md b/docs/neocode-coding-agent-mvp-architecture.md deleted file mode 100644 index e6f13128..00000000 --- a/docs/neocode-coding-agent-mvp-architecture.md +++ /dev/null @@ -1,521 +0,0 @@ -# ⚠️ 已过时:旧版 API 文档(DEPRECATED) - - -# NeoCode Coding Agent MVP 架构设计 - -## 1. 目标 - -本文定义一个基于 Go + Bubble Tea 的本地 Coding Agent MVP,目标是先跑通最小闭环: - -`用户输入 -> Agent 推理 -> 调用工具 -> 获取结果 -> 继续推理 -> UI 展示` - -MVP 聚焦六个模块: - -1. provider:统一不同模型/API 的调用方式 -2. TUI:用户交互入口,承载输入、对话、侧边栏、会话 -3. tools:统一工具定义、参数校验、执行与结果封装 -4. config:管理本地配置、provider 切换、模型选择 -5. context:负责 system prompt、显式上下文源与历史消息裁剪 -6. agent runtime:驱动整个 agent loop,是系统核心 - ---- - -## 2. 设计原则 - -- 模块职责清晰,避免 UI、模型调用、工具执行互相耦合 -- 面向接口设计,方便后续增加 provider 和工具 -- MVP 先保证主链路可用,不追求一次做全 -- 所有副作用操作统一收敛到 provider/tools/config 等边界层 -- Runtime 作为唯一编排中心,TUI 不直接调用 provider 和 tools - ---- - -## 3. 总体架构 - -```mermaid -flowchart LR - U["User"] --> TUI["Bubble Tea TUI"] - TUI --> APP["Application / Bootstrap"] - APP --> CFG["Config"] - APP --> RT["Agent Runtime"] - RT --> PR["Provider"] - RT --> TM["Tool Manager"] - TM --> FS["Filesystem Tool"] - TM --> SH["Bash Tool"] - TM --> WF["WebFetch Tool"] -``` - -系统分层: - -- TUI:负责交互和渲染 -- Application:负责启动和依赖注入 -- Runtime:负责 Agent Loop 和状态编排 -- Context:负责模型请求前的上下文构建 -- Provider:负责模型调用抽象 -- Tool Manager:负责工具注册、校验、执行 -- Config:负责配置加载与选择 - ---- - -## 4. 模块设计 - -### 4.1 Provider - -职责: - -- 屏蔽 OpenAI / Anthropic / Gemini 的协议差异 -- 统一暴露聊天、工具调用、流式输出能力 -- 管理 endpoint、model、api key、超时、重试 - -建议接口: - -```go -type Provider interface { - Name() string - Chat(ctx context.Context, req ChatRequest) (ChatResponse, error) -} - -type ChatRequest struct { - Model string - SystemPrompt string - Messages []Message - Tools []ToolSpec - Stream bool -} - -type ChatResponse struct { - Message Message - FinishReason string - Usage Usage -} - -type Message struct { - Role string - Content string - ToolCalls []ToolCall -} - -type ToolCall struct { - ID string - Name string - Arguments string -} -``` - -MVP 建议: - -- 第一阶段先实现一个 provider,例如 OpenAI 兼容接口 -- Provider 层只关心“模型协议”,不关心 UI 和工具执行 -- Runtime 把 Tool schema 传给 Provider,Provider 把 ToolCall 返回给 Runtime - ---- - -### 4.2 TUI - -职责: - -- 用户输入和结果展示 -- 展示会话列表、当前会话、工具执行状态 -- 接收快捷键和命令 -- 通过事件与 Runtime 通信 - -建议布局: - -- 左侧:会话列表 Sidebar -- 中间:对话消息区 -- 底部:输入框 -- 顶部/状态栏:provider、model、workdir、运行状态 - -建议状态: - -```go -type UIState struct { - Sessions []SessionSummary - ActiveSessionID string - InputText string - IsAgentRunning bool - StatusText string - CurrentProvider string - CurrentModel string -} -``` - -边界原则: - -- TUI 不直接处理模型协议 -- TUI 不直接执行工具 -- TUI 只发送事件,例如“提交输入”“切换会话” -- Runtime 回传事件,例如“开始响应”“工具开始/结束”“最终完成” - ---- - -### 4.3 Tools - -职责: - -- 定义统一工具协议 -- 管理工具注册、查找、schema、执行和结果格式 -- 为 Runtime 提供统一调用入口 - -MVP 工具: - -- filesystem -- bash -- webfetch - -建议接口: - -```go -type Tool interface { - Name() string - Description() string - Schema() any - Execute(ctx context.Context, call ToolCallInput) (ToolResult, error) -} - -type ToolCallInput struct { - ID string - Name string - Arguments []byte - SessionID string - Workdir string -} - -type ToolResult struct { - ToolCallID string - Name string - Content string - IsError bool - Metadata map[string]any -} -``` - -建议增加 `Registry` / `Manager`: - -- 注册所有内置工具 -- 暴露 `ListSchemas()` 给 Provider/Runtime -- 负责参数校验和统一错误封装 -- 把工具输出转成模型可消费的结果消息 - -各工具 MVP 建议: - -- Filesystem:读文件、写文件、列目录、搜索文件 -- Bash:执行命令,限制超时、输出长度、工作目录 -- WebFetch:抓取网页文本内容,限制响应大小 - ---- - -### 4.4 Config - -职责: - -- 从 `~/.neocode/config.yaml` 加载配置 -- 管理 provider 列表、当前 provider、当前 model -- 校验配置完整性并提供默认值 - -示例: - -```yaml -providers: - - name: openai - type: openai - base_url: https://api.openai.com/v1 - model: gpt-4.1 - api_key_env: OPENAI_API_KEY - - - name: anthropic - type: anthropic - base_url: https://api.anthropic.com - model: claude-3-7-sonnet-latest - api_key_env: ANTHROPIC_API_KEY - -selected_provider: openai -current_model: gpt-4.1 -workdir: . -shell: bash -``` - -建议结构: - -```go -type Config struct { - Providers []ProviderConfig `yaml:"providers"` - SelectedProvider string `yaml:"selected_provider"` - CurrentModel string `yaml:"current_model"` - Workdir string `yaml:"workdir"` - Shell string `yaml:"shell"` -} - -type ProviderConfig struct { - Name string `yaml:"name"` - Type string `yaml:"type"` - BaseURL string `yaml:"base_url"` - Model string `yaml:"model"` - APIKeyEnv string `yaml:"api_key_env"` -} -``` - -建议: - -- API Key 不直接写配置文件,只引用环境变量名 -- 加载后生成运行时只读配置对象 -- 启动时立即校验 selected provider 是否存在 - ---- - -### 4.5 Agent Runtime - -职责: - -- 管理会话上下文 -- 调用 Provider 获取模型响应 -- 识别并执行 ToolCall -- 将 ToolResult 回灌模型 -- 持续循环直到得到最终答案或触发停止条件 - -MVP 推荐使用简化版 ReAct / Tool-Calling Loop: - -1. 接收用户输入 -2. 组装 system prompt + 历史消息 + tools -3. 调用 provider -4. 若返回普通文本,则输出 -5. 若返回 tool calls,则执行工具 -6. 将工具结果追加到上下文 -7. 再次调用 provider -8. 重复直到结束 - -建议接口: - -```go -type Runtime interface { - Run(ctx context.Context, input UserInput) error -} - -type UserInput struct { - SessionID string - Content string -} -``` - -Runtime 内部建议拆分: - -- `SessionStore`:管理会话和消息历史 -- `context.Builder`:组装核心 prompt、显式上下文源与裁剪后的消息 -- `Executor`:执行 loop -- `EventBus`:向 TUI 推送运行事件 - -停止条件: - -- provider 返回最终文本 -- 超过最大轮数 -- 工具执行失败且不可恢复 -- 用户取消 - ---- - -## 5. 核心数据模型 - -```go -type Session struct { - ID string - Title string - Messages []Message - CreatedAt time.Time - UpdatedAt time.Time -} - -type RuntimeEvent struct { - Type string - Payload any -} -``` - -建议事件类型: - -- `user_message` -- `agent_chunk` -- `tool_started` -- `tool_finished` -- `agent_completed` -- `error` - ---- - -## 6. 启动与依赖注入 - -Application 层负责把所有模块组装起来: - -```mermaid -sequenceDiagram - participant Main - participant Config - participant Provider - participant Tools - participant Runtime - participant TUI - - Main->>Config: Load() - Main->>Provider: Build() - Main->>Tools: Register builtin tools - Main->>Runtime: New() - Main->>TUI: Start() - - TUI->>Runtime: Submit input - Runtime->>Provider: Chat() - Provider-->>Runtime: ToolCall / Final Answer - Runtime->>Tools: Execute() - Tools-->>Runtime: ToolResult - Runtime-->>TUI: Events -``` - ---- - -## 7. 建议目录结构 - -```text -. -├── cmd/ -│ └── neocode/ -│ └── main.go -├── internal/ -│ ├── app/ -│ │ └── bootstrap.go -│ ├── config/ -│ │ ├── loader.go -│ │ ├── model.go -│ │ └── validate.go -│ ├── context/ -│ │ ├── builder.go -│ │ ├── metadata.go -│ │ ├── prompt.go -│ │ ├── source_rules.go -│ │ ├── source_system.go -│ │ └── trim.go -│ ├── provider/ -│ │ ├── provider.go -│ │ ├── openai/ -│ │ ├── anthropic/ -│ │ └── gemini/ -│ ├── runtime/ -│ │ ├── runtime.go -│ │ ├── executor.go -│ │ ├── prompt_builder.go -│ │ ├── session_store.go -│ │ └── events.go -│ ├── tools/ -│ │ ├── registry.go -│ │ ├── types.go -│ │ ├── filesystem/ -│ │ ├── bash/ -│ │ └── webfetch/ -│ └── tui/ -│ ├── app.go -│ ├── state.go -│ ├── keymap.go -│ ├── views/ -│ └── components/ -└── docs/ - └── mvp-architecture.md -``` - ---- - -## 8. MVP 时序示例 - -场景:用户提问后触发一次工具调用 - -1. 用户在 TUI 输入问题 -2. TUI 将输入发送给 Runtime -3. Runtime 读取 Session 历史 -4. Runtime 获取工具 schema -5. Runtime 调用 Provider -6. Provider 返回 tool call,例如 `filesystem.read_file` -7. Runtime 调用 Tool Manager 执行 -8. Tool Result 写回上下文 -9. Runtime 再次调用 Provider -10. Provider 返回最终回答 -11. Runtime 把结果事件发送给 TUI -12. TUI 刷新界面 - ---- - -## 9. 错误处理 - -Provider 错误: - -- 网络错误:有限重试 -- 认证错误:提示配置问题 -- 限流错误:提示稍后重试 -- 非法响应:记录日志并返回用户可读错误 - -Tool 错误: - -- 参数错误:返回结构化错误 -- 执行失败:不中断程序,作为 tool error 回灌 -- 超时:统一包装 timeout - -Runtime 错误: - -- 超过最大轮数立即停止 -- 构造上下文失败则结束当前请求 -- 通过事件通知 TUI 展示错误 - ---- - -## 10. 安全边界 - -MVP 建议先加基础约束: - -- Filesystem 默认限制在工作目录内 -- Bash 限制超时、输出长度、禁止交互式阻塞命令 -- WebFetch 限制协议和响应大小 -- 配置文件不保存明文 API Key - ---- - -## 11. 开发顺序 - -### Phase 1:先跑通闭环 - -- config -- provider 抽象 + 一个 provider 实现 -- tools registry + 一个 filesystem 工具 -- runtime loop -- tui 单会话输入输出 - -### Phase 2:增强可用性 - -- 会话侧边栏 -- bash / webfetch 工具 -- 流式输出 -- 状态栏和错误展示 - -### Phase 3:增强扩展性 - -- 多 provider 切换 -- session 持久化 -- 更完整的权限控制 -- 更丰富的工具生态 - ---- - -## 12. MVP 成功标准 - -满足以下条件即可认为 MVP 完成: - -- 用户可在 TUI 中输入问题 -- Agent 可调用至少一个模型 provider -- Agent 可调用至少一个工具 -- 工具结果可回灌给模型继续推理 -- UI 可展示基本会话历史和运行状态 -- 配置可从 `~/.neocode/config.yaml` 加载 - ---- - -## 13. 总结 - -这个架构的关键是先把主链路做干净: - -`TUI -> Runtime -> Provider -> Tool Manager -> Runtime -> TUI` - -只要这条链路稳定,后面无论加更多 provider、更多工具,还是把 Runtime 升级成更复杂的 Agent,都不需要推翻当前设计。 From daec7444ef0badb527ffd73c57ec553da330b292 Mon Sep 17 00:00:00 2001 From: phantom5099 <1011668688@qq.com> Date: Mon, 6 Apr 2026 23:08:30 +0800 Subject: [PATCH 24/55] =?UTF-8?q?pref(provider)=EF=BC=9A=E5=88=A0=E6=8E=89?= =?UTF-8?q?=E9=80=8F=E6=98=8E=E9=87=8D=E4=BC=A0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/provider/errors.go | 62 --- internal/provider/errors_test.go | 217 ---------- internal/provider/openai/openai.go | 117 +---- internal/provider/openai/openai_test.go | 553 +----------------------- 4 files changed, 19 insertions(+), 930 deletions(-) diff --git a/internal/provider/errors.go b/internal/provider/errors.go index a2fee38f..daff0da6 100644 --- a/internal/provider/errors.go +++ b/internal/provider/errors.go @@ -1,10 +1,8 @@ package provider import ( - "context" "errors" "fmt" - "net" "net/http" ) @@ -107,63 +105,3 @@ func NewTimeoutProviderError(message string) *ProviderError { Retryable: true, // 超时默认可重试 } } - -// MarkNonRetryable 将错误标记为不可重试,用于防止上层重试叠加放大。 -// -// 若错误链中包含 *ProviderError,返回其 Retryable=false 的副本; -// 否则将原始错误包装为 *ProviderError{Code: ErrorCodeUnknown, Retryable: false}。 -// 原始错误通过 Unwrap 保留,不影响 errors.Is/As 对原始哨兵的匹配。 -func MarkNonRetryable(err error) error { - var pErr *ProviderError - if errors.As(err, &pErr) { - clone := *pErr - clone.Retryable = false - return &clone - } - return &ProviderError{ - StatusCode: 0, - Code: ErrorCodeUnknown, - Message: err.Error(), - Retryable: false, - } -} - -// IsRecoverableStreamError 判断流读取错误是否可通过透明重连恢复。 -// -// 不可恢复的情况: -// - context 取消/超时(调用方主动终止) -// - 缓冲区溢出(重连只会再次溢出) -// - 认证失败等业务错误(重连无意义) -// -// 可恢复的情况: -// - ProviderError 且 Retryable=true(5xx、429 等) -// - 网络层临时错误(*net.OpError) -// - ErrStreamInterrupted(通用流中断标记) -func IsRecoverableStreamError(err error) bool { - if err == nil { - return false - } - // context 取消 → 不可恢复 - if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { - return false - } - // 缓冲区溢出 → 不可恢复(重连同样会溢出) - if errors.Is(err, ErrLineTooLong) || errors.Is(err, ErrStreamTooLarge) { - return false - } - // 流中断标记 → 可恢复 - if errors.Is(err, ErrStreamInterrupted) { - return true - } - // ProviderError → 依据 Retryable 字段 - var pErr *ProviderError - if errors.As(err, &pErr) { - return pErr.Retryable - } - // 网络层临时故障(连接重置、超时等)→ 可恢复 - var netErr *net.OpError - if errors.As(err, &netErr) { - return true - } - return false -} diff --git a/internal/provider/errors_test.go b/internal/provider/errors_test.go index eed23978..e6e9ac4d 100644 --- a/internal/provider/errors_test.go +++ b/internal/provider/errors_test.go @@ -1,10 +1,8 @@ package provider import ( - "context" "errors" "fmt" - "io" "net/http" "strings" "testing" @@ -172,218 +170,3 @@ func TestProviderError_As(t *testing.T) { t.Fatalf("expected retryable") } } - -// --- IsRecoverableStreamError 全分支覆盖 --- - -func TestIsRecoverableStreamError_Nil(t *testing.T) { - t.Parallel() - if IsRecoverableStreamError(nil) { - t.Fatal("nil error should not be recoverable") - } -} - -func TestIsRecoverableStreamError_ContextErrors_NotRecoverable(t *testing.T) { - t.Parallel() - - tests := []struct { - name string - err error - }{ - {"context.Canceled", context.Canceled}, - {"context.DeadlineExceeded", context.DeadlineExceeded}, - {"wrapped Canceled", fmt.Errorf("wrap: %w", context.Canceled)}, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - t.Parallel() - if IsRecoverableStreamError(tt.err) { - t.Fatalf("%v should not be recoverable", tt.err) - } - }) - } -} - -func TestIsRecoverableStreamError_BufferOverflow_NotRecoverable(t *testing.T) { - t.Parallel() - - tests := []struct { - name string - sentinel error - }{ - {"ErrLineTooLong", ErrLineTooLong}, - {"ErrStreamTooLarge", ErrStreamTooLarge}, - {"wrapped ErrLineTooLong", fmt.Errorf("read: %w", ErrLineTooLong)}, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - t.Parallel() - if IsRecoverableStreamError(tt.sentinel) { - t.Fatalf("%v should not be recoverable", tt.sentinel) - } - }) - } -} - -func TestIsRecoverableStreamError_StreamInterrupted_Recoverable(t *testing.T) { - t.Parallel() - if !IsRecoverableStreamError(ErrStreamInterrupted) { - t.Fatal("ErrStreamInterrupted should be recoverable") - } - wrapped := fmt.Errorf("stream broken: %w", ErrStreamInterrupted) - if !IsRecoverableStreamError(wrapped) { - t.Fatal("wrapped ErrStreamInterrupted should be recoverable") - } -} - -func TestIsRecoverableStreamError_ProviderError_ByRetryableField(t *testing.T) { - t.Parallel() - - retryable := NewProviderErrorFromStatus(http.StatusTooManyRequests, "rate limit") - if !IsRecoverableStreamError(retryable) { - t.Fatal("429 ProviderError should be recoverable") - } - - serverErr := NewProviderErrorFromStatus(http.StatusInternalServerError, "internal") - if !IsRecoverableStreamError(serverErr) { - t.Fatal("5xx ProviderError should be recoverable") - } - - authErr := NewProviderErrorFromStatus(http.StatusUnauthorized, "bad key") - if IsRecoverableStreamError(authErr) { - t.Fatal("401 ProviderError should NOT be recoverable") - } - - clientErr := NewProviderErrorFromStatus(http.StatusBadRequest, "bad request") - if IsRecoverableStreamError(clientErr) { - t.Fatal("400 ProviderError should NOT be recoverable") - } - - wrappedRetryable := fmt.Errorf("layer1: %w", retryable) - if !IsRecoverableStreamError(wrappedRetryable) { - t.Fatal("wrapped retryable ProviderError should be recoverable") - } -} - -func TestIsRecoverableStreamError_NetOpError_Recoverable(t *testing.T) { - t.Parallel() - - // 模拟网络错误:使用一个包含 "connection reset" 的通用 error - // net.OpError 需要真实网络操作才能产生,这里用包装方式模拟 - genericNetErr := fmt.Errorf("net error: connection reset by peer") - // 注意:真实的 *net.OpError 需要 errors.As 匹配 - // 此处验证非上述已知不可恢复类型时默认返回 false - if IsRecoverableStreamError(genericNetErr) { - // 通用 error(非 OpError/ProviderError/哨兵)默认不恢复 - t.Fatal("generic non-net error should not be recoverable") - } -} - -func TestIsRecoverableStreamError_UnknownError_NotRecoverable(t *testing.T) { - t.Parallel() - - tests := []struct { - name string - err error - }{ - {"generic error", errors.New("something went wrong")}, - {"io.EOF", io.EOF}, - {"io.ErrClosedPipe", io.ErrClosedPipe}, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - t.Parallel() - if IsRecoverableStreamError(tt.err) { - t.Fatalf("%v should not be recoverable", tt.err) - } - }) - } -} - -// --- MarkNonRetryable 测试 --- - -func TestMarkNonRetryable_ProviderError(t *testing.T) { - t.Parallel() - - // Retryable=true 的 ProviderError → Retryable=false - retryable := NewProviderErrorFromStatus(http.StatusInternalServerError, "internal") - if !retryable.Retryable { - t.Fatal("setup: expected retryable") - } - - marked := MarkNonRetryable(retryable) - var pErr *ProviderError - if !errors.As(marked, &pErr) { - t.Fatal("marked error should be *ProviderError") - } - if pErr.Retryable { - t.Fatal("MarkNonRetryable should set Retryable=false") - } - if pErr.StatusCode != 500 || pErr.Code != ErrorCodeServer { - t.Fatalf("MarkNonRetryable should preserve StatusCode and Code, got status=%d code=%s", pErr.StatusCode, pErr.Code) - } - - // 原始对象不受影响 - if !retryable.Retryable { - t.Fatal("original ProviderError should not be mutated") - } -} - -func TestMarkNonRetryable_WrappedProviderError(t *testing.T) { - t.Parallel() - - inner := NewProviderErrorFromStatus(http.StatusTooManyRequests, "rate limited") - wrapped := fmt.Errorf("layer: %w", inner) - - marked := MarkNonRetryable(wrapped) - var pErr *ProviderError - if !errors.As(marked, &pErr) { - t.Fatal("marked error should contain *ProviderError") - } - if pErr.Retryable { - t.Fatal("MarkNonRetryable on wrapped should set Retryable=false") - } -} - -func TestMarkNonRetryable_NonProviderError(t *testing.T) { - t.Parallel() - - // 非 ProviderError → 包装为 ProviderError{Retryable: false} - generic := errors.New("some error") - marked := MarkNonRetryable(generic) - - var pErr *ProviderError - if !errors.As(marked, &pErr) { - t.Fatal("marked error should be *ProviderError") - } - if pErr.Retryable { - t.Fatal("should be non-retryable") - } - if pErr.Code != ErrorCodeUnknown { - t.Fatalf("expected code %s, got %s", ErrorCodeUnknown, pErr.Code) - } - if !strings.Contains(pErr.Message, "some error") { - t.Fatalf("message should contain original error text, got: %s", pErr.Message) - } -} - -func TestMarkNonRetryable_AlreadyNonRetryable(t *testing.T) { - t.Parallel() - - // 已经 Retryable=false 的 ProviderError → 保持 false - nonRetryable := NewProviderErrorFromStatus(http.StatusUnauthorized, "bad key") - if nonRetryable.Retryable { - t.Fatal("setup: 401 should be non-retryable") - } - - marked := MarkNonRetryable(nonRetryable) - var pErr *ProviderError - if !errors.As(marked, &pErr) { - t.Fatal("marked error should be *ProviderError") - } - if pErr.Retryable { - t.Fatal("should still be non-retryable") - } -} diff --git a/internal/provider/openai/openai.go b/internal/provider/openai/openai.go index 6e3a59c7..2e2cbaf8 100644 --- a/internal/provider/openai/openai.go +++ b/internal/provider/openai/openai.go @@ -101,69 +101,9 @@ func (p *Provider) DiscoverModels(ctx context.Context) ([]config.ModelDescriptor return config.MergeModelDescriptors(descriptors), nil } -// Chat 发起 SSE 流式对话请求,支持透明重连。 -// -// 流中途断连时,将已累积的 assistant 消息(文本 + tool call)注入请求上下文, -// 利用 OpenAI 多轮对话语义实现断点续传,对上层调用方透明。 -// 最多重连 maxReconnects 次;不可恢复错误直接返回。 +// Chat 发起 SSE 流式对话请求。 +// 流中途断连或协议错误时直接返回错误,由上层调用方决定重试策略。 func (p *Provider) Chat(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { - const maxReconnects = 3 - - // 保存原始消息列表的副本,避免重连时反复 append 到同一个切片导致上下文污染 - originalMessages := make([]provider.Message, len(req.Messages)) - copy(originalMessages, req.Messages) - - // 跨重连周期持久化的累积状态:已收到的文本和 tool call - var ( - accumText strings.Builder - accumCalls map[int]*provider.ToolCall - ) - - for attempt := 0; attempt <= maxReconnects; attempt++ { - if attempt > 0 { - // 从原始消息出发构造本次请求的完整消息列表 - req.Messages = make([]provider.Message, len(originalMessages), len(originalMessages)+1) - copy(req.Messages, originalMessages) - - // 仅在有实际累积内容时注入 assistant 快照,避免插入空消息 - if accumText.Len() > 0 || len(accumCalls) > 0 { - req.Messages = append(req.Messages, - p.buildAssistantMsg(&accumText, accumCalls)) - } - - // 指数退避等待 - backoff := time.Duration(1< 0 { - calls := make([]provider.ToolCall, 0, len(accumCalls)) - for _, c := range accumCalls { - calls = append(calls, *c) - } - msg.ToolCalls = calls - } - return msg -} - -// mergeToolCallDeltaWithAccum 在 mergeToolCallDelta 的基础上, -// 同步将 tool call 累积状态写入跨周期的 accumCalls(*map[int]*ToolCall)。 -func mergeToolCallDeltaWithAccum( - ctx context.Context, - events chan<- provider.StreamEvent, - accumCalls *map[int]*provider.ToolCall, - delta toolCallDelta, -) error { - if *accumCalls == nil { - *accumCalls = make(map[int]*provider.ToolCall) - } - - // 先确保 accumCalls 中有对应条目 - call, exists := (*accumCalls)[delta.Index] - if !exists { - call = &provider.ToolCall{} - (*accumCalls)[delta.Index] = call - } - - // 复用原有逻辑处理事件发送和局部累积 - return mergeToolCallDelta(ctx, events, *accumCalls, delta) -} - func emitStreamEvent(ctx context.Context, events chan<- provider.StreamEvent, event provider.StreamEvent) error { if events == nil { return nil diff --git a/internal/provider/openai/openai_test.go b/internal/provider/openai/openai_test.go index 9085e54b..bfbca884 100644 --- a/internal/provider/openai/openai_test.go +++ b/internal/provider/openai/openai_test.go @@ -1,11 +1,9 @@ package openai import ( - "bytes" "context" "encoding/json" "errors" - "fmt" "io" "net/http" "net/http/httptest" @@ -618,7 +616,7 @@ func TestProviderConsumeStreamRejectsDirtyJSON(t *testing.T) { t.Fatalf("New() error = %v", err) } - err = provider.consumeStream(context.Background(), strings.NewReader("data: {not-json}\n\n"), make(chan domain.StreamEvent, 1), &strings.Builder{}, new(map[int]*domain.ToolCall)) + err = provider.consumeStream(context.Background(), strings.NewReader("data: {not-json}\n\n"), make(chan domain.StreamEvent, 1)) if err == nil || !strings.Contains(err.Error(), "decode stream chunk") { t.Fatalf("expected dirty JSON decode error, got %v", err) } @@ -1032,318 +1030,8 @@ func TestProviderChatEmitsFullEventStream(t *testing.T) { } } -// --- 透明重连测试 --- - -func TestProviderChatReconnect_OnRecoverableError(t *testing.T) { - t.Setenv(config.OpenAIDefaultAPIKeyEnv, "test-key") - - attempt := 0 - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - attempt++ - w.Header().Set("Content-Type", "text/event-stream") - - if attempt == 1 { - // 第一次请求:返回 5xx(可恢复) - w.WriteHeader(http.StatusInternalServerError) - _, _ = w.Write([]byte(`{"error":{"message":"temporarily unavailable"}}`)) - return - } - - // 第二次请求:正常返回 - writeSSEChunk(t, w, map[string]any{ - "choices": []map[string]any{ - {"index": 0, "delta": map[string]any{"content": "recovered"}}, - }, - }) - _, _ = w.Write([]byte("data: [DONE]\n\n")) - })) - defer server.Close() - - p, err := New(resolvedConfig(server.URL, config.OpenAIDefaultModel)) - if err != nil { - t.Fatalf("New() error = %v", err) - } - p.client = server.Client() - - events := make(chan domain.StreamEvent, 8) - err = p.Chat(context.Background(), domain.ChatRequest{ - Model: config.OpenAIDefaultModel, - Messages: []domain.Message{{Role: "user", Content: "hello"}}, - }, events) - if err != nil { - t.Fatalf("Chat() should succeed after reconnect, got: %v", err) - } - - drained := drainStreamEvents(events) - var foundText bool - for _, evt := range drained { - if evt.Type == domain.StreamEventTextDelta { - foundText = true - } - } - if !foundText { - t.Fatal("expected text_delta event after reconnect") - } - if attempt < 2 { - t.Fatalf("expected at least 2 attempts, got %d", attempt) - } -} - -func TestProviderChatReconnect_NonRecoverableError_StopsImmediately(t *testing.T) { - t.Setenv(config.OpenAIDefaultAPIKeyEnv, "test-key") - - attempt := 0 - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - attempt++ - // 返回 401(不可恢复)→ 应立即停止,不重试 - w.WriteHeader(http.StatusUnauthorized) - _, _ = w.Write([]byte(`{"error":{"message":"invalid key"}}`)) - })) - defer server.Close() - - p, err := New(resolvedConfig(server.URL, config.OpenAIDefaultModel)) - if err != nil { - t.Fatalf("New() error = %v", err) - } - p.client = server.Client() - - err = p.Chat(context.Background(), domain.ChatRequest{ - Model: config.OpenAIDefaultModel, - Messages: []domain.Message{{Role: "user", Content: "hello"}}, - }, make(chan domain.StreamEvent, 1)) - if err == nil { - t.Fatal("expected error for 401") - } - if !strings.Contains(err.Error(), "invalid key") { - t.Fatalf("expected auth error, got: %v", err) - } - if attempt > 1 { - t.Fatalf("non-recoverable error should stop immediately, but got %d attempts", attempt) - } -} - -func TestProviderChatReconnect_MaxRetriesExhausted(t *testing.T) { - t.Setenv(config.OpenAIDefaultAPIKeyEnv, "test-key") - - attempt := 0 - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - attempt++ - w.WriteHeader(http.StatusBadGateway) - _, _ = w.Write([]byte(`bad gateway`)) - })) - defer server.Close() - - p, err := New(resolvedConfig(server.URL, config.OpenAIDefaultModel)) - if err != nil { - t.Fatalf("New() error = %v", err) - } - p.client = server.Client() - - err = p.Chat(context.Background(), domain.ChatRequest{ - Model: config.OpenAIDefaultModel, - Messages: []domain.Message{{Role: "user", Content: "hello"}}, - }, make(chan domain.StreamEvent, 1)) - if err == nil { - t.Fatal("expected error after exhausting retries") - } - // 重连耗尽后,错误应被标记为不可重试,防止上层 runtime 再次重试叠加放大。 - var pErr *domain.ProviderError - if !errors.As(err, &pErr) { - t.Fatalf("expected *ProviderError after retry exhaustion, got: %T: %v", err, err) - } - if pErr.Retryable { - t.Fatal("error should be non-retryable after reconnect exhaustion") - } - // 初始1次 + 最大3次重连 = 最多4次尝试 - if attempt > 4 { - t.Fatalf("too many attempts: %d (max should be 4)", attempt) - } -} - -func TestProviderChatReconnect_InjectsAccumulatedContext(t *testing.T) { - t.Setenv(config.OpenAIDefaultAPIKeyEnv, "test-key") - - p, err := New(resolvedConfig(config.OpenAIDefaultBaseURL, config.OpenAIDefaultModel)) - if err != nil { - t.Fatalf("New() error = %v", err) - } - - // 构造一个先返回有效 SSE 数据再中断的 reader,验证 consumeStream - // 在中断前正确累积 accumText 和 accumCalls,且错误为 ErrStreamInterrupted。 - sseData := "data: {\"choices\":[{\"index\":0,\"delta\":{\"content\":\"partial \"}}]}\n\n" + - "data: {\"choices\":[{\"index\":0,\"delta\":{\"tool_calls\":[{\"index\":0,\"id\":\"c1\",\"type\":\"function\",\"function\":{\"name\":\"bash\",\"arguments\":\"run\"}}]}}]}\n" - - reader := io.MultiReader(strings.NewReader(sseData), &errReader{err: io.ErrClosedPipe}) - events := make(chan domain.StreamEvent, 8) - accumText := &strings.Builder{} - accumCalls := make(map[int]*domain.ToolCall) - - err = p.consumeStream(context.Background(), reader, events, accumText, &accumCalls) - if err == nil { - t.Fatal("expected error from interrupted stream") - } - if !errors.Is(err, domain.ErrStreamInterrupted) { - t.Fatalf("expected ErrStreamInterrupted, got: %v", err) - } - - // 验证累积状态:文本和 tool call 都应已保留 - if accumText.String() != "partial " { - t.Fatalf("expected accumText %q, got %q", "partial ", accumText.String()) - } - call, ok := accumCalls[0] - if !ok { - t.Fatal("expected accumCalls[0] to exist") - } - if call.ID != "c1" || call.Name != "bash" { - t.Fatalf("expected tool call c1/bash, got %+v", call) - } - if call.Arguments != "run" { - t.Fatalf("expected arguments %q, got %q", "run", call.Arguments) - } - - // 验证累积状态可用于 buildAssistantMsg - msg := p.buildAssistantMsg(accumText, accumCalls) - if msg.Role != domain.RoleAssistant { - t.Fatalf("expected role assistant, got %q", msg.Role) - } - if !strings.Contains(msg.Content, "partial ") { - t.Fatalf("expected assistant content to contain 'partial', got %q", msg.Content) - } - if len(msg.ToolCalls) != 1 || msg.ToolCalls[0].ID != "c1" { - t.Fatalf("expected assistant tool calls to contain c1, got %+v", msg.ToolCalls) - } -} - // --- 辅助方法测试 --- -func TestBuildAssistantMsg_TextOnly(t *testing.T) { - t.Parallel() - - p, _ := New(resolvedConfig("", "")) - var accumText strings.Builder - accumText.WriteString("hello world") - - msg := p.buildAssistantMsg(&accumText, nil) - if msg.Role != domain.RoleAssistant { - t.Fatalf("expected role assistant, got %q", msg.Role) - } - if msg.Content != "hello world" { - t.Fatalf("expected content %q, got %q", "hello world", msg.Content) - } - if len(msg.ToolCalls) != 0 { - t.Fatalf("expected no tool calls, got %+v", msg.ToolCalls) - } -} - -func TestBuildAssistantMsg_WithToolCalls(t *testing.T) { - t.Parallel() - - p, _ := New(resolvedConfig("", "")) - var accumText strings.Builder - accumText.WriteString("done") - accumCalls := map[int]*domain.ToolCall{ - 0: {ID: "call_1", Name: "edit", Arguments: `{"path":"f.go"}`}, - 1: {ID: "call_2", Name: "read", Arguments: `{"path":"f.go"}`}, - } - - msg := p.buildAssistantMsg(&accumText, accumCalls) - if msg.Content != "done" { - t.Fatalf("content mismatch") - } - if len(msg.ToolCalls) != 2 { - t.Fatalf("expected 2 tool calls, got %d", len(msg.ToolCalls)) - } - // 使用 map 检查,避免依赖 Go map 迭代的不确定顺序 - names := make(map[string]bool) - for _, tc := range msg.ToolCalls { - names[tc.Name] = true - } - if !names["edit"] || !names["read"] { - t.Fatalf("expected tool calls 'edit' and 'read', got %+v", msg.ToolCalls) - } -} - -func TestBuildAssistantMsg_EmptyAccum(t *testing.T) { - t.Parallel() - - p, _ := New(resolvedConfig("", "")) - var accumText strings.Builder - - msg := p.buildAssistantMsg(&accumText, nil) - if msg.Content != "" { - t.Fatalf("expected empty content, got %q", msg.Content) - } - if msg.ToolCalls != nil { - t.Fatal("expected nil ToolCalls when accum is nil") - } -} - -func TestMergeToolCallDeltaWithAccum_SyncsExternalState(t *testing.T) { - t.Parallel() - - events := make(chan domain.StreamEvent, 4) - accumCalls := make(map[int]*domain.ToolCall) - - delta1 := toolCallDelta{ - Index: 0, - ID: "call_acc", - Function: openAIFunctionCall{ - Name: "bash", - Arguments: `{"cmd":"ls"`, - }, - } - if err := mergeToolCallDeltaWithAccum(context.Background(), events, &accumCalls, delta1); err != nil { - t.Fatalf("first delta error = %v", err) - } - - delta2 := toolCallDelta{ - Index: 0, - Function: openAIFunctionCall{ - Arguments: `"}`, - }, - } - if err := mergeToolCallDeltaWithAccum(context.Background(), events, &accumCalls, delta2); err != nil { - t.Fatalf("second delta error = %v", err) - } - - // 验证外部 accumCalls 状态已同步 - call, ok := accumCalls[0] - if !ok { - t.Fatal("expected accumCalls[0] to exist") - } - if call.ID != "call_acc" || call.Name != "bash" { - t.Fatalf("unexpected call state: %+v", call) - } - if call.Arguments != `{"cmd":"ls""}` { - t.Fatalf("expected arguments %q, got %q", `{"cmd":"ls""}`, call.Arguments) - } -} - -func TestMergeToolCallDeltaWithAccum_NilMapInitializes(t *testing.T) { - t.Parallel() - - var accumCalls map[int]*domain.ToolCall // nil map(非指针) - - delta := toolCallDelta{ - Index: 2, - ID: "call_nil", - Function: openAIFunctionCall{ - Name: "read", - }, - } - events := make(chan domain.StreamEvent, 2) - if err := mergeToolCallDeltaWithAccum(context.Background(), events, &accumCalls, delta); err != nil { - t.Fatalf("error = %v", err) - } - - if accumCalls == nil { - t.Fatal("expected accumCalls to be initialized from nil") - } - if accumCalls[2] == nil || accumCalls[2].Name != "read" { - t.Fatalf("unexpected accumCalls[2]: %+v", accumCalls[2]) - } -} - // --- consumeStream 错误包装测试 --- func TestConsumeStream_WrapsNonEOFAsInterrupted(t *testing.T) { @@ -1356,10 +1044,8 @@ func TestConsumeStream_WrapsNonEOFAsInterrupted(t *testing.T) { // 使用一个会触发读取错误的 source(模拟网络断开) errReader := &errReader{err: io.ErrClosedPipe} - accumText := &strings.Builder{} - accums := make(map[int]*domain.ToolCall) - err = p.consumeStream(context.Background(), errReader, make(chan domain.StreamEvent, 1), accumText, &accums) + err = p.consumeStream(context.Background(), errReader, make(chan domain.StreamEvent, 1)) if err == nil { t.Fatal("expected error for broken reader") } @@ -1385,10 +1071,8 @@ func TestConsumeStream_FlushesPendingDataOnNonEOFError(t *testing.T) { body := io.MultiReader(strings.NewReader(sseData), &errReader{err: io.ErrClosedPipe}) events := make(chan domain.StreamEvent, 10) - accumText := &strings.Builder{} - accums := make(map[int]*domain.ToolCall) - err = p.consumeStream(context.Background(), body, events, accumText, &accums) + err = p.consumeStream(context.Background(), body, events) if err == nil { t.Fatal("expected error for broken reader") } @@ -1396,9 +1080,16 @@ func TestConsumeStream_FlushesPendingDataOnNonEOFError(t *testing.T) { t.Fatalf("expected ErrStreamInterrupted, got: %v", err) } - // 关键断言:中断前的 data 行必须已被刷新处理,文本累积不为空。 - if accumText.String() != "hello" { - t.Fatalf("expected accumText 'hello', got %q", accumText.String()) + // 关键断言:中断前的 data 行必须已被刷新处理,应有 text_delta 事件。 + drained := drainStreamEvents(events) + var foundText bool + for _, evt := range drained { + if evt.Type == domain.StreamEventTextDelta { + foundText = true + } + } + if !foundText { + t.Fatal("expected text_delta event from flushed pending data") } } @@ -1410,221 +1101,3 @@ type errReader struct { func (e *errReader) Read(p []byte) (int, error) { return 0, e.err } - -// roundTripperFunc 将函数适配为 http.RoundTripper 接口,用于测试中 mock HTTP 行为。 -type roundTripperFunc func(*http.Request) (*http.Response, error) - -func (f roundTripperFunc) RoundTrip(req *http.Request) (*http.Response, error) { - return f(req) -} - -// --- 重连消息完整性测试 --- - -// TestReconnect_NoEmptyAssistantOnFirstFailure 验证首次请求失败(未收到任何 SSE 数据) -// 时,重连不会向消息列表注入空的 assistant 消息。 -func TestReconnect_NoEmptyAssistantOnFirstFailure(t *testing.T) { - t.Setenv(config.OpenAIDefaultAPIKeyEnv, "test-key") - - attempt := 0 - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - attempt++ - - // 解码请求中的消息,验证消息结构 - var payload chatCompletionRequest - if err := json.NewDecoder(r.Body).Decode(&payload); err != nil { - t.Errorf("decode request: %v", err) - w.WriteHeader(http.StatusBadRequest) - return - } - - if attempt == 1 { - // 首次请求:返回 500(无 SSE 数据),不应注入空 assistant 消息 - if len(payload.Messages) != 1 { - t.Errorf("first attempt: expected 1 message, got %d", len(payload.Messages)) - } - w.WriteHeader(http.StatusInternalServerError) - _, _ = w.Write([]byte(`{"error":{"message":"temporarily unavailable"}}`)) - return - } - - // 第二次请求:仍应只有原始 1 条消息(无空 assistant 注入) - if len(payload.Messages) != 1 { - t.Errorf("second attempt: expected 1 message (no empty assistant), got %d; messages: %+v", - len(payload.Messages), payload.Messages) - } - for _, msg := range payload.Messages { - if msg.Role == "assistant" { - t.Errorf("second attempt: unexpected assistant message injected: %+v", msg) - } - } - - w.Header().Set("Content-Type", "text/event-stream") - writeSSEChunk(t, w, map[string]any{ - "choices": []map[string]any{ - {"index": 0, "delta": map[string]any{"content": "ok"}}, - }, - }) - _, _ = w.Write([]byte("data: [DONE]\n\n")) - })) - defer server.Close() - - p, err := New(resolvedConfig(server.URL, config.OpenAIDefaultModel)) - if err != nil { - t.Fatalf("New() error = %v", err) - } - p.client = server.Client() - - events := make(chan domain.StreamEvent, 4) - err = p.Chat(context.Background(), domain.ChatRequest{ - Model: config.OpenAIDefaultModel, - Messages: []domain.Message{{Role: "user", Content: "hello"}}, - }, events) - if err != nil { - t.Fatalf("Chat() should succeed after reconnect, got: %v", err) - } - if attempt != 2 { - t.Fatalf("expected exactly 2 attempts, got %d", attempt) - } -} - -// TestReconnect_SingleAssistantSnapshotNotDuplicated 验证多次重连时,每次请求 -// 只包含原始消息 + 恰好 1 条 assistant 快照,不会出现旧快照残留。 -// 使用自定义 RoundTripper 模拟流中途中断(非 EOF 错误)。 -func TestReconnect_SingleAssistantSnapshotNotDuplicated(t *testing.T) { - t.Setenv(config.OpenAIDefaultAPIKeyEnv, "test-key") - - attempt := 0 - - // 使用 httptest.Server 作为成功响应的代理,RoundTripper 控制中断 - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("Content-Type", "text/event-stream") - writeSSEChunk(t, w, map[string]any{ - "choices": []map[string]any{ - {"index": 0, "delta": map[string]any{"content": " done"}}, - }, - }) - _, _ = w.Write([]byte("data: [DONE]\n\n")) - })) - defer server.Close() - - rt := roundTripperFunc(func(req *http.Request) (*http.Response, error) { - attempt++ - - // 读取并缓存请求体,以便解码后仍可转发给真实 transport - bodyBytes, err := io.ReadAll(req.Body) - _ = req.Body.Close() - if err != nil { - return nil, fmt.Errorf("read request body: %w", err) - } - - // 解码请求消息,用于断言 - var payload chatCompletionRequest - if err := json.Unmarshal(bodyBytes, &payload); err != nil { - return nil, fmt.Errorf("decode request: %w", err) - } - - // 恢复请求体,确保后续 transport 可以读取 - req.Body = io.NopCloser(bytes.NewReader(bodyBytes)) - req.ContentLength = int64(len(bodyBytes)) - - switch attempt { - case 1: - // 首次请求:正常消息(system + user),发送部分 SSE 后流中断 - if len(payload.Messages) != 2 { - t.Errorf("attempt 1: expected 2 messages, got %d", len(payload.Messages)) - } - sseData := "data: {\"choices\":[{\"index\":0,\"delta\":{\"content\":\"partial\"}}]}\n\n" - body := io.MultiReader(strings.NewReader(sseData), &errReader{err: io.ErrClosedPipe}) - return &http.Response{ - StatusCode: http.StatusOK, - Header: http.Header{"Content-Type": []string{"text/event-stream"}}, - Body: io.NopCloser(body), - }, nil - - case 2: - // 第二次请求:应包含原始 2 条 + 1 条 assistant("partial"),再次中断 - if len(payload.Messages) != 3 { - t.Errorf("attempt 2: expected 3 messages, got %d; messages: %+v", - len(payload.Messages), payload.Messages) - } - assistMsg := payload.Messages[2] - if assistMsg.Role != "assistant" || assistMsg.Content != "partial" { - t.Errorf("attempt 2: expected assistant content 'partial', got role=%q content=%q", - assistMsg.Role, assistMsg.Content) - } - assistCount := countMessagesByRole(payload.Messages, "assistant") - if assistCount != 1 { - t.Errorf("attempt 2: expected 1 assistant message, got %d", assistCount) - } - - sseData := "data: {\"choices\":[{\"index\":0,\"delta\":{\"content\":\" more\"}}]}\n\n" - body := io.MultiReader(strings.NewReader(sseData), &errReader{err: io.ErrClosedPipe}) - return &http.Response{ - StatusCode: http.StatusOK, - Header: http.Header{"Content-Type": []string{"text/event-stream"}}, - Body: io.NopCloser(body), - }, nil - - default: - // 第三次请求:应包含原始 2 条 + 1 条 assistant("partial more") - if len(payload.Messages) != 3 { - t.Errorf("attempt 3: expected 3 messages, got %d; messages: %+v", - len(payload.Messages), payload.Messages) - } - assistMsg := payload.Messages[2] - if assistMsg.Role != "assistant" || assistMsg.Content != "partial more" { - t.Errorf("attempt 3: expected assistant content 'partial more', got role=%q content=%q", - assistMsg.Role, assistMsg.Content) - } - // 确认仍然只有 1 条 assistant 消息(旧快照未残留) - assistCount := countMessagesByRole(payload.Messages, "assistant") - if assistCount != 1 { - t.Errorf("attempt 3: expected 1 assistant message, got %d", assistCount) - } - - // 委托给 test server 返回完整响应 - return http.DefaultTransport.RoundTrip(req) - } - }) - - p, err := New(resolvedConfig(server.URL, config.OpenAIDefaultModel), withTransport(rt)) - if err != nil { - t.Fatalf("New() error = %v", err) - } - - events := make(chan domain.StreamEvent, 16) - err = p.Chat(context.Background(), domain.ChatRequest{ - SystemPrompt: "you are helpful", - Model: config.OpenAIDefaultModel, - Messages: []domain.Message{{Role: "user", Content: "hello"}}, - }, events) - if err != nil { - t.Fatalf("Chat() should succeed after reconnect, got: %v", err) - } - if attempt != 3 { - t.Fatalf("expected exactly 3 attempts, got %d", attempt) - } - - // 验证最终累积的文本(三次请求的增量合并) - var fullText strings.Builder - for _, evt := range drainStreamEvents(events) { - if evt.Type == domain.StreamEventTextDelta { - fullText.WriteString(requireTextDeltaPayload(t, evt).Text) - } - } - expectedText := "partial more done" - if fullText.String() != expectedText { - t.Fatalf("expected full text %q, got %q", expectedText, fullText.String()) - } -} - -// countMessagesByRole 统计消息列表中指定角色的消息数量。 -func countMessagesByRole(messages []openAIMessage, role string) int { - count := 0 - for _, msg := range messages { - if msg.Role == role { - count++ - } - } - return count -} From 2c45505efb93df43284186c2a87f71cddd91d25c Mon Sep 17 00:00:00 2001 From: creatang Date: Mon, 6 Apr 2026 14:08:23 +0800 Subject: [PATCH 25/55] refactor(tui): introduce layered state package and boundary docs --- internal/tui/docs/LAYERING.md | 94 +++++++++++++++++++++++++++++ internal/tui/docs/SKILL.md | 55 +++++++++++++++++ internal/tui/state/.gitkeep | 0 internal/tui/state/chat_state.go | 17 ++++++ internal/tui/state/constants.go | 15 +++++ internal/tui/state/messages.go | 53 ++++++++++++++++ internal/tui/state/runtime_state.go | 43 +++++++++++++ internal/tui/state/state_test.go | 25 ++++++++ internal/tui/state/ui_state.go | 47 +++++++++++++++ 9 files changed, 349 insertions(+) create mode 100644 internal/tui/docs/LAYERING.md create mode 100644 internal/tui/docs/SKILL.md create mode 100644 internal/tui/state/.gitkeep create mode 100644 internal/tui/state/chat_state.go create mode 100644 internal/tui/state/constants.go create mode 100644 internal/tui/state/messages.go create mode 100644 internal/tui/state/runtime_state.go create mode 100644 internal/tui/state/state_test.go create mode 100644 internal/tui/state/ui_state.go diff --git a/internal/tui/docs/LAYERING.md b/internal/tui/docs/LAYERING.md new file mode 100644 index 00000000..b15d1833 --- /dev/null +++ b/internal/tui/docs/LAYERING.md @@ -0,0 +1,94 @@ +# TUI 分层约束(Iteration 0) + +本文档用于约束 `internal/tui` 的分层职责与依赖方向,确保后续迭代按层收敛,不跨层扩散。 + +## 改造范围 + +- 本轮只处理 `internal/tui`。 +- 入口层 `cmd/tui` 暂不处理。 + +## 分层定义 + +### L1 - Entry(暂缓) + +- 位置:`cmd/tui/` +- 职责:参数解析、终端初始化、启动 Program。 +- 本轮状态:暂不纳入改造。 + +### L2 - Bootstrap + +- 位置:`internal/tui/bootstrap/` +- 职责:依赖注入(DI)与初始化编排。 +- 负责:工作区/配置初始化、服务装配、Offline/Mock 注入切换。 + +### L3 - App/Core + +- 位置:`internal/tui/core/` +- 职责:Bubble Tea 状态机中枢(ELM 单向数据流)。 +- 负责:消息路由、状态变更、布局调度。 + +### L4 - State + +- 位置:`internal/tui/state/` +- 职责:纯数据容器。 +- 约束:只放结构体和常量,不放方法与副作用。 + +### L5 - Component Adapter + +- 位置:`internal/tui/components/` +- 职责:原子渲染组件。 +- 输入:基础数据或 state。 +- 输出:渲染字符串。 + +### L6 - Services + +- 位置:`internal/tui/services/` +- 职责:对接 runtime/provider/本地系统能力。 +- 约束:统一返回 `tea.Cmd` 或异步产出 `tea.Msg`。 + +### L7 - Infrastructure + +- 位置:`internal/tui/infra/` +- 职责:底层 I/O 与系统能力。 +- 范围:shell 执行、文件扫描、终端 I/O、渲染器、剪贴板等。 + +## 依赖方向(允许) + +- `core` -> `state` +- `core` -> `components` +- `core` -> `services` +- `services` -> `infra` + +## 禁止项 + +- 禁止 `components` 直接访问 runtime/provider 或执行外部 I/O。 +- 禁止 `core` 直接调用底层系统能力(应经 `services`)。 +- 禁止 `state` 承载业务逻辑、网络调用或文件操作。 +- 禁止新增跨层直连(例如 `core` 直接依赖 `infra`)。 +- 禁止在本轮引入行为变更;Iteration 0 只做骨架与规则。 + +## Iteration 0 验收 + +- 目录骨架已创建:`bootstrap/core/state/components/services/infra` +- 分层约束文档已建立 +- `go test ./internal/tui/...` 通过 + +## Iteration 6 补充(Bootstrap 落地) + +- `internal/tui/bootstrap` 已提供 `Build` 装配入口,统一完成 `ConfigManager + Runtime + ProviderService` 注入。 +- 支持 `Mode`(`live/offline/mock`)与 `ServiceFactory` 扩展点,可在不修改 `core` 的情况下替换注入实现。 +- `internal/tui.New(...)` 保持兼容签名,对外作为薄封装;实际装配路径为 `New -> bootstrap.Build -> newApp`。 + +## Iteration 7 补充(Runtime Source 收敛) + +- Runtime 事件新增并接入 UI 桥接: + - `EventToolStatus` + - `EventRunContext` + - `EventUsage` +- Runtime 查询接口已落地: + - `GetRunSnapshot(runID)` + - `GetSessionContext(sessionID)` + - `GetSessionUsage(sessionID)` + - `GetRunUsage(runID)` +- `internal/tui/core/runtime_bridge.go` 统一处理 payload -> VM 映射与 Tool 状态去重合并(覆盖重复/乱序事件场景)。 +- TUI 在会话刷新时优先通过 runtime 查询回填 context/token 快照,避免由 UI 本地推导。 diff --git a/internal/tui/docs/SKILL.md b/internal/tui/docs/SKILL.md new file mode 100644 index 00000000..361d8f57 --- /dev/null +++ b/internal/tui/docs/SKILL.md @@ -0,0 +1,55 @@ +--- +name: bubbletea +description: Browse Bubbletea TUI framework documentation and examples. Use when working with Bubbletea components, models, commands, or building terminal user interfaces in Go. +--- + +# Bubbletea Documentation + +Bubbletea is a Go framework for building terminal user interfaces based on The Elm Architecture. + +## Key Resources + +When you need to understand Bubbletea patterns or find examples: + +1. **Examples README** - Overview of all available examples: + https://github.com/charmbracelet/bubbletea/blob/main/examples/README.md + +2. **Examples Directory** - Full source code for all examples: + https://github.com/charmbracelet/bubbletea/tree/main/examples + +## How to Use + +1. First, fetch the examples README to get an overview of available examples: + + ``` + WebFetch https://github.com/charmbracelet/bubbletea/blob/main/examples/README.md + ``` + +2. Once you identify a relevant example, fetch its source code from the examples directory. + +## Common Examples to Reference + +- `list` - List component with filtering +- `table` - Table component +- `textinput` - Text input handling +- `textarea` - Multi-line text input +- `viewport` - Scrollable content +- `paginator` - Pagination +- `spinner` - Loading spinners +- `progress` - Progress bars +- `tabs` - Tab navigation +- `help` - Help text/keybindings display + +## Core Concepts + +- **Model**: Application state +- **Update**: Handles messages and returns updated model + commands +- **View**: Renders the model to a string +- **Cmd**: Side effects that produce messages +- **Msg**: Events that trigger updates + +## Related Charm Libraries + +- **Bubbles**: Pre-built components (github.com/charmbracelet/bubbles) +- **Lipgloss**: Styling and layout (github.com/charmbracelet/lipgloss) +- **Glamour**: Markdown rendering (github.com/charmbracelet/glamour) diff --git a/internal/tui/state/.gitkeep b/internal/tui/state/.gitkeep new file mode 100644 index 00000000..e69de29b diff --git a/internal/tui/state/chat_state.go b/internal/tui/state/chat_state.go new file mode 100644 index 00000000..18440412 --- /dev/null +++ b/internal/tui/state/chat_state.go @@ -0,0 +1,17 @@ +package state + +import "time" + +// ActivityEntry 表示 Activity 面板中的单条事件记录。 +type ActivityEntry struct { + Time time.Time + Kind string + Title string + Detail string + IsError bool +} + +// CommandMenuMeta 表示命令建议菜单的标题等元信息。 +type CommandMenuMeta struct { + Title string +} diff --git a/internal/tui/state/constants.go b/internal/tui/state/constants.go new file mode 100644 index 00000000..3205ff91 --- /dev/null +++ b/internal/tui/state/constants.go @@ -0,0 +1,15 @@ +package state + +import "time" + +// 这些常量定义了输入框与粘贴检测在 Update 流程中的基础行为阈值。 +const ( + ComposerMinHeight = 1 + ComposerMaxHeight = 5 + ComposerPromptWidth = 2 + MouseWheelStepLines = 3 + PasteBurstWindow = 120 * time.Millisecond + PasteEnterGuard = 180 * time.Millisecond + PasteSessionGuard = 5 * time.Second + PasteBurstThreshold = 12 +) diff --git a/internal/tui/state/messages.go b/internal/tui/state/messages.go new file mode 100644 index 00000000..9016281d --- /dev/null +++ b/internal/tui/state/messages.go @@ -0,0 +1,53 @@ +package state + +import ( + "neo-code/internal/config" + agentruntime "neo-code/internal/runtime" +) + +// RuntimeMsg 封装 runtime 事件流消息。 +type RuntimeMsg struct { + Event agentruntime.RuntimeEvent +} + +// RuntimeClosedMsg 表示 runtime 事件通道已关闭。 +type RuntimeClosedMsg struct{} + +// RunFinishedMsg 表示一次 Run 调用结束。 +type RunFinishedMsg struct { + Err error +} + +// ModelCatalogRefreshMsg 表示模型目录刷新结果。 +type ModelCatalogRefreshMsg struct { + ProviderID string + Models []config.ModelDescriptor + Err error +} + +// CompactFinishedMsg 表示 compact 调用结束。 +type CompactFinishedMsg struct { + Err error +} + +// LocalCommandResultMsg 表示本地命令执行结果。 +type LocalCommandResultMsg struct { + Notice string + Err error + ProviderChanged bool + ModelChanged bool +} + +// SessionWorkdirResultMsg 表示会话工作目录命令结果。 +type SessionWorkdirResultMsg struct { + Notice string + Workdir string + Err error +} + +// WorkspaceCommandResultMsg 表示工作区命令执行结果。 +type WorkspaceCommandResultMsg struct { + Command string + Output string + Err error +} diff --git a/internal/tui/state/runtime_state.go b/internal/tui/state/runtime_state.go new file mode 100644 index 00000000..c3dda088 --- /dev/null +++ b/internal/tui/state/runtime_state.go @@ -0,0 +1,43 @@ +package state + +import "time" + +// ToolLifecycleStatus 描述工具执行生命周期状态。 +type ToolLifecycleStatus string + +const ( + ToolLifecyclePlanned ToolLifecycleStatus = "planned" + ToolLifecycleRunning ToolLifecycleStatus = "running" + ToolLifecycleSucceeded ToolLifecycleStatus = "succeeded" + ToolLifecycleFailed ToolLifecycleStatus = "failed" +) + +// ToolState 记录单个工具调用在 UI 中展示的状态。 +type ToolState struct { + ToolCallID string + ToolName string + Status ToolLifecycleStatus + Message string + DurationMS int64 + UpdatedAt time.Time +} + +// ContextWindowState 描述 runtime 透出的上下文窗口信息。 +type ContextWindowState struct { + RunID string + SessionID string + Provider string + Model string + Workdir string + Mode string +} + +// TokenUsageState 描述 token 统计在 UI 的展示结构。 +type TokenUsageState struct { + RunInputTokens int + RunOutputTokens int + RunTotalTokens int + SessionInputTokens int + SessionOutputTokens int + SessionTotalTokens int +} diff --git a/internal/tui/state/state_test.go b/internal/tui/state/state_test.go new file mode 100644 index 00000000..73e34cda --- /dev/null +++ b/internal/tui/state/state_test.go @@ -0,0 +1,25 @@ +package state + +import "testing" + +func TestPanelAndPickerConstants(t *testing.T) { + if PanelSessions != 0 || PanelTranscript != 1 || PanelActivity != 2 || PanelInput != 3 { + t.Fatalf("unexpected panel constants: %d %d %d %d", PanelSessions, PanelTranscript, PanelActivity, PanelInput) + } + if PickerNone != 0 || PickerProvider != 1 || PickerModel != 2 || PickerFile != 3 { + t.Fatalf("unexpected picker constants: %d %d %d %d", PickerNone, PickerProvider, PickerModel, PickerFile) + } +} + +func TestUIStateCarriesFocusAndPicker(t *testing.T) { + s := UIState{ + Focus: PanelInput, + ActivePicker: PickerModel, + } + if s.Focus != PanelInput { + t.Fatalf("expected focus panel input, got %v", s.Focus) + } + if s.ActivePicker != PickerModel { + t.Fatalf("expected model picker, got %v", s.ActivePicker) + } +} diff --git a/internal/tui/state/ui_state.go b/internal/tui/state/ui_state.go new file mode 100644 index 00000000..ac6c4611 --- /dev/null +++ b/internal/tui/state/ui_state.go @@ -0,0 +1,47 @@ +package state + +import agentruntime "neo-code/internal/runtime" + +// Panel 定义 TUI 中可聚焦的主面板。 +type Panel int + +const ( + PanelSessions Panel = iota + PanelTranscript + PanelActivity + PanelInput +) + +// PickerMode 定义当前激活的选择器类型。 +type PickerMode int + +const ( + PickerNone PickerMode = iota + PickerProvider + PickerModel + PickerFile +) + +// UIState 保存顶层界面状态快照,仅作为数据容器使用。 +type UIState struct { + Sessions []agentruntime.SessionSummary + ActiveSessionID string + ActiveSessionTitle string + ActiveRunID string + InputText string + IsAgentRunning bool + IsCompacting bool + StreamingReply bool + CurrentTool string + ToolStates []ToolState + RunContext ContextWindowState + TokenUsage TokenUsageState + ExecutionError string + StatusText string + CurrentProvider string + CurrentModel string + CurrentWorkdir string + ShowHelp bool + ActivePicker PickerMode + Focus Panel +} From aa7a75517ada4f8532d5c0a115ff65ec75f7e57b Mon Sep 17 00:00:00 2001 From: creatang Date: Mon, 6 Apr 2026 15:26:40 +0800 Subject: [PATCH 26/55] fix(tui/state): remove stray BOM to unblock CI build --- internal/tui/state/runtime_state.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/internal/tui/state/runtime_state.go b/internal/tui/state/runtime_state.go index c3dda088..7909ed11 100644 --- a/internal/tui/state/runtime_state.go +++ b/internal/tui/state/runtime_state.go @@ -1,4 +1,4 @@ -package state +package state import "time" From fe0eb640a3d82cabeb972dea1da6ea63f6ff396d Mon Sep 17 00:00:00 2001 From: creatang Date: Mon, 6 Apr 2026 21:02:29 +0800 Subject: [PATCH 27/55] docs(tui): fix entry path and untrack local SKILL guide --- internal/tui/docs/LAYERING.md | 4 +-- internal/tui/docs/SKILL.md | 55 ----------------------------------- 2 files changed, 2 insertions(+), 57 deletions(-) delete mode 100644 internal/tui/docs/SKILL.md diff --git a/internal/tui/docs/LAYERING.md b/internal/tui/docs/LAYERING.md index b15d1833..4ff7b040 100644 --- a/internal/tui/docs/LAYERING.md +++ b/internal/tui/docs/LAYERING.md @@ -5,13 +5,13 @@ ## 改造范围 - 本轮只处理 `internal/tui`。 -- 入口层 `cmd/tui` 暂不处理。 +- 入口层 `cmd/neocode` 暂不处理。 ## 分层定义 ### L1 - Entry(暂缓) -- 位置:`cmd/tui/` +- 位置:`cmd/neocode/` - 职责:参数解析、终端初始化、启动 Program。 - 本轮状态:暂不纳入改造。 diff --git a/internal/tui/docs/SKILL.md b/internal/tui/docs/SKILL.md deleted file mode 100644 index 361d8f57..00000000 --- a/internal/tui/docs/SKILL.md +++ /dev/null @@ -1,55 +0,0 @@ ---- -name: bubbletea -description: Browse Bubbletea TUI framework documentation and examples. Use when working with Bubbletea components, models, commands, or building terminal user interfaces in Go. ---- - -# Bubbletea Documentation - -Bubbletea is a Go framework for building terminal user interfaces based on The Elm Architecture. - -## Key Resources - -When you need to understand Bubbletea patterns or find examples: - -1. **Examples README** - Overview of all available examples: - https://github.com/charmbracelet/bubbletea/blob/main/examples/README.md - -2. **Examples Directory** - Full source code for all examples: - https://github.com/charmbracelet/bubbletea/tree/main/examples - -## How to Use - -1. First, fetch the examples README to get an overview of available examples: - - ``` - WebFetch https://github.com/charmbracelet/bubbletea/blob/main/examples/README.md - ``` - -2. Once you identify a relevant example, fetch its source code from the examples directory. - -## Common Examples to Reference - -- `list` - List component with filtering -- `table` - Table component -- `textinput` - Text input handling -- `textarea` - Multi-line text input -- `viewport` - Scrollable content -- `paginator` - Pagination -- `spinner` - Loading spinners -- `progress` - Progress bars -- `tabs` - Tab navigation -- `help` - Help text/keybindings display - -## Core Concepts - -- **Model**: Application state -- **Update**: Handles messages and returns updated model + commands -- **View**: Renders the model to a string -- **Cmd**: Side effects that produce messages -- **Msg**: Events that trigger updates - -## Related Charm Libraries - -- **Bubbles**: Pre-built components (github.com/charmbracelet/bubbles) -- **Lipgloss**: Styling and layout (github.com/charmbracelet/lipgloss) -- **Glamour**: Markdown rendering (github.com/charmbracelet/glamour) From 9b5249286470ca64539dd490ab320704dd67379c Mon Sep 17 00:00:00 2001 From: creatang Date: Mon, 6 Apr 2026 21:17:19 +0800 Subject: [PATCH 28/55] test(runtime): align scriptedProvider callers with stream-based chat --- internal/runtime/permission_test.go | 14 ++++---------- internal/runtime/runtime_test.go | 14 ++++---------- 2 files changed, 8 insertions(+), 20 deletions(-) diff --git a/internal/runtime/permission_test.go b/internal/runtime/permission_test.go index 61516346..23102aea 100644 --- a/internal/runtime/permission_test.go +++ b/internal/runtime/permission_test.go @@ -170,19 +170,13 @@ func TestServiceRunPermissionRejectFlow(t *testing.T) { } scripted := &scriptedProvider{ - responses: []scriptedResponse{ + streams: [][]provider.StreamEvent{ { - Message: provider.Message{ - Role: "assistant", - ToolCalls: []provider.ToolCall{ - {ID: "call-ask-reject", Name: "webfetch", Arguments: `{"url":"https://example.com/private"}`}, - }, - }, - FinishReason: "tool_calls", + provider.NewToolCallStartStreamEvent(0, "call-ask-reject", "webfetch"), + provider.NewToolCallDeltaStreamEvent(0, "call-ask-reject", `{"url":"https://example.com/private"}`), }, { - Message: provider.Message{Role: "assistant", Content: "done"}, - FinishReason: "stop", + provider.NewTextDeltaStreamEvent("done"), }, }, } diff --git a/internal/runtime/runtime_test.go b/internal/runtime/runtime_test.go index 25980b8d..64f32bb5 100644 --- a/internal/runtime/runtime_test.go +++ b/internal/runtime/runtime_test.go @@ -1207,19 +1207,13 @@ func TestServiceRunEmitsRememberScopeWhenSessionRejectMemoryHits(t *testing.T) { } scripted := &scriptedProvider{ - responses: []scriptedResponse{ + streams: [][]provider.StreamEvent{ { - Message: provider.Message{ - Role: "assistant", - ToolCalls: []provider.ToolCall{ - {ID: "call-memory-reject", Name: "webfetch", Arguments: `{"url":"https://example.com/private"}`}, - }, - }, - FinishReason: "tool_calls", + provider.NewToolCallStartStreamEvent(0, "call-memory-reject", "webfetch"), + provider.NewToolCallDeltaStreamEvent(0, "call-memory-reject", `{"url":"https://example.com/private"}`), }, { - Message: provider.Message{Role: "assistant", Content: "done"}, - FinishReason: "stop", + provider.NewTextDeltaStreamEvent("done"), }, }, } From f40048edded305ac7162190edac7ba5f45a3f03c Mon Sep 17 00:00:00 2001 From: phantom5099 <1011668688@qq.com> Date: Tue, 7 Apr 2026 14:41:26 +0800 Subject: [PATCH 29/55] =?UTF-8?q?refactor(provider):=E6=8B=86=E5=88=86type?= =?UTF-8?q?=E5=92=8Copenai?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/session-persistence-design.md | 20 - docs/tool-execution-toctou.md | 50 - internal/context/builder.go | 4 +- internal/context/builder_test.go | 226 +-- internal/context/compact/helpers.go | 10 +- internal/context/compact/planner.go | 14 +- internal/context/compact/planner_test.go | 42 +- internal/context/compact/runner.go | 24 +- internal/context/compact/runner_test.go | 184 +-- internal/context/compact/transcript_store.go | 20 +- .../context/compact/transcript_store_test.go | 8 +- internal/context/compact_prompt.go | 10 +- internal/context/compact_prompt_test.go | 14 +- internal/context/internalcompact/messages.go | 12 +- .../context/internalcompact/messages_test.go | 18 +- internal/context/microcompact.go | 24 +- internal/context/microcompact_test.go | 186 +-- internal/context/trim.go | 8 +- internal/context/trim_policy.go | 8 +- internal/context/types.go | 6 +- internal/provider/catalog/service_test.go | 3 +- internal/provider/discovery/discovery.go | 50 - internal/provider/discovery/discovery_test.go | 51 - internal/provider/openai/discovery.go | 51 + internal/provider/openai/driver.go | 35 + internal/provider/openai/events.go | 63 + internal/provider/openai/openai.go | 501 ------ internal/provider/openai/openai_test.go | 1385 +++++++++++------ internal/provider/openai/provider.go | 121 ++ internal/provider/openai/request.go | 105 ++ internal/provider/openai/response.go | 135 ++ internal/provider/openai/toolcall.go | 43 + internal/provider/openai/types.go | 90 ++ internal/provider/provider.go | 9 +- internal/provider/registry_test.go | 3 +- .../provider/{types.go => types/event.go} | 51 +- internal/provider/types/message.go | 36 + internal/provider/types/request.go | 16 + internal/provider/{ => types}/types_test.go | 2 +- internal/runtime/compact.go | 6 +- internal/runtime/compact_generator.go | 10 +- internal/runtime/compact_generator_test.go | 40 +- internal/runtime/permission.go | 6 +- internal/runtime/permission_test.go | 22 +- internal/runtime/runtime.go | 51 +- internal/runtime/runtime_test.go | 327 ++-- internal/runtime/session.go | 14 +- internal/runtime/session_test.go | 8 +- internal/tools/manager.go | 8 +- internal/tools/manager_test.go | 4 +- internal/tools/registry.go | 12 +- internal/tools/types.go | 4 +- internal/tui/app.go | 4 +- internal/tui/copy_code_test.go | 10 +- internal/tui/update.go | 16 +- internal/tui/update_test.go | 21 +- internal/tui/view.go | 4 +- 57 files changed, 2284 insertions(+), 1921 deletions(-) delete mode 100644 docs/session-persistence-design.md delete mode 100644 docs/tool-execution-toctou.md delete mode 100644 internal/provider/discovery/discovery.go delete mode 100644 internal/provider/discovery/discovery_test.go create mode 100644 internal/provider/openai/discovery.go create mode 100644 internal/provider/openai/driver.go create mode 100644 internal/provider/openai/events.go delete mode 100644 internal/provider/openai/openai.go create mode 100644 internal/provider/openai/provider.go create mode 100644 internal/provider/openai/request.go create mode 100644 internal/provider/openai/response.go create mode 100644 internal/provider/openai/toolcall.go create mode 100644 internal/provider/openai/types.go rename internal/provider/{types.go => types/event.go} (76%) create mode 100644 internal/provider/types/message.go create mode 100644 internal/provider/types/request.go rename internal/provider/{ => types}/types_test.go (99%) diff --git a/docs/session-persistence-design.md b/docs/session-persistence-design.md deleted file mode 100644 index 0c5d1887..00000000 --- a/docs/session-persistence-design.md +++ /dev/null @@ -1,20 +0,0 @@ -# Session 持久化设计 -## 存储策略 -NeoCode 在 MVP 阶段使用 JSON 文件持久化 Session,以保持本地优先、易于调试和跨平台可移植。 - -## 数据模型 -- `Session`:完整消息历史以及 `id`、`title`、`updated_at` 等元信息 -- `SessionSummary`:用于侧边栏的轻量摘要结构 - -## 加载策略 -- `ListSummaries` 只读取渲染侧边栏所需的基础信息 -- `Load` 仅在用户真正进入某个会话时读取完整消息历史 -- `Save` 通过临时文件原子写入完整 Session - -## 命名策略 -- 新会话默认展示为 `Draft` -- 一旦持久化,runtime 会根据首轮用户消息生成简短标题 - -## 并发约束 -- SessionStore 实现必须自行保护共享访问 -- 真正的保存时机由 runtime 决定,TUI 不负责直接触发磁盘写入 diff --git a/docs/tool-execution-toctou.md b/docs/tool-execution-toctou.md deleted file mode 100644 index f4ce8347..00000000 --- a/docs/tool-execution-toctou.md +++ /dev/null @@ -1,50 +0,0 @@ -# 工具执行期 TOCTOU 防护设计 - -## 背景 -`WorkspaceSandbox` 在此前版本主要完成“路径边界校验”,执行链为: - -`Permission + Sandbox Check -> Tool.Execute(path string)` - -这会留下检查与执行之间的 TOCTOU(Time-of-Check to Time-of-Use)窗口:路径在校验后被替换,工具仍可能访问到非预期对象。 - -## 本次实现 -本次将链路升级为: - -`Permission + Sandbox Check -> WorkspaceExecutionPlan -> Tool.Execute(plan + args)` - -核心变化如下: - -1. `WorkspaceSandbox.Check` 不再只返回 `error`,而是返回 `*WorkspaceExecutionPlan`。 -2. `ToolManager` 将 plan 透传到 `ToolCallInput.WorkspacePlan`。 -3. `filesystem_read_file` / `filesystem_write_file` / `filesystem_edit` / `bash` 在真实执行前调用 `plan.ValidateForExecution()` 复验锚点状态。 -4. 工具使用 `tools.ResolveWorkspaceTarget` 统一消费 plan,避免再次仅依赖字符串路径解析。 - -## 执行期绑定机制 -`WorkspaceExecutionPlan` 在 sandbox 阶段记录: - -- 规范化后的 workspace root -- 规范化后的最终 target -- 最近存在路径锚点(anchor) -- 锚点快照(模式、大小、修改时间、符号链接目标) - -工具执行前会复验: - -1. 当前锚点是否仍为同一路径; -2. 锚点快照是否与校验阶段一致; -3. 锚点解析后的真实路径是否仍在 workspace 内。 - -若任一步失败,返回稳定错误:`workspace target changed before execution` 或 `escapes workspace root via symlink`。 - -## 当前防护边界 -已覆盖: - -- 校验后 symlink 被替换导致 read 越界 -- 校验后父目录被替换导致 write 越界 -- 校验后 bash workdir 被替换导致 cwd 漂移 - -仍存在限制(已显式记录): - -- 未实现跨平台 `openat/no-follow/dirfd` 级别的系统调用原子封装 -- 仍非容器级隔离,不替代系统沙箱 - -本次目标是把“校验结果”显式带入执行期,显著缩小 TOCTOU 窗口,并为后续更强执行器(含 MCP)预留统一接口。 diff --git a/internal/context/builder.go b/internal/context/builder.go index db2e77c1..96e3c437 100644 --- a/internal/context/builder.go +++ b/internal/context/builder.go @@ -3,7 +3,7 @@ package context import ( "context" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" ) // DefaultBuilder preserves the current runtime context-building behavior. @@ -59,7 +59,7 @@ func (b *DefaultBuilder) Build(ctx context.Context, input BuildInput) (BuildResu } // applyReadTimeContextProjection 负责在 provider 请求前按开关应用只读上下文投影,避免改写原始会话消息。 -func applyReadTimeContextProjection(messages []provider.Message, options CompactOptions, policies MicroCompactPolicySource) []provider.Message { +func applyReadTimeContextProjection(messages []providertypes.Message, options CompactOptions, policies MicroCompactPolicySource) []providertypes.Message { if options.DisableMicroCompact { return cloneContextMessages(messages) } diff --git a/internal/context/builder_test.go b/internal/context/builder_test.go index e76d4389..dfe2b838 100644 --- a/internal/context/builder_test.go +++ b/internal/context/builder_test.go @@ -10,7 +10,7 @@ import ( "testing" "neo-code/internal/context/internalcompact" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" "neo-code/internal/tools" ) @@ -31,7 +31,7 @@ func TestDefaultBuilderBuild(t *testing.T) { builder := NewBuilder() input := BuildInput{ - Messages: []provider.Message{ + Messages: []providertypes.Message{ {Role: "user", Content: "hello"}, }, Metadata: testMetadata(t.TempDir()), @@ -90,7 +90,7 @@ func TestDefaultBuilderBuildComposesPromptSectionsInOrder(t *testing.T) { builder := NewBuilder() got, err := builder.Build(stdcontext.Background(), BuildInput{ - Messages: []provider.Message{{Role: "user", Content: "hello"}}, + Messages: []providertypes.Message{{Role: "user", Content: "hello"}}, Metadata: testMetadata(root), }) if err != nil { @@ -111,10 +111,10 @@ func TestDefaultBuilderBuildComposesPromptSectionsInOrder(t *testing.T) { func TestDefaultBuilderBuildUsesSpanTrimPolicyWhenTrimPolicyIsUnset(t *testing.T) { t.Parallel() - messages := make([]provider.Message, 0, maxRetainedMessageSpans+2) + messages := make([]providertypes.Message, 0, maxRetainedMessageSpans+2) for i := 0; i < maxRetainedMessageSpans+2; i++ { - messages = append(messages, provider.Message{ - Role: provider.RoleUser, + messages = append(messages, providertypes.Message{ + Role: providertypes.RoleUser, Content: fmt.Sprintf("u-%d", i), }) } @@ -161,31 +161,31 @@ func TestDefaultBuilderBuildAppliesMicroCompactAfterTrim(t *testing.T) { }, } - messages := []provider.Message{ - {Role: provider.RoleUser, Content: "older user"}, + messages := []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older user"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-1", Content: "old read result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "old read result"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-2", Name: "bash", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-3", Name: "webfetch", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, - {Role: provider.RoleUser, Content: "latest explicit instruction"}, - {Role: provider.RoleAssistant, Content: "current reply"}, + {Role: providertypes.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, + {Role: providertypes.RoleUser, Content: "latest explicit instruction"}, + {Role: providertypes.RoleAssistant, Content: "current reply"}, } got, err := builder.Build(stdcontext.Background(), BuildInput{Messages: messages}) @@ -215,31 +215,31 @@ func TestDefaultBuilderBuildSkipsMicroCompactWhenDisabled(t *testing.T) { }, } - messages := []provider.Message{ - {Role: provider.RoleUser, Content: "older user"}, + messages := []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older user"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-1", Content: "old read result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "old read result"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-2", Name: "bash", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-3", Name: "webfetch", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, - {Role: provider.RoleUser, Content: "latest explicit instruction"}, - {Role: provider.RoleAssistant, Content: "current reply"}, + {Role: providertypes.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, + {Role: providertypes.RoleUser, Content: "latest explicit instruction"}, + {Role: providertypes.RoleAssistant, Content: "current reply"}, } got, err := builder.Build(stdcontext.Background(), BuildInput{ @@ -271,30 +271,30 @@ func TestDefaultBuilderBuildHonorsToolMicroCompactPolicies(t *testing.T) { }, } - messages := []provider.Message{ - {Role: provider.RoleUser, Content: "older user"}, + messages := []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older user"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-1", Name: "custom_tool", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-1", Content: "old custom result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "old custom result"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-2", Name: "bash", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-3", Name: "webfetch", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, - {Role: provider.RoleUser, Content: "latest explicit instruction"}, + {Role: providertypes.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, + {Role: providertypes.RoleUser, Content: "latest explicit instruction"}, } got, err := builder.Build(stdcontext.Background(), BuildInput{Messages: messages}) @@ -313,30 +313,30 @@ func TestNewBuilderWithToolPoliciesUsesProvidedPolicySource(t *testing.T) { "custom_tool": tools.MicroCompactPolicyPreserveHistory, }) - messages := []provider.Message{ - {Role: provider.RoleUser, Content: "older user"}, + messages := []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older user"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-1", Name: "custom_tool", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-1", Content: "old custom result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "old custom result"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-2", Name: "bash", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-3", Name: "webfetch", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, - {Role: provider.RoleUser, Content: "latest explicit instruction"}, + {Role: providertypes.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, + {Role: providertypes.RoleUser, Content: "latest explicit instruction"}, } got, err := builder.Build(stdcontext.Background(), BuildInput{Messages: messages}) @@ -351,20 +351,20 @@ func TestNewBuilderWithToolPoliciesUsesProvidedPolicySource(t *testing.T) { func TestTrimMessagesPreservesToolPairs(t *testing.T) { t.Parallel() - messages := make([]provider.Message, 0, maxRetainedMessageSpans+4) + messages := make([]providertypes.Message, 0, maxRetainedMessageSpans+4) for i := 0; i < 8; i++ { - messages = append(messages, provider.Message{Role: "user", Content: fmt.Sprintf("u-%d", i)}) + messages = append(messages, providertypes.Message{Role: "user", Content: fmt.Sprintf("u-%d", i)}) } messages = append(messages, - provider.Message{ + providertypes.Message{ Role: "assistant", - ToolCalls: []provider.ToolCall{ + ToolCalls: []providertypes.ToolCall{ {ID: "call-1", Name: "filesystem_edit", Arguments: "{}"}, }, }, - provider.Message{Role: "tool", ToolCallID: "call-1", Content: "tool-result"}, - provider.Message{Role: "assistant", Content: "after-tool"}, - provider.Message{Role: "user", Content: "latest"}, + providertypes.Message{Role: "tool", ToolCallID: "call-1", Content: "tool-result"}, + providertypes.Message{Role: "assistant", Content: "after-tool"}, + providertypes.Message{Role: "user", Content: "latest"}, ) trimmed := trimMessages(messages) @@ -390,27 +390,27 @@ func TestTrimMessagesPreservesToolPairs(t *testing.T) { func TestTrimMessagesProtectsLatestExplicitUserInstructionTail(t *testing.T) { t.Parallel() - messages := make([]provider.Message, 0, maxRetainedMessageSpans+5) + messages := make([]providertypes.Message, 0, maxRetainedMessageSpans+5) for i := 0; i < 2; i++ { - messages = append(messages, provider.Message{Role: provider.RoleUser, Content: fmt.Sprintf("old-%d", i)}) + messages = append(messages, providertypes.Message{Role: providertypes.RoleUser, Content: fmt.Sprintf("old-%d", i)}) } messages = append(messages, - provider.Message{Role: provider.RoleUser, Content: "latest explicit instruction"}, - provider.Message{Role: provider.RoleAssistant, Content: "follow-up-1"}, - provider.Message{Role: provider.RoleAssistant, Content: "follow-up-2"}, - provider.Message{Role: provider.RoleAssistant, Content: "follow-up-3"}, - provider.Message{Role: provider.RoleAssistant, Content: "follow-up-4"}, - provider.Message{Role: provider.RoleAssistant, Content: "follow-up-5"}, - provider.Message{Role: provider.RoleAssistant, Content: "follow-up-6"}, - provider.Message{Role: provider.RoleAssistant, Content: "follow-up-7"}, - provider.Message{Role: provider.RoleAssistant, Content: "follow-up-8"}, - provider.Message{Role: provider.RoleAssistant, Content: "follow-up-9"}, - provider.Message{Role: provider.RoleAssistant, Content: "follow-up-10"}, - provider.Message{Role: provider.RoleAssistant, Content: "follow-up-11"}, + providertypes.Message{Role: providertypes.RoleUser, Content: "latest explicit instruction"}, + providertypes.Message{Role: providertypes.RoleAssistant, Content: "follow-up-1"}, + providertypes.Message{Role: providertypes.RoleAssistant, Content: "follow-up-2"}, + providertypes.Message{Role: providertypes.RoleAssistant, Content: "follow-up-3"}, + providertypes.Message{Role: providertypes.RoleAssistant, Content: "follow-up-4"}, + providertypes.Message{Role: providertypes.RoleAssistant, Content: "follow-up-5"}, + providertypes.Message{Role: providertypes.RoleAssistant, Content: "follow-up-6"}, + providertypes.Message{Role: providertypes.RoleAssistant, Content: "follow-up-7"}, + providertypes.Message{Role: providertypes.RoleAssistant, Content: "follow-up-8"}, + providertypes.Message{Role: providertypes.RoleAssistant, Content: "follow-up-9"}, + providertypes.Message{Role: providertypes.RoleAssistant, Content: "follow-up-10"}, + providertypes.Message{Role: providertypes.RoleAssistant, Content: "follow-up-11"}, ) trimmed := trimMessages(messages) - if trimmed[0].Role != provider.RoleUser || trimmed[0].Content != "latest explicit instruction" { + if trimmed[0].Role != providertypes.RoleUser || trimmed[0].Content != "latest explicit instruction" { t.Fatalf("expected protected tail to keep latest explicit user instruction, got %+v", trimmed[0]) } if len(trimmed) != 12 { @@ -421,27 +421,27 @@ func TestTrimMessagesProtectsLatestExplicitUserInstructionTail(t *testing.T) { func TestTrimMessagesUsesSharedSpanModel(t *testing.T) { t.Parallel() - messages := make([]provider.Message, 0, maxRetainedMessageSpans+6) + messages := make([]providertypes.Message, 0, maxRetainedMessageSpans+6) for i := 0; i < 3; i++ { - messages = append(messages, provider.Message{Role: provider.RoleUser, Content: fmt.Sprintf("u-%d", i)}) + messages = append(messages, providertypes.Message{Role: providertypes.RoleUser, Content: fmt.Sprintf("u-%d", i)}) } messages = append(messages, - provider.Message{ - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + providertypes.Message{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-2", Name: "filesystem_read_file", Arguments: "{}"}, }, }, - provider.Message{Role: provider.RoleTool, ToolCallID: "call-2", Content: "tool-result"}, - provider.Message{Role: provider.RoleAssistant, Content: "after tool"}, - provider.Message{Role: provider.RoleUser, Content: "u-4"}, - provider.Message{Role: provider.RoleAssistant, Content: "a-5"}, - provider.Message{Role: provider.RoleUser, Content: "u-6"}, - provider.Message{Role: provider.RoleAssistant, Content: "a-7"}, - provider.Message{Role: provider.RoleUser, Content: "u-8"}, - provider.Message{Role: provider.RoleAssistant, Content: "a-9"}, - provider.Message{Role: provider.RoleUser, Content: "u-10"}, - provider.Message{Role: provider.RoleAssistant, Content: "a-11"}, + providertypes.Message{Role: providertypes.RoleTool, ToolCallID: "call-2", Content: "tool-result"}, + providertypes.Message{Role: providertypes.RoleAssistant, Content: "after tool"}, + providertypes.Message{Role: providertypes.RoleUser, Content: "u-4"}, + providertypes.Message{Role: providertypes.RoleAssistant, Content: "a-5"}, + providertypes.Message{Role: providertypes.RoleUser, Content: "u-6"}, + providertypes.Message{Role: providertypes.RoleAssistant, Content: "a-7"}, + providertypes.Message{Role: providertypes.RoleUser, Content: "u-8"}, + providertypes.Message{Role: providertypes.RoleAssistant, Content: "a-9"}, + providertypes.Message{Role: providertypes.RoleUser, Content: "u-10"}, + providertypes.Message{Role: providertypes.RoleAssistant, Content: "a-11"}, ) spans := internalcompact.BuildMessageSpans(messages) @@ -451,7 +451,7 @@ func TestTrimMessagesUsesSharedSpanModel(t *testing.T) { if len(trimmed) == 0 || trimmed[0].Content != messages[start].Content { t.Fatalf("expected trim to start from shared span boundary %d, got %+v", start, trimmed) } - if trimmed[0].Role != provider.RoleAssistant || len(trimmed[0].ToolCalls) != 1 { + if trimmed[0].Role != providertypes.RoleAssistant || len(trimmed[0].ToolCalls) != 1 { t.Fatalf("expected trim to keep whole tool block at shared boundary, got %+v", trimmed[0]) } } @@ -461,18 +461,18 @@ func TestTrimMessagesBoundaries(t *testing.T) { tests := []struct { name string - input []provider.Message + input []providertypes.Message wantLen int - assert func(t *testing.T, original []provider.Message, trimmed []provider.Message) + assert func(t *testing.T, original []providertypes.Message, trimmed []providertypes.Message) }{ { name: "within max turns returns full cloned slice", - input: []provider.Message{ + input: []providertypes.Message{ {Role: "user", Content: "one"}, {Role: "assistant", Content: "two"}, }, wantLen: 2, - assert: func(t *testing.T, original []provider.Message, trimmed []provider.Message) { + assert: func(t *testing.T, original []providertypes.Message, trimmed []providertypes.Message) { t.Helper() if &trimmed[0] == &original[0] { t.Fatalf("expected trimmed slice to be cloned") @@ -481,25 +481,25 @@ func TestTrimMessagesBoundaries(t *testing.T) { }, { name: "long message list with limited spans keeps full history", - input: func() []provider.Message { - messages := make([]provider.Message, 0, maxRetainedMessageSpans+3) + input: func() []providertypes.Message { + messages := make([]providertypes.Message, 0, maxRetainedMessageSpans+3) for i := 0; i < maxRetainedMessageSpans-1; i++ { - messages = append(messages, provider.Message{Role: "user", Content: fmt.Sprintf("u-%d", i)}) + messages = append(messages, providertypes.Message{Role: "user", Content: fmt.Sprintf("u-%d", i)}) } messages = append(messages, - provider.Message{ + providertypes.Message{ Role: "assistant", - ToolCalls: []provider.ToolCall{ + ToolCalls: []providertypes.ToolCall{ {ID: "call-1", Name: "filesystem_edit", Arguments: "{}"}, }, }, - provider.Message{Role: "tool", ToolCallID: "call-1", Content: "tool-1"}, - provider.Message{Role: "tool", ToolCallID: "call-1", Content: "tool-2"}, + providertypes.Message{Role: "tool", ToolCallID: "call-1", Content: "tool-1"}, + providertypes.Message{Role: "tool", ToolCallID: "call-1", Content: "tool-2"}, ) return messages }(), wantLen: maxRetainedMessageSpans + 2, - assert: func(t *testing.T, original []provider.Message, trimmed []provider.Message) { + assert: func(t *testing.T, original []providertypes.Message, trimmed []providertypes.Message) { t.Helper() if len(trimmed) != len(original) { t.Fatalf("expected full history to remain, got %d want %d", len(trimmed), len(original)) @@ -508,24 +508,24 @@ func TestTrimMessagesBoundaries(t *testing.T) { }, { name: "message count beyond limit trims by span count", - input: func() []provider.Message { - messages := make([]provider.Message, 0, maxRetainedMessageSpans+5) + input: func() []providertypes.Message { + messages := make([]providertypes.Message, 0, maxRetainedMessageSpans+5) for i := 0; i < maxRetainedMessageSpans+1; i++ { - messages = append(messages, provider.Message{Role: "user", Content: fmt.Sprintf("u-%d", i)}) + messages = append(messages, providertypes.Message{Role: "user", Content: fmt.Sprintf("u-%d", i)}) } messages = append(messages, - provider.Message{ + providertypes.Message{ Role: "assistant", - ToolCalls: []provider.ToolCall{ + ToolCalls: []providertypes.ToolCall{ {ID: "call-2", Name: "filesystem_edit", Arguments: "{}"}, }, }, - provider.Message{Role: "tool", ToolCallID: "call-2", Content: "tool-result"}, + providertypes.Message{Role: "tool", ToolCallID: "call-2", Content: "tool-result"}, ) return messages }(), wantLen: maxRetainedMessageSpans + 1, - assert: func(t *testing.T, original []provider.Message, trimmed []provider.Message) { + assert: func(t *testing.T, original []providertypes.Message, trimmed []providertypes.Message) { t.Helper() if trimmed[0].Content != "u-2" { t.Fatalf("expected oldest spans to be removed, got first message %+v", trimmed[0]) diff --git a/internal/context/compact/helpers.go b/internal/context/compact/helpers.go index edbd8420..a66ce845 100644 --- a/internal/context/compact/helpers.go +++ b/internal/context/compact/helpers.go @@ -3,25 +3,25 @@ package compact import ( "unicode/utf8" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" ) // cloneMessages 深拷贝消息切片,避免后续规划或摘要阶段共享底层数据。 -func cloneMessages(messages []provider.Message) []provider.Message { +func cloneMessages(messages []providertypes.Message) []providertypes.Message { if len(messages) == 0 { return nil } - out := make([]provider.Message, 0, len(messages)) + out := make([]providertypes.Message, 0, len(messages)) for _, message := range messages { next := message - next.ToolCalls = append([]provider.ToolCall(nil), message.ToolCalls...) + next.ToolCalls = append([]providertypes.ToolCall(nil), message.ToolCalls...) out = append(out, next) } return out } // countMessageChars 以 rune 数量统计消息体积,用于 compact 前后指标计算。 -func countMessageChars(messages []provider.Message) int { +func countMessageChars(messages []providertypes.Message) int { total := 0 for _, message := range messages { total += utf8.RuneCountInString(message.Role) diff --git a/internal/context/compact/planner.go b/internal/context/compact/planner.go index 7f06c043..7ecb2204 100644 --- a/internal/context/compact/planner.go +++ b/internal/context/compact/planner.go @@ -6,13 +6,13 @@ import ( "neo-code/internal/config" "neo-code/internal/context/internalcompact" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" ) // compactionPlan 描述一次 compact 在摘要生成前的归档与保留结果。 type compactionPlan struct { - Archived []provider.Message - Retained []provider.Message + Archived []providertypes.Message + Retained []providertypes.Message ArchivedMessageCount int Applied bool } @@ -21,7 +21,7 @@ type compactionPlan struct { type compactionPlanner struct{} // Plan 根据 mode 与配置返回摘要前的裁剪规划结果。 -func (compactionPlanner) Plan(mode Mode, messages []provider.Message, cfg config.CompactConfig) (compactionPlan, error) { +func (compactionPlanner) Plan(mode Mode, messages []providertypes.Message, cfg config.CompactConfig) (compactionPlan, error) { if mode == ModeReactive { return planKeepRecent(messages, cfg.ManualKeepRecentMessages), nil } @@ -37,7 +37,7 @@ func (compactionPlanner) Plan(mode Mode, messages []provider.Message, cfg config } // planKeepRecent 计算 keep_recent 策略下需要摘要与保留的消息集合。 -func planKeepRecent(messages []provider.Message, keepMessages int) compactionPlan { +func planKeepRecent(messages []providertypes.Message, keepMessages int) compactionPlan { spans := internalcompact.BuildMessageSpans(messages) retainedStart := internalcompact.RetainedStartForKeepRecentMessages(spans, keepMessages) if retainedStart <= 0 { @@ -57,7 +57,7 @@ func planKeepRecent(messages []provider.Message, keepMessages int) compactionPla } // planFullReplace 计算 full_replace 策略下需要摘要与保留的消息集合。 -func planFullReplace(messages []provider.Message) compactionPlan { +func planFullReplace(messages []providertypes.Message) compactionPlan { if len(messages) == 0 { return compactionPlan{} } @@ -78,7 +78,7 @@ func planFullReplace(messages []provider.Message) compactionPlan { } // splitMessagesAt 按 retained 起点切分 archived 与 retained,并返回深拷贝结果。 -func splitMessagesAt(messages []provider.Message, retainedStart int) ([]provider.Message, []provider.Message) { +func splitMessagesAt(messages []providertypes.Message, retainedStart int) ([]providertypes.Message, []providertypes.Message) { if retainedStart <= 0 { return nil, cloneMessages(messages) } diff --git a/internal/context/compact/planner_test.go b/internal/context/compact/planner_test.go index be243445..bcce24fd 100644 --- a/internal/context/compact/planner_test.go +++ b/internal/context/compact/planner_test.go @@ -4,20 +4,20 @@ import ( "testing" "neo-code/internal/config" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" ) func TestCompactionPlannerKeepRecentPlan(t *testing.T) { t.Parallel() planner := compactionPlanner{} - plan, err := planner.Plan(ModeManual, []provider.Message{ - {Role: provider.RoleUser, Content: "old request"}, - {Role: provider.RoleAssistant, Content: "old answer"}, - {Role: provider.RoleAssistant, ToolCalls: []provider.ToolCall{{ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}}}, - {Role: provider.RoleTool, ToolCallID: "call-1", Content: "tool result"}, - {Role: provider.RoleUser, Content: "latest instruction"}, - {Role: provider.RoleAssistant, Content: "latest answer"}, + plan, err := planner.Plan(ModeManual, []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "old request"}, + {Role: providertypes.RoleAssistant, Content: "old answer"}, + {Role: providertypes.RoleAssistant, ToolCalls: []providertypes.ToolCall{{ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}}}, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "tool result"}, + {Role: providertypes.RoleUser, Content: "latest instruction"}, + {Role: providertypes.RoleAssistant, Content: "latest answer"}, }, config.CompactConfig{ ManualStrategy: config.CompactManualStrategyKeepRecent, ManualKeepRecentMessages: 3, @@ -31,10 +31,10 @@ func TestCompactionPlannerKeepRecentPlan(t *testing.T) { if len(plan.Archived) != 2 || len(plan.Retained) != 4 { t.Fatalf("unexpected keep_recent plan: %+v", plan) } - if plan.Retained[0].Role != provider.RoleAssistant || len(plan.Retained[0].ToolCalls) != 1 { + if plan.Retained[0].Role != providertypes.RoleAssistant || len(plan.Retained[0].ToolCalls) != 1 { t.Fatalf("expected retained tool block start, got %+v", plan.Retained[0]) } - if plan.Retained[1].Role != provider.RoleTool { + if plan.Retained[1].Role != providertypes.RoleTool { t.Fatalf("expected retained tool result, got %+v", plan.Retained[1]) } } @@ -43,11 +43,11 @@ func TestCompactionPlannerFullReplaceProtectsLatestExplicitUserInstruction(t *te t.Parallel() planner := compactionPlanner{} - plan, err := planner.Plan(ModeManual, []provider.Message{ - {Role: provider.RoleUser, Content: "old request"}, - {Role: provider.RoleAssistant, Content: "old answer"}, - {Role: provider.RoleUser, Content: "latest instruction"}, - {Role: provider.RoleAssistant, Content: "latest answer"}, + plan, err := planner.Plan(ModeManual, []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "old request"}, + {Role: providertypes.RoleAssistant, Content: "old answer"}, + {Role: providertypes.RoleUser, Content: "latest instruction"}, + {Role: providertypes.RoleAssistant, Content: "latest answer"}, }, config.CompactConfig{ ManualStrategy: config.CompactManualStrategyFullReplace, }) @@ -60,7 +60,7 @@ func TestCompactionPlannerFullReplaceProtectsLatestExplicitUserInstruction(t *te if len(plan.Archived) != 2 || len(plan.Retained) != 2 { t.Fatalf("unexpected full_replace plan: %+v", plan) } - if plan.Retained[0].Role != provider.RoleUser || plan.Retained[0].Content != "latest instruction" { + if plan.Retained[0].Role != providertypes.RoleUser || plan.Retained[0].Content != "latest instruction" { t.Fatalf("expected latest explicit user instruction to stay retained, got %+v", plan.Retained) } } @@ -77,11 +77,11 @@ func TestCompactionPlannerRejectsUnsupportedStrategy(t *testing.T) { func TestCompactionPlannerReactiveModeAlwaysUsesKeepRecentStrategy(t *testing.T) { t.Parallel() - plan, err := (compactionPlanner{}).Plan(ModeReactive, []provider.Message{ - {Role: provider.RoleUser, Content: "old request"}, - {Role: provider.RoleAssistant, Content: "old answer"}, - {Role: provider.RoleUser, Content: "latest request"}, - {Role: provider.RoleAssistant, Content: "latest answer"}, + plan, err := (compactionPlanner{}).Plan(ModeReactive, []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "old request"}, + {Role: providertypes.RoleAssistant, Content: "old answer"}, + {Role: providertypes.RoleUser, Content: "latest request"}, + {Role: providertypes.RoleAssistant, Content: "latest answer"}, }, config.CompactConfig{ ManualStrategy: "unsupported", ManualKeepRecentMessages: 2, diff --git a/internal/context/compact/runner.go b/internal/context/compact/runner.go index 1896c03b..6ae8dc9c 100644 --- a/internal/context/compact/runner.go +++ b/internal/context/compact/runner.go @@ -9,7 +9,7 @@ import ( "time" "neo-code/internal/config" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" ) // Mode identifies the compact execution mode. @@ -34,15 +34,15 @@ type Input struct { Mode Mode SessionID string Workdir string - Messages []provider.Message + Messages []providertypes.Message Config config.CompactConfig } // SummaryInput describes the historical context that must be summarized. type SummaryInput struct { Mode Mode - ArchivedMessages []provider.Message - RetainedMessages []provider.Message + ArchivedMessages []providertypes.Message + RetainedMessages []providertypes.Message ArchivedMessageCount int Config config.CompactConfig } @@ -57,12 +57,12 @@ type Metrics struct { // Result is the compact execution result. type Result struct { - Messages []provider.Message `json:"messages"` - Metrics Metrics `json:"metrics"` - TranscriptID string `json:"transcript_id"` - TranscriptPath string `json:"transcript_path"` - Applied bool `json:"applied"` - ErrorMode ErrorMode `json:"error_mode"` + Messages []providertypes.Message `json:"messages"` + Metrics Metrics `json:"metrics"` + TranscriptID string `json:"transcript_id"` + TranscriptPath string `json:"transcript_path"` + Applied bool `json:"applied"` + ErrorMode ErrorMode `json:"error_mode"` } // SummaryGenerator produces the semantic compact summary. @@ -155,8 +155,8 @@ func (s *Service) Run(ctx context.Context, input Input) (Result, error) { return Result{}, err } - next := make([]provider.Message, 0, len(plan.Retained)+1) - next = append(next, provider.Message{Role: provider.RoleAssistant, Content: summary}) + next := make([]providertypes.Message, 0, len(plan.Retained)+1) + next = append(next, providertypes.Message{Role: providertypes.RoleAssistant, Content: summary}) next = append(next, plan.Retained...) afterChars := countMessageChars(next) diff --git a/internal/context/compact/runner_test.go b/internal/context/compact/runner_test.go index ccf92451..4ab50825 100644 --- a/internal/context/compact/runner_test.go +++ b/internal/context/compact/runner_test.go @@ -11,7 +11,7 @@ import ( "neo-code/internal/config" "neo-code/internal/context/internalcompact" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" ) type stubSummaryGenerator struct { @@ -56,19 +56,19 @@ func TestManualCompactKeepRecentRetainsRecentMessagesAndWholeToolBlock(t *testin home := t.TempDir() runner.userHomeDir = func() (string, error) { return home, nil } - messages := []provider.Message{ - {Role: provider.RoleUser, Content: "old requirement"}, - {Role: provider.RoleAssistant, Content: "old answer"}, - {Role: provider.RoleUser, Content: "middle request"}, - {Role: provider.RoleAssistant, ToolCalls: []provider.ToolCall{{ID: "call-old", Name: "filesystem_grep", Arguments: "{}"}}}, - {Role: provider.RoleTool, ToolCallID: "call-old", Content: "old result"}, - {Role: provider.RoleAssistant, Content: "after tool"}, - {Role: provider.RoleUser, Content: "instruction to keep"}, - {Role: provider.RoleAssistant, Content: "ack"}, - {Role: provider.RoleUser, Content: "recent follow up"}, - {Role: provider.RoleAssistant, Content: "recent answer"}, - {Role: provider.RoleUser, Content: "latest explicit instruction"}, - {Role: provider.RoleAssistant, Content: "latest result"}, + messages := []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "old requirement"}, + {Role: providertypes.RoleAssistant, Content: "old answer"}, + {Role: providertypes.RoleUser, Content: "middle request"}, + {Role: providertypes.RoleAssistant, ToolCalls: []providertypes.ToolCall{{ID: "call-old", Name: "filesystem_grep", Arguments: "{}"}}}, + {Role: providertypes.RoleTool, ToolCallID: "call-old", Content: "old result"}, + {Role: providertypes.RoleAssistant, Content: "after tool"}, + {Role: providertypes.RoleUser, Content: "instruction to keep"}, + {Role: providertypes.RoleAssistant, Content: "ack"}, + {Role: providertypes.RoleUser, Content: "recent follow up"}, + {Role: providertypes.RoleAssistant, Content: "recent answer"}, + {Role: providertypes.RoleUser, Content: "latest explicit instruction"}, + {Role: providertypes.RoleAssistant, Content: "latest result"}, } result, err := runner.Run(context.Background(), Input{ @@ -91,7 +91,7 @@ func TestManualCompactKeepRecentRetainsRecentMessagesAndWholeToolBlock(t *testin if len(result.Messages) != 10 { t.Fatalf("expected summary + 9 retained messages, got %d", len(result.Messages)) } - if result.Messages[0].Role != provider.RoleAssistant { + if result.Messages[0].Role != providertypes.RoleAssistant { t.Fatalf("expected summary role assistant, got %q", result.Messages[0].Role) } for _, section := range []string{"done:", "in_progress:", "decisions:", "code_changes:", "constraints:"} { @@ -99,10 +99,10 @@ func TestManualCompactKeepRecentRetainsRecentMessagesAndWholeToolBlock(t *testin t.Fatalf("expected summary to include section %q, got %q", section, result.Messages[0].Content) } } - if result.Messages[1].Role != provider.RoleAssistant || len(result.Messages[1].ToolCalls) != 1 { + if result.Messages[1].Role != providertypes.RoleAssistant || len(result.Messages[1].ToolCalls) != 1 { t.Fatalf("expected retained tool call block start, got %+v", result.Messages[1]) } - if result.Messages[2].Role != provider.RoleTool || result.Messages[2].ToolCallID != "call-old" { + if result.Messages[2].Role != providertypes.RoleTool || result.Messages[2].ToolCallID != "call-old" { t.Fatalf("expected retained tool result, got %+v", result.Messages[2]) } if len(generator.calls) != 1 { @@ -124,19 +124,19 @@ func TestReactiveCompactUsesKeepRecentAndReportsReactiveMode(t *testing.T) { home := t.TempDir() runner.userHomeDir = func() (string, error) { return home, nil } - messages := []provider.Message{ - {Role: provider.RoleUser, Content: "old requirement"}, - {Role: provider.RoleAssistant, Content: "old answer"}, - {Role: provider.RoleUser, Content: "middle request"}, - {Role: provider.RoleAssistant, ToolCalls: []provider.ToolCall{{ID: "call-old", Name: "filesystem_grep", Arguments: "{}"}}}, - {Role: provider.RoleTool, ToolCallID: "call-old", Content: "old result"}, - {Role: provider.RoleAssistant, Content: "after tool"}, - {Role: provider.RoleUser, Content: "instruction to keep"}, - {Role: provider.RoleAssistant, Content: "ack"}, - {Role: provider.RoleUser, Content: "recent follow up"}, - {Role: provider.RoleAssistant, Content: "recent answer"}, - {Role: provider.RoleUser, Content: "latest explicit instruction"}, - {Role: provider.RoleAssistant, Content: "latest result"}, + messages := []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "old requirement"}, + {Role: providertypes.RoleAssistant, Content: "old answer"}, + {Role: providertypes.RoleUser, Content: "middle request"}, + {Role: providertypes.RoleAssistant, ToolCalls: []providertypes.ToolCall{{ID: "call-old", Name: "filesystem_grep", Arguments: "{}"}}}, + {Role: providertypes.RoleTool, ToolCallID: "call-old", Content: "old result"}, + {Role: providertypes.RoleAssistant, Content: "after tool"}, + {Role: providertypes.RoleUser, Content: "instruction to keep"}, + {Role: providertypes.RoleAssistant, Content: "ack"}, + {Role: providertypes.RoleUser, Content: "recent follow up"}, + {Role: providertypes.RoleAssistant, Content: "recent answer"}, + {Role: providertypes.RoleUser, Content: "latest explicit instruction"}, + {Role: providertypes.RoleAssistant, Content: "latest result"}, } result, err := runner.Run(context.Background(), Input{ @@ -165,10 +165,10 @@ func TestReactiveCompactUsesKeepRecentAndReportsReactiveMode(t *testing.T) { if len(result.Messages) != 10 { t.Fatalf("expected summary + 9 retained messages, got %d", len(result.Messages)) } - if result.Messages[1].Role != provider.RoleAssistant || len(result.Messages[1].ToolCalls) != 1 { + if result.Messages[1].Role != providertypes.RoleAssistant || len(result.Messages[1].ToolCalls) != 1 { t.Fatalf("expected retained tool call block start, got %+v", result.Messages[1]) } - if result.Messages[2].Role != provider.RoleTool || result.Messages[2].ToolCallID != "call-old" { + if result.Messages[2].Role != providertypes.RoleTool || result.Messages[2].ToolCallID != "call-old" { t.Fatalf("expected retained tool result, got %+v", result.Messages[2]) } if len(generator.calls) != 1 { @@ -192,16 +192,16 @@ func TestManualCompactKeepRecentProtectsLatestExplicitUserInstruction(t *testing runner := NewRunner(generator) runner.userHomeDir = func() (string, error) { return t.TempDir(), nil } - messages := []provider.Message{ - {Role: provider.RoleUser, Content: "old requirement"}, - {Role: provider.RoleAssistant, Content: "old answer"}, - {Role: provider.RoleUser, Content: "latest explicit instruction"}, - {Role: provider.RoleAssistant, Content: "ack"}, - {Role: provider.RoleAssistant, Content: "follow up 1"}, - {Role: provider.RoleAssistant, Content: "follow up 2"}, - {Role: provider.RoleAssistant, Content: "follow up 3"}, - {Role: provider.RoleAssistant, Content: "follow up 4"}, - {Role: provider.RoleAssistant, Content: "follow up 5"}, + messages := []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "old requirement"}, + {Role: providertypes.RoleAssistant, Content: "old answer"}, + {Role: providertypes.RoleUser, Content: "latest explicit instruction"}, + {Role: providertypes.RoleAssistant, Content: "ack"}, + {Role: providertypes.RoleAssistant, Content: "follow up 1"}, + {Role: providertypes.RoleAssistant, Content: "follow up 2"}, + {Role: providertypes.RoleAssistant, Content: "follow up 3"}, + {Role: providertypes.RoleAssistant, Content: "follow up 4"}, + {Role: providertypes.RoleAssistant, Content: "follow up 5"}, } result, err := runner.Run(context.Background(), Input{ @@ -227,7 +227,7 @@ func TestManualCompactKeepRecentProtectsLatestExplicitUserInstruction(t *testing if len(generator.calls[0].ArchivedMessages) != 2 || len(generator.calls[0].RetainedMessages) != 7 { t.Fatalf("expected protected tail to start at latest user instruction, got %+v", generator.calls[0]) } - if result.Messages[1].Role != provider.RoleUser || result.Messages[1].Content != "latest explicit instruction" { + if result.Messages[1].Role != providertypes.RoleUser || result.Messages[1].Content != "latest explicit instruction" { t.Fatalf("expected retained latest explicit instruction, got %+v", result.Messages[1]) } } @@ -243,8 +243,8 @@ func TestManualCompactWritesTranscriptJSONL(t *testing.T) { Mode: ModeManual, SessionID: "session-jsonl", Workdir: filepath.Join(home, "workspace"), - Messages: []provider.Message{ - {Role: provider.RoleUser, Content: "hello"}, + Messages: []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "hello"}, }, Config: config.CompactConfig{ ManualStrategy: config.CompactManualStrategyKeepRecent, @@ -284,7 +284,7 @@ func TestManualCompactFailsWhenTranscriptWriteFails(t *testing.T) { Mode: ModeManual, SessionID: "session-fail", Workdir: t.TempDir(), - Messages: []provider.Message{{Role: provider.RoleUser, Content: "hello"}}, + Messages: []providertypes.Message{{Role: providertypes.RoleUser, Content: "hello"}}, Config: config.CompactConfig{ ManualStrategy: config.CompactManualStrategyKeepRecent, ManualKeepRecentMessages: 10, @@ -304,13 +304,13 @@ func TestManualCompactFullReplaceKeepsProtectedTail(t *testing.T) { home := t.TempDir() runner.userHomeDir = func() (string, error) { return home, nil } - messages := []provider.Message{ - {Role: provider.RoleUser, Content: "old requirement"}, - {Role: provider.RoleAssistant, Content: "old answer"}, - {Role: provider.RoleUser, Content: "latest explicit instruction"}, - {Role: provider.RoleAssistant, ToolCalls: []provider.ToolCall{{ID: "call-old", Name: "filesystem_grep", Arguments: "{}"}}}, - {Role: provider.RoleTool, ToolCallID: "call-old", Content: "old result"}, - {Role: provider.RoleAssistant, Content: "latest answer"}, + messages := []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "old requirement"}, + {Role: providertypes.RoleAssistant, Content: "old answer"}, + {Role: providertypes.RoleUser, Content: "latest explicit instruction"}, + {Role: providertypes.RoleAssistant, ToolCalls: []providertypes.ToolCall{{ID: "call-old", Name: "filesystem_grep", Arguments: "{}"}}}, + {Role: providertypes.RoleTool, ToolCallID: "call-old", Content: "old result"}, + {Role: providertypes.RoleAssistant, Content: "latest answer"}, } result, err := runner.Run(context.Background(), Input{ @@ -333,7 +333,7 @@ func TestManualCompactFullReplaceKeepsProtectedTail(t *testing.T) { if len(result.Messages) != 5 { t.Fatalf("expected summary plus protected tail, got %d", len(result.Messages)) } - if result.Messages[0].Role != provider.RoleAssistant { + if result.Messages[0].Role != providertypes.RoleAssistant { t.Fatalf("expected summary role assistant, got %q", result.Messages[0].Role) } if len(generator.calls) != 1 || len(generator.calls[0].RetainedMessages) != 4 { @@ -351,9 +351,9 @@ func TestManualCompactFullReplaceWithoutArchivableMessagesSkipsGenerator(t *test runner := NewRunner(generator) runner.userHomeDir = func() (string, error) { return t.TempDir(), nil } - messages := []provider.Message{ - {Role: provider.RoleUser, Content: "latest explicit instruction"}, - {Role: provider.RoleAssistant, Content: "latest answer"}, + messages := []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "latest explicit instruction"}, + {Role: providertypes.RoleAssistant, Content: "latest answer"}, } result, err := runner.Run(context.Background(), Input{ @@ -393,7 +393,7 @@ func TestRunManualRejectsUnsupportedStrategy(t *testing.T) { Mode: ModeManual, SessionID: "session-invalid-strategy", Workdir: t.TempDir(), - Messages: []provider.Message{{Role: provider.RoleUser, Content: "hello"}}, + Messages: []providertypes.Message{{Role: providertypes.RoleUser, Content: "hello"}}, Config: config.CompactConfig{ ManualStrategy: "unknown_strategy", ManualKeepRecentMessages: 10, @@ -415,7 +415,7 @@ func TestRunRejectsUnsupportedMode(t *testing.T) { Mode: Mode("unexpected"), SessionID: "session-invalid-mode", Workdir: t.TempDir(), - Messages: []provider.Message{{Role: provider.RoleUser, Content: "hello"}}, + Messages: []providertypes.Message{{Role: providertypes.RoleUser, Content: "hello"}}, Config: config.CompactConfig{ ManualStrategy: config.CompactManualStrategyKeepRecent, ManualKeepRecentMessages: 10, @@ -430,12 +430,12 @@ func TestRunRejectsUnsupportedMode(t *testing.T) { func TestCountMessageCharsUsesRunes(t *testing.T) { t.Parallel() - messages := []provider.Message{ + messages := []providertypes.Message{ {Role: "用户", Content: "你好"}, - {Role: provider.RoleAssistant, Content: "done"}, + {Role: providertypes.RoleAssistant, Content: "done"}, } got := countMessageChars(messages) - want := len([]rune("用户")) + len([]rune("你好")) + len([]rune(provider.RoleAssistant)) + len([]rune("done")) + want := len([]rune("用户")) + len([]rune("你好")) + len([]rune(providertypes.RoleAssistant)) + len([]rune("done")) if got != want { t.Fatalf("countMessageChars() = %d, want %d", got, want) } @@ -460,9 +460,9 @@ func TestSaveTranscriptUsesUniqueIDWithinSameTimestamp(t *testing.T) { Mode: ModeManual, SessionID: "session-dup-safe", Workdir: t.TempDir(), - Messages: []provider.Message{ - {Role: provider.RoleUser, Content: "hello"}, - {Role: provider.RoleAssistant, Content: "world"}, + Messages: []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "hello"}, + {Role: providertypes.RoleAssistant, Content: "world"}, }, Config: config.CompactConfig{ ManualStrategy: config.CompactManualStrategyFullReplace, @@ -506,11 +506,11 @@ func TestManualCompactGeneratorInvalidSummaryFails(t *testing.T) { Mode: ModeManual, SessionID: "session-invalid-summary", Workdir: t.TempDir(), - Messages: []provider.Message{ - {Role: provider.RoleUser, Content: "older"}, - {Role: provider.RoleAssistant, Content: "older answer"}, - {Role: provider.RoleUser, Content: "latest explicit instruction"}, - {Role: provider.RoleAssistant, Content: "newer"}, + Messages: []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older"}, + {Role: providertypes.RoleAssistant, Content: "older answer"}, + {Role: providertypes.RoleUser, Content: "latest explicit instruction"}, + {Role: providertypes.RoleAssistant, Content: "newer"}, }, Config: config.CompactConfig{ ManualStrategy: config.CompactManualStrategyFullReplace, @@ -551,11 +551,11 @@ func TestManualCompactGeneratorEmptyBulletFails(t *testing.T) { Mode: ModeManual, SessionID: "session-empty-bullet", Workdir: t.TempDir(), - Messages: []provider.Message{ - {Role: provider.RoleUser, Content: "older"}, - {Role: provider.RoleAssistant, Content: "older answer"}, - {Role: provider.RoleUser, Content: "latest explicit instruction"}, - {Role: provider.RoleAssistant, Content: "newer"}, + Messages: []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older"}, + {Role: providertypes.RoleAssistant, Content: "older answer"}, + {Role: providertypes.RoleUser, Content: "latest explicit instruction"}, + {Role: providertypes.RoleAssistant, Content: "newer"}, }, Config: config.CompactConfig{ ManualStrategy: config.CompactManualStrategyFullReplace, @@ -579,11 +579,11 @@ func TestManualCompactTruncationFailsWhenStructureBreaks(t *testing.T) { Mode: ModeManual, SessionID: "session-truncate-fail", Workdir: t.TempDir(), - Messages: []provider.Message{ - {Role: provider.RoleUser, Content: "older"}, - {Role: provider.RoleAssistant, Content: "older answer"}, - {Role: provider.RoleUser, Content: "latest explicit instruction"}, - {Role: provider.RoleAssistant, Content: "newer"}, + Messages: []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older"}, + {Role: providertypes.RoleAssistant, Content: "older answer"}, + {Role: providertypes.RoleUser, Content: "latest explicit instruction"}, + {Role: providertypes.RoleAssistant, Content: "newer"}, }, Config: config.CompactConfig{ ManualStrategy: config.CompactManualStrategyFullReplace, @@ -607,8 +607,8 @@ func TestManualCompactKeepRecentWithoutEnoughMessagesSkipsGenerator(t *testing.T Mode: ModeManual, SessionID: "session-no-compact", Workdir: t.TempDir(), - Messages: []provider.Message{ - {Role: provider.RoleUser, Content: "single message"}, + Messages: []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "single message"}, }, Config: config.CompactConfig{ ManualStrategy: config.CompactManualStrategyKeepRecent, @@ -637,11 +637,11 @@ func TestManualCompactReturnsErrorWhenSummaryGeneratorIsMissing(t *testing.T) { Mode: ModeManual, SessionID: "session-missing-generator", Workdir: t.TempDir(), - Messages: []provider.Message{ - {Role: provider.RoleUser, Content: "older"}, - {Role: provider.RoleAssistant, Content: "older answer"}, - {Role: provider.RoleUser, Content: "latest explicit instruction"}, - {Role: provider.RoleAssistant, Content: "newer"}, + Messages: []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older"}, + {Role: providertypes.RoleAssistant, Content: "older answer"}, + {Role: providertypes.RoleUser, Content: "latest explicit instruction"}, + {Role: providertypes.RoleAssistant, Content: "newer"}, }, Config: config.CompactConfig{ ManualStrategy: config.CompactManualStrategyFullReplace, @@ -665,11 +665,11 @@ func TestManualCompactDefaultsToKeepRecentStrategyWhenManualStrategyIsEmpty(t *t Mode: ModeManual, SessionID: "session-default-strategy", Workdir: t.TempDir(), - Messages: []provider.Message{ - {Role: provider.RoleUser, Content: "old request"}, - {Role: provider.RoleAssistant, Content: "old answer"}, - {Role: provider.RoleUser, Content: "latest request"}, - {Role: provider.RoleAssistant, Content: "latest answer"}, + Messages: []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "old request"}, + {Role: providertypes.RoleAssistant, Content: "old answer"}, + {Role: providertypes.RoleUser, Content: "latest request"}, + {Role: providertypes.RoleAssistant, Content: "latest answer"}, }, Config: config.CompactConfig{ ManualStrategy: "", diff --git a/internal/context/compact/transcript_store.go b/internal/context/compact/transcript_store.go index 7ff88b44..568f9b9f 100644 --- a/internal/context/compact/transcript_store.go +++ b/internal/context/compact/transcript_store.go @@ -13,7 +13,7 @@ import ( "strings" "time" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" ) const ( @@ -26,13 +26,13 @@ const ( ) type transcriptLine struct { - Index int `json:"index"` - Timestamp string `json:"timestamp"` - Role string `json:"role"` - Content string `json:"content"` - ToolCalls []provider.ToolCall `json:"tool_calls,omitempty"` - ToolCallID string `json:"tool_call_id,omitempty"` - IsError bool `json:"is_error,omitempty"` + Index int `json:"index"` + Timestamp string `json:"timestamp"` + Role string `json:"role"` + Content string `json:"content"` + ToolCalls []providertypes.ToolCall `json:"tool_calls,omitempty"` + ToolCallID string `json:"tool_call_id,omitempty"` + IsError bool `json:"is_error,omitempty"` } // transcriptStore 只负责 compact 原始 transcript 的目录规划与安全落盘。 @@ -47,7 +47,7 @@ type transcriptStore struct { } // Save 按项目维度持久化当前 compact 前的 transcript,并返回 ID 与路径。 -func (s transcriptStore) Save(messages []provider.Message, sessionID string, workdir string) (string, string, error) { +func (s transcriptStore) Save(messages []providertypes.Message, sessionID string, workdir string) (string, string, error) { home, err := s.userHomeDir() if err != nil { return "", "", fmt.Errorf("compact: resolve user home: %w", err) @@ -85,7 +85,7 @@ func (s transcriptStore) Save(messages []provider.Message, sessionID string, wor Timestamp: now, Role: message.Role, Content: message.Content, - ToolCalls: append([]provider.ToolCall(nil), message.ToolCalls...), + ToolCalls: append([]providertypes.ToolCall(nil), message.ToolCalls...), ToolCallID: message.ToolCallID, IsError: message.IsError, } diff --git a/internal/context/compact/transcript_store_test.go b/internal/context/compact/transcript_store_test.go index b3616250..9425182d 100644 --- a/internal/context/compact/transcript_store_test.go +++ b/internal/context/compact/transcript_store_test.go @@ -8,7 +8,7 @@ import ( "testing" "time" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" ) func TestTranscriptStoreSaveSanitizesSessionIDAndWritesJSONL(t *testing.T) { @@ -25,8 +25,8 @@ func TestTranscriptStoreSaveSanitizesSessionIDAndWritesJSONL(t *testing.T) { remove: os.Remove, } - id, path, err := store.Save([]provider.Message{ - {Role: provider.RoleUser, Content: "hello"}, + id, path, err := store.Save([]providertypes.Message{ + {Role: providertypes.RoleUser, Content: "hello"}, }, "session with spaces", filepath.Join(home, "workspace")) if err != nil { t.Fatalf("Save() error = %v", err) @@ -113,7 +113,7 @@ func TestTranscriptStoreSaveRemovesTemporaryFileWhenRenameFails(t *testing.T) { }, } - _, _, err := store.Save([]provider.Message{{Role: provider.RoleUser, Content: "hello"}}, "session", filepath.Join(home, "workspace")) + _, _, err := store.Save([]providertypes.Message{{Role: providertypes.RoleUser, Content: "hello"}}, "session", filepath.Join(home, "workspace")) if err == nil || !strings.Contains(err.Error(), "rename boom") { t.Fatalf("expected rename error, got %v", err) } diff --git a/internal/context/compact_prompt.go b/internal/context/compact_prompt.go index d79e56f4..4cbd8ce2 100644 --- a/internal/context/compact_prompt.go +++ b/internal/context/compact_prompt.go @@ -5,7 +5,7 @@ import ( "strings" "neo-code/internal/context/internalcompact" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" ) var compactSummarySystemPrompt = buildCompactSummarySystemPrompt() @@ -17,8 +17,8 @@ type CompactPromptInput struct { ManualKeepRecentMessages int ArchivedMessageCount int MaxSummaryChars int - ArchivedMessages []provider.Message - RetainedMessages []provider.Message + ArchivedMessages []providertypes.Message + RetainedMessages []providertypes.Message } // CompactPrompt is the provider-facing prompt pair for compact summaries. @@ -87,7 +87,7 @@ func buildCompactSummarySystemPrompt() string { } // renderCompactPromptMessages 将消息渲染为紧凑的 transcript 视图,减少冗余 JSON 噪音。 -func renderCompactPromptMessages(messages []provider.Message) string { +func renderCompactPromptMessages(messages []providertypes.Message) string { if len(messages) == 0 { return "[]" } @@ -120,7 +120,7 @@ func renderCompactPromptMessages(messages []provider.Message) string { } // renderCompactPromptToolCall 以单行形式渲染工具调用元信息,压缩摘要输入体积。 -func renderCompactPromptToolCall(call provider.ToolCall) string { +func renderCompactPromptToolCall(call providertypes.ToolCall) string { line := fmt.Sprintf( "tool_call id=%s name=%s arguments=%s", strings.TrimSpace(call.ID), diff --git a/internal/context/compact_prompt_test.go b/internal/context/compact_prompt_test.go index 5d38c006..676f02ab 100644 --- a/internal/context/compact_prompt_test.go +++ b/internal/context/compact_prompt_test.go @@ -5,7 +5,7 @@ import ( "testing" "neo-code/internal/context/internalcompact" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" ) func TestBuildCompactPromptIncludesFixedInstructionsAndBoundaries(t *testing.T) { @@ -17,20 +17,20 @@ func TestBuildCompactPromptIncludesFixedInstructionsAndBoundaries(t *testing.T) ManualKeepRecentMessages: 10, ArchivedMessageCount: 3, MaxSummaryChars: 1200, - ArchivedMessages: []provider.Message{ + ArchivedMessages: []providertypes.Message{ { - Role: provider.RoleUser, + Role: providertypes.RoleUser, Content: "legacy request\nwith details", }, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-1", Name: "filesystem_read_file", Arguments: "{\n \"path\": \"a.txt\"\n}"}, }, }, }, - RetainedMessages: []provider.Message{ - {Role: provider.RoleAssistant, Content: "recent answer"}, + RetainedMessages: []providertypes.Message{ + {Role: providertypes.RoleAssistant, Content: "recent answer"}, }, }) diff --git a/internal/context/internalcompact/messages.go b/internal/context/internalcompact/messages.go index a81e891b..3574707b 100644 --- a/internal/context/internalcompact/messages.go +++ b/internal/context/internalcompact/messages.go @@ -3,7 +3,7 @@ package internalcompact import ( "strings" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" ) // MessageSpan 描述一段不可拆分的消息区间,并携带是否需要保护的尾部语义。 @@ -15,13 +15,13 @@ type MessageSpan struct { } // BuildMessageSpans 按工具调用原子块构建消息分段,并保护最后一条明确用户指令所在分段。 -func BuildMessageSpans(messages []provider.Message) []MessageSpan { +func BuildMessageSpans(messages []providertypes.Message) []MessageSpan { spans := make([]MessageSpan, 0, len(messages)) for i := 0; i < len(messages); { start := i end := i + 1 - if messages[start].Role == provider.RoleAssistant && len(messages[start].ToolCalls) > 0 { - for end < len(messages) && messages[end].Role == provider.RoleTool { + if messages[start].Role == providertypes.RoleAssistant && len(messages[start].ToolCalls) > 0 { + for end < len(messages) && messages[end].Role == providertypes.RoleTool { end++ } } @@ -73,9 +73,9 @@ func RetainedStartForKeepRecentMessages(spans []MessageSpan, keepMessages int) i } // lastExplicitUserMessageIndex 返回最后一条非空用户消息的位置,用于保护最近明确指令。 -func lastExplicitUserMessageIndex(messages []provider.Message) int { +func lastExplicitUserMessageIndex(messages []providertypes.Message) int { for index := len(messages) - 1; index >= 0; index-- { - if messages[index].Role == provider.RoleUser && strings.TrimSpace(messages[index].Content) != "" { + if messages[index].Role == providertypes.RoleUser && strings.TrimSpace(messages[index].Content) != "" { return index } } diff --git a/internal/context/internalcompact/messages_test.go b/internal/context/internalcompact/messages_test.go index 09c73181..b561ccec 100644 --- a/internal/context/internalcompact/messages_test.go +++ b/internal/context/internalcompact/messages_test.go @@ -3,24 +3,24 @@ package internalcompact import ( "testing" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" ) func TestBuildMessageSpansPreservesToolBlocksAndProtectedTail(t *testing.T) { t.Parallel() - messages := []provider.Message{ - {Role: provider.RoleUser, Content: "old"}, + messages := []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "old"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-1", Content: "result"}, - {Role: provider.RoleAssistant, Content: "after tool"}, - {Role: provider.RoleUser, Content: "latest explicit instruction"}, - {Role: provider.RoleAssistant, Content: "latest answer"}, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "result"}, + {Role: providertypes.RoleAssistant, Content: "after tool"}, + {Role: providertypes.RoleUser, Content: "latest explicit instruction"}, + {Role: providertypes.RoleAssistant, Content: "latest answer"}, } spans := BuildMessageSpans(messages) diff --git a/internal/context/microcompact.go b/internal/context/microcompact.go index e90a5d24..a881a8b9 100644 --- a/internal/context/microcompact.go +++ b/internal/context/microcompact.go @@ -4,7 +4,7 @@ import ( "strings" "neo-code/internal/context/internalcompact" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" "neo-code/internal/tools" ) @@ -16,12 +16,12 @@ const ( ) // microCompactMessages 对裁剪后的消息做只读投影式微压缩,仅清理旧工具结果内容。 -func microCompactMessages(messages []provider.Message) []provider.Message { +func microCompactMessages(messages []providertypes.Message) []providertypes.Message { return microCompactMessagesWithPolicies(messages, nil) } // microCompactMessagesWithPolicies 按工具策略对裁剪后的消息做只读投影式微压缩。 -func microCompactMessagesWithPolicies(messages []provider.Message, policies MicroCompactPolicySource) []provider.Message { +func microCompactMessagesWithPolicies(messages []providertypes.Message, policies MicroCompactPolicySource) []providertypes.Message { cloned := cloneContextMessages(messages) if len(cloned) == 0 { return cloned @@ -63,31 +63,31 @@ func microCompactMessagesWithPolicies(messages []provider.Message, policies Micr } // cloneContextMessages 深拷贝消息切片,避免读时投影污染 runtime 持有的原始会话消息。 -func cloneContextMessages(messages []provider.Message) []provider.Message { +func cloneContextMessages(messages []providertypes.Message) []providertypes.Message { if len(messages) == 0 { return nil } - cloned := make([]provider.Message, 0, len(messages)) + cloned := make([]providertypes.Message, 0, len(messages)) for _, message := range messages { next := message - next.ToolCalls = append([]provider.ToolCall(nil), message.ToolCalls...) + next.ToolCalls = append([]providertypes.ToolCall(nil), message.ToolCalls...) cloned = append(cloned, next) } return cloned } // isToolCallSpan 判断当前 span 是否是由 assistant tool call 起始的原子工具块。 -func isToolCallSpan(messages []provider.Message, span internalcompact.MessageSpan) bool { +func isToolCallSpan(messages []providertypes.Message, span internalcompact.MessageSpan) bool { if span.Start < 0 || span.Start >= len(messages) { return false } message := messages[span.Start] - return message.Role == provider.RoleAssistant && len(message.ToolCalls) > 0 + return message.Role == providertypes.RoleAssistant && len(message.ToolCalls) > 0 } // compactableToolCallIDs 返回 assistant tool call 中可参与微压缩的调用 ID 集合。 -func compactableToolCallIDs(calls []provider.ToolCall, policies MicroCompactPolicySource) map[string]struct{} { +func compactableToolCallIDs(calls []providertypes.ToolCall, policies MicroCompactPolicySource) map[string]struct{} { if len(calls) == 0 { return nil } @@ -119,7 +119,7 @@ func toolParticipatesInMicroCompact(toolName string, policies MicroCompactPolicy } // hasCompactableToolContent 判断工具块中是否存在会影响保留预算的有效工具结果内容。 -func hasCompactableToolContent(messages []provider.Message, span internalcompact.MessageSpan, compactableIDs map[string]struct{}) bool { +func hasCompactableToolContent(messages []providertypes.Message, span internalcompact.MessageSpan, compactableIDs map[string]struct{}) bool { for messageIndex := span.Start + 1; messageIndex < span.End; messageIndex++ { if shouldClearToolMessage(messages[messageIndex], compactableIDs) { return true @@ -129,8 +129,8 @@ func hasCompactableToolContent(messages []provider.Message, span internalcompact } // shouldClearToolMessage 判断一条 tool 消息是否满足旧结果清理条件。 -func shouldClearToolMessage(message provider.Message, compactableIDs map[string]struct{}) bool { - if message.Role != provider.RoleTool || message.IsError { +func shouldClearToolMessage(message providertypes.Message, compactableIDs map[string]struct{}) bool { + if message.Role != providertypes.RoleTool || message.IsError { return false } if compactableIDs == nil { diff --git a/internal/context/microcompact_test.go b/internal/context/microcompact_test.go index a690b4c6..24e1e0e4 100644 --- a/internal/context/microcompact_test.go +++ b/internal/context/microcompact_test.go @@ -3,7 +3,7 @@ package context import ( "testing" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" "neo-code/internal/tools" ) @@ -19,31 +19,31 @@ func (s stubMicroCompactPolicySource) MicroCompactPolicy(name string) tools.Micr func TestMicroCompactMessagesClearsOlderCompactableToolResults(t *testing.T) { t.Parallel() - messages := []provider.Message{ - {Role: provider.RoleUser, Content: "older user"}, + messages := []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older user"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-1", Content: "old read result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "old read result"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-2", Name: "bash", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-3", Name: "webfetch", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, - {Role: provider.RoleUser, Content: "latest explicit instruction"}, - {Role: provider.RoleAssistant, Content: "current working reply"}, + {Role: providertypes.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, + {Role: providertypes.RoleUser, Content: "latest explicit instruction"}, + {Role: providertypes.RoleAssistant, Content: "current working reply"}, } got := microCompactMessages(messages) @@ -71,10 +71,10 @@ func TestMicroCompactMessagesHandlesEmptyAndInvalidSpanInputs(t *testing.T) { t.Fatalf("expected nil input to remain nil, got %+v", got) } - assistantOnly := []provider.Message{ + assistantOnly := []providertypes.Message{ { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "", Name: "bash", Arguments: "{}"}, }, }, @@ -88,37 +88,37 @@ func TestMicroCompactMessagesHandlesEmptyAndInvalidSpanInputs(t *testing.T) { func TestMicroCompactMessagesKeepsProtectedTailUntouched(t *testing.T) { t.Parallel() - messages := []provider.Message{ - {Role: provider.RoleUser, Content: "older user"}, + messages := []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older user"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-0", Name: "filesystem_grep", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-0", Content: "old grep result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-0", Content: "old grep result"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-1", Content: "recent read result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "recent read result"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-2", Name: "bash", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, - {Role: provider.RoleUser, Content: "latest explicit instruction"}, + {Role: providertypes.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + {Role: providertypes.RoleUser, Content: "latest explicit instruction"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-3", Name: "bash", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-3", Content: "tail bash result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-3", Content: "tail bash result"}, } got := microCompactMessages(messages) @@ -139,36 +139,36 @@ func TestMicroCompactMessagesKeepsProtectedTailUntouched(t *testing.T) { func TestMicroCompactMessagesKeepsPreservedToolsErrorsAndOrphans(t *testing.T) { t.Parallel() - messages := []provider.Message{ + messages := []providertypes.Message{ { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-1", Name: "custom_tool", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-1", Content: "custom result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "custom result"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-2", Name: "filesystem_edit", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-2", Content: "edit failed", IsError: true}, - {Role: provider.RoleTool, ToolCallID: "orphan", Content: "orphan result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-2", Content: "edit failed", IsError: true}, + {Role: providertypes.RoleTool, ToolCallID: "orphan", Content: "orphan result"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-3", Name: "filesystem_write_file", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-3", Content: microCompactClearedMessage}, + {Role: providertypes.RoleTool, ToolCallID: "call-3", Content: microCompactClearedMessage}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-4", Name: "filesystem_grep", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-4", Content: ""}, + {Role: providertypes.RoleTool, ToolCallID: "call-4", Content: ""}, } got := microCompactMessagesWithPolicies(messages, stubMicroCompactPolicySource{ @@ -194,33 +194,33 @@ func TestMicroCompactMessagesKeepsPreservedToolsErrorsAndOrphans(t *testing.T) { func TestMicroCompactMessagesClearsOnlyNonPreservedResultsInMixedToolSpan(t *testing.T) { t.Parallel() - messages := []provider.Message{ - {Role: provider.RoleUser, Content: "older user"}, + messages := []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older user"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, {ID: "call-2", Name: "custom_tool", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-1", Content: "read result"}, - {Role: provider.RoleTool, ToolCallID: "call-2", Content: "custom result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "read result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-2", Content: "custom result"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-3", Name: "bash", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-3", Content: "recent bash result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-3", Content: "recent bash result"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-4", Name: "webfetch", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-4", Content: "latest webfetch result"}, - {Role: provider.RoleUser, Content: "latest explicit instruction"}, - {Role: provider.RoleAssistant, Content: "current reply"}, + {Role: providertypes.RoleTool, ToolCallID: "call-4", Content: "latest webfetch result"}, + {Role: providertypes.RoleUser, Content: "latest explicit instruction"}, + {Role: providertypes.RoleAssistant, Content: "current reply"}, } got := microCompactMessagesWithPolicies(messages, stubMicroCompactPolicySource{ @@ -240,30 +240,30 @@ func TestMicroCompactMessagesClearsOnlyNonPreservedResultsInMixedToolSpan(t *tes func TestMicroCompactMessagesTreatsNewToolsAsCompactableByDefault(t *testing.T) { t.Parallel() - messages := []provider.Message{ - {Role: provider.RoleUser, Content: "older user"}, + messages := []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older user"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-1", Name: "repo_search", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-1", Content: "old repo search result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "old repo search result"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-2", Name: "bash", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-3", Name: "webfetch", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, - {Role: provider.RoleUser, Content: "latest explicit instruction"}, + {Role: providertypes.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, + {Role: providertypes.RoleUser, Content: "latest explicit instruction"}, } got := microCompactMessagesWithPolicies(messages, stubMicroCompactPolicySource{}) @@ -275,45 +275,45 @@ func TestMicroCompactMessagesTreatsNewToolsAsCompactableByDefault(t *testing.T) func TestMicroCompactMessagesSkipsEmptyRecentSpansWhenCountingRetainedBudget(t *testing.T) { t.Parallel() - messages := []provider.Message{ - {Role: provider.RoleUser, Content: "older user"}, + messages := []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older user"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-1", Content: "older read result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "older read result"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-2", Name: "filesystem_grep", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-2", Content: "middle grep result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-2", Content: "middle grep result"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-3", Name: "filesystem_edit", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-3", Content: "near edit result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-3", Content: "near edit result"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-4", Name: "bash", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-4", Content: "", IsError: true}, + {Role: providertypes.RoleTool, ToolCallID: "call-4", Content: "", IsError: true}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-5", Name: "webfetch", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-5", Content: ""}, - {Role: provider.RoleUser, Content: "latest explicit instruction"}, - {Role: provider.RoleAssistant, Content: "current reply"}, + {Role: providertypes.RoleTool, ToolCallID: "call-5", Content: ""}, + {Role: providertypes.RoleUser, Content: "latest explicit instruction"}, + {Role: providertypes.RoleAssistant, Content: "current reply"}, } got := microCompactMessages(messages) @@ -337,8 +337,8 @@ func TestMicroCompactMessagesSkipsEmptyRecentSpansWhenCountingRetainedBudget(t * func TestMicroCompactMessagesSkipsToolMessagesWhenCompactableIDsMissing(t *testing.T) { t.Parallel() - messages := []provider.Message{ - {Role: provider.RoleTool, ToolCallID: "orphan", Content: "orphan result"}, + messages := []providertypes.Message{ + {Role: providertypes.RoleTool, ToolCallID: "orphan", Content: "orphan result"}, } got := microCompactMessagesWithPolicies(messages, stubMicroCompactPolicySource{}) diff --git a/internal/context/trim.go b/internal/context/trim.go index 91d04e8f..d818a6fe 100644 --- a/internal/context/trim.go +++ b/internal/context/trim.go @@ -2,21 +2,21 @@ package context import ( "neo-code/internal/context/internalcompact" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" ) const maxRetainedMessageSpans = 10 // trimMessages 按消息分段裁剪上下文,并始终保护最近一条明确用户指令所在尾部。 -func trimMessages(messages []provider.Message) []provider.Message { +func trimMessages(messages []providertypes.Message) []providertypes.Message { spans := internalcompact.BuildMessageSpans(messages) if len(spans) <= maxRetainedMessageSpans { - return append([]provider.Message(nil), messages...) + return append([]providertypes.Message(nil), messages...) } start := spans[len(spans)-maxRetainedMessageSpans].Start if protectedStart, ok := internalcompact.ProtectedTailStart(spans); ok && protectedStart < start { start = protectedStart } - return append([]provider.Message(nil), messages[start:]...) + return append([]providertypes.Message(nil), messages[start:]...) } diff --git a/internal/context/trim_policy.go b/internal/context/trim_policy.go index e04da9a4..6e7c73b1 100644 --- a/internal/context/trim_policy.go +++ b/internal/context/trim_policy.go @@ -1,16 +1,18 @@ package context -import "neo-code/internal/provider" +import ( + providertypes "neo-code/internal/provider/types" +) // messageTrimPolicy 约束消息裁剪策略的最小接口,避免 Builder 直接持有裁剪细节。 type messageTrimPolicy interface { - Trim(messages []provider.Message) []provider.Message + Trim(messages []providertypes.Message) []providertypes.Message } // spanMessageTrimPolicy 以消息 span 为单位裁剪历史,确保 tool block 不被拆散。 type spanMessageTrimPolicy struct{} // Trim 返回保留关键 tool block 原子性的裁剪后消息副本。 -func (spanMessageTrimPolicy) Trim(messages []provider.Message) []provider.Message { +func (spanMessageTrimPolicy) Trim(messages []providertypes.Message) []providertypes.Message { return trimMessages(messages) } diff --git a/internal/context/types.go b/internal/context/types.go index e1bb0c93..53af2bfa 100644 --- a/internal/context/types.go +++ b/internal/context/types.go @@ -3,7 +3,7 @@ package context import ( "context" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" "neo-code/internal/tools" ) @@ -14,7 +14,7 @@ type Builder interface { // BuildInput contains the runtime state needed to assemble model context. type BuildInput struct { - Messages []provider.Message + Messages []providertypes.Message Metadata Metadata Compact CompactOptions } @@ -22,7 +22,7 @@ type BuildInput struct { // BuildResult is the provider-facing context produced for a single round. type BuildResult struct { SystemPrompt string - Messages []provider.Message + Messages []providertypes.Message } // MicroCompactPolicySource 定义 context 读取工具 micro compact 策略的最小依赖。 diff --git a/internal/provider/catalog/service_test.go b/internal/provider/catalog/service_test.go index 3d5d0b05..da3218dc 100644 --- a/internal/provider/catalog/service_test.go +++ b/internal/provider/catalog/service_test.go @@ -8,6 +8,7 @@ import ( "neo-code/internal/config" "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" ) func TestNewService(t *testing.T) { @@ -234,7 +235,7 @@ func containsModelDescriptorID(models []config.ModelDescriptor, modelID string) type catalogTestProvider struct{} -func (catalogTestProvider) Chat(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { +func (catalogTestProvider) Chat(ctx context.Context, req providertypes.ChatRequest, events chan<- providertypes.StreamEvent) error { return nil } diff --git a/internal/provider/discovery/discovery.go b/internal/provider/discovery/discovery.go deleted file mode 100644 index 4311117c..00000000 --- a/internal/provider/discovery/discovery.go +++ /dev/null @@ -1,50 +0,0 @@ -package discovery - -import ( - "context" - "encoding/json" - "fmt" - "io" - "net/http" - "strings" -) - -type openAIModelsResponse struct { - Data []map[string]any `json:"data"` -} - -// FetchOpenAICompatibleModels fetches raw model objects from an OpenAI-compatible /models endpoint. -func FetchOpenAICompatibleModels(ctx context.Context, client *http.Client, baseURL string, apiKey string) ([]map[string]any, error) { - endpoint := strings.TrimRight(strings.TrimSpace(baseURL), "/") + "/models" - req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil) - if err != nil { - return nil, fmt.Errorf("provider discovery: build request: %w", err) - } - req.Header.Set("Accept", "application/json") - if strings.TrimSpace(apiKey) != "" { - req.Header.Set("Authorization", "Bearer "+strings.TrimSpace(apiKey)) - } - - resp, err := client.Do(req) - if err != nil { - return nil, fmt.Errorf("provider discovery: send request: %w", err) - } - defer func(body io.ReadCloser) { - _ = body.Close() - }(resp.Body) - - if resp.StatusCode >= http.StatusBadRequest { - data, _ := io.ReadAll(resp.Body) - body := strings.TrimSpace(string(data)) - if body == "" { - body = resp.Status - } - return nil, fmt.Errorf("provider discovery: %s", body) - } - - var payload openAIModelsResponse - if err := json.NewDecoder(resp.Body).Decode(&payload); err != nil { - return nil, fmt.Errorf("provider discovery: decode response: %w", err) - } - return payload.Data, nil -} diff --git a/internal/provider/discovery/discovery_test.go b/internal/provider/discovery/discovery_test.go deleted file mode 100644 index 50380a4c..00000000 --- a/internal/provider/discovery/discovery_test.go +++ /dev/null @@ -1,51 +0,0 @@ -package discovery - -import ( - "context" - "net/http" - "net/http/httptest" - "strings" - "testing" -) - -func TestFetchOpenAICompatibleModels(t *testing.T) { - t.Parallel() - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if r.URL.Path != "/models" { - t.Fatalf("unexpected path %q", r.URL.Path) - } - if got := r.Header.Get("Authorization"); got != "Bearer test-key" { - t.Fatalf("unexpected auth header %q", got) - } - w.Header().Set("Content-Type", "application/json") - _, _ = w.Write([]byte(`{"data":[{"id":"gpt-test","context_window":128000,"extra":"kept"}]}`)) - })) - defer server.Close() - - models, err := FetchOpenAICompatibleModels(context.Background(), server.Client(), server.URL, "test-key") - if err != nil { - t.Fatalf("FetchOpenAICompatibleModels() error = %v", err) - } - if len(models) != 1 || models[0]["id"] != "gpt-test" { - t.Fatalf("unexpected models payload: %+v", models) - } - if models[0]["extra"] != "kept" { - t.Fatalf("expected unknown fields to remain in raw payload, got %+v", models[0]) - } -} - -func TestFetchOpenAICompatibleModelsHTTPError(t *testing.T) { - t.Parallel() - - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusBadGateway) - _, _ = w.Write([]byte("gateway failed")) - })) - defer server.Close() - - _, err := FetchOpenAICompatibleModels(context.Background(), server.Client(), server.URL, "test-key") - if err == nil || !strings.Contains(err.Error(), "gateway failed") { - t.Fatalf("expected gateway failure error, got %v", err) - } -} diff --git a/internal/provider/openai/discovery.go b/internal/provider/openai/discovery.go new file mode 100644 index 00000000..b20929d8 --- /dev/null +++ b/internal/provider/openai/discovery.go @@ -0,0 +1,51 @@ +package openai + +import ( + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "strings" +) + +// openAIModelsResponse 表示 /models 端点的响应结构。 +type openAIModelsResponse struct { + Data []map[string]any `json:"data"` +} + +// fetchModels 从 OpenAI 兼容的 /models 端点获取原始模型列表。 +func (p *Provider) fetchModels(ctx context.Context) ([]map[string]any, error) { + endpoint := strings.TrimRight(strings.TrimSpace(p.cfg.BaseURL), "/") + "/models" + req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil) + if err != nil { + return nil, fmt.Errorf("openai provider: build models request: %w", err) + } + req.Header.Set("Accept", "application/json") + if strings.TrimSpace(p.cfg.APIKey) != "" { + req.Header.Set("Authorization", "Bearer "+strings.TrimSpace(p.cfg.APIKey)) + } + + resp, err := p.client.Do(req) + if err != nil { + return nil, fmt.Errorf("openai provider: send models request: %w", err) + } + defer func(body io.ReadCloser) { + _ = body.Close() + }(resp.Body) + + if resp.StatusCode >= http.StatusBadRequest { + data, _ := io.ReadAll(resp.Body) + body := strings.TrimSpace(string(data)) + if body == "" { + body = resp.Status + } + return nil, fmt.Errorf("openai provider: models endpoint %s", body) + } + + var payload openAIModelsResponse + if err := json.NewDecoder(resp.Body).Decode(&payload); err != nil { + return nil, fmt.Errorf("openai provider: decode models response: %w", err) + } + return payload.Data, nil +} diff --git a/internal/provider/openai/driver.go b/internal/provider/openai/driver.go new file mode 100644 index 00000000..fbde6820 --- /dev/null +++ b/internal/provider/openai/driver.go @@ -0,0 +1,35 @@ +package openai + +import ( + "context" + "net/http" + + "neo-code/internal/config" + "neo-code/internal/provider" + "neo-code/internal/provider/transport" +) + +// DriverName 是 OpenAI 驱动的注册标识。 +const DriverName = "openai" + +// defaultRetryTransport 返回内置的带重试的 HTTP Transport。 +func defaultRetryTransport() http.RoundTripper { + return transport.NewRetryTransport(http.DefaultTransport, transport.DefaultRetryConfig()) +} + +// Driver 返回 OpenAI 协议驱动的定义,供 Registry 注册使用。 +func Driver() provider.DriverDefinition { + return provider.DriverDefinition{ + Name: DriverName, + Build: func(ctx context.Context, cfg config.ResolvedProviderConfig) (provider.Provider, error) { + return New(cfg, withTransport(defaultRetryTransport())) + }, + Discover: func(ctx context.Context, cfg config.ResolvedProviderConfig) ([]config.ModelDescriptor, error) { + p, err := New(cfg, withTransport(defaultRetryTransport())) + if err != nil { + return nil, err + } + return p.DiscoverModels(ctx) + }, + } +} diff --git a/internal/provider/openai/events.go b/internal/provider/openai/events.go new file mode 100644 index 00000000..50e13338 --- /dev/null +++ b/internal/provider/openai/events.go @@ -0,0 +1,63 @@ +package openai + +import ( + "context" + + providertypes "neo-code/internal/provider/types" +) + +// emitTextDelta 发送文本增量事件,空文本时跳过。 +func emitTextDelta(ctx context.Context, events chan<- providertypes.StreamEvent, text string) error { + if text == "" { + return nil + } + return emitStreamEvent(ctx, events, providertypes.NewTextDeltaStreamEvent(text)) +} + +// emitToolCallStart 发送工具调用开始事件,空名称时跳过。 +func emitToolCallStart(ctx context.Context, events chan<- providertypes.StreamEvent, index int, id, name string) error { + if name == "" { + return nil + } + return emitStreamEvent(ctx, events, providertypes.NewToolCallStartStreamEvent(index, id, name)) +} + +// emitToolCallDelta 发送工具调用参数增量事件。 +// id 为工具调用 ID,由上游 mergeToolCallDelta 从累积状态中传入。 +func emitToolCallDelta(ctx context.Context, events chan<- providertypes.StreamEvent, index int, id, argumentsDelta string) error { + if argumentsDelta == "" { + return nil + } + return emitStreamEvent(ctx, events, providertypes.NewToolCallDeltaStreamEvent(index, id, argumentsDelta)) +} + +// emitMessageDone 发送消息完成事件。 +func emitMessageDone(ctx context.Context, events chan<- providertypes.StreamEvent, finishReason string, usage *providertypes.Usage) error { + return emitStreamEvent(ctx, events, providertypes.NewMessageDoneStreamEvent(finishReason, usage)) +} + +// emitStreamEvent 通过 channel 安全发送流式事件,支持上下文取消和 nil channel 保护。 +func emitStreamEvent(ctx context.Context, events chan<- providertypes.StreamEvent, event providertypes.StreamEvent) error { + if events == nil { + return nil + } + + select { + case events <- event: + return nil + case <-ctx.Done(): + return ctx.Err() + } +} + +// flushDataLines 逐行处理缓冲的 data lines,每行作为独立 payload 通过 processChunk 处理。 +// SSE 规范允许同一事件内多行 data 拼接,但 OpenAI 实际行为是每行 data 为独立 JSON, +// 因此逐行处理更可靠,避免拼接产生无效 JSON。 +func flushDataLines(dataLines []string, processChunk func(string) error) error { + for _, line := range dataLines { + if err := processChunk(line); err != nil { + return err + } + } + return nil +} diff --git a/internal/provider/openai/openai.go b/internal/provider/openai/openai.go deleted file mode 100644 index 2e2cbaf8..00000000 --- a/internal/provider/openai/openai.go +++ /dev/null @@ -1,501 +0,0 @@ -package openai - -import ( - "bytes" - "context" - "encoding/json" - "errors" - "fmt" - "io" - "log" - "net/http" - "strings" - "time" - - "neo-code/internal/config" - "neo-code/internal/provider" - modeldiscovery "neo-code/internal/provider/discovery" - "neo-code/internal/provider/transport" -) - -type Provider struct { - cfg config.ResolvedProviderConfig - client *http.Client -} - -type buildOptions struct { - transport http.RoundTripper -} - -type buildOption func(*buildOptions) - -// withTransport 注入自定义 HTTP Transport(如 RetryTransport)。 -func withTransport(rt http.RoundTripper) buildOption { - return func(o *buildOptions) { - o.transport = rt - } -} - -const DriverName = "openai" - -// defaultRetryTransport 返回内置的带重试的 HTTP Transport。 -func defaultRetryTransport() http.RoundTripper { - return transport.NewRetryTransport(http.DefaultTransport, transport.DefaultRetryConfig()) -} - -// Driver 返回 OpenAI 协议驱动的定义。 -func Driver() provider.DriverDefinition { - return provider.DriverDefinition{ - Name: DriverName, - Build: func(ctx context.Context, cfg config.ResolvedProviderConfig) (provider.Provider, error) { - return New(cfg, withTransport(defaultRetryTransport())) - }, - Discover: func(ctx context.Context, cfg config.ResolvedProviderConfig) ([]config.ModelDescriptor, error) { - p, err := New(cfg, withTransport(defaultRetryTransport())) - if err != nil { - return nil, err - } - return p.DiscoverModels(ctx) - }, - } -} - -func New(cfg config.ResolvedProviderConfig, opts ...buildOption) (*Provider, error) { - if err := cfg.Validate(); err != nil { - return nil, fmt.Errorf("openai provider: %w", err) - } - if strings.TrimSpace(cfg.APIKey) == "" { - return nil, errors.New("openai provider: api key is empty") - } - - o := &buildOptions{ - transport: http.DefaultTransport, - } - for _, apply := range opts { - apply(o) - } - - return &Provider{ - cfg: cfg, - client: &http.Client{ - Timeout: 90 * time.Second, - Transport: o.transport, - }, - }, nil -} - -func (p *Provider) DiscoverModels(ctx context.Context) ([]config.ModelDescriptor, error) { - rawModels, err := modeldiscovery.FetchOpenAICompatibleModels(ctx, p.client, p.cfg.BaseURL, p.cfg.APIKey) - if err != nil { - return nil, err - } - - descriptors := make([]config.ModelDescriptor, 0, len(rawModels)) - for _, raw := range rawModels { - descriptor, ok := config.DescriptorFromRawModel(raw) - if !ok { - continue - } - descriptors = append(descriptors, descriptor) - } - return config.MergeModelDescriptors(descriptors), nil -} - -// Chat 发起 SSE 流式对话请求。 -// 流中途断连或协议错误时直接返回错误,由上层调用方决定重试策略。 -func (p *Provider) Chat(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { - payload, err := p.buildRequest(req) - if err != nil { - return err - } - - body, err := json.Marshal(payload) - if err != nil { - return fmt.Errorf("openai provider: marshal request: %w", err) - } - - endpoint := strings.TrimRight(p.cfg.BaseURL, "/") + "/chat/completions" - httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(body)) - if err != nil { - return fmt.Errorf("openai provider: build request: %w", err) - } - httpReq.Header.Set("Authorization", "Bearer "+p.cfg.APIKey) - httpReq.Header.Set("Content-Type", "application/json") - httpReq.Header.Set("Accept", "text/event-stream") - - resp, err := p.client.Do(httpReq) - if err != nil { - return fmt.Errorf("openai provider: send request: %w", err) - } - defer func(Body io.ReadCloser) { - err := Body.Close() - if err != nil { - log.Printf("openai provider: close response body: %v", err) - } - }(resp.Body) - - if resp.StatusCode >= http.StatusBadRequest { - return p.parseError(resp) - } - - return p.consumeStream(ctx, resp.Body, events) -} - -func (p *Provider) buildRequest(req provider.ChatRequest) (chatCompletionRequest, error) { - model := strings.TrimSpace(req.Model) - if model == "" { - model = strings.TrimSpace(p.cfg.Model) - } - if model == "" { - return chatCompletionRequest{}, errors.New("openai provider: model is empty") - } - - payload := chatCompletionRequest{ - Model: model, - Stream: true, - Messages: make([]openAIMessage, 0, len(req.Messages)+1), - } - - if strings.TrimSpace(req.SystemPrompt) != "" { - payload.Messages = append(payload.Messages, openAIMessage{ - Role: provider.RoleSystem, - Content: req.SystemPrompt, - }) - } - - for _, message := range req.Messages { - payload.Messages = append(payload.Messages, toOpenAIMessage(message)) - } - - if len(req.Tools) > 0 { - payload.ToolChoice = "auto" - payload.Tools = make([]openAIToolDefinition, 0, len(req.Tools)) - for _, spec := range req.Tools { - payload.Tools = append(payload.Tools, openAIToolDefinition{ - Type: "function", - Function: openAIFunctionDefinition{ - Name: spec.Name, - Description: spec.Description, - Parameters: spec.Schema, - }, - }) - } - } - - return payload, nil -} - -// consumeStream 消费 SSE 响应流,使用有界读取器防止缓冲区溢出。 -func (p *Provider) consumeStream( - ctx context.Context, - body io.Reader, - events chan<- provider.StreamEvent, -) error { - reader := newBoundedSSEReader(body) - - var ( - finishReason string - usage provider.Usage - done bool - toolCalls = make(map[int]*provider.ToolCall) - ) - - dataLines := make([]string, 0, 4) - - // processChunk 解析单个 SSE data payload,发送事件。 - processChunk := func(payload string) error { - if strings.TrimSpace(payload) == "[DONE]" { - done = true - return nil - } - - var chunk chatCompletionChunk - if err := json.Unmarshal([]byte(payload), &chunk); err != nil { - return fmt.Errorf("openai provider: decode stream chunk: %w", err) - } - - if chunk.Error != nil && strings.TrimSpace(chunk.Error.Message) != "" { - return errors.New(chunk.Error.Message) - } - - extractStreamUsage(&usage, chunk.Usage) - - for _, choice := range chunk.Choices { - if choice.FinishReason != "" { - finishReason = choice.FinishReason - } - if choice.Delta.Content != "" { - if err := emitTextDelta(ctx, events, choice.Delta.Content); err != nil { - return err - } - } - for _, delta := range choice.Delta.ToolCalls { - if err := mergeToolCallDelta(ctx, events, toolCalls, delta); err != nil { - return err - } - } - } - return nil - } - - // finishStream 统一的流结束处理:发送 message_done 事件。 - finishStream := func() error { - return emitMessageDone(ctx, events, finishReason, &usage) - } - - flushPendingData := func() error { - defer func() { dataLines = dataLines[:0] }() - return flushDataLines(dataLines, processChunk) - } - - for { - line, err := reader.ReadLine() - - if err != nil && !errors.Is(err, io.EOF) { - // 非 EOF 的读取错误:先刷新缓冲的 data 行,再包装为流中断, - // 避免中断前最后一段数据丢失。 - if flushErr := flushPendingData(); flushErr != nil { - return flushErr - } - return fmt.Errorf("%w: %w", provider.ErrStreamInterrupted, err) - } - - trimmed := line - - switch { - case strings.HasPrefix(trimmed, "data:"): - dataLines = append(dataLines, strings.TrimSpace(strings.TrimPrefix(trimmed, "data:"))) - case trimmed == "": - if flushErr := flushPendingData(); flushErr != nil { - return flushErr - } - if done { - return finishStream() - } - case strings.HasPrefix(trimmed, ":"): - // SSE comment/heartbeat; ignore. - } - - if errors.Is(err, io.EOF) { - if flushErr := flushPendingData(); flushErr != nil { - return flushErr - } - return finishStream() - } - } -} - -func (p *Provider) parseError(resp *http.Response) error { - data, readErr := io.ReadAll(resp.Body) - if readErr != nil { - return provider.NewProviderErrorFromStatus(resp.StatusCode, - fmt.Sprintf("openai provider: read error response: %v", readErr)) - } - - var parsed openAIErrorResponse - if err := json.Unmarshal(data, &parsed); err == nil && strings.TrimSpace(parsed.Error.Message) != "" { - return provider.NewProviderErrorFromStatus(resp.StatusCode, parsed.Error.Message) - } - - bodyText := strings.TrimSpace(string(data)) - if bodyText == "" { - return provider.NewProviderErrorFromStatus(resp.StatusCode, resp.Status) - } - - return provider.NewProviderErrorFromStatus(resp.StatusCode, bodyText) -} - -func toOpenAIMessage(message provider.Message) openAIMessage { - out := openAIMessage{ - Role: message.Role, - Content: message.Content, - ToolCallID: message.ToolCallID, - } - - if len(message.ToolCalls) > 0 { - out.ToolCalls = make([]openAIToolCall, 0, len(message.ToolCalls)) - for _, call := range message.ToolCalls { - out.ToolCalls = append(out.ToolCalls, openAIToolCall{ - ID: call.ID, - Type: "function", - Function: openAIFunctionCall{ - Name: call.Name, - Arguments: call.Arguments, - }, - }) - } - } - - return out -} - -func emitTextDelta(ctx context.Context, events chan<- provider.StreamEvent, text string) error { - if text == "" { - return nil - } - return emitStreamEvent(ctx, events, provider.NewTextDeltaStreamEvent(text)) -} - -func emitToolCallStart(ctx context.Context, events chan<- provider.StreamEvent, index int, id, name string) error { - if name == "" { - return nil - } - return emitStreamEvent(ctx, events, provider.NewToolCallStartStreamEvent(index, id, name)) -} - -// emitToolCallDelta 发送工具调用参数增量事件。 -// id 为工具调用 ID,由上游 mergeToolCallDelta 从累积状态中传入。 -func emitToolCallDelta(ctx context.Context, events chan<- provider.StreamEvent, index int, id, argumentsDelta string) error { - if argumentsDelta == "" { - return nil - } - return emitStreamEvent(ctx, events, provider.NewToolCallDeltaStreamEvent(index, id, argumentsDelta)) -} - -// emitMessageDone 发送消息完成事件。 -func emitMessageDone(ctx context.Context, events chan<- provider.StreamEvent, finishReason string, usage *provider.Usage) error { - return emitStreamEvent(ctx, events, provider.NewMessageDoneStreamEvent(finishReason, usage)) -} - -// extractStreamUsage 从 OpenAI usage 响应提取并覆盖累积的 token 统计。 -func extractStreamUsage(usage *provider.Usage, raw *openAIUsage) { - if raw == nil { - return - } - *usage = provider.Usage{ - InputTokens: raw.PromptTokens, - OutputTokens: raw.CompletionTokens, - TotalTokens: raw.TotalTokens, - } -} - -// mergeToolCallDelta 将单个 tool call delta 累积到 toolCalls map 中。 -// 首次发现带名称的 delta 时发送 tool_call_start 事件; -// 每次收到 arguments 增量时发送 tool_call_delta 事件。 -func mergeToolCallDelta(ctx context.Context, events chan<- provider.StreamEvent, toolCalls map[int]*provider.ToolCall, delta toolCallDelta) error { - call, exists := toolCalls[delta.Index] - if !exists { - call = &provider.ToolCall{} - toolCalls[delta.Index] = call - } - - hadName := strings.TrimSpace(call.Name) != "" - - if id := strings.TrimSpace(delta.ID); id != "" { - call.ID = id - } - if name := strings.TrimSpace(delta.Function.Name); name != "" { - call.Name = name - } - - if !hadName && strings.TrimSpace(call.Name) != "" { - if err := emitToolCallStart(ctx, events, delta.Index, call.ID, call.Name); err != nil { - return err - } - } - - // 发送参数增量事件(同一 chunk 可能同时携带 name 和 arguments) - if args := delta.Function.Arguments; args != "" { - call.Arguments += args - if err := emitToolCallDelta(ctx, events, delta.Index, call.ID, args); err != nil { - return err - } - } - return nil -} - -func emitStreamEvent(ctx context.Context, events chan<- provider.StreamEvent, event provider.StreamEvent) error { - if events == nil { - return nil - } - - select { - case events <- event: - return nil - case <-ctx.Done(): - return ctx.Err() - } -} - -// flushDataLines 将缓冲的 data lines 合并为单个 payload 并通过 processChunk 处理。 -func flushDataLines(dataLines []string, processChunk func(string) error) error { - if len(dataLines) == 0 { - return nil - } - return processChunk(strings.Join(dataLines, "\n")) -} - -type chatCompletionRequest struct { - Model string `json:"model"` - Messages []openAIMessage `json:"messages"` - Tools []openAIToolDefinition `json:"tools,omitempty"` - ToolChoice string `json:"tool_choice,omitempty"` - Stream bool `json:"stream"` -} - -type openAIMessage struct { - Role string `json:"role"` - Content string `json:"content,omitempty"` - ToolCallID string `json:"tool_call_id,omitempty"` - ToolCalls []openAIToolCall `json:"tool_calls,omitempty"` -} - -type openAIToolDefinition struct { - Type string `json:"type"` - Function openAIFunctionDefinition `json:"function"` -} - -type openAIFunctionDefinition struct { - Name string `json:"name"` - Description string `json:"description,omitempty"` - Parameters map[string]any `json:"parameters,omitempty"` -} - -type openAIToolCall struct { - ID string `json:"id,omitempty"` - Type string `json:"type,omitempty"` - Function openAIFunctionCall `json:"function"` -} - -type openAIFunctionCall struct { - Name string `json:"name,omitempty"` - Arguments string `json:"arguments,omitempty"` -} - -type chatCompletionChunk struct { - Choices []struct { - Index int `json:"index"` - Delta chunkDelta `json:"delta"` - FinishReason string `json:"finish_reason"` - } `json:"choices"` - Usage *openAIUsage `json:"usage,omitempty"` - Error *struct { - Message string `json:"message"` - } `json:"error,omitempty"` -} - -type chunkDelta struct { - Role string `json:"role,omitempty"` - Content string `json:"content,omitempty"` - ToolCalls []toolCallDelta `json:"tool_calls,omitempty"` -} - -type toolCallDelta struct { - Index int `json:"index"` - ID string `json:"id,omitempty"` - Type string `json:"type,omitempty"` - Function openAIFunctionCall `json:"function"` -} - -type openAIUsage struct { - PromptTokens int `json:"prompt_tokens"` - CompletionTokens int `json:"completion_tokens"` - TotalTokens int `json:"total_tokens"` -} - -type openAIErrorResponse struct { - Error struct { - Message string `json:"message"` - Code string `json:"code,omitempty"` - } `json:"error"` -} diff --git a/internal/provider/openai/openai_test.go b/internal/provider/openai/openai_test.go index bfbca884..33179a2e 100644 --- a/internal/provider/openai/openai_test.go +++ b/internal/provider/openai/openai_test.go @@ -9,9 +9,11 @@ import ( "net/http/httptest" "strings" "testing" + "time" "neo-code/internal/config" - domain "neo-code/internal/provider" + "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" ) func TestDriver(t *testing.T) { @@ -73,7 +75,6 @@ func TestNewValidationErrors(t *testing.T) { t.Run("invalid config validate fails", func(t *testing.T) { t.Parallel() - // 空字符串的 BaseURL 和 Model 会导致 Validate 失败(取决于 config 实现) cfg := config.ResolvedProviderConfig{ ProviderConfig: config.ProviderConfig{ Driver: DriverName, @@ -84,12 +85,9 @@ func TestNewValidationErrors(t *testing.T) { APIKey: "test-key", } _, err := New(cfg) - // 验证失败时应该返回错误 if err != nil { - // 预期行为:config 校验不通过 return } - // 如果校验通过了,也接受(取决于具体实现) }) } @@ -97,7 +95,7 @@ func TestNewDefaultTransportWhenNoOption(t *testing.T) { t.Parallel() cfg := resolvedConfig("", "") - provider, err := New(cfg) // 不传任何 buildOption + provider, err := New(cfg) if err != nil { t.Fatalf("New() error = %v", err) } @@ -150,106 +148,726 @@ func TestDiscoverModels(t *testing.T) { } } -func TestEmitToolCallDelta(t *testing.T) { +// --- toOpenAIMessage 转换测试 --- + +func TestToOpenAIMessage_BasicMessage(t *testing.T) { + t.Parallel() + + msg := providertypes.Message{ + Role: "user", + Content: "hello world", + } + result := toOpenAIMessage(msg) + + if result.Role != "user" || result.Content != "hello world" { + t.Fatalf("unexpected basic message: role=%q content=%q", result.Role, result.Content) + } + if result.ToolCallID != "" || len(result.ToolCalls) > 0 { + t.Fatal("basic message should not have tool call fields") + } +} + +func TestToOpenAIMessage_ToolRoleMessage(t *testing.T) { + t.Parallel() + + msg := providertypes.Message{ + Role: "tool", + Content: "result data", + ToolCallID: "call_123", + } + result := toOpenAIMessage(msg) + + if result.Role != "tool" || result.ToolCallID != "call_123" { + t.Fatalf("unexpected tool message: role=%q toolCallID=%q", result.Role, result.ToolCallID) + } +} + +func TestToOpenAIMessage_AssistantWithToolCalls(t *testing.T) { + t.Parallel() + + msg := providertypes.Message{ + Role: "assistant", + ToolCalls: []providertypes.ToolCall{ + {ID: "call_1", Name: "read_file", Arguments: `{"path":"main.go"}`}, + {ID: "call_2", Name: "write_file", Arguments: `{"path":"test.go","content":"..."}`}, + }, + } + result := toOpenAIMessage(msg) + + if len(result.ToolCalls) != 2 { + t.Fatalf("expected 2 tool calls, got %d", len(result.ToolCalls)) + } + tc1 := result.ToolCalls[0] + if tc1.ID != "call_1" || tc1.Type != "function" { + t.Fatalf("unexpected first tool call: id=%q type=%q", tc1.ID, tc1.Type) + } + if tc1.Function.Name != "read_file" || tc1.Function.Arguments != `{"path":"main.go"}` { + t.Fatalf("unexpected first function: name=%q args=%q", tc1.Function.Name, tc1.Function.Arguments) + } + tc2 := result.ToolCalls[1] + if tc2.Function.Name != "write_file" { + t.Fatalf("unexpected second function name: %q", tc2.Function.Name) + } +} + +func TestToOpenAIMessage_EmptyToolCalls(t *testing.T) { + t.Parallel() + + msg := providertypes.Message{Role: "user", Content: "test"} + result := toOpenAIMessage(msg) + if len(result.ToolCalls) != 0 { + t.Fatalf("expected no tool calls for user message, got %d", len(result.ToolCalls)) + } +} + +// --- extractStreamUsage 测试 --- + +func TestExtractStreamUsage_NilInput(t *testing.T) { + t.Parallel() + + var usage providertypes.Usage + extractStreamUsage(&usage, nil) + if usage.InputTokens != 0 || usage.OutputTokens != 0 || usage.TotalTokens != 0 { + t.Fatalf("expected zero values for nil input, got %+v", usage) + } +} + +func TestExtractStreamUsage_NormalValues(t *testing.T) { + t.Parallel() + + var usage providertypes.Usage + raw := &openAIUsage{PromptTokens: 100, CompletionTokens: 50, TotalTokens: 150} + extractStreamUsage(&usage, raw) + if usage.InputTokens != 100 || usage.OutputTokens != 50 || usage.TotalTokens != 150 { + t.Fatalf("unexpected usage values: %+v", usage) + } +} + +func TestExtractStreamUsage_ZeroValues(t *testing.T) { + t.Parallel() + + var usage providertypes.Usage + usage.InputTokens = 999 + raw := &openAIUsage{} + extractStreamUsage(&usage, raw) + if usage.InputTokens != 0 || usage.OutputTokens != 0 || usage.TotalTokens != 0 { + t.Fatalf("expected zero values to overwrite previous, got %+v", usage) + } +} + +func TestExtractStreamUsage_MultipleOverwrites(t *testing.T) { + t.Parallel() + + var usage providertypes.Usage + extractStreamUsage(&usage, &openAIUsage{PromptTokens: 10, CompletionTokens: 5, TotalTokens: 15}) + extractStreamUsage(&usage, &openAIUsage{PromptTokens: 20, CompletionTokens: 10, TotalTokens: 30}) + if usage.TotalTokens != 30 { + t.Fatalf("expected last write to win (total=30), got %d", usage.TotalTokens) + } +} + +// --- buildRequest 边界测试 --- + +func TestBuildRequest_EmptyModelReturnsError(t *testing.T) { + t.Parallel() + + // 直接构造 Provider 跳过 New() 的 Validate 校验, + // 以便测试 buildRequest 对空 model 的独立校验。 + p := &Provider{ + cfg: config.ResolvedProviderConfig{ + ProviderConfig: config.ProviderConfig{ + Name: DriverName, + Driver: DriverName, + BaseURL: config.OpenAIDefaultBaseURL, + Model: "", + APIKeyEnv: config.OpenAIDefaultAPIKeyEnv, + }, + APIKey: "test-key", + }, + client: &http.Client{}, + } + + _, buildErr := p.buildRequest(providertypes.ChatRequest{}) + if buildErr == nil { + t.Fatal("expected error for empty model") + } + if !strings.Contains(buildErr.Error(), "model is empty") { + t.Fatalf("unexpected error message: %v", buildErr) + } +} + +func TestBuildRequest_FallsBackToConfigModel(t *testing.T) { + t.Parallel() + + p, err := New(resolvedConfig(config.OpenAIDefaultBaseURL, config.OpenAIDefaultModel)) + if err != nil { + t.Fatalf("New() error = %v", err) + } + + payload, err := p.buildRequest(providertypes.ChatRequest{Messages: []providertypes.Message{{Role: "user", Content: "hi"}}}) + if err != nil { + t.Fatalf("buildRequest() error = %v", err) + } + if payload.Model != config.OpenAIDefaultModel { + t.Fatalf("expected model %q, got %q", config.OpenAIDefaultModel, payload.Model) + } +} + +func TestBuildRequest_RequestModelTakesPrecedence(t *testing.T) { + t.Parallel() + + p, err := New(resolvedConfig(config.OpenAIDefaultBaseURL, config.OpenAIDefaultModel)) + if err != nil { + t.Fatalf("New() error = %v", err) + } + + payload, err := p.buildRequest(providertypes.ChatRequest{Model: "gpt-4-custom", Messages: []providertypes.Message{{Role: "user", Content: "hi"}}}) + if err != nil { + t.Fatalf("buildRequest() error = %v", err) + } + if payload.Model != "gpt-4-custom" { + t.Fatalf("expected model %q, got %q", "gpt-4-custom", payload.Model) + } +} + +func TestBuildRequest_NoSystemPrompt(t *testing.T) { + t.Parallel() + + p, err := New(resolvedConfig(config.OpenAIDefaultBaseURL, config.OpenAIDefaultModel)) + if err != nil { + t.Fatalf("New() error = %v", err) + } + + payload, err := p.buildRequest(providertypes.ChatRequest{SystemPrompt: "", Messages: []providertypes.Message{{Role: "user", Content: "hi"}}}) + if err != nil { + t.Fatalf("buildRequest() error = %v", err) + } + for _, msg := range payload.Messages { + if msg.Role == "system" { + t.Fatal("expected no system message when SystemPrompt is empty") + } + } +} + +func TestBuildRequest_NoTools(t *testing.T) { + t.Parallel() + + p, err := New(resolvedConfig(config.OpenAIDefaultBaseURL, config.OpenAIDefaultModel)) + if err != nil { + t.Fatalf("New() error = %v", err) + } + + payload, err := p.buildRequest(providertypes.ChatRequest{Messages: []providertypes.Message{{Role: "user", Content: "hi"}}, Tools: nil}) + if err != nil { + t.Fatalf("buildRequest() error = %v", err) + } + if payload.ToolChoice != "" || len(payload.Tools) != 0 { + t.Fatalf("expected no tools, got choice=%q tools=%d", payload.ToolChoice, len(payload.Tools)) + } +} + +func TestBuildRequest_EmptyToolsSlice(t *testing.T) { + t.Parallel() + + p, err := New(resolvedConfig(config.OpenAIDefaultBaseURL, config.OpenAIDefaultModel)) + if err != nil { + t.Fatalf("New() error = %v", err) + } + + payload, err := p.buildRequest(providertypes.ChatRequest{Messages: []providertypes.Message{{Role: "user", Content: "hi"}}, Tools: []providertypes.ToolSpec{}}) + if err != nil { + t.Fatalf("buildRequest() error = %v", err) + } + if payload.ToolChoice != "" || len(payload.Tools) != 0 { + t.Fatalf("expected empty tools for empty slice, got choice=%q tools=%d", payload.ToolChoice, len(payload.Tools)) + } +} + +func TestBuildRequest_MultipleTools(t *testing.T) { + t.Parallel() + + p, err := New(resolvedConfig(config.OpenAIDefaultBaseURL, config.OpenAIDefaultModel)) + if err != nil { + t.Fatalf("New() error = %v", err) + } + + payload, err := p.buildRequest(providertypes.ChatRequest{ + Messages: []providertypes.Message{{Role: "user", Content: "use tools"}}, + Tools: []providertypes.ToolSpec{ + {Name: "tool_a", Description: "Tool A", Schema: map[string]any{"type": "object"}}, + {Name: "tool_b", Description: "Tool B", Schema: map[string]any{"type": "object"}}, + }, + }) + if err != nil { + t.Fatalf("buildRequest() error = %v", err) + } + if len(payload.Tools) != 2 { + t.Fatalf("expected 2 tools, got %d", len(payload.Tools)) + } + if payload.ToolChoice != "auto" { + t.Fatalf("expected tool_choice=auto, got %q", payload.ToolChoice) + } +} + +func TestBuildRequest_WhitespaceSystemPromptSkipped(t *testing.T) { + t.Parallel() + + p, err := New(resolvedConfig(config.OpenAIDefaultBaseURL, config.OpenAIDefaultModel)) + if err != nil { + t.Fatalf("New() error = %v", err) + } + + payload, err := p.buildRequest(providertypes.ChatRequest{SystemPrompt: " ", Messages: []providertypes.Message{{Role: "user", Content: "hi"}}}) + if err != nil { + t.Fatalf("buildRequest() error = %v", err) + } + for _, msg := range payload.Messages { + if msg.Role == "system" { + t.Fatal("expected no system message for whitespace-only system prompt") + } + } +} + +// --- consumeStream SSE 场景测试 --- + +func TestConsumeStream_SSECommentIgnored(t *testing.T) { + t.Setenv(config.OpenAIDefaultAPIKeyEnv, "test-key") + + p, err := New(resolvedConfig("", "")) + if err != nil { + t.Fatalf("New() error = %v", err) + } + + sseData := `: heartbeat +data: {"id":"a","choices":[{"delta":{"content":"ok"},"finish_reason":""}]} +data: [DONE] + +` + events := make(chan providertypes.StreamEvent, 4) + err = p.consumeStream(context.Background(), strings.NewReader(sseData), events) + if err != nil { + t.Fatalf("consumeStream() error = %v", err) + } + drained := drainStreamEvents(events) + var foundText bool + for _, evt := range drained { + if evt.Type == providertypes.StreamEventTextDelta && requireTextDeltaPayload(t, evt).Text == "ok" { + foundText = true + } + } + if !foundText { + t.Fatal("expected text_delta event after SSE comment") + } +} + +func TestConsumeStream_ChunkErrorInPayload(t *testing.T) { + t.Setenv(config.OpenAIDefaultAPIKeyEnv, "test-key") + + p, err := New(resolvedConfig("", "")) + if err != nil { + t.Fatalf("New() error = %v", err) + } + + sseData := `data: {"error":{"message":"rate limit exceeded"}} +` + events := make(chan providertypes.StreamEvent, 1) + err = p.consumeStream(context.Background(), strings.NewReader(sseData), events) + if err == nil { + t.Fatal("expected error for chunk with error field") + } + if !strings.Contains(err.Error(), "rate limit exceeded") { + t.Fatalf("expected rate limit error, got: %v", err) + } +} + +func TestConsumeStream_MultiLineDataPayload(t *testing.T) { + t.Setenv(config.OpenAIDefaultAPIKeyEnv, "test-key") + + p, err := New(resolvedConfig("", "")) + if err != nil { + t.Fatalf("New() error = %v", err) + } + + sseData := `data: {"id":"a","choices":[{"delta":{"content":"part1"},"finish_reason":""}]} +data: {"id":"b","choices":[{"delta":{"content":"part2"},"finish_reason":"stop"}]} + +` + events := make(chan providertypes.StreamEvent, 8) + err = p.consumeStream(context.Background(), strings.NewReader(sseData), events) + if err != nil { + t.Fatalf("consumeStream() error = %v", err) + } + drained := drainStreamEvents(events) + if len(drained) == 0 { + t.Fatal("expected events from multi-line data payload") + } +} + +func TestConsumeStream_EOFWithoutDone(t *testing.T) { + t.Setenv(config.OpenAIDefaultAPIKeyEnv, "test-key") + + p, err := New(resolvedConfig("", "")) + if err != nil { + t.Fatalf("New() error = %v", err) + } + + sseData := `data: {"id":"a","choices":[{"delta":{"content":"partial"},"finish_reason":""}]} +` + events := make(chan providertypes.StreamEvent, 4) + err = p.consumeStream(context.Background(), strings.NewReader(sseData), events) + if err != nil { + t.Fatalf("consumeStream() error = %v", err) + } + drained := drainStreamEvents(events) + var foundText bool + for _, evt := range drained { + if evt.Type == providertypes.StreamEventTextDelta { + foundText = true + } + } + if !foundText { + t.Fatal("expected text_delta event before EOF") + } +} + +func TestConsumeStream_ContextCancellation(t *testing.T) { + t.Setenv(config.OpenAIDefaultAPIKeyEnv, "test-key") + + p, err := New(resolvedConfig("", "")) + if err != nil { + t.Fatalf("New() error = %v", err) + } + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + sseData := `data: {"id":"a","choices":[{"delta":{"content":"should not emit"}}]} + +` + events := make(chan providertypes.StreamEvent, 1) + err = p.consumeStream(ctx, strings.NewReader(sseData), events) + if err != nil && !errors.Is(err, context.Canceled) { + t.Fatalf("expected context.Canceled or nil, got: %v", err) + } +} + +func TestConsumeStream_FinishReasonAccumulation(t *testing.T) { + t.Setenv(config.OpenAIDefaultAPIKeyEnv, "test-key") + + p, err := New(resolvedConfig("", "")) + if err != nil { + t.Fatalf("New() error = %v", err) + } + + sseData := `data: {"id":"a","choices":[{"delta":{"content":"text"},"finish_reason":""}]} +data: {"id":"b","choices":[{"index":0,"finish_reason":"stop"}]} +data: [DONE] + +` + events := make(chan providertypes.StreamEvent, 8) + err = p.consumeStream(context.Background(), strings.NewReader(sseData), events) + if err != nil { + t.Fatalf("consumeStream() error = %v", err) + } + + drained := drainStreamEvents(events) + var donePayload *providertypes.MessageDonePayload + for _, evt := range drained { + if evt.Type == providertypes.StreamEventMessageDone { + p := requireMessageDonePayload(t, evt) + donePayload = &p + } + } + if donePayload == nil { + t.Fatal("expected message_done event") + } + if donePayload.FinishReason != "stop" { + t.Fatalf("expected finish_reason=stop, got %q", donePayload.FinishReason) + } +} + +// --- emit 函数守卫和边界测试 --- + +func TestEmitTextDelta_NilEventsGuard(t *testing.T) { + t.Parallel() + if err := emitTextDelta(context.Background(), nil, "some text"); err != nil { + t.Fatalf("expected nil events guard to return nil, got %v", err) + } +} + +func TestEmitTextDelta_EmptyTextGuard(t *testing.T) { + t.Parallel() + events := make(chan providertypes.StreamEvent, 1) + if err := emitTextDelta(context.Background(), events, ""); err != nil { + t.Fatalf("expected empty text guard to return nil, got %v", err) + } + select { + case <-events: + t.Fatal("expected no event for empty text") + default: + } +} + +func TestEmitStreamEvent_NilEventsChannel(t *testing.T) { + t.Parallel() + event := providertypes.NewTextDeltaStreamEvent("test") + if err := emitStreamEvent(context.Background(), nil, event); err != nil { + t.Fatalf("expected nil channel to return nil, got %v", err) + } +} + +func TestEmitStreamEvent_NormalSend(t *testing.T) { + t.Parallel() + events := make(chan providertypes.StreamEvent, 1) + event := providertypes.NewTextDeltaStreamEvent("hello") + if err := emitStreamEvent(context.Background(), events, event); err != nil { + t.Fatalf("emitStreamEvent() error = %v", err) + } + got := <-events + if got.Type != providertypes.StreamEventTextDelta { + t.Fatalf("unexpected event type: %s", got.Type) + } +} + +func TestEmitStreamEvent_ContextCancelled(t *testing.T) { + t.Parallel() + ctx, cancel := context.WithCancel(context.Background()) + cancel() + events := make(chan providertypes.StreamEvent) + event := providertypes.NewTextDeltaStreamEvent("test") + + err := emitStreamEvent(ctx, events, event) + if err == nil { + t.Fatal("expected context cancellation error") + } + if !errors.Is(err, context.Canceled) { + t.Fatalf("expected context.Canceled, got: %v", err) + } +} + +// --- flushDataLines 测试 --- + +func TestFlushDataLines_EmptyLines(t *testing.T) { + t.Parallel() + called := false + err := flushDataLines([]string{}, func(string) error { called = true; return nil }) + if err != nil { + t.Fatalf("flushDataLines() error = %v", err) + } + if called { + t.Fatal("processChunk should not be called for empty lines") + } +} + +func TestFlushDataLines_SingleLine(t *testing.T) { + t.Parallel() + var received string + err := flushDataLines([]string{"line1"}, func(p string) error { received = p; return nil }) + if err != nil { + t.Fatalf("flushDataLines() error = %v", err) + } + if received != "line1" { + t.Fatalf("expected %q, got %q", "line1", received) + } +} + +func TestFlushDataLines_MultipleLinesProcessedIndividually(t *testing.T) { + t.Parallel() + var received []string + err := flushDataLines([]string{"a", "b", "c"}, func(p string) error { received = append(received, p); return nil }) + if err != nil { + t.Fatalf("flushDataLines() error = %v", err) + } + if len(received) != 3 || received[0] != "a" || received[1] != "b" || received[2] != "c" { + t.Fatalf("expected each line processed individually, got %v", received) + } +} + +func TestFlushDataLines_ProcessChunkError(t *testing.T) { + t.Parallel() + expectedErr := errors.New("process error") + err := flushDataLines([]string{"data"}, func(string) error { return expectedErr }) + if err != expectedErr { + t.Fatalf("expected processChunk error, got %v", err) + } +} + +// --- DiscoverModels 错误场景测试 --- + +func TestDiscoverModels_HTTPError(t *testing.T) { + t.Parallel() + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + _, _ = w.Write([]byte("internal error")) + })) + defer server.Close() + + p, err := New(resolvedConfig(server.URL, "")) + if err != nil { + t.Fatalf("New() error = %v", err) + } + p.client = server.Client() + + models, err := p.DiscoverModels(context.Background()) + if err == nil { + t.Fatal("expected error for HTTP 500 response") + } + if models != nil { + t.Fatalf("expected nil models on error, got %d models", len(models)) + } +} + +func TestDiscoverModels_NetworkError(t *testing.T) { + t.Parallel() + + p, err := New(resolvedConfig("http://127.0.0.1:1", "")) + if err != nil { + t.Fatalf("New() error = %v", err) + } + p.client = &http.Client{Timeout: time.Millisecond * 10} + + _, err = p.DiscoverModels(context.Background()) + if err == nil { + t.Fatal("expected error for unreachable server") + } +} + +// --- mergeToolCallDelta 边界测试 --- + +func TestMergeToolCallDelta_MultipleIndices(t *testing.T) { t.Parallel() - t.Run("nil events guard", func(t *testing.T) { - t.Parallel() - if err := emitToolCallDelta(context.Background(), nil, 0, "", "args"); err != nil { - t.Fatalf("expected nil events guard to return nil, got %v", err) - } - }) + events := make(chan providertypes.StreamEvent, 8) + toolCalls := make(map[int]*providertypes.ToolCall) - t.Run("empty arguments guard", func(t *testing.T) { - t.Parallel() - events := make(chan domain.StreamEvent, 1) - if err := emitToolCallDelta(context.Background(), events, 0, "", ""); err != nil { - t.Fatalf("expected empty arguments guard to return nil, got %v", err) - } - select { - case <-events: - t.Fatal("expected no event for empty arguments") - default: - } + mergeToolCallDelta(context.Background(), events, toolCalls, toolCallDelta{ + Index: 0, ID: "call_0", + Function: openAIFunctionCall{Name: "tool_a", Arguments: `{"arg":"a"`}, }) - - t.Run("normal send", func(t *testing.T) { - t.Parallel() - events := make(chan domain.StreamEvent, 1) - if err := emitToolCallDelta(context.Background(), events, 3, "call_123", `{"path":"main.go"}`); err != nil { - t.Fatalf("emitToolCallDelta() error = %v", err) - } - got := <-events - payload := requireToolCallDeltaPayload(t, got) - if got.Type != domain.StreamEventToolCallDelta || payload.Index != 3 || payload.ArgumentsDelta != `{"path":"main.go"}` || payload.ID != "call_123" { - t.Fatalf("unexpected event: %+v", got) - } + mergeToolCallDelta(context.Background(), events, toolCalls, toolCallDelta{ + Index: 1, ID: "call_1", + Function: openAIFunctionCall{Name: "tool_b", Arguments: `{"arg":"b"}`}, }) - - t.Run("context cancellation", func(t *testing.T) { - t.Parallel() - cancelledCtx, cancel := context.WithCancel(context.Background()) - cancel() - if err := emitToolCallDelta(cancelledCtx, make(chan domain.StreamEvent), 0, "", "args"); err == nil { - t.Fatal("expected cancellation error") - } + mergeToolCallDelta(context.Background(), events, toolCalls, toolCallDelta{ + Index: 0, + Function: openAIFunctionCall{Arguments: `,"more":"data"}`}, }) + + if len(toolCalls) != 2 { + t.Fatalf("expected 2 tool calls, got %d", len(toolCalls)) + } + call0 := toolCalls[0] + if call0.Name != "tool_a" || call0.ID != "call_0" { + t.Fatalf("unexpected tool call 0: %+v", call0) + } + expectedArgs0 := `{"arg":"a","more":"data"}` + if call0.Arguments != expectedArgs0 { + t.Fatalf("expected arguments %q for call 0, got %q", expectedArgs0, call0.Arguments) + } + call1 := toolCalls[1] + if call1.Name != "tool_b" || call1.ID != "call_1" { + t.Fatalf("unexpected tool call 1: %+v", call1) + } } -func TestEmitMessageDone(t *testing.T) { +func TestMergeToolCallDelta_IDUpdateOnly(t *testing.T) { t.Parallel() - t.Run("nil events guard", func(t *testing.T) { - t.Parallel() - if err := emitMessageDone(context.Background(), nil, "stop", nil); err != nil { - t.Fatalf("expected nil events guard to return nil, got %v", err) - } - }) + events := make(chan providertypes.StreamEvent, 4) + toolCalls := make(map[int]*providertypes.ToolCall) - t.Run("normal send", func(t *testing.T) { - t.Parallel() - events := make(chan domain.StreamEvent, 1) - usage := &domain.Usage{TotalTokens: 100} - if err := emitMessageDone(context.Background(), events, "stop", usage); err != nil { - t.Fatalf("emitMessageDone() error = %v", err) - } - got := <-events - payload := requireMessageDonePayload(t, got) - if got.Type != domain.StreamEventMessageDone || payload.FinishReason != "stop" || payload.Usage == nil || payload.Usage.TotalTokens != 100 { - t.Fatalf("unexpected event: %+v", got) - } - }) + mergeToolCallDelta(context.Background(), events, toolCalls, toolCallDelta{Index: 0, ID: "call_only_id"}) - t.Run("context cancellation", func(t *testing.T) { - t.Parallel() - cancelledCtx, cancel := context.WithCancel(context.Background()) - cancel() - if err := emitMessageDone(cancelledCtx, make(chan domain.StreamEvent), "stop", nil); err == nil { - t.Fatal("expected cancellation error") + call := toolCalls[0] + if call == nil { + t.Fatal("expected tool call entry to be created") + } + if call.ID != "call_only_id" { + t.Fatalf("expected ID %q, got %q", "call_only_id", call.ID) + } + select { + case <-events: + t.Fatal("expected no event for ID-only delta") + default: + } +} + +// --- Chat 集成测试 --- + +func TestChat_BaseURLTrailingSlashHandled(t *testing.T) { + t.Setenv(config.OpenAIDefaultAPIKeyEnv, "test-key") + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/chat/completions" { + t.Fatalf("unexpected path: %s", r.URL.Path) } - }) + w.Header().Set("Content-Type", "text/event-stream") + _, _ = w.Write([]byte(`data: {"choices":[{"delta":{"content":"ok"},"finish_reason":"stop"}],"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}} +data: [DONE] + +`)) + })) + defer server.Close() + + p, err := New(resolvedConfig(server.URL+"/", config.OpenAIDefaultModel)) + if err != nil { + t.Fatalf("New() error = %v", err) + } + p.client = server.Client() + + events := make(chan providertypes.StreamEvent, 4) + err = p.Chat(context.Background(), providertypes.ChatRequest{ + Model: config.OpenAIDefaultModel, + Messages: []providertypes.Message{{Role: "user", Content: "hi"}}, + }, events) + if err != nil { + t.Fatalf("Chat() error = %v", err) + } } -func resolvedConfig(baseURL string, model string) config.ResolvedProviderConfig { - if strings.TrimSpace(baseURL) == "" { - baseURL = config.OpenAIDefaultBaseURL +// --- parseError 边界测试 --- + +func TestParseError_ReadBodyFailure(t *testing.T) { + t.Parallel() + + p, err := New(resolvedConfig(config.OpenAIDefaultBaseURL, config.OpenAIDefaultModel)) + if err != nil { + t.Fatalf("New() error = %v", err) } - if strings.TrimSpace(model) == "" { - model = config.OpenAIDefaultModel + + readErr := errors.New("simulated read failure") + resp := &http.Response{Status: "400 Bad Request", StatusCode: 400, Body: &failingReadCloser{err: readErr}} + + err = p.parseError(resp) + if err == nil { + t.Fatal("expected error when body read fails") } + if !strings.Contains(err.Error(), "read error response") { + t.Fatalf("expected read error in message, got: %v", err) + } +} - return config.ResolvedProviderConfig{ - ProviderConfig: config.ProviderConfig{ - Name: DriverName, - Driver: DriverName, - BaseURL: baseURL, - Model: model, - APIKeyEnv: config.OpenAIDefaultAPIKeyEnv, - }, - APIKey: "test-key", +func TestParseError_InvalidJSONBody(t *testing.T) { + t.Parallel() + + p, err := New(resolvedConfig(config.OpenAIDefaultBaseURL, config.OpenAIDefaultModel)) + if err != nil { + t.Fatalf("New() error = %v", err) + } + + resp := &http.Response{Status: "400 Bad Request", StatusCode: 400, Body: ioNopCloser("this is not json at all")} + err = p.parseError(resp) + if err == nil { + t.Fatal("expected error for non-JSON body") + } + if !strings.Contains(err.Error(), "this is not json at all") { + t.Fatalf("expected plain text fallback, got: %v", err) } } +// --- 原有保留的集成测试(保持兼容) --- + func TestProviderChatConsumesSSEAndMergesToolCalls(t *testing.T) { t.Setenv(config.OpenAIDefaultAPIKeyEnv, "test-key") @@ -268,118 +886,32 @@ func TestProviderChatConsumesSSEAndMergesToolCalls(t *testing.T) { if payload.Model != "gpt-5.4" { t.Fatalf("expected model gpt-5.4, got %q", payload.Model) } - if !containsToolRoleMessage(payload.Messages, "call_1", "tool finished") { - t.Fatalf("expected tool role message with tool_call_id in payload: %+v", payload.Messages) - } w.Header().Set("Content-Type", "text/event-stream") - writeSSEChunk(t, w, map[string]any{ - "choices": []map[string]any{ - { - "index": 0, - "delta": map[string]any{ - "content": "Hello ", - }, - }, - }, - }) - writeSSEChunk(t, w, map[string]any{ - "choices": []map[string]any{ - { - "index": 0, - "delta": map[string]any{ - "tool_calls": []map[string]any{ - { - "index": 0, - "id": "call_1", - "type": "function", - "function": map[string]any{ - "name": "filesystem_edit", - "arguments": `{"path":"main.go",`, - }, - }, - }, - }, - }, - }, - }) - writeSSEChunk(t, w, map[string]any{ - "choices": []map[string]any{ - { - "index": 0, - "delta": map[string]any{ - "content": "world", - "tool_calls": []map[string]any{ - { - "index": 0, - "function": map[string]any{ - "arguments": `"search_string":"old",`, - }, - }, - }, - }, - }, - }, - }) - writeSSEChunk(t, w, map[string]any{ - "choices": []map[string]any{ - { - "index": 0, - "finish_reason": "tool_calls", - "delta": map[string]any{ - "tool_calls": []map[string]any{ - { - "index": 0, - "function": map[string]any{ - "arguments": `"replace_string":"new"}`, - }, - }, - }, - }, - }, - }, - "usage": map[string]any{ - "prompt_tokens": 10, - "completion_tokens": 5, - "total_tokens": 15, - }, - }) + writeSSEChunk(t, w, map[string]any{"choices": []map[string]any{{"index": 0, "delta": map[string]any{"content": "Hello "}}}}) + writeSSEChunk(t, w, map[string]any{"choices": []map[string]any{{"index": 0, "delta": map[string]any{"tool_calls": []map[string]any{{"index": 0, "id": "call_1", "type": "function", "function": map[string]any{"name": "filesystem_edit", "arguments": `{"path":"main.go","search_string":"old"`}}}}}}}) + writeSSEChunk(t, w, map[string]any{"choices": []map[string]any{{"index": 0, "delta": map[string]any{"content": "world"}}}}) + writeSSEChunk(t, w, map[string]any{"choices": []map[string]any{{"index": 0, "delta": map[string]any{"tool_calls": []map[string]any{{"index": 0, "function": map[string]any{"arguments": `,"replace_string":"new"}`}}}}}}}) + writeSSEChunk(t, w, map[string]any{"choices": []map[string]any{{"index": 0, "finish_reason": "tool_calls"}}, "usage": map[string]any{"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}}) _, _ = w.Write([]byte("data: [DONE]\n\n")) })) defer server.Close() - provider, err := New(resolvedConfig(server.URL, "gpt-5.4")) + p, err := New(resolvedConfig(server.URL, "gpt-5.4")) if err != nil { t.Fatalf("New() error = %v", err) } - provider.client = server.Client() + p.client = server.Client() - events := make(chan domain.StreamEvent, 8) - err = provider.Chat(context.Background(), domain.ChatRequest{ + events := make(chan providertypes.StreamEvent, 8) + err = p.Chat(context.Background(), providertypes.ChatRequest{ Model: "gpt-5.4", - Messages: []domain.Message{ + Messages: []providertypes.Message{ {Role: "user", Content: "please edit the file"}, - { - Role: "assistant", - ToolCalls: []domain.ToolCall{ - { - ID: "call_1", - Name: "filesystem_edit", - Arguments: `{"path":"main.go","search_string":"old","replace_string":"new"}`, - }, - }, - }, + {Role: "assistant", ToolCalls: []providertypes.ToolCall{{ID: "call_1", Name: "filesystem_edit", Arguments: `{"path":"main.go","search_string":"old","replace_string":"new"}`}}}, {Role: "tool", ToolCallID: "call_1", Content: "tool finished"}, }, - Tools: []domain.ToolSpec{ - { - Name: "filesystem_edit", - Description: "Edit one matching block in a file", - Schema: map[string]any{ - "type": "object", - }, - }, - }, + Tools: []providertypes.ToolSpec{{Name: "filesystem_edit", Description: "Edit one matching block in a file", Schema: map[string]any{"type": "object"}}}, }, events) if err != nil { t.Fatalf("Chat() error = %v", err) @@ -390,32 +922,30 @@ func TestProviderChatConsumesSSEAndMergesToolCalls(t *testing.T) { t.Fatal("expected streamed events") } - var ( - chunks []string - toolCallStartSeen bool - toolCallArgs strings.Builder - messageDone *domain.MessageDonePayload - ) + var chunks []string + var toolCallStartSeen bool + var toolCallArgs strings.Builder + var messageDone *providertypes.MessageDonePayload for _, event := range streamEvents { switch event.Type { - case domain.StreamEventTextDelta: + case providertypes.StreamEventTextDelta: chunks = append(chunks, requireTextDeltaPayload(t, event).Text) - case domain.StreamEventToolCallStart: + case providertypes.StreamEventToolCallStart: payload := requireToolCallStartPayload(t, event) toolCallStartSeen = true if payload.Index != 0 || payload.ID != "call_1" || payload.Name != "filesystem_edit" { t.Fatalf("unexpected tool_call_start payload: %+v", payload) } - case domain.StreamEventToolCallDelta: + case providertypes.StreamEventToolCallDelta: payload := requireToolCallDeltaPayload(t, event) if payload.Index != 0 || payload.ID != "call_1" { t.Fatalf("unexpected tool_call_delta payload: %+v", payload) } toolCallArgs.WriteString(payload.ArgumentsDelta) - case domain.StreamEventMessageDone: - payload := requireMessageDonePayload(t, event) - messageDone = &payload + case providertypes.StreamEventMessageDone: + p := requireMessageDonePayload(t, event) + messageDone = &p } } @@ -448,18 +978,8 @@ func TestProviderChatHTTPErrorResponses(t *testing.T) { body string expectErr string }{ - { - name: "http 401 json error", - status: http.StatusUnauthorized, - body: `{"error":{"message":"invalid api key"}}`, - expectErr: "invalid api key", - }, - { - name: "http 500 empty body falls back to status", - status: http.StatusInternalServerError, - body: ``, - expectErr: "500 Internal Server Error", - }, + {name: "http 401 json error", status: http.StatusUnauthorized, body: `{"error":{"message":"invalid api key"}}`, expectErr: "invalid api key"}, + {name: "http 500 empty body falls back to status", status: http.StatusInternalServerError, body: ``, expectErr: "500 Internal Server Error"}, } for _, tt := range tests { @@ -473,15 +993,13 @@ func TestProviderChatHTTPErrorResponses(t *testing.T) { })) defer server.Close() - provider, err := New(resolvedConfig(server.URL, config.OpenAIDefaultModel)) + p, err := New(resolvedConfig(server.URL, config.OpenAIDefaultModel)) if err != nil { t.Fatalf("New() error = %v", err) } - provider.client = server.Client() + p.client = server.Client() - err = provider.Chat(context.Background(), domain.ChatRequest{ - Model: config.OpenAIDefaultModel, - }, make(chan domain.StreamEvent, 1)) + err = p.Chat(context.Background(), providertypes.ChatRequest{Model: config.OpenAIDefaultModel}, make(chan providertypes.StreamEvent, 1)) if err == nil || !strings.Contains(err.Error(), tt.expectErr) { t.Fatalf("expected error containing %q, got %v", tt.expectErr, err) } @@ -492,36 +1010,19 @@ func TestProviderChatHTTPErrorResponses(t *testing.T) { func TestBuildRequestIncludesSystemPromptToolsAndToolMessages(t *testing.T) { t.Parallel() - provider, err := New(resolvedConfig(config.OpenAIDefaultBaseURL, config.OpenAIDefaultModel)) + p, err := New(resolvedConfig(config.OpenAIDefaultBaseURL, config.OpenAIDefaultModel)) if err != nil { t.Fatalf("New() error = %v", err) } - payload, err := provider.buildRequest(domain.ChatRequest{ + payload, err := p.buildRequest(providertypes.ChatRequest{ SystemPrompt: "system prompt", - Messages: []domain.Message{ + Messages: []providertypes.Message{ {Role: "user", Content: "hello"}, - { - Role: "assistant", - ToolCalls: []domain.ToolCall{ - { - ID: "call_1", - Name: "filesystem_edit", - Arguments: `{"path":"main.go"}`, - }, - }, - }, + {Role: "assistant", ToolCalls: []providertypes.ToolCall{{ID: "call_1", Name: "filesystem_edit", Arguments: `{"path":"main.go"}`}}}, {Role: "tool", ToolCallID: "call_1", Content: "tool finished"}, }, - Tools: []domain.ToolSpec{ - { - Name: "filesystem_edit", - Description: "Edit file content", - Schema: map[string]any{ - "type": "object", - }, - }, - }, + Tools: []providertypes.ToolSpec{{Name: "filesystem_edit", Description: "Edit file content", Schema: map[string]any{"type": "object"}}}, }) if err != nil { t.Fatalf("buildRequest() error = %v", err) @@ -553,7 +1054,7 @@ func TestBuildRequestIncludesSystemPromptToolsAndToolMessages(t *testing.T) { func TestParseErrorAndEmitTextDelta(t *testing.T) { t.Parallel() - provider, err := New(resolvedConfig(config.OpenAIDefaultBaseURL, config.OpenAIDefaultModel)) + p, err := New(resolvedConfig(config.OpenAIDefaultBaseURL, config.OpenAIDefaultModel)) if err != nil { t.Fatalf("New() error = %v", err) } @@ -564,46 +1065,33 @@ func TestParseErrorAndEmitTextDelta(t *testing.T) { body string expectErr string }{ - { - name: "json error payload", - status: "400 Bad Request", - body: `{"error":{"message":"invalid request"}}`, - expectErr: "invalid request", - }, - { - name: "plain text fallback", - status: "502 Bad Gateway", - body: `gateway timeout`, - expectErr: "gateway timeout", - }, + {"json error payload", "400 Bad Request", `{"error":{"message":"invalid request"}}`, "invalid request"}, + {"plain text fallback", "502 Bad Gateway", `gateway timeout`, "gateway timeout"}, } for _, tt := range tests { tt := tt t.Run(tt.name, func(t *testing.T) { t.Parallel() - resp := &http.Response{ - Status: tt.status, - Body: ioNopCloser(tt.body), - } - err := provider.parseError(resp) + resp := &http.Response{Status: tt.status, Body: ioNopCloser(tt.body)} + err := p.parseError(resp) if err == nil || !strings.Contains(err.Error(), tt.expectErr) { t.Fatalf("expected error containing %q, got %v", tt.expectErr, err) } }) } - eventCh := make(chan domain.StreamEvent, 1) + eventCh := make(chan providertypes.StreamEvent, 1) if err := emitTextDelta(context.Background(), eventCh, "chunk"); err != nil { t.Fatalf("emitTextDelta() error = %v", err) } - if got := <-eventCh; got.Type != domain.StreamEventTextDelta || requireTextDeltaPayload(t, got).Text != "chunk" { + if got := <-eventCh; got.Type != providertypes.StreamEventTextDelta || requireTextDeltaPayload(t, got).Text != "chunk" { t.Fatalf("unexpected stream event: %+v", got) } cancelledCtx, cancel := context.WithCancel(context.Background()) cancel() - if err := emitTextDelta(cancelledCtx, make(chan domain.StreamEvent), "chunk"); err == nil { + if err := emitTextDelta(cancelledCtx, make(chan providertypes.StreamEvent), "chunk"); err == nil { t.Fatalf("expected cancellation error") } } @@ -611,19 +1099,34 @@ func TestParseErrorAndEmitTextDelta(t *testing.T) { func TestProviderConsumeStreamRejectsDirtyJSON(t *testing.T) { t.Parallel() - provider, err := New(resolvedConfig(config.OpenAIDefaultBaseURL, config.OpenAIDefaultModel)) + p, err := New(resolvedConfig(config.OpenAIDefaultBaseURL, config.OpenAIDefaultModel)) if err != nil { t.Fatalf("New() error = %v", err) } - err = provider.consumeStream(context.Background(), strings.NewReader("data: {not-json}\n\n"), make(chan domain.StreamEvent, 1)) + err = p.consumeStream(context.Background(), strings.NewReader("data: {not-json}\n\n"), make(chan providertypes.StreamEvent, 1)) if err == nil || !strings.Contains(err.Error(), "decode stream chunk") { t.Fatalf("expected dirty JSON decode error, got %v", err) } } -func drainStreamEvents(events <-chan domain.StreamEvent) []domain.StreamEvent { - drained := make([]domain.StreamEvent, 0) +// --- 辅助函数 --- + +func resolvedConfig(baseURL string, model string) config.ResolvedProviderConfig { + if strings.TrimSpace(baseURL) == "" { + baseURL = config.OpenAIDefaultBaseURL + } + if strings.TrimSpace(model) == "" { + model = config.OpenAIDefaultModel + } + return config.ResolvedProviderConfig{ + ProviderConfig: config.ProviderConfig{Name: DriverName, Driver: DriverName, BaseURL: baseURL, Model: model, APIKeyEnv: config.OpenAIDefaultAPIKeyEnv}, + APIKey: "test-key", + } +} + +func drainStreamEvents(events <-chan providertypes.StreamEvent) []providertypes.StreamEvent { + drained := make([]providertypes.StreamEvent, 0) for { select { case evt, ok := <-events: @@ -637,7 +1140,7 @@ func drainStreamEvents(events <-chan domain.StreamEvent) []domain.StreamEvent { } } -func requireTextDeltaPayload(t *testing.T, event domain.StreamEvent) domain.TextDeltaPayload { +func requireTextDeltaPayload(t *testing.T, event providertypes.StreamEvent) providertypes.TextDeltaPayload { t.Helper() payload, err := event.TextDeltaValue() if err != nil { @@ -646,7 +1149,7 @@ func requireTextDeltaPayload(t *testing.T, event domain.StreamEvent) domain.Text return payload } -func requireToolCallStartPayload(t *testing.T, event domain.StreamEvent) domain.ToolCallStartPayload { +func requireToolCallStartPayload(t *testing.T, event providertypes.StreamEvent) providertypes.ToolCallStartPayload { t.Helper() payload, err := event.ToolCallStartValue() if err != nil { @@ -655,7 +1158,7 @@ func requireToolCallStartPayload(t *testing.T, event domain.StreamEvent) domain. return payload } -func requireToolCallDeltaPayload(t *testing.T, event domain.StreamEvent) domain.ToolCallDeltaPayload { +func requireToolCallDeltaPayload(t *testing.T, event providertypes.StreamEvent) providertypes.ToolCallDeltaPayload { t.Helper() payload, err := event.ToolCallDeltaValue() if err != nil { @@ -664,7 +1167,7 @@ func requireToolCallDeltaPayload(t *testing.T, event domain.StreamEvent) domain. return payload } -func requireMessageDonePayload(t *testing.T, event domain.StreamEvent) domain.MessageDonePayload { +func requireMessageDonePayload(t *testing.T, event providertypes.StreamEvent) providertypes.MessageDonePayload { t.Helper() payload, err := event.MessageDoneValue() if err != nil { @@ -674,8 +1177,8 @@ func requireMessageDonePayload(t *testing.T, event domain.StreamEvent) domain.Me } func containsToolRoleMessage(messages []openAIMessage, toolCallID string, content string) bool { - for _, message := range messages { - if message.Role == "tool" && message.ToolCallID == toolCallID && message.Content == content { + for _, m := range messages { + if m.Role == "tool" && m.ToolCallID == toolCallID && m.Content == content { return true } } @@ -691,37 +1194,90 @@ func writeSSEChunk(t *testing.T, w http.ResponseWriter, payload any) { if _, err := w.Write([]byte("data: " + string(data) + "\n\n")); err != nil { t.Fatalf("write SSE payload: %v", err) } - if flusher, ok := w.(http.Flusher); ok { - flusher.Flush() + if f, ok := w.(http.Flusher); ok { + f.Flush() } } -func ioNopCloser(body string) *readCloser { - return &readCloser{Reader: strings.NewReader(body)} -} +func ioNopCloser(body string) *readCloser { return &readCloser{Reader: strings.NewReader(body)} } + +type readCloser struct{ *strings.Reader } + +func (r *readCloser) Close() error { return nil } + +// --- 错误包装测试 --- + +func TestConsumeStream_WrapsNonEOFAsInterrupted(t *testing.T) { + t.Setenv(config.OpenAIDefaultAPIKeyEnv, "test-key") -type readCloser struct { - *strings.Reader + p, err := New(resolvedConfig("", "")) + if err != nil { + t.Fatalf("New() error = %v", err) + } + + errReader := &errReader{err: io.ErrClosedPipe} + err = p.consumeStream(context.Background(), errReader, make(chan providertypes.StreamEvent, 1)) + if err == nil { + t.Fatal("expected error for broken reader") + } + if !errors.Is(err, provider.ErrStreamInterrupted) { + t.Fatalf("expected ErrStreamInterrupted wrapping, got: %v", err) + } } -func (r *readCloser) Close() error { - return nil +func TestConsumeStream_FlushesPendingDataOnNonEOFError(t *testing.T) { + t.Setenv(config.OpenAIDefaultAPIKeyEnv, "test-key") + + p, err := New(resolvedConfig("", "")) + if err != nil { + t.Fatalf("New() error = %v", err) + } + + sseData := `data: {"id":"a","choices":[{"delta":{"content":"hello"},"finish_reason":""}]} +` + body := io.MultiReader(strings.NewReader(sseData), &errReader{err: io.ErrClosedPipe}) + events := make(chan providertypes.StreamEvent, 10) + + err = p.consumeStream(context.Background(), body, events) + if err == nil { + t.Fatal("expected error for broken reader") + } + if !errors.Is(err, provider.ErrStreamInterrupted) { + t.Fatalf("expected ErrStreamInterrupted, got: %v", err) + } + + drained := drainStreamEvents(events) + var foundText bool + for _, evt := range drained { + if evt.Type == providertypes.StreamEventTextDelta { + foundText = true + } + } + if !foundText { + t.Fatal("expected text_delta event from flushed pending data") + } } -// --- emitToolCallStart 边界测试 --- +type errReader struct{ err error } + +func (e *errReader) Read(_ []byte) (int, error) { return 0, e.err } + +type failingReadCloser struct{ err error } + +func (f *failingReadCloser) Read(_ []byte) (int, error) { return 0, f.err } +func (f *failingReadCloser) Close() error { return f.err } + +// --- emitToolCallStart 和 mergeToolCallDelta 保留测试 --- func TestEmitToolCallStartGuards(t *testing.T) { t.Parallel() - ctx := context.Background() - // nil events 守卫 if err := emitToolCallStart(ctx, nil, 0, "call-1", "filesystem_edit"); err != nil { t.Fatalf("expected nil events guard to return nil, got %v", err) } - // 空 name 守卫 - events := make(chan domain.StreamEvent, 1) + events := make(chan providertypes.StreamEvent, 1) if err := emitToolCallStart(ctx, events, 0, "call-1", ""); err != nil { t.Fatalf("expected empty name guard to return nil, got %v", err) } @@ -731,55 +1287,41 @@ func TestEmitToolCallStartGuards(t *testing.T) { default: } - // 正常发送 if err := emitToolCallStart(ctx, events, 2, "call-1", "filesystem_edit"); err != nil { t.Fatalf("emitToolCallStart() error = %v", err) } got := <-events payload := requireToolCallStartPayload(t, got) - if got.Type != domain.StreamEventToolCallStart || payload.Name != "filesystem_edit" || payload.ID != "call-1" || payload.Index != 2 { + if got.Type != providertypes.StreamEventToolCallStart || payload.Name != "filesystem_edit" || payload.ID != "call-1" || payload.Index != 2 { t.Fatalf("unexpected event: %+v", got) } - // context 取消 cancelledCtx, cancel := context.WithCancel(context.Background()) cancel() - if err := emitToolCallStart(cancelledCtx, make(chan domain.StreamEvent), 0, "call-1", "filesystem_edit"); err == nil { - t.Fatalf("expected cancellation error") + if err := emitToolCallStart(cancelledCtx, make(chan providertypes.StreamEvent), 0, "call-1", "filesystem_edit"); err == nil { + t.Fatal("expected cancellation error") } } func TestMergeToolCallDeltaEmitsStartWhenNameArrivesLater(t *testing.T) { t.Parallel() - events := make(chan domain.StreamEvent, 4) - toolCalls := make(map[int]*domain.ToolCall) - - if err := mergeToolCallDelta(context.Background(), events, toolCalls, toolCallDelta{ - Index: 0, - ID: "call_late_name", - }); err != nil { - t.Fatalf("mergeToolCallDelta() first delta error = %v", err) - } + events := make(chan providertypes.StreamEvent, 4) + toolCalls := make(map[int]*providertypes.ToolCall) + mergeToolCallDelta(context.Background(), events, toolCalls, toolCallDelta{Index: 0, ID: "call_late_name"}) select { case evt := <-events: t.Fatalf("expected no event before tool name arrives, got %+v", evt) default: } - if err := mergeToolCallDelta(context.Background(), events, toolCalls, toolCallDelta{ - Index: 0, - Function: openAIFunctionCall{ - Name: "filesystem_edit", - Arguments: `{"path":"main.go"}`, - }, - }); err != nil { - t.Fatalf("mergeToolCallDelta() late-name delta error = %v", err) - } + mergeToolCallDelta(context.Background(), events, toolCalls, toolCallDelta{ + Index: 0, Function: openAIFunctionCall{Name: "filesystem_edit", Arguments: `{"path":"main.go"}`}, + }) start := <-events - if start.Type != domain.StreamEventToolCallStart { + if start.Type != providertypes.StreamEventToolCallStart { t.Fatalf("expected tool_call_start event, got %+v", start) } startPayload := requireToolCallStartPayload(t, start) @@ -788,7 +1330,7 @@ func TestMergeToolCallDeltaEmitsStartWhenNameArrivesLater(t *testing.T) { } delta := <-events - if delta.Type != domain.StreamEventToolCallDelta { + if delta.Type != providertypes.StreamEventToolCallDelta { t.Fatalf("expected tool_call_delta event, got %+v", delta) } deltaPayload := requireToolCallDeltaPayload(t, delta) @@ -810,43 +1352,21 @@ func TestProviderChatEmitsToolCallStartEvent(t *testing.T) { server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "text/event-stream") - writeSSEChunk(t, w, map[string]any{ - "choices": []map[string]any{ - { - "index": 0, - "delta": map[string]any{ - "tool_calls": []map[string]any{ - { - "index": 0, - "id": "call_tool", - "type": "function", - "function": map[string]any{ - "name": "filesystem_edit", - "arguments": `{}`, - }, - }, - }, - }, - }, - }, - }) + writeSSEChunk(t, w, map[string]any{"choices": []map[string]any{{"index": 0, "delta": map[string]any{"tool_calls": []map[string]any{{"index": 0, "id": "call_tool", "type": "function", "function": map[string]any{"name": "filesystem_edit", "arguments": `{}`}}}}}}}) _, _ = w.Write([]byte("data: [DONE]\n\n")) })) defer server.Close() - provider, err := New(resolvedConfig(server.URL, config.OpenAIDefaultModel)) + p, err := New(resolvedConfig(server.URL, config.OpenAIDefaultModel)) if err != nil { t.Fatalf("New() error = %v", err) } - provider.client = server.Client() + p.client = server.Client() - events := make(chan domain.StreamEvent, 8) - err = provider.Chat(context.Background(), domain.ChatRequest{ - Model: config.OpenAIDefaultModel, - Messages: []domain.Message{{Role: "user", Content: "edit"}}, - Tools: []domain.ToolSpec{ - {Name: "filesystem_edit", Description: "edit", Schema: map[string]any{"type": "object"}}, - }, + events := make(chan providertypes.StreamEvent, 8) + err = p.Chat(context.Background(), providertypes.ChatRequest{ + Model: config.OpenAIDefaultModel, Messages: []providertypes.Message{{Role: "user", Content: "edit"}}, + Tools: []providertypes.ToolSpec{{Name: "filesystem_edit", Description: "edit", Schema: map[string]any{"type": "object"}}}, }, events) if err != nil { t.Fatalf("Chat() error = %v", err) @@ -854,7 +1374,7 @@ func TestProviderChatEmitsToolCallStartEvent(t *testing.T) { var foundToolCallStart bool for _, evt := range drainStreamEvents(events) { - if evt.Type == domain.StreamEventToolCallStart { + if evt.Type == providertypes.StreamEventToolCallStart { foundToolCallStart = true payload := requireToolCallStartPayload(t, evt) if payload.Name != "filesystem_edit" { @@ -870,129 +1390,51 @@ func TestProviderChatEmitsToolCallStartEvent(t *testing.T) { } } -// TestProviderChatEmitsFullEventStream 测试完整的事件流(包括 tool_call_delta 和 message_done)。 func TestProviderChatEmitsFullEventStream(t *testing.T) { t.Setenv(config.OpenAIDefaultAPIKeyEnv, "test-key") server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "text/event-stream") - - // 发送文本 delta - writeSSEChunk(t, w, map[string]any{ - "choices": []map[string]any{ - { - "index": 0, - "delta": map[string]any{ - "content": "Hello", - }, - }, - }, - }) - - // 发送 tool call start 和 delta - writeSSEChunk(t, w, map[string]any{ - "choices": []map[string]any{ - { - "index": 0, - "delta": map[string]any{ - "tool_calls": []map[string]any{ - { - "index": 0, - "id": "call_tool_1", - "type": "function", - "function": map[string]any{ - "name": "filesystem_edit", - "arguments": `{"path":"a.`, - }, - }, - }, - }, - }, - }, - }) - - // 发送 tool call delta(参数增量) - writeSSEChunk(t, w, map[string]any{ - "choices": []map[string]any{ - { - "index": 0, - "delta": map[string]any{ - "tool_calls": []map[string]any{ - { - "index": 0, - "function": map[string]any{ - "arguments": `go"}`, - }, - }, - }, - }, - }, - }, - }) - - // 发送 usage 和 finish_reason - writeSSEChunk(t, w, map[string]any{ - "choices": []map[string]any{ - { - "index": 0, - "finish_reason": "tool_calls", - }, - }, - "usage": map[string]any{ - "prompt_tokens": 100, - "completion_tokens": 50, - "total_tokens": 150, - }, - }) - + writeSSEChunk(t, w, map[string]any{"choices": []map[string]any{{"index": 0, "delta": map[string]any{"content": "Hello"}}}}) + writeSSEChunk(t, w, map[string]any{"choices": []map[string]any{{"index": 0, "delta": map[string]any{"tool_calls": []map[string]any{{"index": 0, "id": "call_tool_1", "type": "function", "function": map[string]any{"name": "filesystem_edit", "arguments": `{"path":"a.go"}`}}}}}}}) + writeSSEChunk(t, w, map[string]any{"choices": []map[string]any{{"index": 0, "finish_reason": "tool_calls"}}, "usage": map[string]any{"prompt_tokens": 100, "completion_tokens": 50, "total_tokens": 150}}) _, _ = w.Write([]byte("data: [DONE]\n\n")) })) defer server.Close() - provider, err := New(resolvedConfig(server.URL, config.OpenAIDefaultModel)) + p, err := New(resolvedConfig(server.URL, config.OpenAIDefaultModel)) if err != nil { t.Fatalf("New() error = %v", err) } - provider.client = server.Client() + p.client = server.Client() - events := make(chan domain.StreamEvent, 16) - err = provider.Chat(context.Background(), domain.ChatRequest{ - Model: config.OpenAIDefaultModel, - Messages: []domain.Message{{Role: "user", Content: "test"}}, - }, events) + events := make(chan providertypes.StreamEvent, 16) + err = p.Chat(context.Background(), providertypes.ChatRequest{Model: config.OpenAIDefaultModel, Messages: []providertypes.Message{{Role: "user", Content: "test"}}}, events) if err != nil { t.Fatalf("Chat() error = %v", err) } - var ( - foundTextDelta bool - foundToolCallStart bool - foundToolCallDelta bool - foundMessageDone bool - toolCallDeltaContent string - messageDonePayload *domain.MessageDonePayload - ) + var foundTextDelta, foundToolCallStart, foundToolCallDelta, foundMessageDone bool + var toolCallDeltaContent string + var messageDonePayload *providertypes.MessageDonePayload for _, evt := range drainStreamEvents(events) { switch evt.Type { - case domain.StreamEventTextDelta: + case providertypes.StreamEventTextDelta: foundTextDelta = true - case domain.StreamEventToolCallStart: + case providertypes.StreamEventToolCallStart: foundToolCallStart = true - payload := requireToolCallStartPayload(t, evt) - if payload.Name != "filesystem_edit" { - t.Fatalf("expected ToolName %q, got %q", "filesystem_edit", payload.Name) - } - if payload.Index != 0 { - t.Fatalf("expected ToolCallIndex %d for tool_call_start, got %d", 0, payload.Index) + p := requireToolCallStartPayload(t, evt) + if p.Name != "filesystem_edit" { + t.Fatalf("expected ToolName %q, got %q", "filesystem_edit", p.Name) } - case domain.StreamEventToolCallDelta: + case providertypes.StreamEventToolCallDelta: foundToolCallDelta = true toolCallDeltaContent += requireToolCallDeltaPayload(t, evt).ArgumentsDelta - case domain.StreamEventMessageDone: + case providertypes.StreamEventMessageDone: foundMessageDone = true - payload := requireMessageDonePayload(t, evt) - messageDonePayload = &payload + p := requireMessageDonePayload(t, evt) + messageDonePayload = &p } } @@ -1008,14 +1450,9 @@ func TestProviderChatEmitsFullEventStream(t *testing.T) { if !foundMessageDone { t.Fatal("expected StreamEventMessageDone event") } - - // 验证 tool_call_delta 内容被正确累加 - expectedDelta := `{"path":"a.go"}` - if toolCallDeltaContent != expectedDelta { - t.Fatalf("expected tool call delta content %q, got %q", expectedDelta, toolCallDeltaContent) + if toolCallDeltaContent != `{"path":"a.go"}` { + t.Fatalf("expected tool call delta content %q, got %q", `{"path":"a.go"}`, toolCallDeltaContent) } - - // 验证 message_done 事件包含正确的字段 if messageDonePayload == nil { t.Fatal("message_done event is nil") } @@ -1029,75 +1466,3 @@ func TestProviderChatEmitsFullEventStream(t *testing.T) { t.Fatalf("expected TotalTokens %d, got %d", 150, messageDonePayload.Usage.TotalTokens) } } - -// --- 辅助方法测试 --- - -// --- consumeStream 错误包装测试 --- - -func TestConsumeStream_WrapsNonEOFAsInterrupted(t *testing.T) { - t.Setenv(config.OpenAIDefaultAPIKeyEnv, "test-key") - - p, err := New(resolvedConfig("", "")) - if err != nil { - t.Fatalf("New() error = %v", err) - } - - // 使用一个会触发读取错误的 source(模拟网络断开) - errReader := &errReader{err: io.ErrClosedPipe} - - err = p.consumeStream(context.Background(), errReader, make(chan domain.StreamEvent, 1)) - if err == nil { - t.Fatal("expected error for broken reader") - } - if !errors.Is(err, domain.ErrStreamInterrupted) { - t.Fatalf("expected ErrStreamInterrupted wrapping, got: %v", err) - } -} - -// TestConsumeStream_FlushesPendingDataOnNonEOFError 验证非 EOF 读取错误发生前, -// 已缓冲但尚未刷新的 data: 行仍会被处理(不会因中断而丢失)。 -func TestConsumeStream_FlushesPendingDataOnNonEOFError(t *testing.T) { - t.Setenv(config.OpenAIDefaultAPIKeyEnv, "test-key") - - p, err := New(resolvedConfig("", "")) - if err != nil { - t.Fatalf("New() error = %v", err) - } - - // 构造一段 SSE 数据:包含一个有效 data 行,但紧跟一个错误而非空行。 - // 这模拟了流中断前最后一帧数据还没来得及被空行触发刷新的场景。 - sseData := `data: {"id":"a","object":"chat.completion.chunk","choices":[{"delta":{"content":"hello"},"finish_reason":""}]} -` // 注意:这里有换行符,但由于紧接着是 error,不会被空行刷新 - body := io.MultiReader(strings.NewReader(sseData), &errReader{err: io.ErrClosedPipe}) - - events := make(chan domain.StreamEvent, 10) - - err = p.consumeStream(context.Background(), body, events) - if err == nil { - t.Fatal("expected error for broken reader") - } - if !errors.Is(err, domain.ErrStreamInterrupted) { - t.Fatalf("expected ErrStreamInterrupted, got: %v", err) - } - - // 关键断言:中断前的 data 行必须已被刷新处理,应有 text_delta 事件。 - drained := drainStreamEvents(events) - var foundText bool - for _, evt := range drained { - if evt.Type == domain.StreamEventTextDelta { - foundText = true - } - } - if !foundText { - t.Fatal("expected text_delta event from flushed pending data") - } -} - -// errReader 是一个每次 ReadLine 都返回指定错误的测试辅助类型。 -type errReader struct { - err error -} - -func (e *errReader) Read(p []byte) (int, error) { - return 0, e.err -} diff --git a/internal/provider/openai/provider.go b/internal/provider/openai/provider.go new file mode 100644 index 00000000..f6db0d37 --- /dev/null +++ b/internal/provider/openai/provider.go @@ -0,0 +1,121 @@ +package openai + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "log" + "net/http" + "strings" + "time" + + "neo-code/internal/config" + providertypes "neo-code/internal/provider/types" +) + +// Provider 封装 OpenAI 兼容 API 的客户端配置和 HTTP 连接。 +type Provider struct { + cfg config.ResolvedProviderConfig + client *http.Client +} + +// buildOptions 控制构造行为,用于注入自定义 Transport 等选项。 +type buildOptions struct { + transport http.RoundTripper +} + +// buildOption 是 New() 的函数式配置选项。 +type buildOption func(*buildOptions) + +// withTransport 注入自定义 HTTP Transport(如 RetryTransport)。 +func withTransport(rt http.RoundTripper) buildOption { + return func(o *buildOptions) { + o.transport = rt + } +} + +// New 创建 OpenAI provider 实例。cfg 必须 Validate 通过且包含有效 API Key。 +func New(cfg config.ResolvedProviderConfig, opts ...buildOption) (*Provider, error) { + if err := cfg.Validate(); err != nil { + return nil, fmt.Errorf("openai provider: %w", err) + } + if strings.TrimSpace(cfg.APIKey) == "" { + return nil, errors.New("openai provider: api key is empty") + } + + o := &buildOptions{ + transport: http.DefaultTransport, + } + for _, apply := range opts { + apply(o) + } + + return &Provider{ + cfg: cfg, + client: &http.Client{ + Timeout: 90 * time.Second, + Transport: o.transport, + }, + }, nil +} + +// DiscoverModels 通过 /models 端点查询可用模型列表。 +func (p *Provider) DiscoverModels(ctx context.Context) ([]config.ModelDescriptor, error) { + rawModels, err := p.fetchModels(ctx) + if err != nil { + return nil, err + } + + descriptors := make([]config.ModelDescriptor, 0, len(rawModels)) + for _, raw := range rawModels { + descriptor, ok := config.DescriptorFromRawModel(raw) + if !ok { + continue + } + descriptors = append(descriptors, descriptor) + } + return config.MergeModelDescriptors(descriptors), nil +} + +// Chat 发起 SSE 流式对话请求。 +// 流中途断连或协议错误时直接返回错误,由上层调用方决定重试策略。 +func (p *Provider) Chat(ctx context.Context, req providertypes.ChatRequest, events chan<- providertypes.StreamEvent) error { + payload, err := p.buildRequest(req) + if err != nil { + return err + } + + body, err := json.Marshal(payload) + if err != nil { + return fmt.Errorf("openai provider: marshal request: %w", err) + } + + endpoint := strings.TrimRight(p.cfg.BaseURL, "/") + "/chat/completions" + httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(body)) + if err != nil { + return fmt.Errorf("openai provider: build request: %w", err) + } + httpReq.Header.Set("Authorization", "Bearer "+p.cfg.APIKey) + httpReq.Header.Set("Content-Type", "application/json") + httpReq.Header.Set("Accept", "text/event-stream") + + resp, err := p.client.Do(httpReq) + if err != nil { + return fmt.Errorf("openai provider: send request: %w", err) + } + defer func(Body io.ReadCloser) { + err := Body.Close() + if err != nil { + log.Printf("openai provider: close response body: %v", err) + } + }(resp.Body) + + if resp.StatusCode >= http.StatusBadRequest { + return p.parseError(resp) + } + + return p.consumeStream(ctx, resp.Body, events) +} diff --git a/internal/provider/openai/request.go b/internal/provider/openai/request.go new file mode 100644 index 00000000..02f5398b --- /dev/null +++ b/internal/provider/openai/request.go @@ -0,0 +1,105 @@ +package openai + +import ( + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "strings" + + "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" +) + +// buildRequest 将 provider.ChatRequest 转换为 OpenAI API 请求结构。 +// 模型优先取 req.Model,其次使用配置中的默认模型。 +func (p *Provider) buildRequest(req providertypes.ChatRequest) (chatCompletionRequest, error) { + model := strings.TrimSpace(req.Model) + if model == "" { + model = strings.TrimSpace(p.cfg.Model) + } + if model == "" { + return chatCompletionRequest{}, errors.New("openai provider: model is empty") + } + + payload := chatCompletionRequest{ + Model: model, + Stream: true, + Messages: make([]openAIMessage, 0, len(req.Messages)+1), + } + + if strings.TrimSpace(req.SystemPrompt) != "" { + payload.Messages = append(payload.Messages, openAIMessage{ + Role: providertypes.RoleSystem, + Content: req.SystemPrompt, + }) + } + + for _, message := range req.Messages { + payload.Messages = append(payload.Messages, toOpenAIMessage(message)) + } + + if len(req.Tools) > 0 { + payload.ToolChoice = "auto" + payload.Tools = make([]openAIToolDefinition, 0, len(req.Tools)) + for _, spec := range req.Tools { + payload.Tools = append(payload.Tools, openAIToolDefinition{ + Type: "function", + Function: openAIFunctionDefinition{ + Name: spec.Name, + Description: spec.Description, + Parameters: spec.Schema, + }, + }) + } + } + + return payload, nil +} + +// toOpenAIMessage 将通用 Message 转换为 OpenAI 协议消息格式。 +func toOpenAIMessage(message providertypes.Message) openAIMessage { + out := openAIMessage{ + Role: message.Role, + Content: message.Content, + ToolCallID: message.ToolCallID, + } + + if len(message.ToolCalls) > 0 { + out.ToolCalls = make([]openAIToolCall, 0, len(message.ToolCalls)) + for _, call := range message.ToolCalls { + out.ToolCalls = append(out.ToolCalls, openAIToolCall{ + ID: call.ID, + Type: "function", + Function: openAIFunctionCall{ + Name: call.Name, + Arguments: call.Arguments, + }, + }) + } + } + + return out +} + +// parseError 解析 HTTP 错误响应并包装为 ProviderError。 +func (p *Provider) parseError(resp *http.Response) error { + data, readErr := io.ReadAll(resp.Body) + if readErr != nil { + return provider.NewProviderErrorFromStatus(resp.StatusCode, + fmt.Sprintf("openai provider: read error response: %v", readErr)) + } + + var parsed openAIErrorResponse + if err := json.Unmarshal(data, &parsed); err == nil && strings.TrimSpace(parsed.Error.Message) != "" { + return provider.NewProviderErrorFromStatus(resp.StatusCode, parsed.Error.Message) + } + + bodyText := strings.TrimSpace(string(data)) + if bodyText == "" { + return provider.NewProviderErrorFromStatus(resp.StatusCode, resp.Status) + } + + return provider.NewProviderErrorFromStatus(resp.StatusCode, bodyText) +} diff --git a/internal/provider/openai/response.go b/internal/provider/openai/response.go new file mode 100644 index 00000000..02b7b9b5 --- /dev/null +++ b/internal/provider/openai/response.go @@ -0,0 +1,135 @@ +package openai + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "strings" + + "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" +) + +// consumeStream 消费 SSE 响应流,使用有界读取器防止缓冲区溢出。 +func (p *Provider) consumeStream( + ctx context.Context, + body io.Reader, + events chan<- providertypes.StreamEvent, +) error { + reader := newBoundedSSEReader(body) + + var ( + finishReason string + usage providertypes.Usage + done bool + toolCalls = make(map[int]*providertypes.ToolCall) + ) + + dataLines := make([]string, 0, 4) + + // processChunk 解析单个 SSE data payload,发送事件。 + processChunk := func(payload string) error { + if strings.TrimSpace(payload) == "[DONE]" { + done = true + return nil + } + + var chunk chatCompletionChunk + if err := json.Unmarshal([]byte(payload), &chunk); err != nil { + return fmt.Errorf("openai provider: decode stream chunk: %w", err) + } + + if chunk.Error != nil && strings.TrimSpace(chunk.Error.Message) != "" { + return errors.New(chunk.Error.Message) + } + + extractStreamUsage(&usage, chunk.Usage) + + for _, choice := range chunk.Choices { + if choice.FinishReason != "" { + finishReason = choice.FinishReason + } + if choice.Delta.Content != "" { + if err := emitTextDelta(ctx, events, choice.Delta.Content); err != nil { + return err + } + } + for _, delta := range choice.Delta.ToolCalls { + if err := mergeToolCallDelta(ctx, events, toolCalls, delta); err != nil { + return err + } + } + } + return nil + } + + // finishStream 统一的流结束处理:发送 message_done 事件。 + finishStream := func() error { + return emitMessageDone(ctx, events, finishReason, &usage) + } + + flushPendingData := func() error { + defer func() { dataLines = dataLines[:0] }() + return flushDataLines(dataLines, processChunk) + } + + for { + line, err := reader.ReadLine() + + if err != nil && !errors.Is(err, io.EOF) { + // 非 EOF 的读取错误:先刷新缓冲的 data 行,再包装为流中断, + // 避免中断前最后一段数据丢失。 + if flushErr := flushPendingData(); flushErr != nil { + return flushErr + } + return fmt.Errorf("%w: %w", provider.ErrStreamInterrupted, err) + } + + trimmed := line + + switch { + case strings.HasPrefix(trimmed, "data:"): + data := strings.TrimSpace(strings.TrimPrefix(trimmed, "data:")) + // data: [DONE] 需要立即处理:先刷新已缓冲的 data 行,再标记结束, + // 避免与前面的合法 JSON 拼接后导致 json.Unmarshal 失败。 + if data == "[DONE]" { + if flushErr := flushPendingData(); flushErr != nil { + return flushErr + } + done = true + } else { + dataLines = append(dataLines, data) + } + case trimmed == "": + if flushErr := flushPendingData(); flushErr != nil { + return flushErr + } + if done { + return finishStream() + } + case strings.HasPrefix(trimmed, ":"): + // SSE comment/heartbeat; ignore. + } + + if errors.Is(err, io.EOF) { + if flushErr := flushPendingData(); flushErr != nil { + return flushErr + } + return finishStream() + } + } +} + +// extractStreamUsage 从 OpenAI usage 响应提取并覆盖累积的 token 统计。 +func extractStreamUsage(usage *providertypes.Usage, raw *openAIUsage) { + if raw == nil { + return + } + *usage = providertypes.Usage{ + InputTokens: raw.PromptTokens, + OutputTokens: raw.CompletionTokens, + TotalTokens: raw.TotalTokens, + } +} diff --git a/internal/provider/openai/toolcall.go b/internal/provider/openai/toolcall.go new file mode 100644 index 00000000..339d86c9 --- /dev/null +++ b/internal/provider/openai/toolcall.go @@ -0,0 +1,43 @@ +package openai + +import ( + "context" + "strings" + + providertypes "neo-code/internal/provider/types" +) + +// mergeToolCallDelta 将单个 tool call delta 累积到 toolCalls map 中。 +// 首次发现带名称的 delta 时发送 tool_call_start 事件; +// 每次收到 arguments 增量时发送 tool_call_delta 事件。 +func mergeToolCallDelta(ctx context.Context, events chan<- providertypes.StreamEvent, toolCalls map[int]*providertypes.ToolCall, delta toolCallDelta) error { + call, exists := toolCalls[delta.Index] + if !exists { + call = &providertypes.ToolCall{} + toolCalls[delta.Index] = call + } + + hadName := strings.TrimSpace(call.Name) != "" + + if id := strings.TrimSpace(delta.ID); id != "" { + call.ID = id + } + if name := strings.TrimSpace(delta.Function.Name); name != "" { + call.Name = name + } + + if !hadName && strings.TrimSpace(call.Name) != "" { + if err := emitToolCallStart(ctx, events, delta.Index, call.ID, call.Name); err != nil { + return err + } + } + + // 发送参数增量事件(同一 chunk 可能同时携带 name 和 arguments) + if args := delta.Function.Arguments; args != "" { + call.Arguments += args + if err := emitToolCallDelta(ctx, events, delta.Index, call.ID, args); err != nil { + return err + } + } + return nil +} diff --git a/internal/provider/openai/types.go b/internal/provider/openai/types.go new file mode 100644 index 00000000..f943e1e3 --- /dev/null +++ b/internal/provider/openai/types.go @@ -0,0 +1,90 @@ +package openai + +// 以下类型定义了 OpenAI Chat Completions API 的请求和响应结构体, +// 仅在 openai 子包内部使用,不对外暴露。 + +// chatCompletionRequest 表示 /chat/completions 端点的请求体。 +type chatCompletionRequest struct { + Model string `json:"model"` + Messages []openAIMessage `json:"messages"` + Tools []openAIToolDefinition `json:"tools,omitempty"` + ToolChoice string `json:"tool_choice,omitempty"` + Stream bool `json:"stream"` +} + +// openAIMessage 表示 OpenAI 协议中的消息格式。 +type openAIMessage struct { + Role string `json:"role"` + Content string `json:"content,omitempty"` + ToolCallID string `json:"tool_call_id,omitempty"` + ToolCalls []openAIToolCall `json:"tool_calls,omitempty"` +} + +// openAIToolDefinition 表示工具定义的 OpenAI 格式。 +type openAIToolDefinition struct { + Type string `json:"type"` + Function openAIFunctionDefinition `json:"function"` +} + +// openAIFunctionDefinition 表示函数描述的 OpenAI 格式。 +type openAIFunctionDefinition struct { + Name string `json:"name"` + Description string `json:"description,omitempty"` + Parameters map[string]any `json:"parameters,omitempty"` +} + +// openAIToolCall 表示响应中工具调用的 OpenAI 格式。 +type openAIToolCall struct { + ID string `json:"id,omitempty"` + Type string `json:"type,omitempty"` + Function openAIFunctionCall `json:"function"` +} + +// openAIFunctionCall 表示函数调用参数的 OpenAI 格式。 +type openAIFunctionCall struct { + Name string `json:"name,omitempty"` + Arguments string `json:"arguments,omitempty"` +} + +// chatCompletionChunk 表示 SSE 流式响应中的单个 chunk(内部使用,非导出)。 +type chatCompletionChunk struct { + Choices []struct { + Index int `json:"index"` + Delta chunkDelta `json:"delta"` + FinishReason string `json:"finish_reason"` + } `json:"choices"` + Usage *openAIUsage `json:"usage,omitempty"` + Error *struct { + Message string `json:"message"` + } `json:"error,omitempty"` +} + +// chunkDelta 表示流式 chunk 中的增量内容。 +type chunkDelta struct { + Role string `json:"role,omitempty"` + Content string `json:"content,omitempty"` + ToolCalls []toolCallDelta `json:"tool_calls,omitempty"` +} + +// toolCallDelta 表示流式 tool call 增量。 +type toolCallDelta struct { + Index int `json:"index"` + ID string `json:"id,omitempty"` + Type string `json:"type,omitempty"` + Function openAIFunctionCall `json:"function"` +} + +// openAIUsage 表示 token 使用统计。 +type openAIUsage struct { + PromptTokens int `json:"prompt_tokens"` + CompletionTokens int `json:"completion_tokens"` + TotalTokens int `json:"total_tokens"` +} + +// openAIErrorResponse 表示 API 错误响应。 +type openAIErrorResponse struct { + Error struct { + Message string `json:"message"` + Code string `json:"code,omitempty"` + } `json:"error"` +} diff --git a/internal/provider/provider.go b/internal/provider/provider.go index 91a99b52..f8aff036 100644 --- a/internal/provider/provider.go +++ b/internal/provider/provider.go @@ -1,7 +1,12 @@ package provider -import "context" +import ( + "context" + "neo-code/internal/provider/types" +) + +// Provider 定义模型对话能力,通过 channel 推送流式事件给上层消费。 type Provider interface { - Chat(ctx context.Context, req ChatRequest, events chan<- StreamEvent) error + Chat(ctx context.Context, req types.ChatRequest, events chan<- types.StreamEvent) error } diff --git a/internal/provider/registry_test.go b/internal/provider/registry_test.go index df57c9da..5ddc5b31 100644 --- a/internal/provider/registry_test.go +++ b/internal/provider/registry_test.go @@ -8,11 +8,12 @@ import ( "neo-code/internal/config" "neo-code/internal/provider" "neo-code/internal/provider/openai" + providertypes "neo-code/internal/provider/types" ) type stubProvider struct{} -func (stubProvider) Chat(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { +func (stubProvider) Chat(ctx context.Context, req providertypes.ChatRequest, events chan<- providertypes.StreamEvent) error { return nil } diff --git a/internal/provider/types.go b/internal/provider/types/event.go similarity index 76% rename from internal/provider/types.go rename to internal/provider/types/event.go index 1f689c12..72df4bc1 100644 --- a/internal/provider/types.go +++ b/internal/provider/types/event.go @@ -1,56 +1,7 @@ -package provider +package types import "fmt" -const ( - // RoleSystem 标识系统消息。 - RoleSystem = "system" - // RoleUser 标识用户消息。 - RoleUser = "user" - // RoleAssistant 标识助手消息。 - RoleAssistant = "assistant" - // RoleTool 标识工具结果消息。 - RoleTool = "tool" -) - -// Message 表示对话中的单条消息。 -type Message struct { - Role string `json:"role"` - Content string `json:"content"` - ToolCalls []ToolCall `json:"tool_calls,omitempty"` - ToolCallID string `json:"tool_call_id,omitempty"` - IsError bool `json:"is_error,omitempty"` -} - -// ToolCall 表示模型发起的工具调用请求。 -type ToolCall struct { - ID string `json:"id"` - Name string `json:"name"` - Arguments string `json:"arguments"` -} - -// ToolSpec 表示暴露给模型的可调用工具描述。 -type ToolSpec struct { - Name string `json:"name"` - Description string `json:"description"` - Schema map[string]any `json:"schema"` -} - -// ChatRequest 是 provider.Chat() 的请求参数。 -type ChatRequest struct { - Model string `json:"model"` - SystemPrompt string `json:"system_prompt"` - Messages []Message `json:"messages"` - Tools []ToolSpec `json:"tools,omitempty"` -} - -// Usage 记录本次请求的 token 使用统计。 -type Usage struct { - InputTokens int `json:"input_tokens"` - OutputTokens int `json:"output_tokens"` - TotalTokens int `json:"total_tokens"` -} - // StreamEventType 定义流式事件类型。 type StreamEventType string diff --git a/internal/provider/types/message.go b/internal/provider/types/message.go new file mode 100644 index 00000000..3758c99f --- /dev/null +++ b/internal/provider/types/message.go @@ -0,0 +1,36 @@ +package types + +// RoleSystem 标识系统消息。 +const RoleSystem = "system" + +// RoleUser 标识用户消息。 +const RoleUser = "user" + +// RoleAssistant 标识助手消息。 +const RoleAssistant = "assistant" + +// RoleTool 标识工具结果消息。 +const RoleTool = "tool" + +// Message 表示对话中的单条消息。 +type Message struct { + Role string `json:"role"` + Content string `json:"content"` + ToolCalls []ToolCall `json:"tool_calls,omitempty"` + ToolCallID string `json:"tool_call_id,omitempty"` + IsError bool `json:"is_error,omitempty"` +} + +// ToolCall 表示模型发起的工具调用请求。 +type ToolCall struct { + ID string `json:"id"` + Name string `json:"name"` + Arguments string `json:"arguments"` +} + +// ToolSpec 表示暴露给模型的可调用工具描述。 +type ToolSpec struct { + Name string `json:"name"` + Description string `json:"description"` + Schema map[string]any `json:"schema"` +} diff --git a/internal/provider/types/request.go b/internal/provider/types/request.go new file mode 100644 index 00000000..1171f6b0 --- /dev/null +++ b/internal/provider/types/request.go @@ -0,0 +1,16 @@ +package types + +// ChatRequest 是 provider.Chat() 的请求参数。 +type ChatRequest struct { + Model string `json:"model"` + SystemPrompt string `json:"system_prompt"` + Messages []Message `json:"messages"` + Tools []ToolSpec `json:"tools,omitempty"` +} + +// Usage 记录本次请求的 token 使用统计。 +type Usage struct { + InputTokens int `json:"input_tokens"` + OutputTokens int `json:"output_tokens"` + TotalTokens int `json:"total_tokens"` +} diff --git a/internal/provider/types_test.go b/internal/provider/types/types_test.go similarity index 99% rename from internal/provider/types_test.go rename to internal/provider/types/types_test.go index 5c822ee8..958ebb16 100644 --- a/internal/provider/types_test.go +++ b/internal/provider/types/types_test.go @@ -1,4 +1,4 @@ -package provider +package types import ( "encoding/json" diff --git a/internal/runtime/compact.go b/internal/runtime/compact.go index 4939ca62..9798d250 100644 --- a/internal/runtime/compact.go +++ b/internal/runtime/compact.go @@ -8,7 +8,7 @@ import ( "neo-code/internal/config" contextcompact "neo-code/internal/context/compact" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" ) // CompactInput 描述一次手动 compact 请求所需的最小输入。 @@ -103,7 +103,7 @@ func (s *Service) runCompactForSession( } } - originalMessages := append([]provider.Message(nil), session.Messages...) + originalMessages := append([]providertypes.Message(nil), session.Messages...) s.emit(ctx, EventCompactStart, runID, session.ID, string(contextcompact.ModeManual)) result, err := runner.Run(ctx, contextcompact.Input{ @@ -125,7 +125,7 @@ func (s *Service) runCompactForSession( } if result.Applied { - session.Messages = append([]provider.Message(nil), result.Messages...) + session.Messages = append([]providertypes.Message(nil), result.Messages...) session.UpdatedAt = time.Now() if err := s.sessionStore.Save(ctx, &session); err != nil { s.emit(ctx, EventCompactError, runID, session.ID, CompactErrorPayload{ diff --git a/internal/runtime/compact_generator.go b/internal/runtime/compact_generator.go index 20be11a1..af62c13a 100644 --- a/internal/runtime/compact_generator.go +++ b/internal/runtime/compact_generator.go @@ -8,7 +8,7 @@ import ( "neo-code/internal/config" agentcontext "neo-code/internal/context" contextcompact "neo-code/internal/context/compact" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" ) type compactSummaryGenerator struct { @@ -58,7 +58,7 @@ func (g *compactSummaryGenerator) Generate(ctx context.Context, input contextcom } // 使用流式事件通道收集 compact 摘要响应。 - streamEvents := make(chan provider.StreamEvent, 32) + streamEvents := make(chan providertypes.StreamEvent, 32) streamDone := make(chan error, 1) acc := newStreamAccumulator() @@ -84,11 +84,11 @@ func (g *compactSummaryGenerator) Generate(ctx context.Context, input contextcom } }() - err = modelProvider.Chat(ctx, provider.ChatRequest{ + err = modelProvider.Chat(ctx, providertypes.ChatRequest{ Model: g.model, SystemPrompt: prompt.SystemPrompt, - Messages: []provider.Message{{ - Role: provider.RoleUser, + Messages: []providertypes.Message{{ + Role: providertypes.RoleUser, Content: prompt.UserPrompt, }}, }, streamEvents) diff --git a/internal/runtime/compact_generator_test.go b/internal/runtime/compact_generator_test.go index e82c625f..00a96ec9 100644 --- a/internal/runtime/compact_generator_test.go +++ b/internal/runtime/compact_generator_test.go @@ -8,7 +8,7 @@ import ( "neo-code/internal/config" contextcompact "neo-code/internal/context/compact" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" ) func TestCompactSummaryGeneratorBuildsProviderRequestWithoutTools(t *testing.T) { @@ -21,8 +21,8 @@ func TestCompactSummaryGeneratorBuildsProviderRequestWithoutTools(t *testing.T) } scripted := &scriptedProvider{ - streams: [][]provider.StreamEvent{ - {provider.NewTextDeltaStreamEvent(strings.Join([]string{ + streams: [][]providertypes.StreamEvent{ + {providertypes.NewTextDeltaStreamEvent(strings.Join([]string{ "[compact_summary]", "done:", "- Completed the historical task and kept the final result.", @@ -46,17 +46,17 @@ func TestCompactSummaryGeneratorBuildsProviderRequestWithoutTools(t *testing.T) summary, err := generator.Generate(context.Background(), contextcompact.SummaryInput{ Mode: contextcompact.ModeManual, - ArchivedMessages: []provider.Message{ - {Role: provider.RoleUser, Content: "legacy request"}, + ArchivedMessages: []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "legacy request"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, }, }, }, - RetainedMessages: []provider.Message{ - {Role: provider.RoleAssistant, Content: "recent answer"}, + RetainedMessages: []providertypes.Message{ + {Role: providertypes.RoleAssistant, Content: "recent answer"}, }, ArchivedMessageCount: 2, Config: manager.Get().Context.Compact, @@ -87,7 +87,7 @@ func TestCompactSummaryGeneratorBuildsProviderRequestWithoutTools(t *testing.T) if !strings.Contains(req.SystemPrompt, "[compact_summary]") { t.Fatalf("expected compact system prompt, got %q", req.SystemPrompt) } - if len(req.Messages) != 1 || req.Messages[0].Role != provider.RoleUser { + if len(req.Messages) != 1 || req.Messages[0].Role != providertypes.RoleUser { t.Fatalf("expected a single user prompt, got %+v", req.Messages) } if !strings.Contains(req.Messages[0].Content, "") { @@ -114,10 +114,10 @@ func TestCompactSummaryGeneratorRejectsToolCalls(t *testing.T) { } scripted := &scriptedProvider{ - streams: [][]provider.StreamEvent{ + streams: [][]providertypes.StreamEvent{ { - provider.NewToolCallStartStreamEvent(0, "call-1", "filesystem_read_file"), - provider.NewToolCallDeltaStreamEvent(0, "call-1", "{}"), + providertypes.NewToolCallStartStreamEvent(0, "call-1", "filesystem_read_file"), + providertypes.NewToolCallDeltaStreamEvent(0, "call-1", "{}"), }, }, } @@ -125,8 +125,8 @@ func TestCompactSummaryGeneratorRejectsToolCalls(t *testing.T) { _, err = generator.Generate(context.Background(), contextcompact.SummaryInput{ Mode: contextcompact.ModeManual, - ArchivedMessages: []provider.Message{ - {Role: provider.RoleUser, Content: "legacy request"}, + ArchivedMessages: []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "legacy request"}, }, Config: manager.Get().Context.Compact, }) @@ -145,9 +145,9 @@ func TestCompactSummaryGeneratorRejectsMalformedStreamEvent(t *testing.T) { } scripted := &scriptedProvider{ - streams: [][]provider.StreamEvent{ + streams: [][]providertypes.StreamEvent{ { - {Type: provider.StreamEventTextDelta}, + {Type: providertypes.StreamEventTextDelta}, }, }, } @@ -171,12 +171,12 @@ func TestCompactSummaryGeneratorMalformedStreamEventDoesNotDeadlock(t *testing.T t.Fatalf("resolve provider: %v", err) } - stream := []provider.StreamEvent{{Type: provider.StreamEventTextDelta}} + stream := []providertypes.StreamEvent{{Type: providertypes.StreamEventTextDelta}} for i := 0; i < 40; i++ { - stream = append(stream, provider.NewTextDeltaStreamEvent("ignored")) + stream = append(stream, providertypes.NewTextDeltaStreamEvent("ignored")) } scripted := &scriptedProvider{ - streams: [][]provider.StreamEvent{stream}, + streams: [][]providertypes.StreamEvent{stream}, } generator := newCompactSummaryGenerator(&scriptedProviderFactory{provider: scripted}, resolvedProvider, "session-model") diff --git a/internal/runtime/permission.go b/internal/runtime/permission.go index 6782e9be..0ffe62ee 100644 --- a/internal/runtime/permission.go +++ b/internal/runtime/permission.go @@ -8,7 +8,7 @@ import ( "sync" "time" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" "neo-code/internal/security" "neo-code/internal/tools" ) @@ -31,7 +31,7 @@ const ( type permissionExecutionInput struct { RunID string SessionID string - Call provider.ToolCall + Call providertypes.ToolCall Workdir string ToolTimeout time.Duration } @@ -40,7 +40,7 @@ type pendingPermissionRequest struct { RequestID string RunID string SessionID string - Call provider.ToolCall + Call providertypes.ToolCall Action security.Action ResultCh chan PermissionResolutionDecision Submitted bool diff --git a/internal/runtime/permission_test.go b/internal/runtime/permission_test.go index 23102aea..c82a9526 100644 --- a/internal/runtime/permission_test.go +++ b/internal/runtime/permission_test.go @@ -6,7 +6,7 @@ import ( "testing" "time" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" "neo-code/internal/security" "neo-code/internal/tools" ) @@ -53,7 +53,7 @@ func TestResolvePermissionSuccess(t *testing.T) { request := registerPendingPermission(service, permissionExecutionInput{ RunID: "run-permission", SessionID: "session-permission", - Call: provider.ToolCall{ + Call: providertypes.ToolCall{ ID: "call-1", Name: "webfetch", }, @@ -105,7 +105,7 @@ func TestResolvePermissionDuplicateSubmissionIsNonBlocking(t *testing.T) { request := registerPendingPermission(service, permissionExecutionInput{ RunID: "run-permission-dup", SessionID: "session-permission-dup", - Call: provider.ToolCall{ + Call: providertypes.ToolCall{ ID: "call-dup", Name: "webfetch", }, @@ -170,13 +170,19 @@ func TestServiceRunPermissionRejectFlow(t *testing.T) { } scripted := &scriptedProvider{ - streams: [][]provider.StreamEvent{ + responses: []scriptedResponse{ { - provider.NewToolCallStartStreamEvent(0, "call-ask-reject", "webfetch"), - provider.NewToolCallDeltaStreamEvent(0, "call-ask-reject", `{"url":"https://example.com/private"}`), + Message: providertypes.Message{ + Role: "assistant", + ToolCalls: []providertypes.ToolCall{ + {ID: "call-ask-reject", Name: "webfetch", Arguments: `{"url":"https://example.com/private"}`}, + }, + }, + FinishReason: "tool_calls", }, { - provider.NewTextDeltaStreamEvent("done"), + Message: providertypes.Message{Role: "assistant", Content: "done"}, + FinishReason: "stop", }, }, } @@ -306,7 +312,7 @@ func TestResolvePermissionCanceledContext(t *testing.T) { request := registerPendingPermission(service, permissionExecutionInput{ RunID: "run-canceled", SessionID: "session-canceled", - Call: provider.ToolCall{ + Call: providertypes.ToolCall{ ID: "call-canceled", Name: "webfetch", }, diff --git a/internal/runtime/runtime.go b/internal/runtime/runtime.go index b9e97b1d..7b727e08 100644 --- a/internal/runtime/runtime.go +++ b/internal/runtime/runtime.go @@ -17,6 +17,7 @@ import ( agentcontext "neo-code/internal/context" contextcompact "neo-code/internal/context/compact" "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" "neo-code/internal/tools" ) @@ -34,13 +35,13 @@ const ( // 包括文本内容和工具调用列表。 type streamAccumulator struct { content strings.Builder - toolCalls map[int]*provider.ToolCall + toolCalls map[int]*providertypes.ToolCall } // newStreamAccumulator 创建并初始化一个空的流式事件累积器。 func newStreamAccumulator() *streamAccumulator { return &streamAccumulator{ - toolCalls: make(map[int]*provider.ToolCall), + toolCalls: make(map[int]*providertypes.ToolCall), } } @@ -50,10 +51,10 @@ func (a *streamAccumulator) accumulateTextDelta(text string) { } // ensureToolCall 返回指定索引的工具调用条目,不存在时会先创建占位对象。 -func (a *streamAccumulator) ensureToolCall(index int) *provider.ToolCall { +func (a *streamAccumulator) ensureToolCall(index int) *providertypes.ToolCall { call, exists := a.toolCalls[index] if !exists { - call = &provider.ToolCall{} + call = &providertypes.ToolCall{} a.toolCalls[index] = call } return call @@ -80,15 +81,15 @@ func (a *streamAccumulator) accumulateToolCallDelta(index int, id, argumentsDelt } // buildMessage 从累积状态构建最终的 assistant Message 对象,并校验工具调用元数据是否完整。 -func (a *streamAccumulator) buildMessage() (provider.Message, error) { +func (a *streamAccumulator) buildMessage() (providertypes.Message, error) { ordered := make([]int, 0, len(a.toolCalls)) for index := range a.toolCalls { ordered = append(ordered, index) } sort.Ints(ordered) - message := provider.Message{ - Role: provider.RoleAssistant, + message := providertypes.Message{ + Role: providertypes.RoleAssistant, Content: a.content.String(), } for _, index := range ordered { @@ -97,10 +98,10 @@ func (a *streamAccumulator) buildMessage() (provider.Message, error) { continue } if strings.TrimSpace(call.ID) == "" { - return provider.Message{}, fmt.Errorf("runtime: provider emitted tool call %d without id", index) + return providertypes.Message{}, fmt.Errorf("runtime: provider emitted tool call %d without id", index) } if strings.TrimSpace(call.Name) == "" { - return provider.Message{}, fmt.Errorf("runtime: provider emitted tool call %d without name", index) + return providertypes.Message{}, fmt.Errorf("runtime: provider emitted tool call %d without name", index) } message.ToolCalls = append(message.ToolCalls, *call) } @@ -200,8 +201,8 @@ func (s *Service) Run(ctx context.Context, input UserInput) error { return s.handleRunError(ctx, input.RunID, input.SessionID, err) } - userMessage := provider.Message{ - Role: provider.RoleUser, + userMessage := providertypes.Message{ + Role: providertypes.RoleUser, Content: input.Content, } session.Messages = append(session.Messages, userMessage) @@ -251,7 +252,7 @@ func (s *Service) Run(ctx context.Context, input UserInput) error { return s.handleRunError(ctx, input.RunID, session.ID, err) } - acc, err := s.callProviderWithRetry(ctx, input.RunID, session.ID, provider.ChatRequest{ + acc, err := s.callProviderWithRetry(ctx, input.RunID, session.ID, providertypes.ChatRequest{ Model: cfg.CurrentModel, SystemPrompt: builtContext.SystemPrompt, Messages: builtContext.Messages, @@ -273,7 +274,7 @@ func (s *Service) Run(ctx context.Context, input UserInput) error { return s.handleRunError(ctx, input.RunID, session.ID, err) } if strings.TrimSpace(assistant.Role) == "" { - assistant.Role = provider.RoleAssistant + assistant.Role = providertypes.RoleAssistant } if strings.TrimSpace(assistant.Content) != "" || len(assistant.ToolCalls) > 0 { @@ -326,8 +327,8 @@ func (s *Service) Run(ctx context.Context, input UserInput) error { s.emit(ctx, EventPermissionResolved, input.RunID, session.ID, permissionEvent.toResolvedPayload()) } - toolMessage := provider.Message{ - Role: provider.RoleTool, + toolMessage := providertypes.Message{ + Role: providertypes.RoleTool, Content: result.Content, ToolCallID: call.ID, IsError: result.IsError, @@ -496,13 +497,13 @@ func (s *Service) emit(ctx context.Context, kind EventType, runID string, sessio // handleProviderStreamEvent 解析并应用单条 provider 流式事件,缺失载荷或未知类型时返回错误。 func handleProviderStreamEvent( - event provider.StreamEvent, + event providertypes.StreamEvent, acc *streamAccumulator, onTextDelta func(string), - onToolCallStart func(provider.ToolCallStartPayload), + onToolCallStart func(providertypes.ToolCallStartPayload), ) error { switch event.Type { - case provider.StreamEventTextDelta: + case providertypes.StreamEventTextDelta: payload, err := event.TextDeltaValue() if err != nil { return err @@ -513,7 +514,7 @@ func handleProviderStreamEvent( if acc != nil { acc.accumulateTextDelta(payload.Text) } - case provider.StreamEventToolCallStart: + case providertypes.StreamEventToolCallStart: payload, err := event.ToolCallStartValue() if err != nil { return err @@ -524,7 +525,7 @@ func handleProviderStreamEvent( if acc != nil { acc.accumulateToolCallStart(payload.Index, payload.ID, payload.Name) } - case provider.StreamEventToolCallDelta: + case providertypes.StreamEventToolCallDelta: payload, err := event.ToolCallDeltaValue() if err != nil { return err @@ -532,7 +533,7 @@ func handleProviderStreamEvent( if acc != nil { acc.accumulateToolCallDelta(payload.Index, payload.ID, payload.ArgumentsDelta) } - case provider.StreamEventMessageDone: + case providertypes.StreamEventMessageDone: if _, err := event.MessageDoneValue(); err != nil { return err } @@ -548,7 +549,7 @@ func (s *Service) forwardProviderEvents( ctx context.Context, runID string, sessionID string, - input <-chan provider.StreamEvent, + input <-chan providertypes.StreamEvent, done chan<- error, acc *streamAccumulator, ) { @@ -569,7 +570,7 @@ func (s *Service) forwardProviderEvents( func(text string) { s.emit(ctx, EventAgentChunk, runID, sessionID, text) }, - func(payload provider.ToolCallStartPayload) { + func(payload providertypes.ToolCallStartPayload) { s.emit(ctx, EventToolCallThinking, runID, sessionID, payload.Name) }, ) @@ -640,7 +641,7 @@ func (s *Service) callProviderWithRetry( ctx context.Context, runID string, sessionID string, - req provider.ChatRequest, + req providertypes.ChatRequest, ) (*streamAccumulator, error) { acc := newStreamAccumulator() var lastErr error @@ -671,7 +672,7 @@ func (s *Service) callProviderWithRetry( return nil, err } - streamEvents := make(chan provider.StreamEvent, 32) + streamEvents := make(chan providertypes.StreamEvent, 32) streamDone := make(chan error, 1) go s.forwardProviderEvents(ctx, runID, sessionID, streamEvents, streamDone, acc) diff --git a/internal/runtime/runtime_test.go b/internal/runtime/runtime_test.go index 64f32bb5..0fe9f86a 100644 --- a/internal/runtime/runtime_test.go +++ b/internal/runtime/runtime_test.go @@ -14,6 +14,7 @@ import ( agentcontext "neo-code/internal/context" contextcompact "neo-code/internal/context/compact" "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" "neo-code/internal/security" "neo-code/internal/tools" ) @@ -90,19 +91,19 @@ func (s *memoryStore) ListSummaries(ctx context.Context) ([]SessionSummary, erro type scriptedProvider struct { name string - streams [][]provider.StreamEvent + streams [][]providertypes.StreamEvent responses []scriptedResponse - requests []provider.ChatRequest + requests []providertypes.ChatRequest callCount int - chatFn func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error + chatFn func(ctx context.Context, req providertypes.ChatRequest, events chan<- providertypes.StreamEvent) error } type scriptedResponse struct { - Message provider.Message + Message providertypes.Message FinishReason string } -func (p *scriptedProvider) Chat(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { +func (p *scriptedProvider) Chat(ctx context.Context, req providertypes.ChatRequest, events chan<- providertypes.StreamEvent) error { p.requests = append(p.requests, cloneChatRequest(req)) callIndex := p.callCount @@ -125,25 +126,25 @@ func (p *scriptedProvider) Chat(ctx context.Context, req provider.ChatRequest, e response := p.responses[callIndex] for index, toolCall := range response.Message.ToolCalls { select { - case events <- provider.NewToolCallStartStreamEvent(index, toolCall.ID, toolCall.Name): + case events <- providertypes.NewToolCallStartStreamEvent(index, toolCall.ID, toolCall.Name): case <-ctx.Done(): return ctx.Err() } select { - case events <- provider.NewToolCallDeltaStreamEvent(index, toolCall.ID, toolCall.Arguments): + case events <- providertypes.NewToolCallDeltaStreamEvent(index, toolCall.ID, toolCall.Arguments): case <-ctx.Done(): return ctx.Err() } } if response.Message.Content != "" { select { - case events <- provider.NewTextDeltaStreamEvent(response.Message.Content): + case events <- providertypes.NewTextDeltaStreamEvent(response.Message.Content): case <-ctx.Done(): return ctx.Err() } } select { - case events <- provider.NewMessageDoneStreamEvent(response.FinishReason, nil): + case events <- providertypes.NewMessageDoneStreamEvent(response.FinishReason, nil): case <-ctx.Done(): return ctx.Err() } @@ -227,12 +228,12 @@ func (b *stubContextBuilder) Build(ctx context.Context, input agentcontext.Build } return agentcontext.BuildResult{ SystemPrompt: "stub system prompt", - Messages: append([]provider.Message(nil), input.Messages...), + Messages: append([]providertypes.Message(nil), input.Messages...), }, nil } type stubToolManager struct { - specs []provider.ToolSpec + specs []providertypes.ToolSpec result tools.ToolResult err error listErr error @@ -248,7 +249,7 @@ type stubToolManager struct { } } -func (m *stubToolManager) ListAvailableSpecs(ctx context.Context, input tools.SpecListInput) ([]provider.ToolSpec, error) { +func (m *stubToolManager) ListAvailableSpecs(ctx context.Context, input tools.SpecListInput) ([]providertypes.ToolSpec, error) { m.listCalls++ if err := ctx.Err(); err != nil { return nil, err @@ -256,7 +257,7 @@ func (m *stubToolManager) ListAvailableSpecs(ctx context.Context, input tools.Sp if m.listErr != nil { return nil, m.listErr } - return append([]provider.ToolSpec(nil), m.specs...), nil + return append([]providertypes.ToolSpec(nil), m.specs...), nil } func (m *stubToolManager) MicroCompactPolicy(name string) tools.MicroCompactPolicy { @@ -293,7 +294,7 @@ func TestServiceRun(t *testing.T) { tests := []struct { name string input UserInput - providerStreams [][]provider.StreamEvent + providerStreams [][]providertypes.StreamEvent registerTool tools.Tool contextBuilder agentcontext.Builder expectProviderCalls int @@ -305,17 +306,17 @@ func TestServiceRun(t *testing.T) { { name: "normal dialogue exits after final assistant reply", input: UserInput{RunID: "run-normal", Content: "hello"}, - providerStreams: [][]provider.StreamEvent{ + providerStreams: [][]providertypes.StreamEvent{ { - provider.NewTextDeltaStreamEvent("plain "), - provider.NewTextDeltaStreamEvent("answer"), + providertypes.NewTextDeltaStreamEvent("plain "), + providertypes.NewTextDeltaStreamEvent("answer"), }, }, contextBuilder: &stubContextBuilder{ buildFn: func(ctx context.Context, input agentcontext.BuildInput) (agentcontext.BuildResult, error) { return agentcontext.BuildResult{ SystemPrompt: "custom system prompt", - Messages: []provider.Message{ + Messages: []providertypes.Message{ {Role: "user", Content: "trimmed history"}, }, }, nil @@ -346,13 +347,13 @@ func TestServiceRun(t *testing.T) { input: UserInput{RunID: "run-tool", Content: "edit file"}, // 第一轮:工具调用事件流(tool_call_start + tool_call_delta) // 第二轮:普通文本回复 - providerStreams: [][]provider.StreamEvent{ + providerStreams: [][]providertypes.StreamEvent{ { - provider.NewToolCallStartStreamEvent(0, "call-1", "filesystem_edit"), - provider.NewToolCallDeltaStreamEvent(0, "call-1", `{"path":"main.go"}`), + providertypes.NewToolCallStartStreamEvent(0, "call-1", "filesystem_edit"), + providertypes.NewToolCallDeltaStreamEvent(0, "call-1", `{"path":"main.go"}`), }, { - provider.NewTextDeltaStreamEvent("done"), + providertypes.NewTextDeltaStreamEvent("done"), }, }, registerTool: &stubTool{ @@ -458,13 +459,13 @@ func TestServiceRunMergesLateToolCallMetadata(t *testing.T) { registry.Register(tool) scripted := &scriptedProvider{ - streams: [][]provider.StreamEvent{ + streams: [][]providertypes.StreamEvent{ { - provider.NewToolCallDeltaStreamEvent(0, "", `{"path":"main.go"`), - provider.NewToolCallStartStreamEvent(0, "call-late", "filesystem_edit"), - provider.NewToolCallDeltaStreamEvent(0, "call-late", `}`), + providertypes.NewToolCallDeltaStreamEvent(0, "", `{"path":"main.go"`), + providertypes.NewToolCallStartStreamEvent(0, "call-late", "filesystem_edit"), + providertypes.NewToolCallDeltaStreamEvent(0, "call-late", `}`), }, - {provider.NewTextDeltaStreamEvent("done")}, + {providertypes.NewTextDeltaStreamEvent("done")}, }, } @@ -511,10 +512,10 @@ func TestServiceRunRejectsToolCallWithoutID(t *testing.T) { registry.Register(tool) scripted := &scriptedProvider{ - streams: [][]provider.StreamEvent{ + streams: [][]providertypes.StreamEvent{ { - provider.NewToolCallStartStreamEvent(0, "", "filesystem_edit"), - provider.NewToolCallDeltaStreamEvent(0, "", `{}`), + providertypes.NewToolCallStartStreamEvent(0, "", "filesystem_edit"), + providertypes.NewToolCallDeltaStreamEvent(0, "", `{}`), }, }, } @@ -538,9 +539,9 @@ func TestServiceRunRejectsMalformedProviderStreamEvent(t *testing.T) { registry.Register(&stubTool{name: "filesystem_read_file", content: "default"}) scripted := &scriptedProvider{ - streams: [][]provider.StreamEvent{ + streams: [][]providertypes.StreamEvent{ { - {Type: provider.StreamEventTextDelta}, + {Type: providertypes.StreamEventTextDelta}, }, }, } @@ -560,12 +561,12 @@ func TestServiceRunMalformedProviderStreamEventDoesNotDeadlock(t *testing.T) { registry := tools.NewRegistry() registry.Register(&stubTool{name: "filesystem_read_file", content: "default"}) - stream := []provider.StreamEvent{{Type: provider.StreamEventTextDelta}} + stream := []providertypes.StreamEvent{{Type: providertypes.StreamEventTextDelta}} for i := 0; i < 40; i++ { - stream = append(stream, provider.NewTextDeltaStreamEvent("ignored")) + stream = append(stream, providertypes.NewTextDeltaStreamEvent("ignored")) } scripted := &scriptedProvider{ - streams: [][]provider.StreamEvent{stream}, + streams: [][]providertypes.StreamEvent{stream}, } service := NewWithFactory(manager, registry, store, &scriptedProviderFactory{provider: scripted}, nil) @@ -593,7 +594,7 @@ type stubCompactRunner struct { func (r *stubCompactRunner) Run(ctx context.Context, input contextcompact.Input) (contextcompact.Result, error) { cloned := input - cloned.Messages = append([]provider.Message(nil), input.Messages...) + cloned.Messages = append([]providertypes.Message(nil), input.Messages...) r.calls = append(r.calls, cloned) if r.runFn != nil { return r.runFn(ctx, input) @@ -616,7 +617,7 @@ func TestServiceRunDelegatesToContextBuilder(t *testing.T) { buildFn: func(ctx context.Context, input agentcontext.BuildInput) (agentcontext.BuildResult, error) { return agentcontext.BuildResult{ SystemPrompt: "delegated prompt", - Messages: []provider.Message{ + Messages: []providertypes.Message{ {Role: "user", Content: "delegated message"}, }, }, nil @@ -624,8 +625,8 @@ func TestServiceRunDelegatesToContextBuilder(t *testing.T) { } scripted := &scriptedProvider{ - streams: [][]provider.StreamEvent{ - {provider.NewTextDeltaStreamEvent("done")}, + streams: [][]providertypes.StreamEvent{ + {providertypes.NewTextDeltaStreamEvent("done")}, }, } @@ -686,14 +687,14 @@ func TestServiceRunCanDisableMicroCompactViaConfig(t *testing.T) { buildFn: func(ctx context.Context, input agentcontext.BuildInput) (agentcontext.BuildResult, error) { return agentcontext.BuildResult{ SystemPrompt: "delegated prompt", - Messages: append([]provider.Message(nil), input.Messages...), + Messages: append([]providertypes.Message(nil), input.Messages...), }, nil }, } scripted := &scriptedProvider{ responses: []scriptedResponse{{ - Message: provider.Message{Role: provider.RoleAssistant, Content: "done"}, + Message: providertypes.Message{Role: providertypes.RoleAssistant, Content: "done"}, FinishReason: "stop", }}, } @@ -717,8 +718,8 @@ func TestServiceRunPersistsSessionProviderAndModel(t *testing.T) { registry.Register(&stubTool{name: "filesystem_read_file", content: "default"}) scripted := &scriptedProvider{ - streams: [][]provider.StreamEvent{ - {provider.NewTextDeltaStreamEvent("done")}, + streams: [][]providertypes.StreamEvent{ + {providertypes.NewTextDeltaStreamEvent("done")}, }, } @@ -749,35 +750,35 @@ func TestServiceRunDefaultBuilderUsesToolManagerMicroCompactPolicies(t *testing. session := newSession("preserve history") session.ID = "session-preserve-history" - session.Messages = []provider.Message{ - {Role: provider.RoleUser, Content: "older user"}, + session.Messages = []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older user"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-1", Name: "preserve_tool", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-1", Content: "preserved result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "preserved result"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-2", Name: "bash", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-3", Name: "webfetch", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, } store.sessions[session.ID] = cloneSession(session) scripted := &scriptedProvider{ responses: []scriptedResponse{{ - Message: provider.Message{Role: provider.RoleAssistant, Content: "done"}, + Message: providertypes.Message{Role: providertypes.RoleAssistant, Content: "done"}, FinishReason: "stop", }}, } @@ -812,35 +813,35 @@ func TestServiceRunDefaultBuilderUsesGenericToolManagerMicroCompactPolicies(t *t session := newSession("preserve history by manager") session.ID = "session-preserve-history-manager" - session.Messages = []provider.Message{ - {Role: provider.RoleUser, Content: "older user"}, + session.Messages = []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older user"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-1", Name: "preserve_tool", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-1", Content: "preserved result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "preserved result"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-2", Name: "bash", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, { - Role: provider.RoleAssistant, - ToolCalls: []provider.ToolCall{ + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ {ID: "call-3", Name: "webfetch", Arguments: "{}"}, }, }, - {Role: provider.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, } store.sessions[session.ID] = cloneSession(session) scripted := &scriptedProvider{ responses: []scriptedResponse{{ - Message: provider.Message{Role: provider.RoleAssistant, Content: "done"}, + Message: providertypes.Message{Role: providertypes.RoleAssistant, Content: "done"}, FinishReason: "stop", }}, } @@ -883,8 +884,8 @@ func TestServiceRunFailurePreservesExistingSessionProviderAndModel(t *testing.T) session.ID = "session-preserve-metadata" session.Provider = config.OpenAIName session.Model = "openai-original-model" - session.Messages = []provider.Message{ - {Role: provider.RoleUser, Content: "earlier"}, + session.Messages = []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "earlier"}, } store.sessions[session.ID] = cloneSession(session) @@ -924,7 +925,7 @@ func TestServiceRunUsesToolManager(t *testing.T) { manager := newRuntimeConfigManager(t) store := newMemoryStore() toolManager := &stubToolManager{ - specs: []provider.ToolSpec{ + specs: []providertypes.ToolSpec{ {Name: "filesystem_edit", Description: "stub", Schema: map[string]any{"type": "object"}}, }, result: tools.ToolResult{ @@ -934,12 +935,12 @@ func TestServiceRunUsesToolManager(t *testing.T) { } scripted := &scriptedProvider{ - streams: [][]provider.StreamEvent{ + streams: [][]providertypes.StreamEvent{ { - provider.NewToolCallStartStreamEvent(0, "call-manager", "filesystem_edit"), - provider.NewToolCallDeltaStreamEvent(0, "call-manager", `{"path":"main.go"}`), + providertypes.NewToolCallStartStreamEvent(0, "call-manager", "filesystem_edit"), + providertypes.NewToolCallDeltaStreamEvent(0, "call-manager", `{"path":"main.go"}`), }, - {provider.NewTextDeltaStreamEvent("done")}, + {providertypes.NewTextDeltaStreamEvent("done")}, }, } @@ -964,7 +965,7 @@ func TestServiceRunUsesToolManager(t *testing.T) { session := onlySession(t, store) foundToolMessage := false for _, message := range session.Messages { - if message.Role == provider.RoleTool && message.Content == "tool manager output" { + if message.Role == providertypes.RoleTool && message.Content == "tool manager output" { foundToolMessage = true break } @@ -1004,12 +1005,12 @@ func TestServiceRunWaitsForPermissionResolutionAndContinues(t *testing.T) { } scripted := &scriptedProvider{ - streams: [][]provider.StreamEvent{ + streams: [][]providertypes.StreamEvent{ { - provider.NewToolCallStartStreamEvent(0, "call-ask", "webfetch"), - provider.NewToolCallDeltaStreamEvent(0, "call-ask", `{"url":"https://example.com/private"}`), + providertypes.NewToolCallStartStreamEvent(0, "call-ask", "webfetch"), + providertypes.NewToolCallDeltaStreamEvent(0, "call-ask", `{"url":"https://example.com/private"}`), }, - {provider.NewTextDeltaStreamEvent("done")}, + {providertypes.NewTextDeltaStreamEvent("done")}, }, } @@ -1120,12 +1121,12 @@ func TestServiceRunEmitsPermissionResolvedForDeny(t *testing.T) { } scripted := &scriptedProvider{ - streams: [][]provider.StreamEvent{ + streams: [][]providertypes.StreamEvent{ { - provider.NewToolCallStartStreamEvent(0, "call-deny", "bash"), - provider.NewToolCallDeltaStreamEvent(0, "call-deny", `{"command":"echo hi"}`), + providertypes.NewToolCallStartStreamEvent(0, "call-deny", "bash"), + providertypes.NewToolCallDeltaStreamEvent(0, "call-deny", `{"command":"echo hi"}`), }, - {provider.NewTextDeltaStreamEvent("done")}, + {providertypes.NewTextDeltaStreamEvent("done")}, }, } @@ -1207,13 +1208,19 @@ func TestServiceRunEmitsRememberScopeWhenSessionRejectMemoryHits(t *testing.T) { } scripted := &scriptedProvider{ - streams: [][]provider.StreamEvent{ + responses: []scriptedResponse{ { - provider.NewToolCallStartStreamEvent(0, "call-memory-reject", "webfetch"), - provider.NewToolCallDeltaStreamEvent(0, "call-memory-reject", `{"url":"https://example.com/private"}`), + Message: providertypes.Message{ + Role: "assistant", + ToolCalls: []providertypes.ToolCall{ + {ID: "call-memory-reject", Name: "webfetch", Arguments: `{"url":"https://example.com/private"}`}, + }, + }, + FinishReason: "tool_calls", }, { - provider.NewTextDeltaStreamEvent("done"), + Message: providertypes.Message{Role: "assistant", Content: "done"}, + FinishReason: "stop", }, }, } @@ -1272,7 +1279,7 @@ func TestServiceRunHandlesToolManagerSpecError(t *testing.T) { assertEventsRunID(t, events, input.RunID) session := onlySession(t, store) - if len(session.Messages) != 1 || session.Messages[0].Role != provider.RoleUser { + if len(session.Messages) != 1 || session.Messages[0].Role != providertypes.RoleUser { t.Fatalf("expected only user message to persist, got %+v", session.Messages) } } @@ -1282,8 +1289,8 @@ func TestServiceNewWithFactoryDefaultsToolManager(t *testing.T) { store := newMemoryStore() service := NewWithFactory(manager, nil, store, &scriptedProviderFactory{ provider: &scriptedProvider{ - streams: [][]provider.StreamEvent{ - {provider.NewTextDeltaStreamEvent("done")}, + streams: [][]providertypes.StreamEvent{ + {providertypes.NewTextDeltaStreamEvent("done")}, }, }, }, nil) @@ -1325,10 +1332,10 @@ func TestServiceRunErrorPaths(t *testing.T) { input: UserInput{RunID: "run-max-loops", Content: "loop"}, maxLoops: 1, provider: &scriptedProvider{ - streams: [][]provider.StreamEvent{ + streams: [][]providertypes.StreamEvent{ { - provider.NewToolCallStartStreamEvent(0, "loop-call", "filesystem_edit"), - provider.NewToolCallDeltaStreamEvent(0, "loop-call", `{"path":"x"}`), + providertypes.NewToolCallStartStreamEvent(0, "loop-call", "filesystem_edit"), + providertypes.NewToolCallDeltaStreamEvent(0, "loop-call", `{"path":"x"}`), }, }, }, @@ -1364,8 +1371,8 @@ func TestServiceRunErrorPaths(t *testing.T) { Content: "continue", }, provider: &scriptedProvider{ - streams: [][]provider.StreamEvent{ - {provider.NewTextDeltaStreamEvent("resumed")}, + streams: [][]providertypes.StreamEvent{ + {providertypes.NewTextDeltaStreamEvent("resumed")}, }, }, seedSession: &Session{ @@ -1373,7 +1380,7 @@ func TestServiceRunErrorPaths(t *testing.T) { Title: "Resume Me", CreatedAt: newSession("seed").CreatedAt, UpdatedAt: newSession("seed").UpdatedAt, - Messages: []provider.Message{ + Messages: []providertypes.Message{ {Role: "user", Content: "earlier"}, }, }, @@ -1396,7 +1403,7 @@ func TestServiceRunErrorPaths(t *testing.T) { callIdx := 0 return &scriptedProvider{ name: "retry-then-success", - chatFn: func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { + chatFn: func(ctx context.Context, req providertypes.ChatRequest, events chan<- providertypes.StreamEvent) error { callIdx++ if callIdx == 1 { return &provider.ProviderError{ @@ -1406,7 +1413,7 @@ func TestServiceRunErrorPaths(t *testing.T) { Retryable: true, } } - events <- provider.NewTextDeltaStreamEvent("recovered") + events <- providertypes.NewTextDeltaStreamEvent("recovered") return nil }, } @@ -1431,7 +1438,7 @@ func TestServiceRunErrorPaths(t *testing.T) { input: UserInput{RunID: "run-no-retry", Content: "hello"}, provider: &scriptedProvider{ name: "auth-error-no-retry", - chatFn: func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { + chatFn: func(ctx context.Context, req providertypes.ChatRequest, events chan<- providertypes.StreamEvent) error { return &provider.ProviderError{ StatusCode: 401, Code: provider.ErrorCodeAuthFailed, @@ -1454,7 +1461,7 @@ func TestServiceRunErrorPaths(t *testing.T) { input: UserInput{RunID: "run-retry-exhausted", Content: "hello"}, provider: &scriptedProvider{ name: "always-500", - chatFn: func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { + chatFn: func(ctx context.Context, req providertypes.ChatRequest, events chan<- providertypes.StreamEvent) error { return &provider.ProviderError{ StatusCode: 500, Code: provider.ErrorCodeServer, @@ -1538,7 +1545,7 @@ func TestServiceCancelActiveRun(t *testing.T) { started := make(chan struct{}) scripted := &scriptedProvider{ name: "cancel-active-run-provider", - chatFn: func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { + chatFn: func(ctx context.Context, req providertypes.ChatRequest, events chan<- providertypes.StreamEvent) error { close(started) <-ctx.Done() return ctx.Err() @@ -1581,7 +1588,7 @@ func TestServiceRunCanceledByProvider(t *testing.T) { started := make(chan struct{}) scripted := &scriptedProvider{ name: "blocking-provider", - chatFn: func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { + chatFn: func(ctx context.Context, req providertypes.ChatRequest, events chan<- providertypes.StreamEvent) error { close(started) <-ctx.Done() return ctx.Err() @@ -1626,7 +1633,7 @@ func TestServiceRunPreservesProviderErrorAfterCancel(t *testing.T) { providerErr := errors.New("provider failed after cancel") scripted := &scriptedProvider{ name: "provider-error-after-cancel", - chatFn: func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { + chatFn: func(ctx context.Context, req providertypes.ChatRequest, events chan<- providertypes.StreamEvent) error { close(started) <-ctx.Done() return providerErr @@ -1678,10 +1685,10 @@ func TestServiceRunCanceledDuringToolExecution(t *testing.T) { scripted := &scriptedProvider{ name: "tool-cancel-provider", - streams: [][]provider.StreamEvent{ + streams: [][]providertypes.StreamEvent{ { - provider.NewToolCallStartStreamEvent(0, "cancel-call", "filesystem_edit"), - provider.NewToolCallDeltaStreamEvent(0, "cancel-call", `{"path":"main.go"}`), + providertypes.NewToolCallStartStreamEvent(0, "cancel-call", "filesystem_edit"), + providertypes.NewToolCallDeltaStreamEvent(0, "cancel-call", `{"path":"main.go"}`), }, }, } @@ -1738,10 +1745,10 @@ func TestServiceRunPreservesToolErrorAfterCancel(t *testing.T) { scripted := &scriptedProvider{ name: "tool-error-after-cancel-provider", - streams: [][]provider.StreamEvent{ + streams: [][]providertypes.StreamEvent{ { - provider.NewToolCallStartStreamEvent(0, "tool-error-call", "filesystem_edit"), - provider.NewToolCallDeltaStreamEvent(0, "tool-error-call", `{"path":"main.go"}`), + providertypes.NewToolCallStartStreamEvent(0, "tool-error-call", "filesystem_edit"), + providertypes.NewToolCallDeltaStreamEvent(0, "tool-error-call", `{"path":"main.go"}`), }, }, } @@ -1830,12 +1837,12 @@ func TestServiceRunToolTimeoutIsNotCancellation(t *testing.T) { scripted := &scriptedProvider{ name: "timeout-provider", - streams: [][]provider.StreamEvent{ + streams: [][]providertypes.StreamEvent{ { - provider.NewToolCallStartStreamEvent(0, "timeout-call", "filesystem_edit"), - provider.NewToolCallDeltaStreamEvent(0, "timeout-call", `{"path":"main.go"}`), + providertypes.NewToolCallStartStreamEvent(0, "timeout-call", "filesystem_edit"), + providertypes.NewToolCallDeltaStreamEvent(0, "timeout-call", `{"path":"main.go"}`), }, - {provider.NewTextDeltaStreamEvent("done after timeout")}, + {providertypes.NewTextDeltaStreamEvent("done after timeout")}, }, } @@ -1864,10 +1871,10 @@ func TestServiceCompactManualAppliesAndPersists(t *testing.T) { store := newMemoryStore() session := newSession("manual") session.ID = "session-manual" - session.Messages = []provider.Message{ - {Role: provider.RoleUser, Content: "older"}, - {Role: provider.RoleAssistant, Content: "older answer"}, - {Role: provider.RoleUser, Content: "before"}, + session.Messages = []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older"}, + {Role: providertypes.RoleAssistant, Content: "older answer"}, + {Role: providertypes.RoleUser, Content: "before"}, } store.sessions[session.ID] = cloneSession(session) @@ -1877,9 +1884,9 @@ func TestServiceCompactManualAppliesAndPersists(t *testing.T) { service := NewWithFactory(manager, registry, store, &scriptedProviderFactory{provider: &scriptedProvider{}}, nil) service.compactRunner = &stubCompactRunner{ result: contextcompact.Result{ - Messages: []provider.Message{ - {Role: provider.RoleAssistant, Content: "[compact_summary]\ndone:\n- ok\n\nin_progress:\n- continue"}, - {Role: provider.RoleAssistant, Content: "latest"}, + Messages: []providertypes.Message{ + {Role: providertypes.RoleAssistant, Content: "[compact_summary]\ndone:\n- ok\n\nin_progress:\n- continue"}, + {Role: providertypes.RoleAssistant, Content: "latest"}, }, Applied: true, Metrics: contextcompact.Metrics{ @@ -1922,10 +1929,10 @@ func TestServiceCompactManualFailureReturnsError(t *testing.T) { store := newMemoryStore() session := newSession("manual-fail") session.ID = "session-manual-fail" - session.Messages = []provider.Message{ - {Role: provider.RoleUser, Content: "older"}, - {Role: provider.RoleAssistant, Content: "older answer"}, - {Role: provider.RoleUser, Content: "before"}, + session.Messages = []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older"}, + {Role: providertypes.RoleAssistant, Content: "older answer"}, + {Role: providertypes.RoleUser, Content: "before"}, } store.sessions[session.ID] = cloneSession(session) @@ -1978,10 +1985,10 @@ func TestServiceCompactUsesSessionProviderAndModelWhenPresent(t *testing.T) { session.ID = "session-manual-provider" session.Provider = config.OpenAIName session.Model = "session-model" - session.Messages = []provider.Message{ - {Role: provider.RoleUser, Content: "older"}, - {Role: provider.RoleAssistant, Content: "older answer"}, - {Role: provider.RoleUser, Content: "before"}, + session.Messages = []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older"}, + {Role: providertypes.RoleAssistant, Content: "older answer"}, + {Role: providertypes.RoleUser, Content: "before"}, } store.sessions[session.ID] = cloneSession(session) @@ -1989,8 +1996,8 @@ func TestServiceCompactUsesSessionProviderAndModelWhenPresent(t *testing.T) { registry.Register(&stubTool{name: "filesystem_read_file", content: "ok"}) scripted := &scriptedProvider{ - streams: [][]provider.StreamEvent{ - {provider.NewTextDeltaStreamEvent(strings.Join([]string{ + streams: [][]providertypes.StreamEvent{ + {providertypes.NewTextDeltaStreamEvent(strings.Join([]string{ "[compact_summary]", "done:", "- ok", @@ -2050,10 +2057,10 @@ func TestServiceCompactFallsBackToCurrentProviderWhenSessionMetadataMissing(t *t store := newMemoryStore() session := newSession("manual-fallback") session.ID = "session-manual-fallback" - session.Messages = []provider.Message{ - {Role: provider.RoleUser, Content: "older"}, - {Role: provider.RoleAssistant, Content: "older answer"}, - {Role: provider.RoleUser, Content: "before"}, + session.Messages = []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older"}, + {Role: providertypes.RoleAssistant, Content: "older answer"}, + {Role: providertypes.RoleUser, Content: "before"}, } store.sessions[session.ID] = cloneSession(session) @@ -2061,8 +2068,8 @@ func TestServiceCompactFallsBackToCurrentProviderWhenSessionMetadataMissing(t *t registry.Register(&stubTool{name: "filesystem_read_file", content: "ok"}) scripted := &scriptedProvider{ - streams: [][]provider.StreamEvent{ - {provider.NewTextDeltaStreamEvent(strings.Join([]string{ + streams: [][]providertypes.StreamEvent{ + {providertypes.NewTextDeltaStreamEvent(strings.Join([]string{ "[compact_summary]", "done:", "- ok", @@ -2107,9 +2114,9 @@ func TestServiceManualCompactThenRunContinuesToolRound(t *testing.T) { store := newMemoryStore() session := newSession("manual-continue") session.ID = "session-manual-continue" - session.Messages = []provider.Message{ - {Role: provider.RoleUser, Content: "legacy request"}, - {Role: provider.RoleAssistant, Content: "legacy answer"}, + session.Messages = []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "legacy request"}, + {Role: providertypes.RoleAssistant, Content: "legacy answer"}, } store.sessions[session.ID] = cloneSession(session) @@ -2118,12 +2125,12 @@ func TestServiceManualCompactThenRunContinuesToolRound(t *testing.T) { registry.Register(tool) scripted := &scriptedProvider{ - streams: [][]provider.StreamEvent{ + streams: [][]providertypes.StreamEvent{ { - provider.NewToolCallStartStreamEvent(0, "call-1", "filesystem_read_file"), - provider.NewToolCallDeltaStreamEvent(0, "call-1", `{"path":"main.go"}`), + providertypes.NewToolCallStartStreamEvent(0, "call-1", "filesystem_read_file"), + providertypes.NewToolCallDeltaStreamEvent(0, "call-1", `{"path":"main.go"}`), }, - {provider.NewTextDeltaStreamEvent("done")}, + {providertypes.NewTextDeltaStreamEvent("done")}, }, } @@ -2131,9 +2138,9 @@ func TestServiceManualCompactThenRunContinuesToolRound(t *testing.T) { service.compactRunner = &stubCompactRunner{ runFn: func(ctx context.Context, input contextcompact.Input) (contextcompact.Result, error) { return contextcompact.Result{ - Messages: []provider.Message{ - {Role: provider.RoleAssistant, Content: "[compact_summary]\ndone:\n- archived\n\nin_progress:\n- continue"}, - {Role: provider.RoleAssistant, Content: "latest answer"}, + Messages: []providertypes.Message{ + {Role: providertypes.RoleAssistant, Content: "[compact_summary]\ndone:\n- archived\n\nin_progress:\n- continue"}, + {Role: providertypes.RoleAssistant, Content: "latest answer"}, }, Applied: true, Metrics: contextcompact.Metrics{ @@ -2198,14 +2205,14 @@ func TestServiceSerializesRunAndCompact(t *testing.T) { providerStarted := make(chan struct{}) unblockProvider := make(chan struct{}) scripted := &scriptedProvider{ - chatFn: func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { + chatFn: func(ctx context.Context, req providertypes.ChatRequest, events chan<- providertypes.StreamEvent) error { select { case <-providerStarted: default: close(providerStarted) } <-unblockProvider - events <- provider.NewTextDeltaStreamEvent("done") + events <- providertypes.NewTextDeltaStreamEvent("done") return nil }, } @@ -2216,7 +2223,7 @@ func TestServiceSerializesRunAndCompact(t *testing.T) { runFn: func(ctx context.Context, input contextcompact.Input) (contextcompact.Result, error) { compactEntered <- struct{}{} return contextcompact.Result{ - Messages: append([]provider.Message(nil), input.Messages...), + Messages: append([]providertypes.Message(nil), input.Messages...), Metrics: contextcompact.Metrics{ BeforeChars: 1, AfterChars: 1, @@ -2330,12 +2337,12 @@ func TestServiceRunUsesSessionWorkdirForContextAndTools(t *testing.T) { builder := &stubContextBuilder{} scripted := &scriptedProvider{ - streams: [][]provider.StreamEvent{ + streams: [][]providertypes.StreamEvent{ { - provider.NewToolCallStartStreamEvent(0, "call-session-workdir", "filesystem_edit"), - provider.NewToolCallDeltaStreamEvent(0, "call-session-workdir", `{"path":"main.go"}`), + providertypes.NewToolCallStartStreamEvent(0, "call-session-workdir", "filesystem_edit"), + providertypes.NewToolCallDeltaStreamEvent(0, "call-session-workdir", `{"path":"main.go"}`), }, - {provider.NewTextDeltaStreamEvent("done")}, + {providertypes.NewTextDeltaStreamEvent("done")}, }, } @@ -2372,8 +2379,8 @@ func TestServiceRunUsesInputWorkdirForNewSession(t *testing.T) { registry.Register(&stubTool{name: "filesystem_read_file", content: "default"}) builder := &stubContextBuilder{} scripted := &scriptedProvider{ - streams: [][]provider.StreamEvent{ - {provider.NewTextDeltaStreamEvent("done")}, + streams: [][]providertypes.StreamEvent{ + {providertypes.NewTextDeltaStreamEvent("done")}, }, } @@ -2608,20 +2615,20 @@ func assertEventsRunID(t *testing.T, events []RuntimeEvent, runID string) { func cloneSession(session Session) Session { cloned := session - cloned.Messages = append([]provider.Message(nil), session.Messages...) + cloned.Messages = append([]providertypes.Message(nil), session.Messages...) return cloned } -func cloneChatRequest(req provider.ChatRequest) provider.ChatRequest { +func cloneChatRequest(req providertypes.ChatRequest) providertypes.ChatRequest { cloned := req - cloned.Messages = append([]provider.Message(nil), req.Messages...) - cloned.Tools = append([]provider.ToolSpec(nil), req.Tools...) + cloned.Messages = append([]providertypes.Message(nil), req.Messages...) + cloned.Tools = append([]providertypes.ToolSpec(nil), req.Tools...) return cloned } func cloneBuildInput(input agentcontext.BuildInput) agentcontext.BuildInput { cloned := input - cloned.Messages = append([]provider.Message(nil), input.Messages...) + cloned.Messages = append([]providertypes.Message(nil), input.Messages...) return cloned } diff --git a/internal/runtime/session.go b/internal/runtime/session.go index af64317d..24378d28 100644 --- a/internal/runtime/session.go +++ b/internal/runtime/session.go @@ -12,7 +12,7 @@ import ( "sync" "time" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" ) const sessionsDirName = "sessions" @@ -23,11 +23,11 @@ type Session struct { // Provider 记录最近一次成功运行会话时使用的 provider,用于 compact 优先复用历史配置。 Provider string `json:"provider,omitempty"` // Model 记录最近一次成功运行会话时使用的 model,用于 compact 优先复用历史配置。 - Model string `json:"model,omitempty"` - CreatedAt time.Time `json:"created_at"` - UpdatedAt time.Time `json:"updated_at"` - Workdir string `json:"-"` - Messages []provider.Message `json:"messages"` + Model string `json:"model,omitempty"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` + Workdir string `json:"-"` + Messages []providertypes.Message `json:"messages"` } type SessionSummary struct { @@ -181,7 +181,7 @@ func newSessionWithWorkdir(title string, workdir string) Session { CreatedAt: now, UpdatedAt: now, Workdir: strings.TrimSpace(workdir), - Messages: []provider.Message{}, + Messages: []providertypes.Message{}, } } diff --git a/internal/runtime/session_test.go b/internal/runtime/session_test.go index a82f1215..06faa257 100644 --- a/internal/runtime/session_test.go +++ b/internal/runtime/session_test.go @@ -8,7 +8,7 @@ import ( "testing" "time" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" ) func TestJSONSessionStoreSaveLoadAndListSummaries(t *testing.T) { @@ -22,7 +22,7 @@ func TestJSONSessionStoreSaveLoadAndListSummaries(t *testing.T) { Title: "Old Session", CreatedAt: time.Now().Add(-2 * time.Hour), UpdatedAt: time.Now().Add(-1 * time.Hour), - Messages: []provider.Message{ + Messages: []providertypes.Message{ {Role: "user", Content: "hello"}, {Role: "assistant", Content: "world"}, }, @@ -33,7 +33,7 @@ func TestJSONSessionStoreSaveLoadAndListSummaries(t *testing.T) { CreatedAt: time.Now().Add(-30 * time.Minute), UpdatedAt: time.Now(), Workdir: t.TempDir(), - Messages: []provider.Message{ + Messages: []providertypes.Message{ {Role: "user", Content: "new"}, }, } @@ -119,7 +119,7 @@ func TestJSONSessionStoreCorruptedSessionBehaviors(t *testing.T) { Title: "Valid Session", CreatedAt: time.Now().Add(-time.Minute), UpdatedAt: time.Now(), - Messages: []provider.Message{{Role: "user", Content: "hello"}}, + Messages: []providertypes.Message{{Role: "user", Content: "hello"}}, } if err := store.Save(context.Background(), valid); err != nil { t.Fatalf("Save valid session: %v", err) diff --git a/internal/tools/manager.go b/internal/tools/manager.go index b91af823..dab53833 100644 --- a/internal/tools/manager.go +++ b/internal/tools/manager.go @@ -6,7 +6,7 @@ import ( "fmt" "strings" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" "neo-code/internal/security" ) @@ -18,7 +18,7 @@ type SpecListInput struct { // Manager is the runtime-facing tool execution and schema exposure boundary. type Manager interface { - ListAvailableSpecs(ctx context.Context, input SpecListInput) ([]provider.ToolSpec, error) + ListAvailableSpecs(ctx context.Context, input SpecListInput) ([]providertypes.ToolSpec, error) MicroCompactPolicy(name string) MicroCompactPolicy Execute(ctx context.Context, input ToolCallInput) (ToolResult, error) RememberSessionDecision(sessionID string, action security.Action, scope SessionPermissionScope) error @@ -26,7 +26,7 @@ type Manager interface { // Executor is the concrete tool execution layer under the manager. type Executor interface { - ListAvailableSpecs(ctx context.Context, input SpecListInput) ([]provider.ToolSpec, error) + ListAvailableSpecs(ctx context.Context, input SpecListInput) ([]providertypes.ToolSpec, error) Execute(ctx context.Context, input ToolCallInput) (ToolResult, error) Supports(name string) bool } @@ -161,7 +161,7 @@ func NewManager(executor Executor, engine security.PermissionEngine, sandbox Wor } // ListAvailableSpecs returns the currently visible tool specs from the executor. -func (m *DefaultManager) ListAvailableSpecs(ctx context.Context, input SpecListInput) ([]provider.ToolSpec, error) { +func (m *DefaultManager) ListAvailableSpecs(ctx context.Context, input SpecListInput) ([]providertypes.ToolSpec, error) { if m == nil || m.executor == nil { return nil, errors.New("tools: manager executor is nil") } diff --git a/internal/tools/manager_test.go b/internal/tools/manager_test.go index 96d1b75a..45e7de66 100644 --- a/internal/tools/manager_test.go +++ b/internal/tools/manager_test.go @@ -8,7 +8,7 @@ import ( "strings" "testing" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" "neo-code/internal/security" ) @@ -47,7 +47,7 @@ type stubSandbox struct { type executorWithoutMicroCompactPolicy struct{} -func (executorWithoutMicroCompactPolicy) ListAvailableSpecs(ctx context.Context, input SpecListInput) ([]provider.ToolSpec, error) { +func (executorWithoutMicroCompactPolicy) ListAvailableSpecs(ctx context.Context, input SpecListInput) ([]providertypes.ToolSpec, error) { if err := ctx.Err(); err != nil { return nil, err } diff --git a/internal/tools/registry.go b/internal/tools/registry.go index 4ebc8234..90f8a485 100644 --- a/internal/tools/registry.go +++ b/internal/tools/registry.go @@ -6,7 +6,7 @@ import ( "sort" "strings" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" "neo-code/internal/security" ) @@ -65,17 +65,17 @@ func (r *Registry) MicroCompactPolicy(name string) MicroCompactPolicy { return MicroCompactPolicyCompact } -func (r *Registry) GetSpecs() []provider.ToolSpec { +func (r *Registry) GetSpecs() []providertypes.ToolSpec { names := make([]string, 0, len(r.tools)) for name := range r.tools { names = append(names, name) } sort.Strings(names) - specs := make([]provider.ToolSpec, 0, len(names)) + specs := make([]providertypes.ToolSpec, 0, len(names)) for _, name := range names { tool := r.tools[name] - specs = append(specs, provider.ToolSpec{ + specs = append(specs, providertypes.ToolSpec{ Name: tool.Name(), Description: tool.Description(), Schema: tool.Schema(), @@ -84,12 +84,12 @@ func (r *Registry) GetSpecs() []provider.ToolSpec { return specs } -func (r *Registry) ListSchemas() []provider.ToolSpec { +func (r *Registry) ListSchemas() []providertypes.ToolSpec { return r.GetSpecs() } // ListAvailableSpecs returns all registered tool specs. -func (r *Registry) ListAvailableSpecs(ctx context.Context, input SpecListInput) ([]provider.ToolSpec, error) { +func (r *Registry) ListAvailableSpecs(ctx context.Context, input SpecListInput) ([]providertypes.ToolSpec, error) { if err := ctx.Err(); err != nil { return nil, err } diff --git a/internal/tools/types.go b/internal/tools/types.go index 4a0a0a32..77bd75bb 100644 --- a/internal/tools/types.go +++ b/internal/tools/types.go @@ -3,7 +3,7 @@ package tools import ( "context" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" "neo-code/internal/security" ) @@ -35,4 +35,4 @@ type ToolResult struct { Metadata map[string]any } -type ToolSpec = provider.ToolSpec +type ToolSpec = providertypes.ToolSpec diff --git a/internal/tui/app.go b/internal/tui/app.go index eae67f5d..19816f03 100644 --- a/internal/tui/app.go +++ b/internal/tui/app.go @@ -15,7 +15,7 @@ import ( "github.com/charmbracelet/lipgloss" "neo-code/internal/config" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" agentruntime "neo-code/internal/runtime" ) @@ -47,7 +47,7 @@ type App struct { inputBurstCount int pasteMode bool pendingPermission *pendingPermissionPrompt - activeMessages []provider.Message + activeMessages []providertypes.Message activities []activityEntry fileCandidates []string modelRefreshID string diff --git a/internal/tui/copy_code_test.go b/internal/tui/copy_code_test.go index ae9e4305..9cadcee0 100644 --- a/internal/tui/copy_code_test.go +++ b/internal/tui/copy_code_test.go @@ -7,7 +7,7 @@ import ( tea "github.com/charmbracelet/bubbletea" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" ) func TestExtractFencedCodeBlocks(t *testing.T) { @@ -136,7 +136,7 @@ func TestTranscriptMouseClickCopiesCodeBlock(t *testing.T) { app.width = 128 app.height = 40 - app.activeMessages = []provider.Message{ + app.activeMessages = []providertypes.Message{ {Role: roleAssistant, Content: "```go\nfmt.Println(1)\n```"}, } app.applyComponentLayout(true) @@ -223,7 +223,7 @@ func TestTranscriptMouseCopyFailureSetsError(t *testing.T) { app.width = 128 app.height = 40 - app.activeMessages = []provider.Message{ + app.activeMessages = []providertypes.Message{ {Role: roleAssistant, Content: "```txt\nhello\n```"}, } app.applyComponentLayout(true) @@ -266,6 +266,6 @@ func TestTranscriptMouseCopyFailureSetsError(t *testing.T) { } } -func providerMessage(role, content string) provider.Message { - return provider.Message{Role: role, Content: content} +func providerMessage(role, content string) providertypes.Message { + return providertypes.Message{Role: role, Content: content} } diff --git a/internal/tui/update.go b/internal/tui/update.go index 4836fd28..8a2a513c 100644 --- a/internal/tui/update.go +++ b/internal/tui/update.go @@ -16,7 +16,7 @@ import ( "github.com/charmbracelet/lipgloss" "neo-code/internal/config" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" agentruntime "neo-code/internal/runtime" "neo-code/internal/tools" ) @@ -426,7 +426,7 @@ func (a App) updateInputPanel(msg tea.Msg, typed tea.KeyMsg, cmds []tea.Cmd) (te a.state.ExecutionError = "" a.state.StatusText = statusThinking a.state.CurrentTool = "" - a.activeMessages = append(a.activeMessages, provider.Message{Role: roleUser, Content: input}) + a.activeMessages = append(a.activeMessages, providertypes.Message{Role: roleUser, Content: input}) a.rebuildTranscript() requestedWorkdir := "" if strings.TrimSpace(a.state.ActiveSessionID) == "" { @@ -711,7 +711,7 @@ func (a *App) handleRuntimeEvent(event agentruntime.RuntimeEvent) bool { case agentruntime.EventToolStart: a.state.StatusText = statusRunningTool a.state.StreamingReply = false - if payload, ok := event.Payload.(provider.ToolCall); ok { + if payload, ok := event.Payload.(providertypes.ToolCall); ok { a.state.CurrentTool = payload.Name a.setRunProgress(0.6, "Running tool") a.appendActivity("tool", "Running tool", payload.Name, false) @@ -721,7 +721,7 @@ func (a *App) handleRuntimeEvent(event agentruntime.RuntimeEvent) bool { a.state.CurrentTool = "" a.setRunProgress(0.8, "Integrating result") if payload, ok := event.Payload.(tools.ToolResult); ok { - a.activeMessages = append(a.activeMessages, provider.Message{ + a.activeMessages = append(a.activeMessages, providertypes.Message{ Role: roleTool, Content: payload.Content, IsError: payload.IsError, @@ -758,8 +758,8 @@ func (a *App) handleRuntimeEvent(event agentruntime.RuntimeEvent) bool { if strings.TrimSpace(a.state.ExecutionError) == "" { a.state.StatusText = statusReady } - if payload, ok := event.Payload.(provider.Message); ok && strings.TrimSpace(payload.Content) != "" && !a.lastAssistantMatches(payload.Content) { - a.activeMessages = append(a.activeMessages, provider.Message{Role: roleAssistant, Content: payload.Content}) + if payload, ok := event.Payload.(providertypes.Message); ok && strings.TrimSpace(payload.Content) != "" && !a.lastAssistantMatches(payload.Content) { + a.activeMessages = append(a.activeMessages, providertypes.Message{Role: roleAssistant, Content: payload.Content}) transcriptDirty = true } case agentruntime.EventRunCanceled: @@ -871,7 +871,7 @@ func (a *App) appendAssistantChunk(chunk string) { } if !a.state.StreamingReply || len(a.activeMessages) == 0 || a.activeMessages[len(a.activeMessages)-1].Role != roleAssistant { - a.activeMessages = append(a.activeMessages, provider.Message{Role: roleAssistant, Content: chunk}) + a.activeMessages = append(a.activeMessages, providertypes.Message{Role: roleAssistant, Content: chunk}) a.state.StreamingReply = true return } @@ -885,7 +885,7 @@ func (a *App) appendInlineMessage(role string, message string) { return } - a.activeMessages = append(a.activeMessages, provider.Message{Role: role, Content: content}) + a.activeMessages = append(a.activeMessages, providertypes.Message{Role: role, Content: content}) } func (a *App) appendActivity(kind string, title string, detail string, isError bool) { diff --git a/internal/tui/update_test.go b/internal/tui/update_test.go index 6a741a06..a02958f7 100644 --- a/internal/tui/update_test.go +++ b/internal/tui/update_test.go @@ -20,6 +20,7 @@ import ( contextcompact "neo-code/internal/context/compact" "neo-code/internal/provider" providercatalog "neo-code/internal/provider/catalog" + providertypes "neo-code/internal/provider/types" agentruntime "neo-code/internal/runtime" "neo-code/internal/tools" ) @@ -675,7 +676,7 @@ func TestAppUpdateModelPickerAndRuntimeMessages(t *testing.T) { msg: RuntimeMsg{Event: agentruntime.RuntimeEvent{ Type: agentruntime.EventAgentDone, SessionID: "session-2", - Payload: provider.Message{ + Payload: providertypes.Message{ Role: roleAssistant, Content: "final", }, @@ -778,7 +779,7 @@ func TestAppHelpersAndRenderingSmoke(t *testing.T) { now := agentruntime.Session{ ID: "session-1", Title: "Existing Session", - Messages: []provider.Message{ + Messages: []providertypes.Message{ {Role: roleUser, Content: "hi"}, {Role: roleAssistant, Content: "hello"}, }, @@ -932,12 +933,12 @@ func TestAppHelpersAndRenderingSmoke(t *testing.T) { if app.statusBadge("error: boom") == "" || app.statusBadge("running now") == "" { t.Fatalf("expected status badge variants") } - if rendered, _ := app.renderMessageBlockWithCopy(provider.Message{Role: roleError, Content: "boom"}, 80, 1); rendered == "" { + if rendered, _ := app.renderMessageBlockWithCopy(providertypes.Message{Role: roleError, Content: "boom"}, 80, 1); rendered == "" { t.Fatalf("expected error message block") } - if rendered, _ := app.renderMessageBlockWithCopy(provider.Message{ + if rendered, _ := app.renderMessageBlockWithCopy(providertypes.Message{ Role: roleAssistant, - ToolCalls: []provider.ToolCall{ + ToolCalls: []providertypes.ToolCall{ {Name: "filesystem_edit"}, }, }, 80, 1); rendered == "" { @@ -1194,7 +1195,7 @@ func TestAppUpdateAdditionalTransitions(t *testing.T) { setup: func(t *testing.T, app *App, runtime *stubRuntime, manager *config.Manager) { app.state.ActiveSessionID = "existing" app.state.ActiveSessionTitle = "Existing" - app.activeMessages = []provider.Message{{Role: roleUser, Content: "hello"}} + app.activeMessages = []providertypes.Message{{Role: roleUser, Content: "hello"}} }, msg: tea.KeyMsg{Type: tea.KeyCtrlN}, assert: func(t *testing.T, app App, runtime *stubRuntime, manager *config.Manager, msgs []tea.Msg) { @@ -1211,7 +1212,7 @@ func TestAppUpdateAdditionalTransitions(t *testing.T) { runtime.loads["s1"] = agentruntime.Session{ ID: "s1", Title: "One", - Messages: []provider.Message{{Role: roleAssistant, Content: "loaded"}}, + Messages: []providertypes.Message{{Role: roleAssistant, Content: "loaded"}}, } if err := app.refreshSessions(); err != nil { t.Fatalf("refresh sessions: %v", err) @@ -1786,7 +1787,7 @@ func TestAppHandleRuntimeEventAdditionalBranches(t *testing.T) { event: agentruntime.RuntimeEvent{ Type: agentruntime.EventToolStart, SessionID: "s1", - Payload: provider.ToolCall{ + Payload: providertypes.ToolCall{ Name: "filesystem_edit", }, }, @@ -2446,7 +2447,7 @@ func TestRenderMessageBlockUserContentAlignsWithUserTag(t *testing.T) { t.Fatalf("New() error = %v", err) } - renderedMessage, _ := app.renderMessageBlockWithCopy(provider.Message{Role: roleUser, Content: "nihao"}, 80, 1) + renderedMessage, _ := app.renderMessageBlockWithCopy(providertypes.Message{Role: roleUser, Content: "nihao"}, 80, 1) rendered := stripANSI(renderedMessage) lines := strings.Split(rendered, "\n") @@ -2735,7 +2736,7 @@ func newTestProviderService(t *testing.T, manager *config.Manager) *config.Selec type tUItestProvider struct{} -func (tUItestProvider) Chat(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { +func (tUItestProvider) Chat(ctx context.Context, req providertypes.ChatRequest, events chan<- providertypes.StreamEvent) error { return nil } diff --git a/internal/tui/view.go b/internal/tui/view.go index f248f0a7..8398f4ef 100644 --- a/internal/tui/view.go +++ b/internal/tui/view.go @@ -7,7 +7,7 @@ import ( "github.com/charmbracelet/bubbles/list" "github.com/charmbracelet/lipgloss" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" ) type layout struct { @@ -229,7 +229,7 @@ func (a App) renderPanel(title string, subtitle string, body string, width int, return lipgloss.Place(width, height, lipgloss.Left, lipgloss.Top, panel) } -func (a App) renderMessageBlockWithCopy(message provider.Message, width int, startCopyID int) (string, []copyCodeButtonBinding) { +func (a App) renderMessageBlockWithCopy(message providertypes.Message, width int, startCopyID int) (string, []copyCodeButtonBinding) { switch message.Role { case roleEvent: return a.styles.inlineNotice.Width(width).Render(" > " + wrapPlain(message.Content, max(16, width-6))), nil From 04ccc2e84c90c7d8191fa046df81b232a3d461f0 Mon Sep 17 00:00:00 2001 From: Yumiue <229866007@qq.com> Date: Tue, 7 Apr 2026 17:24:57 +0800 Subject: [PATCH 30/55] =?UTF-8?q?refator:=E6=8B=86=E5=88=86=E5=87=BAsessio?= =?UTF-8?q?n=E6=A8=A1=E5=9D=97?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- README.md | 4 +- docs/session-persistence-design.md | 16 ++++- internal/app/bootstrap.go | 3 +- internal/runtime/compact.go | 9 +-- internal/runtime/id.go | 12 ---- internal/runtime/runtime.go | 37 +++++----- internal/runtime/runtime_test.go | 71 ++++++++++--------- internal/runtime/workdir_branch_test.go | 4 +- internal/session/id.go | 18 +++++ .../{runtime/session.go => session/store.go} | 68 +++++++++++------- .../session_test.go => session/store_test.go} | 24 +++---- internal/tui/state.go | 6 +- internal/tui/state/ui_state.go | 4 +- internal/tui/update_test.go | 41 +++++------ 14 files changed, 178 insertions(+), 139 deletions(-) delete mode 100644 internal/runtime/id.go create mode 100644 internal/session/id.go rename internal/{runtime/session.go => session/store.go} (58%) rename internal/{runtime/session_test.go => session/store_test.go} (87%) diff --git a/README.md b/README.md index d710be7c..33fdaa73 100644 --- a/README.md +++ b/README.md @@ -94,7 +94,8 @@ go run ./cmd/neocode - **`internal/provider/catalog`** — 模型发现、catalog 缓存与后台刷新 - **`internal/provider/selection`** — provider/model 选择与配置同步 - **`internal/provider/builtin`** — 内建 driver 注册 -- **`internal/runtime`** — ReAct 主循环、事件流、会话管理 +- **`internal/runtime`** — ReAct 主循环与事件流编排(不直接承载会话存储实现;不再导出会话模型与存储类型) +- **`internal/session`** — 会话模型、会话存储抽象与 JSON 持久化实现(统一对外暴露 `Session` / `Summary` / `Store`) - **`internal/tools`** — 工具注册表与具体工具实现 - **`internal/tui`** — 终端 UI、交互体验、事件桥接 - **`internal/app`** — 应用装配与依赖注入 @@ -116,6 +117,7 @@ go run ./cmd/neocode │ │ ├── catalog # 模型发现与缓存 │ │ └── selection # provider/model 选择服务 │ ├── runtime # ReAct 循环与事件流 +│ ├── session # 会话模型与持久化 │ ├── tools # 工具系统 │ └── tui # 终端 UI └── README.md diff --git a/docs/session-persistence-design.md b/docs/session-persistence-design.md index 0c5d1887..e66f7800 100644 --- a/docs/session-persistence-design.md +++ b/docs/session-persistence-design.md @@ -1,10 +1,16 @@ # Session 持久化设计 + +## 模块职责与收口边界 +- `internal/session`:承载会话领域模型、存储抽象与 JSON 持久化实现,是唯一的会话持久化实现归属层 +- `internal/runtime`:只依赖 `internal/session` 提供的抽象与模型,负责会话保存时机与主循环编排,不再维护会话存储实现细节 +- `internal/tui`:仅消费 runtime 暴露的会话数据,不直接执行会话持久化 + ## 存储策略 NeoCode 在 MVP 阶段使用 JSON 文件持久化 Session,以保持本地优先、易于调试和跨平台可移植。 ## 数据模型 - `Session`:完整消息历史以及 `id`、`title`、`updated_at` 等元信息 -- `SessionSummary`:用于侧边栏的轻量摘要结构 +- `Summary`:用于侧边栏的轻量摘要结构(原 `SessionSummary` 命名已统一收口为 `Summary`) ## 加载策略 - `ListSummaries` 只读取渲染侧边栏所需的基础信息 @@ -12,9 +18,13 @@ NeoCode 在 MVP 阶段使用 JSON 文件持久化 Session,以保持本地优 - `Save` 通过临时文件原子写入完整 Session ## 命名策略 -- 新会话默认展示为 `Draft` +- 新会话默认展示为 `New Session` - 一旦持久化,runtime 会根据首轮用户消息生成简短标题 ## 并发约束 -- SessionStore 实现必须自行保护共享访问 +- `internal/session` 中的 Store 实现必须自行保护共享访问 - 真正的保存时机由 runtime 决定,TUI 不负责直接触发磁盘写入 + +## 兼容性与演进说明 +- 会话持久化能力已从 runtime 侧实现中彻底收口到 `internal/session` +- 新增会话存储实现时,应优先在 `internal/session` 内扩展并通过接口注入 runtime,避免跨层实现 diff --git a/internal/app/bootstrap.go b/internal/app/bootstrap.go index dd0b8436..6d2892b3 100644 --- a/internal/app/bootstrap.go +++ b/internal/app/bootstrap.go @@ -12,6 +12,7 @@ import ( providercatalog "neo-code/internal/provider/catalog" agentruntime "neo-code/internal/runtime" "neo-code/internal/security" + agentsession "neo-code/internal/session" "neo-code/internal/tools" "neo-code/internal/tools/bash" "neo-code/internal/tools/filesystem" @@ -62,7 +63,7 @@ func NewProgram(ctx context.Context) (*tea.Program, error) { return nil, err } - sessionStore := agentruntime.NewSessionStore(loader.BaseDir()) + sessionStore := agentsession.NewStore(loader.BaseDir()) runtimeSvc := agentruntime.NewWithFactory( manager, toolManager, diff --git a/internal/runtime/compact.go b/internal/runtime/compact.go index 4939ca62..3b1040a1 100644 --- a/internal/runtime/compact.go +++ b/internal/runtime/compact.go @@ -9,6 +9,7 @@ import ( "neo-code/internal/config" contextcompact "neo-code/internal/context/compact" "neo-code/internal/provider" + agentsession "neo-code/internal/session" ) // CompactInput 描述一次手动 compact 请求所需的最小输入。 @@ -83,10 +84,10 @@ func (s *Service) Compact(ctx context.Context, input CompactInput) (CompactResul func (s *Service) runCompactForSession( ctx context.Context, runID string, - session Session, + session agentsession.Session, cfg config.Config, failOnError bool, -) (Session, contextcompact.Result, error) { +) (agentsession.Session, contextcompact.Result, error) { runner := s.compactRunner if runner == nil { var err error @@ -155,7 +156,7 @@ func (s *Service) runCompactForSession( } // defaultCompactRunner 为手动 compact 选择摘要生成器并构造默认 runner。 -func (s *Service) defaultCompactRunner(session Session, cfg config.Config) (contextcompact.Runner, error) { +func (s *Service) defaultCompactRunner(session agentsession.Session, cfg config.Config) (contextcompact.Runner, error) { resolvedProvider, model, err := resolveCompactProviderSelection(session, cfg) if err != nil { return nil, err @@ -164,7 +165,7 @@ func (s *Service) defaultCompactRunner(session Session, cfg config.Config) (cont } // resolveCompactProviderSelection 优先复用会话记录的 provider/model,缺失时再回退当前配置。 -func resolveCompactProviderSelection(session Session, cfg config.Config) (config.ResolvedProviderConfig, string, error) { +func resolveCompactProviderSelection(session agentsession.Session, cfg config.Config) (config.ResolvedProviderConfig, string, error) { sessionProvider := strings.TrimSpace(session.Provider) sessionModel := strings.TrimSpace(session.Model) if sessionProvider != "" && sessionModel != "" { diff --git a/internal/runtime/id.go b/internal/runtime/id.go deleted file mode 100644 index b036fbea..00000000 --- a/internal/runtime/id.go +++ /dev/null @@ -1,12 +0,0 @@ -package runtime - -import ( - "crypto/rand" - "encoding/hex" -) - -func newID(prefix string) string { - buf := make([]byte, 8) - _, _ = rand.Read(buf) - return prefix + "_" + hex.EncodeToString(buf) -} diff --git a/internal/runtime/runtime.go b/internal/runtime/runtime.go index b9e97b1d..28a21f42 100644 --- a/internal/runtime/runtime.go +++ b/internal/runtime/runtime.go @@ -17,6 +17,7 @@ import ( agentcontext "neo-code/internal/context" contextcompact "neo-code/internal/context/compact" "neo-code/internal/provider" + agentsession "neo-code/internal/session" "neo-code/internal/tools" ) @@ -120,9 +121,9 @@ type Runtime interface { ResolvePermission(ctx context.Context, input PermissionResolutionInput) error CancelActiveRun() bool Events() <-chan RuntimeEvent - ListSessions(ctx context.Context) ([]SessionSummary, error) - LoadSession(ctx context.Context, id string) (Session, error) - SetSessionWorkdir(ctx context.Context, sessionID string, workdir string) (Session, error) + ListSessions(ctx context.Context) ([]agentsession.Summary, error) + LoadSession(ctx context.Context, id string) (agentsession.Session, error) + SetSessionWorkdir(ctx context.Context, sessionID string, workdir string) (agentsession.Session, error) } type UserInput struct { @@ -138,7 +139,7 @@ type ProviderFactory interface { type Service struct { configManager *config.Manager // 配置管理器,提供当前选中的 provider、model、workdir 等配置读取能力。 - sessionStore Store // 会话持久化接口,负责保存和加载聊天会话。 + sessionStore agentsession.Store // 会话持久化接口,负责保存和加载聊天会话。 toolManager tools.Manager // 工具管理器,统一工具 schema 暴露与执行入口。 providerFactory ProviderFactory // Provider 工厂接口,根据配置动态创建具体的 provider 实例。 contextBuilder agentcontext.Builder // 上下文构建器,负责组装 system prompt 与本轮发给模型的消息上下文。 @@ -154,7 +155,7 @@ type Service struct { func NewWithFactory( configManager *config.Manager, toolManager tools.Manager, - sessionStore Store, + sessionStore agentsession.Store, providerFactory ProviderFactory, contextBuilder agentcontext.Builder, ) *Service { @@ -372,35 +373,35 @@ func (s *Service) Events() <-chan RuntimeEvent { return s.events } -func (s *Service) ListSessions(ctx context.Context) ([]SessionSummary, error) { +func (s *Service) ListSessions(ctx context.Context) ([]agentsession.Summary, error) { return s.sessionStore.ListSummaries(ctx) } -func (s *Service) LoadSession(ctx context.Context, id string) (Session, error) { +func (s *Service) LoadSession(ctx context.Context, id string) (agentsession.Session, error) { session, err := s.sessionStore.Load(ctx, id) if err != nil { - return Session{}, err + return agentsession.Session{}, err } session.Workdir = s.sessionWorkdir(id, session.Workdir) return session, nil } -func (s *Service) SetSessionWorkdir(ctx context.Context, sessionID string, workdir string) (Session, error) { +func (s *Service) SetSessionWorkdir(ctx context.Context, sessionID string, workdir string) (agentsession.Session, error) { sessionID = strings.TrimSpace(sessionID) if sessionID == "" { - return Session{}, errors.New("runtime: session id is empty") + return agentsession.Session{}, errors.New("runtime: session id is empty") } session, err := s.sessionStore.Load(ctx, sessionID) if err != nil { - return Session{}, err + return agentsession.Session{}, err } session.Workdir = s.sessionWorkdir(sessionID, session.Workdir) cfg := s.configManager.Get() resolved, err := resolveWorkdirForSession(cfg.Workdir, session.Workdir, workdir) if err != nil { - return Session{}, err + return agentsession.Session{}, err } if session.Workdir == resolved { return session, nil @@ -439,22 +440,22 @@ func (s *Service) loadOrCreateSession( title string, defaultWorkdir string, requestedWorkdir string, -) (Session, error) { +) (agentsession.Session, error) { if strings.TrimSpace(sessionID) == "" { sessionWorkdir, err := resolveWorkdirForSession(defaultWorkdir, "", requestedWorkdir) if err != nil { - return Session{}, err + return agentsession.Session{}, err } - session := newSessionWithWorkdir(title, sessionWorkdir) + session := agentsession.NewWithWorkdir(title, sessionWorkdir) s.setSessionWorkdir(session.ID, sessionWorkdir) if err := s.sessionStore.Save(ctx, &session); err != nil { - return Session{}, err + return agentsession.Session{}, err } return session, nil } session, err := s.sessionStore.Load(ctx, sessionID) if err != nil { - return Session{}, err + return agentsession.Session{}, err } session.Workdir = s.sessionWorkdir(sessionID, session.Workdir) if strings.TrimSpace(requestedWorkdir) == "" && strings.TrimSpace(session.Workdir) != "" { @@ -463,7 +464,7 @@ func (s *Service) loadOrCreateSession( resolved, err := resolveWorkdirForSession(defaultWorkdir, session.Workdir, requestedWorkdir) if err != nil { - return Session{}, err + return agentsession.Session{}, err } if session.Workdir == resolved { return session, nil diff --git a/internal/runtime/runtime_test.go b/internal/runtime/runtime_test.go index 64f32bb5..d010a908 100644 --- a/internal/runtime/runtime_test.go +++ b/internal/runtime/runtime_test.go @@ -15,16 +15,17 @@ import ( contextcompact "neo-code/internal/context/compact" "neo-code/internal/provider" "neo-code/internal/security" + agentsession "neo-code/internal/session" "neo-code/internal/tools" ) type memoryStore struct { - sessions map[string]Session + sessions map[string]agentsession.Session saves int } type failingStore struct { - Store + agentsession.Store saveErr error failOnSave int saveCalls int @@ -32,10 +33,10 @@ type failingStore struct { } func newMemoryStore() *memoryStore { - return &memoryStore{sessions: map[string]Session{}} + return &memoryStore{sessions: map[string]agentsession.Session{}} } -func (s *failingStore) Save(ctx context.Context, session *Session) error { +func (s *failingStore) Save(ctx context.Context, session *agentsession.Session) error { s.saveCalls++ if s.failOnSave > 0 && s.saveCalls == s.failOnSave { return s.saveErr @@ -49,7 +50,7 @@ func (s *failingStore) Save(ctx context.Context, session *Session) error { return s.Store.Save(ctx, session) } -func (s *memoryStore) Save(ctx context.Context, session *Session) error { +func (s *memoryStore) Save(ctx context.Context, session *agentsession.Session) error { if err := ctx.Err(); err != nil { return err } @@ -61,24 +62,24 @@ func (s *memoryStore) Save(ctx context.Context, session *Session) error { return nil } -func (s *memoryStore) Load(ctx context.Context, id string) (Session, error) { +func (s *memoryStore) Load(ctx context.Context, id string) (agentsession.Session, error) { if err := ctx.Err(); err != nil { - return Session{}, err + return agentsession.Session{}, err } session, ok := s.sessions[id] if !ok { - return Session{}, errors.New("not found") + return agentsession.Session{}, errors.New("not found") } return cloneSession(session), nil } -func (s *memoryStore) ListSummaries(ctx context.Context) ([]SessionSummary, error) { +func (s *memoryStore) ListSummaries(ctx context.Context) ([]agentsession.Summary, error) { if err := ctx.Err(); err != nil { return nil, err } - summaries := make([]SessionSummary, 0, len(s.sessions)) + summaries := make([]agentsession.Summary, 0, len(s.sessions)) for _, session := range s.sessions { - summaries = append(summaries, SessionSummary{ + summaries = append(summaries, agentsession.Summary{ ID: session.ID, Title: session.Title, CreatedAt: session.CreatedAt, @@ -606,7 +607,7 @@ func TestServiceRunDelegatesToContextBuilder(t *testing.T) { manager := newRuntimeConfigManager(t) store := newMemoryStore() - session := newSession("memory reject") + session := agentsession.New("memory reject") session.ID = "session-memory-reject" store.sessions[session.ID] = cloneSession(session) registry := tools.NewRegistry() @@ -747,7 +748,7 @@ func TestServiceRunDefaultBuilderUsesToolManagerMicroCompactPolicies(t *testing. registry.Register(&stubTool{name: "bash", content: "default"}) registry.Register(&stubTool{name: "webfetch", content: "default"}) - session := newSession("preserve history") + session := agentsession.New("preserve history") session.ID = "session-preserve-history" session.Messages = []provider.Message{ {Role: provider.RoleUser, Content: "older user"}, @@ -810,7 +811,7 @@ func TestServiceRunDefaultBuilderUsesGenericToolManagerMicroCompactPolicies(t *t }, } - session := newSession("preserve history by manager") + session := agentsession.New("preserve history by manager") session.ID = "session-preserve-history-manager" session.Messages = []provider.Message{ {Role: provider.RoleUser, Content: "older user"}, @@ -879,7 +880,7 @@ func TestServiceRunFailurePreservesExistingSessionProviderAndModel(t *testing.T) } store := newMemoryStore() - session := newSession("preserve-metadata") + session := agentsession.New("preserve-metadata") session.ID = "session-preserve-metadata" session.Provider = config.OpenAIName session.Model = "openai-original-model" @@ -979,7 +980,7 @@ func TestServiceRunWaitsForPermissionResolutionAndContinues(t *testing.T) { manager := newRuntimeConfigManager(t) store := newMemoryStore() - session := newSession("memory reject") + session := agentsession.New("memory reject") session.ID = "session-memory-reject" store.sessions[session.ID] = cloneSession(session) registry := tools.NewRegistry() @@ -1170,7 +1171,7 @@ func TestServiceRunEmitsRememberScopeWhenSessionRejectMemoryHits(t *testing.T) { manager := newRuntimeConfigManager(t) store := newMemoryStore() - session := newSession("memory reject") + session := agentsession.New("memory reject") session.ID = "session-memory-reject" store.sessions[session.ID] = cloneSession(session) registry := tools.NewRegistry() @@ -1304,7 +1305,7 @@ func TestServiceRunErrorPaths(t *testing.T) { provider *scriptedProvider factoryErr error registerTool *stubTool - seedSession *Session + seedSession *agentsession.Session expectErr string expectEvents []EventType assert func(t *testing.T, store *memoryStore, provider *scriptedProvider, tool *stubTool) @@ -1368,11 +1369,11 @@ func TestServiceRunErrorPaths(t *testing.T) { {provider.NewTextDeltaStreamEvent("resumed")}, }, }, - seedSession: &Session{ + seedSession: &agentsession.Session{ ID: "existing-session", Title: "Resume Me", - CreatedAt: newSession("seed").CreatedAt, - UpdatedAt: newSession("seed").UpdatedAt, + CreatedAt: agentsession.New("seed").CreatedAt, + UpdatedAt: agentsession.New("seed").UpdatedAt, Messages: []provider.Message{ {Role: "user", Content: "earlier"}, }, @@ -1862,7 +1863,7 @@ func TestServiceRunToolTimeoutIsNotCancellation(t *testing.T) { func TestServiceCompactManualAppliesAndPersists(t *testing.T) { manager := newRuntimeConfigManager(t) store := newMemoryStore() - session := newSession("manual") + session := agentsession.New("manual") session.ID = "session-manual" session.Messages = []provider.Message{ {Role: provider.RoleUser, Content: "older"}, @@ -1920,7 +1921,7 @@ func TestServiceCompactManualAppliesAndPersists(t *testing.T) { func TestServiceCompactManualFailureReturnsError(t *testing.T) { manager := newRuntimeConfigManager(t) store := newMemoryStore() - session := newSession("manual-fail") + session := agentsession.New("manual-fail") session.ID = "session-manual-fail" session.Messages = []provider.Message{ {Role: provider.RoleUser, Content: "older"}, @@ -1974,7 +1975,7 @@ func TestServiceCompactUsesSessionProviderAndModelWhenPresent(t *testing.T) { } store := newMemoryStore() - session := newSession("manual-provider") + session := agentsession.New("manual-provider") session.ID = "session-manual-provider" session.Provider = config.OpenAIName session.Model = "session-model" @@ -2048,7 +2049,7 @@ func TestServiceCompactFallsBackToCurrentProviderWhenSessionMetadataMissing(t *t } store := newMemoryStore() - session := newSession("manual-fallback") + session := agentsession.New("manual-fallback") session.ID = "session-manual-fallback" session.Messages = []provider.Message{ {Role: provider.RoleUser, Content: "older"}, @@ -2105,7 +2106,7 @@ func TestServiceCompactFallsBackToCurrentProviderWhenSessionMetadataMissing(t *t func TestServiceManualCompactThenRunContinuesToolRound(t *testing.T) { manager := newRuntimeConfigManager(t) store := newMemoryStore() - session := newSession("manual-continue") + session := agentsession.New("manual-continue") session.ID = "session-manual-continue" session.Messages = []provider.Message{ {Role: provider.RoleUser, Content: "legacy request"}, @@ -2188,7 +2189,7 @@ func TestServiceManualCompactThenRunContinuesToolRound(t *testing.T) { func TestServiceSerializesRunAndCompact(t *testing.T) { manager := newRuntimeConfigManager(t) store := newMemoryStore() - session := newSession("serialized") + session := agentsession.New("serialized") session.ID = "session-serialized" store.sessions[session.ID] = cloneSession(session) @@ -2284,7 +2285,7 @@ func TestServiceConstructorsAndDelegates(t *testing.T) { t.Fatalf("expected events channel") } - session := newSession("List Me") + session := agentsession.New("List Me") store.sessions[session.ID] = cloneSession(session) summaries, err := service.ListSessions(context.Background()) @@ -2303,7 +2304,7 @@ func TestServiceConstructorsAndDelegates(t *testing.T) { t.Fatalf("expected loaded session %q, got %q", session.ID, loaded.ID) } - sessionStore := NewSessionStore(t.TempDir()) + sessionStore := agentsession.NewStore(t.TempDir()) if sessionStore == nil { t.Fatalf("expected JSON session store") } @@ -2321,7 +2322,7 @@ func TestServiceRunUsesSessionWorkdirForContextAndTools(t *testing.T) { } store := newMemoryStore() - session := newSessionWithWorkdir("Session Workdir", sessionWorkdir) + session := agentsession.NewWithWorkdir("Session Workdir", sessionWorkdir) store.sessions[session.ID] = cloneSession(session) tool := &stubTool{name: "filesystem_edit", content: "ok"} @@ -2410,7 +2411,7 @@ func TestServiceSetSessionWorkdir(t *testing.T) { } store := newMemoryStore() - session := newSession("set workdir") + session := agentsession.New("set workdir") store.sessions[session.ID] = cloneSession(session) registry := tools.NewRegistry() registry.Register(&stubTool{name: "filesystem_read_file", content: "default"}) @@ -2541,7 +2542,7 @@ func restoreRuntimeEnv(t *testing.T, key string) { }) } -func onlySession(t *testing.T, store *memoryStore) Session { +func onlySession(t *testing.T, store *memoryStore) agentsession.Session { t.Helper() if len(store.sessions) != 1 { t.Fatalf("expected exactly 1 session, got %d", len(store.sessions)) @@ -2549,7 +2550,7 @@ func onlySession(t *testing.T, store *memoryStore) Session { for _, session := range store.sessions { return session } - return Session{} + return agentsession.Session{} } func resolvedProviderForTests(cfg config.Config, providerName string) (config.ResolvedProviderConfig, error) { @@ -2606,7 +2607,7 @@ func assertEventsRunID(t *testing.T, events []RuntimeEvent, runID string) { } } -func cloneSession(session Session) Session { +func cloneSession(session agentsession.Session) agentsession.Session { cloned := session cloned.Messages = append([]provider.Message(nil), session.Messages...) return cloned @@ -2691,7 +2692,7 @@ func TestServiceSetSessionWorkdirNoopDoesNotSave(t *testing.T) { store := newMemoryStore() target := t.TempDir() - session := newSessionWithWorkdir("noop", target) + session := agentsession.NewWithWorkdir("noop", target) store.sessions[session.ID] = cloneSession(session) registry := tools.NewRegistry() registry.Register(&stubTool{name: "filesystem_read_file", content: "default"}) diff --git a/internal/runtime/workdir_branch_test.go b/internal/runtime/workdir_branch_test.go index 0edc39ef..266fed5e 100644 --- a/internal/runtime/workdir_branch_test.go +++ b/internal/runtime/workdir_branch_test.go @@ -6,6 +6,8 @@ import ( "path/filepath" "strings" "testing" + + agentsession "neo-code/internal/session" ) func TestSessionWorkdirKeyAndMemoryMap(t *testing.T) { @@ -82,7 +84,7 @@ func TestLoadSessionUsesFallbackWorkdirWhenMemoryMissing(t *testing.T) { manager := newRuntimeConfigManager(t) store := newMemoryStore() - session := newSession("fallback") + session := agentsession.New("fallback") session.Workdir = t.TempDir() store.sessions[session.ID] = cloneSession(session) diff --git a/internal/session/id.go b/internal/session/id.go new file mode 100644 index 00000000..dc2e84e9 --- /dev/null +++ b/internal/session/id.go @@ -0,0 +1,18 @@ +package session + +import ( + "crypto/rand" + "encoding/hex" +) + +// NewID 生成带前缀的随机 ID,格式为 "_<16hex>"。 +func NewID(prefix string) string { + buf := make([]byte, 8) + _, _ = rand.Read(buf) + return prefix + "_" + hex.EncodeToString(buf) +} + +// newID 保留为内部兼容入口,后续代码请优先使用 NewID。 +func newID(prefix string) string { + return NewID(prefix) +} diff --git a/internal/runtime/session.go b/internal/session/store.go similarity index 58% rename from internal/runtime/session.go rename to internal/session/store.go index af64317d..4deeffec 100644 --- a/internal/runtime/session.go +++ b/internal/session/store.go @@ -1,4 +1,4 @@ -package runtime +package session import ( "context" @@ -17,6 +17,8 @@ import ( const sessionsDirName = "sessions" +// Session 表示单个会话的持久化模型,包含基础元数据与消息历史。 +// Provider / Model 用于在 compact 等流程中优先复用会话最近一次成功运行的模型配置。 type Session struct { ID string `json:"id"` Title string `json:"title"` @@ -30,71 +32,78 @@ type Session struct { Messages []provider.Message `json:"messages"` } -type SessionSummary struct { +// Summary 表示会话列表视图所需的轻量摘要信息。 +type Summary struct { ID string `json:"id"` Title string `json:"title"` CreatedAt time.Time `json:"created_at"` UpdatedAt time.Time `json:"updated_at"` } +// Store 定义会话持久化抽象。 type Store interface { Save(ctx context.Context, session *Session) error Load(ctx context.Context, id string) (Session, error) - ListSummaries(ctx context.Context) ([]SessionSummary, error) + ListSummaries(ctx context.Context) ([]Summary, error) } -type JSONSessionStore struct { +// JSONStore 是基于 JSON 文件的会话存储实现。 +type JSONStore struct { mu sync.RWMutex baseDir string } -func NewJSONSessionStore(baseDir string) *JSONSessionStore { - return &JSONSessionStore{ +// NewJSONStore 创建 JSONStore,实际会话目录为 {baseDir}/sessions。 +func NewJSONStore(baseDir string) *JSONStore { + return &JSONStore{ baseDir: filepath.Join(baseDir, sessionsDirName), } } -func NewSessionStore(baseDir string) *JSONSessionStore { - return NewJSONSessionStore(baseDir) +// NewStore 返回默认会话存储实现(当前为 JSONStore)。 +func NewStore(baseDir string) *JSONStore { + return NewJSONStore(baseDir) } -func (s *JSONSessionStore) Save(ctx context.Context, session *Session) error { +// Save 持久化会话到 JSON 文件,采用临时文件 + 原子替换策略。 +func (s *JSONStore) Save(ctx context.Context, session *Session) error { if err := ctx.Err(); err != nil { return err } if session == nil { - return errors.New("runtime: session is nil") + return errors.New("session: session is nil") } s.mu.Lock() defer s.mu.Unlock() if err := os.MkdirAll(s.baseDir, 0o755); err != nil { - return fmt.Errorf("runtime: create sessions dir: %w", err) + return fmt.Errorf("session: create sessions dir: %w", err) } payload, err := json.MarshalIndent(session, "", " ") if err != nil { - return fmt.Errorf("runtime: marshal session: %w", err) + return fmt.Errorf("session: marshal session: %w", err) } payload = append(payload, '\n') target := s.filePath(session.ID) temp := target + ".tmp" if err := os.WriteFile(temp, payload, 0o644); err != nil { - return fmt.Errorf("runtime: write temp session: %w", err) + return fmt.Errorf("session: write temp session: %w", err) } if err := os.Remove(target); err != nil && !errors.Is(err, os.ErrNotExist) { - return fmt.Errorf("runtime: replace session file: %w", err) + return fmt.Errorf("session: replace session file: %w", err) } if err := os.Rename(temp, target); err != nil { - return fmt.Errorf("runtime: commit session file: %w", err) + return fmt.Errorf("session: commit session file: %w", err) } return nil } -func (s *JSONSessionStore) Load(ctx context.Context, id string) (Session, error) { +// Load 读取并反序列化指定 ID 的会话文件。 +func (s *JSONStore) Load(ctx context.Context, id string) (Session, error) { if err := ctx.Err(); err != nil { return Session{}, err } @@ -109,12 +118,13 @@ func (s *JSONSessionStore) Load(ctx context.Context, id string) (Session, error) var session Session if err := json.Unmarshal(data, &session); err != nil { - return Session{}, fmt.Errorf("runtime: decode session %s: %w", id, err) + return Session{}, fmt.Errorf("session: decode session %s: %w", id, err) } return session, nil } -func (s *JSONSessionStore) ListSummaries(ctx context.Context) ([]SessionSummary, error) { +// ListSummaries 列出所有会话摘要,并按 UpdatedAt 倒序返回。 +func (s *JSONStore) ListSummaries(ctx context.Context) ([]Summary, error) { if err := ctx.Err(); err != nil { return nil, err } @@ -123,15 +133,15 @@ func (s *JSONSessionStore) ListSummaries(ctx context.Context) ([]SessionSummary, defer s.mu.RUnlock() if err := os.MkdirAll(s.baseDir, 0o755); err != nil { - return nil, fmt.Errorf("runtime: create sessions dir: %w", err) + return nil, fmt.Errorf("session: create sessions dir: %w", err) } entries, err := os.ReadDir(s.baseDir) if err != nil { - return nil, fmt.Errorf("runtime: list sessions dir: %w", err) + return nil, fmt.Errorf("session: list sessions dir: %w", err) } - summaries := make([]SessionSummary, 0, len(entries)) + summaries := make([]Summary, 0, len(entries)) for _, entry := range entries { if entry.IsDir() || filepath.Ext(entry.Name()) != ".json" { continue @@ -148,7 +158,7 @@ func (s *JSONSessionStore) ListSummaries(ctx context.Context) ([]SessionSummary, continue } - var summary SessionSummary + var summary Summary if err := json.Unmarshal(data, &summary); err != nil { continue } @@ -165,18 +175,21 @@ func (s *JSONSessionStore) ListSummaries(ctx context.Context) ([]SessionSummary, return summaries, nil } -func (s *JSONSessionStore) filePath(id string) string { +// filePath 生成会话 ID 对应的 JSON 文件路径。 +func (s *JSONStore) filePath(id string) string { return filepath.Join(s.baseDir, id+".json") } -func newSession(title string) Session { - return newSessionWithWorkdir(title, "") +// New 创建一个默认标题策略的新会话对象。 +func New(title string) Session { + return NewWithWorkdir(title, "") } -func newSessionWithWorkdir(title string, workdir string) Session { +// NewWithWorkdir 创建一个包含运行目录的会话对象。 +func NewWithWorkdir(title string, workdir string) Session { now := time.Now() return Session{ - ID: newID("session"), + ID: NewID("session"), Title: sanitizeTitle(title), CreatedAt: now, UpdatedAt: now, @@ -185,6 +198,7 @@ func newSessionWithWorkdir(title string, workdir string) Session { } } +// sanitizeTitle 规范化会话标题:去空白、空标题回退默认值、超长截断。 func sanitizeTitle(title string) string { title = strings.TrimSpace(title) if title == "" { diff --git a/internal/runtime/session_test.go b/internal/session/store_test.go similarity index 87% rename from internal/runtime/session_test.go rename to internal/session/store_test.go index a82f1215..0e5e02d4 100644 --- a/internal/runtime/session_test.go +++ b/internal/session/store_test.go @@ -1,4 +1,4 @@ -package runtime +package session import ( "context" @@ -11,11 +11,11 @@ import ( "neo-code/internal/provider" ) -func TestJSONSessionStoreSaveLoadAndListSummaries(t *testing.T) { +func TestJSONStoreSaveLoadAndListSummaries(t *testing.T) { t.Parallel() baseDir := t.TempDir() - store := NewJSONSessionStore(baseDir) + store := NewJSONStore(baseDir) older := &Session{ ID: "session-old", @@ -68,7 +68,7 @@ func TestJSONSessionStoreSaveLoadAndListSummaries(t *testing.T) { t.Fatalf("expected persisted session file to exclude workdir, got:\n%s", string(raw)) } - mustWriteRuntimeFile(t, filepath.Join(baseDir, sessionsDirName, "invalid.json"), "{invalid") + mustWriteSessionFile(t, filepath.Join(baseDir, sessionsDirName, "invalid.json"), "{invalid") if err := os.MkdirAll(filepath.Join(baseDir, sessionsDirName, "directory"), 0o755); err != nil { t.Fatalf("mkdir stray directory: %v", err) } @@ -85,11 +85,11 @@ func TestJSONSessionStoreSaveLoadAndListSummaries(t *testing.T) { } } -func TestJSONSessionStoreErrors(t *testing.T) { +func TestJSONStoreErrors(t *testing.T) { t.Parallel() baseDir := t.TempDir() - store := NewJSONSessionStore(baseDir) + store := NewJSONStore(baseDir) cancelledCtx, cancel := context.WithCancel(context.Background()) cancel() @@ -108,11 +108,11 @@ func TestJSONSessionStoreErrors(t *testing.T) { } } -func TestJSONSessionStoreCorruptedSessionBehaviors(t *testing.T) { +func TestJSONStoreCorruptedSessionBehaviors(t *testing.T) { t.Parallel() baseDir := t.TempDir() - store := NewJSONSessionStore(baseDir) + store := NewJSONStore(baseDir) valid := &Session{ ID: "valid-session", @@ -125,7 +125,7 @@ func TestJSONSessionStoreCorruptedSessionBehaviors(t *testing.T) { t.Fatalf("Save valid session: %v", err) } - mustWriteRuntimeFile(t, filepath.Join(baseDir, sessionsDirName, "broken.json"), "{broken") + mustWriteSessionFile(t, filepath.Join(baseDir, sessionsDirName, "broken.json"), "{broken") _, err := store.Load(context.Background(), "broken") if err == nil || !strings.Contains(err.Error(), "decode session broken") { @@ -141,7 +141,7 @@ func TestJSONSessionStoreCorruptedSessionBehaviors(t *testing.T) { } } -func TestJSONSessionStoreSaveInvalidBaseDir(t *testing.T) { +func TestJSONStoreSaveInvalidBaseDir(t *testing.T) { t.Parallel() tempDir := t.TempDir() @@ -150,7 +150,7 @@ func TestJSONSessionStoreSaveInvalidBaseDir(t *testing.T) { t.Fatalf("write base file: %v", err) } - store := NewJSONSessionStore(baseFile) + store := NewJSONStore(baseFile) err := store.Save(context.Background(), &Session{ ID: "session-x", Title: "Broken Save", @@ -162,7 +162,7 @@ func TestJSONSessionStoreSaveInvalidBaseDir(t *testing.T) { } } -func mustWriteRuntimeFile(t *testing.T, path string, content string) { +func mustWriteSessionFile(t *testing.T, path string, content string) { t.Helper() if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil { t.Fatalf("mkdir %s: %v", filepath.Dir(path), err) diff --git a/internal/tui/state.go b/internal/tui/state.go index 1f939dad..c65ecc63 100644 --- a/internal/tui/state.go +++ b/internal/tui/state.go @@ -10,7 +10,7 @@ import ( tea "github.com/charmbracelet/bubbletea" "github.com/charmbracelet/lipgloss" - agentruntime "neo-code/internal/runtime" + agentsession "neo-code/internal/session" ) type panel int @@ -32,7 +32,7 @@ const ( ) type UIState struct { - Sessions []agentruntime.SessionSummary + Sessions []agentsession.Summary ActiveSessionID string ActiveSessionTitle string InputText string @@ -142,7 +142,7 @@ func (d commandMenuDelegate) Render(w io.Writer, m list.Model, index int, item l } type sessionItem struct { - Summary agentruntime.SessionSummary + Summary agentsession.Summary Active bool } diff --git a/internal/tui/state/ui_state.go b/internal/tui/state/ui_state.go index ac6c4611..9fa071a3 100644 --- a/internal/tui/state/ui_state.go +++ b/internal/tui/state/ui_state.go @@ -1,6 +1,6 @@ package state -import agentruntime "neo-code/internal/runtime" +import agentsession "neo-code/internal/session" // Panel 定义 TUI 中可聚焦的主面板。 type Panel int @@ -24,7 +24,7 @@ const ( // UIState 保存顶层界面状态快照,仅作为数据容器使用。 type UIState struct { - Sessions []agentruntime.SessionSummary + Sessions []agentsession.Summary ActiveSessionID string ActiveSessionTitle string ActiveRunID string diff --git a/internal/tui/update_test.go b/internal/tui/update_test.go index 6a741a06..255786b8 100644 --- a/internal/tui/update_test.go +++ b/internal/tui/update_test.go @@ -21,6 +21,7 @@ import ( "neo-code/internal/provider" providercatalog "neo-code/internal/provider/catalog" agentruntime "neo-code/internal/runtime" + agentsession "neo-code/internal/session" "neo-code/internal/tools" ) @@ -28,15 +29,15 @@ type stubRuntime struct { runInputs []agentruntime.UserInput compactInputs []agentruntime.CompactInput events chan agentruntime.RuntimeEvent - sessions []agentruntime.SessionSummary - loads map[string]agentruntime.Session + sessions []agentsession.Summary + loads map[string]agentsession.Session runErr error compactErr error compactResult agentruntime.CompactResult listErr error loadErr error setWorkdirErr error - setResult *agentruntime.Session + setResult *agentsession.Session setCalls int resolveInputs []agentruntime.PermissionResolutionInput resolveErr error @@ -66,7 +67,7 @@ func (r *stubMarkdownRenderer) Render(content string, width int) (string, error) func newStubRuntime() *stubRuntime { return &stubRuntime{ events: make(chan agentruntime.RuntimeEvent, 16), - loads: map[string]agentruntime.Session{}, + loads: map[string]agentsession.Session{}, } } @@ -94,34 +95,34 @@ func (r *stubRuntime) CancelActiveRun() bool { return r.cancelResult } -func (r *stubRuntime) ListSessions(ctx context.Context) ([]agentruntime.SessionSummary, error) { +func (r *stubRuntime) ListSessions(ctx context.Context) ([]agentsession.Summary, error) { if r.listErr != nil { return nil, r.listErr } - return append([]agentruntime.SessionSummary(nil), r.sessions...), nil + return append([]agentsession.Summary(nil), r.sessions...), nil } -func (r *stubRuntime) LoadSession(ctx context.Context, id string) (agentruntime.Session, error) { +func (r *stubRuntime) LoadSession(ctx context.Context, id string) (agentsession.Session, error) { if r.loadErr != nil { - return agentruntime.Session{}, r.loadErr + return agentsession.Session{}, r.loadErr } if session, ok := r.loads[id]; ok { return session, nil } - return agentruntime.Session{}, nil + return agentsession.Session{}, nil } -func (r *stubRuntime) SetSessionWorkdir(ctx context.Context, sessionID string, workdir string) (agentruntime.Session, error) { +func (r *stubRuntime) SetSessionWorkdir(ctx context.Context, sessionID string, workdir string) (agentsession.Session, error) { r.setCalls++ if r.setWorkdirErr != nil { - return agentruntime.Session{}, r.setWorkdirErr + return agentsession.Session{}, r.setWorkdirErr } if r.setResult != nil { return *r.setResult, nil } session, ok := r.loads[sessionID] if !ok { - session = agentruntime.Session{ID: sessionID} + session = agentsession.Session{ID: sessionID} } session.Workdir = strings.TrimSpace(workdir) r.loads[sessionID] = session @@ -250,7 +251,7 @@ func TestAppUpdateWorkspaceSlashCommands(t *testing.T) { manager := newTestConfigManager(t) runtime := newStubRuntime() sessionID := "session-workdir" - runtime.loads[sessionID] = agentruntime.Session{ID: sessionID, Workdir: t.TempDir()} + runtime.loads[sessionID] = agentsession.Session{ID: sessionID, Workdir: t.TempDir()} app, err := New(nil, manager, runtime, newTestProviderService(t, manager)) if err != nil { @@ -327,7 +328,7 @@ func TestRunSessionWorkdirCommandBranches(t *testing.T) { t.Run("session workdir fallback uses current workdir when runtime returns empty", func(t *testing.T) { current := t.TempDir() runtime := newStubRuntime() - runtime.setResult = &agentruntime.Session{ID: "session-1", Workdir: ""} + runtime.setResult = &agentsession.Session{ID: "session-1", Workdir: ""} msg := runSessionWorkdirCommand(runtime, "session-1", current, "/cwd ./subdir")() result := msg.(sessionWorkdirResultMsg) if result.err != nil { @@ -342,7 +343,7 @@ func TestRunSessionWorkdirCommandBranches(t *testing.T) { current := t.TempDir() target := t.TempDir() runtime := newStubRuntime() - runtime.setResult = &agentruntime.Session{ID: "session-1", Workdir: target} + runtime.setResult = &agentsession.Session{ID: "session-1", Workdir: target} msg := runSessionWorkdirCommand(runtime, "session-1", current, "/cwd ./subdir")() result := msg.(sessionWorkdirResultMsg) if result.err != nil { @@ -775,7 +776,7 @@ func TestAppUpdateModelPickerAndRuntimeMessages(t *testing.T) { func TestAppHelpersAndRenderingSmoke(t *testing.T) { manager := newTestConfigManager(t) runtime := newStubRuntime() - now := agentruntime.Session{ + now := agentsession.Session{ ID: "session-1", Title: "Existing Session", Messages: []provider.Message{ @@ -783,7 +784,7 @@ func TestAppHelpersAndRenderingSmoke(t *testing.T) { {Role: roleAssistant, Content: "hello"}, }, } - runtime.sessions = []agentruntime.SessionSummary{ + runtime.sessions = []agentsession.Summary{ {ID: now.ID, Title: now.Title, UpdatedAt: now.UpdatedAt}, } runtime.loads[now.ID] = now @@ -985,7 +986,7 @@ func TestTUIStandaloneHelpers(t *testing.T) { t.Fatalf("expected numeric helpers to work") } - sItem := sessionItem{Summary: agentruntime.SessionSummary{Title: "My Session"}} + sItem := sessionItem{Summary: agentsession.Summary{Title: "My Session"}} if sItem.FilterValue() != "my session" { t.Fatalf("unexpected session item filter value") } @@ -1207,8 +1208,8 @@ func TestAppUpdateAdditionalTransitions(t *testing.T) { { name: "session enter activates selected session", setup: func(t *testing.T, app *App, runtime *stubRuntime, manager *config.Manager) { - runtime.sessions = []agentruntime.SessionSummary{{ID: "s1", Title: "One"}} - runtime.loads["s1"] = agentruntime.Session{ + runtime.sessions = []agentsession.Summary{{ID: "s1", Title: "One"}} + runtime.loads["s1"] = agentsession.Session{ ID: "s1", Title: "One", Messages: []provider.Message{{Role: roleAssistant, Content: "loaded"}}, From ffce7c3814e2280504081ee74a68deccb55ff1da Mon Sep 17 00:00:00 2001 From: Yumiue <229866007@qq.com> Date: Tue, 7 Apr 2026 18:04:04 +0800 Subject: [PATCH 31/55] =?UTF-8?q?test:=E8=A1=A5=E5=85=85=E6=B5=8B=E8=AF=95?= =?UTF-8?q?=E6=96=87=E4=BB=B6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/session/id_test.go | 41 ++++++++++++++++++++++++++ internal/session/store_test.go | 54 ++++++++++++++++++++++++++++++++++ 2 files changed, 95 insertions(+) create mode 100644 internal/session/id_test.go diff --git a/internal/session/id_test.go b/internal/session/id_test.go new file mode 100644 index 00000000..8ef1bb51 --- /dev/null +++ b/internal/session/id_test.go @@ -0,0 +1,41 @@ +package session + +import ( + "strings" + "testing" +) + +func TestNewIDFormatAndUniqueness(t *testing.T) { + t.Parallel() + + id1 := NewID("session") + id2 := NewID("session") + + if !strings.HasPrefix(id1, "session_") || !strings.HasPrefix(id2, "session_") { + t.Fatalf("expected prefix session_, got %q and %q", id1, id2) + } + + hex1 := strings.TrimPrefix(id1, "session_") + hex2 := strings.TrimPrefix(id2, "session_") + + if len(hex1) != 16 || len(hex2) != 16 { + t.Fatalf("expected 16 hex chars, got %d and %d", len(hex1), len(hex2)) + } + for _, ch := range hex1 + hex2 { + if !((ch >= '0' && ch <= '9') || (ch >= 'a' && ch <= 'f')) { + t.Fatalf("expected lowercase hex, got %q in ids %q %q", ch, id1, id2) + } + } + if id1 == id2 { + t.Fatalf("expected different ids, got identical %q", id1) + } +} + +func TestNewIDAllowsEmptyPrefix(t *testing.T) { + t.Parallel() + + id := NewID("") + if len(id) != 17 || id[0] != '_' { + t.Fatalf("expected format _<16hex>, got %q", id) + } +} diff --git a/internal/session/store_test.go b/internal/session/store_test.go index 0e5e02d4..75161164 100644 --- a/internal/session/store_test.go +++ b/internal/session/store_test.go @@ -162,6 +162,60 @@ func TestJSONStoreSaveInvalidBaseDir(t *testing.T) { } } +func TestNewUsesDefaultWorkdirAndEmptyMessages(t *testing.T) { + t.Parallel() + + session := New("hello title") + + if session.ID == "" { + t.Fatalf("expected non-empty id") + } + if !strings.HasPrefix(session.ID, "session_") { + t.Fatalf("expected id with session_ prefix, got %q", session.ID) + } + if session.Title != "hello title" { + t.Fatalf("expected title %q, got %q", "hello title", session.Title) + } + if session.Workdir != "" { + t.Fatalf("expected empty workdir, got %q", session.Workdir) + } + if len(session.Messages) != 0 { + t.Fatalf("expected empty messages, got %+v", session.Messages) + } + if session.CreatedAt.IsZero() || session.UpdatedAt.IsZero() { + t.Fatalf("expected non-zero timestamps, got created=%v updated=%v", session.CreatedAt, session.UpdatedAt) + } + if session.UpdatedAt.Before(session.CreatedAt) { + t.Fatalf("expected UpdatedAt >= CreatedAt, got created=%v updated=%v", session.CreatedAt, session.UpdatedAt) + } +} + +func TestNewWithWorkdirTrimAndTitleSanitize(t *testing.T) { + t.Parallel() + + tooLong := strings.Repeat("中", 45) // rune 长度 > 40 + workdir := " /tmp/workdir " + + session := NewWithWorkdir(tooLong, workdir) + + if session.Workdir != "/tmp/workdir" { + t.Fatalf("expected trimmed workdir %q, got %q", "/tmp/workdir", session.Workdir) + } + if got := len([]rune(session.Title)); got != 40 { + t.Fatalf("expected title rune length 40, got %d (title=%q)", got, session.Title) + } +} + +func TestNewWithWorkdirFallsBackDefaultTitle(t *testing.T) { + t.Parallel() + + session := NewWithWorkdir(" \n\t ", "") + + if session.Title != "New Session" { + t.Fatalf("expected default title %q, got %q", "New Session", session.Title) + } +} + func mustWriteSessionFile(t *testing.T, path string, content string) { t.Helper() if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil { From 48ba7cf823e6b5a9b33baa9d50fb4aaa46d5f768 Mon Sep 17 00:00:00 2001 From: Cai_Tang <106404101+Cai-Tang-www@users.noreply.github.com> Date: Tue, 7 Apr 2026 12:05:49 +0800 Subject: [PATCH 32/55] =?UTF-8?q?feat:=20=E6=96=B0=E5=A2=9E=20MCP=20Regist?= =?UTF-8?q?ry=20=E4=B8=8E=20Adapter=20=E5=9F=BA=E7=A1=80=E8=83=BD=E5=8A=9B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/tools/mcp/adapter.go | 171 +++++++++++++++ internal/tools/mcp/adapter_test.go | 135 ++++++++++++ internal/tools/mcp/registry.go | 319 ++++++++++++++++++++++++++++ internal/tools/mcp/registry_test.go | 163 ++++++++++++++ 4 files changed, 788 insertions(+) create mode 100644 internal/tools/mcp/adapter.go create mode 100644 internal/tools/mcp/adapter_test.go create mode 100644 internal/tools/mcp/registry.go create mode 100644 internal/tools/mcp/registry_test.go diff --git a/internal/tools/mcp/adapter.go b/internal/tools/mcp/adapter.go new file mode 100644 index 00000000..e238e07e --- /dev/null +++ b/internal/tools/mcp/adapter.go @@ -0,0 +1,171 @@ +package mcp + +import ( + "context" + "errors" + "fmt" + "strings" + + "neo-code/internal/tools" +) + +const mcpToolNamePrefix = "mcp." + +// AdapterFactory 基于 registry 快照构造 MCP tool 适配器集合。 +type AdapterFactory struct { + registry *Registry +} + +// NewAdapterFactory 创建 MCP adapter 工厂。 +func NewAdapterFactory(registry *Registry) *AdapterFactory { + return &AdapterFactory{registry: registry} +} + +// BuildTools 将当前所有 MCP tool 快照转换为统一 tools.Tool 列表。 +func (f *AdapterFactory) BuildTools(ctx context.Context) ([]tools.Tool, error) { + if f == nil || f.registry == nil { + return nil, errors.New("mcp: adapter factory registry is nil") + } + if err := ctx.Err(); err != nil { + return nil, err + } + + snapshots := f.registry.Snapshot() + if len(snapshots) == 0 { + return nil, nil + } + + result := make([]tools.Tool, 0, len(snapshots)*2) + for _, snapshot := range snapshots { + for _, descriptor := range snapshot.Tools { + adapter, err := NewAdapter(f.registry, snapshot.ServerID, descriptor) + if err != nil { + return nil, err + } + result = append(result, adapter) + } + } + return result, nil +} + +// Adapter 将单个 MCP tool 适配为统一 tools.Tool 接口。 +type Adapter struct { + registry *Registry + serverID string + toolName string + description string + schema map[string]any +} + +// NewAdapter 创建指定 server/tool 的 MCP 适配器。 +func NewAdapter(registry *Registry, serverID string, descriptor ToolDescriptor) (*Adapter, error) { + if registry == nil { + return nil, errors.New("mcp: registry is nil") + } + normalizedServerID := normalizeServerID(serverID) + if normalizedServerID == "" { + return nil, errors.New("mcp: server id is empty") + } + normalizedToolName := strings.TrimSpace(descriptor.Name) + if normalizedToolName == "" { + return nil, errors.New("mcp: descriptor tool name is empty") + } + + return &Adapter{ + registry: registry, + serverID: normalizedServerID, + toolName: normalizedToolName, + description: strings.TrimSpace(descriptor.Description), + schema: ensureObjectSchema(descriptor.InputSchema), + }, nil +} + +// Name 返回统一的 MCP tool 名称:mcp..。 +func (a *Adapter) Name() string { + return composeToolName(a.serverID, a.toolName) +} + +// Description 返回工具描述,不存在时回退到稳定默认文案。 +func (a *Adapter) Description() string { + if strings.TrimSpace(a.description) != "" { + return a.description + } + return fmt.Sprintf("MCP tool %s from server %s", a.toolName, a.serverID) +} + +// Schema 返回 MCP 工具输入 schema 的标准对象结构。 +func (a *Adapter) Schema() map[string]any { + return cloneSchema(a.schema) +} + +// MicroCompactPolicy 返回 MCP tool 历史结果默认 micro compact 策略。 +func (a *Adapter) MicroCompactPolicy() tools.MicroCompactPolicy { + return tools.MicroCompactPolicyCompact +} + +// Execute 分发 MCP tool 调用并收敛为统一 ToolResult。 +func (a *Adapter) Execute(ctx context.Context, call tools.ToolCallInput) (tools.ToolResult, error) { + if a == nil || a.registry == nil { + err := errors.New("mcp: adapter is not initialized") + return tools.NewErrorResult("mcp", tools.NormalizeErrorReason("mcp", err), "", nil), err + } + if err := ctx.Err(); err != nil { + return tools.NewErrorResult(a.Name(), tools.NormalizeErrorReason(a.Name(), err), "", adapterMetadata(a.serverID, a.toolName)), err + } + + result, err := a.registry.Call(ctx, a.serverID, a.toolName, call.Arguments) + if err != nil { + errorResult := tools.NewErrorResult(a.Name(), tools.NormalizeErrorReason(a.Name(), err), "", adapterMetadata(a.serverID, a.toolName)) + errorResult.ToolCallID = call.ID + return errorResult, err + } + + metadata := adapterMetadata(a.serverID, a.toolName) + for key, value := range result.Metadata { + metadata[key] = value + } + + toolResult := tools.ToolResult{ + ToolCallID: call.ID, + Name: a.Name(), + Content: strings.TrimSpace(result.Content), + IsError: result.IsError, + Metadata: metadata, + } + if strings.TrimSpace(toolResult.Content) == "" { + toolResult.Content = "ok" + } + return tools.ApplyOutputLimit(toolResult, tools.DefaultOutputLimitBytes), nil +} + +// composeToolName 组装统一的 MCP tool 名称,保持权限映射可预测。 +func composeToolName(serverID string, toolName string) string { + return mcpToolNamePrefix + normalizeServerID(serverID) + "." + strings.TrimSpace(toolName) +} + +// ensureObjectSchema 确保 schema 至少是 object,避免上层 provider 解析异常。 +func ensureObjectSchema(schema map[string]any) map[string]any { + cloned := cloneSchema(schema) + if len(cloned) == 0 { + return map[string]any{ + "type": "object", + "properties": map[string]any{}, + } + } + + if strings.TrimSpace(fmt.Sprintf("%v", cloned["type"])) == "" { + cloned["type"] = "object" + } + if _, ok := cloned["properties"]; !ok { + cloned["properties"] = map[string]any{} + } + return cloned +} + +// adapterMetadata 生成 MCP 调用结果的基础元信息。 +func adapterMetadata(serverID string, toolName string) map[string]any { + return map[string]any{ + "mcp_server_id": normalizeServerID(serverID), + "mcp_tool_name": strings.TrimSpace(toolName), + } +} diff --git a/internal/tools/mcp/adapter_test.go b/internal/tools/mcp/adapter_test.go new file mode 100644 index 00000000..689075b3 --- /dev/null +++ b/internal/tools/mcp/adapter_test.go @@ -0,0 +1,135 @@ +package mcp + +import ( + "context" + "errors" + "testing" + + "neo-code/internal/tools" +) + +func TestAdapterFactoryBuildTools(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + client := &stubServerClient{ + tools: []ToolDescriptor{ + { + Name: "search", + Description: "search docs", + InputSchema: map[string]any{"type": "object"}, + }, + }, + } + if err := registry.RegisterServer("docs", "stdio", "v1", client); err != nil { + t.Fatalf("register server: %v", err) + } + if err := registry.RefreshServerTools(context.Background(), "docs"); err != nil { + t.Fatalf("refresh tools: %v", err) + } + + factory := NewAdapterFactory(registry) + toolsList, err := factory.BuildTools(context.Background()) + if err != nil { + t.Fatalf("BuildTools() error = %v", err) + } + if len(toolsList) != 1 { + t.Fatalf("expected one adapter tool, got %d", len(toolsList)) + } + if toolsList[0].Name() != "mcp.docs.search" { + t.Fatalf("unexpected adapter tool name: %q", toolsList[0].Name()) + } +} + +func TestAdapterExecute(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + client := &stubServerClient{ + tools: []ToolDescriptor{ + {Name: "search", InputSchema: map[string]any{"type": "object"}}, + }, + callResult: CallResult{ + Content: "result body", + Metadata: map[string]any{ + "latency_ms": 20, + }, + }, + } + if err := registry.RegisterServer("docs", "stdio", "v1", client); err != nil { + t.Fatalf("register server: %v", err) + } + if err := registry.RefreshServerTools(context.Background(), "docs"); err != nil { + t.Fatalf("refresh tools: %v", err) + } + + adapter, err := NewAdapter(registry, "docs", ToolDescriptor{ + Name: "search", + Description: "search docs", + InputSchema: map[string]any{"type": "object"}, + }) + if err != nil { + t.Fatalf("NewAdapter() error = %v", err) + } + + result, err := adapter.Execute(context.Background(), tools.ToolCallInput{ + ID: "tool-call-1", + Name: adapter.Name(), + Arguments: []byte(`{"q":"mcp"}`), + }) + if err != nil { + t.Fatalf("Execute() error = %v", err) + } + if result.ToolCallID != "tool-call-1" { + t.Fatalf("expected tool call id tool-call-1, got %q", result.ToolCallID) + } + if result.Name != "mcp.docs.search" { + t.Fatalf("expected tool name mcp.docs.search, got %q", result.Name) + } + if result.Content != "result body" { + t.Fatalf("expected result content, got %q", result.Content) + } + if result.Metadata["mcp_server_id"] != "docs" || result.Metadata["mcp_tool_name"] != "search" { + t.Fatalf("unexpected metadata: %+v", result.Metadata) + } +} + +func TestAdapterExecuteErrorMapping(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + client := &stubServerClient{ + tools: []ToolDescriptor{ + {Name: "search", InputSchema: map[string]any{"type": "object"}}, + }, + callErr: errors.New("transport timeout"), + } + if err := registry.RegisterServer("docs", "stdio", "v1", client); err != nil { + t.Fatalf("register server: %v", err) + } + if err := registry.RefreshServerTools(context.Background(), "docs"); err != nil { + t.Fatalf("refresh tools: %v", err) + } + + adapter, err := NewAdapter(registry, "docs", ToolDescriptor{ + Name: "search", + }) + if err != nil { + t.Fatalf("NewAdapter() error = %v", err) + } + + result, execErr := adapter.Execute(context.Background(), tools.ToolCallInput{ + ID: "tool-call-error", + Name: adapter.Name(), + Arguments: []byte(`{"q":"mcp"}`), + }) + if execErr == nil { + t.Fatalf("expected execute error") + } + if !result.IsError { + t.Fatalf("expected error result, got %+v", result) + } + if result.Metadata["mcp_server_id"] != "docs" { + t.Fatalf("unexpected metadata for error result: %+v", result.Metadata) + } +} diff --git a/internal/tools/mcp/registry.go b/internal/tools/mcp/registry.go new file mode 100644 index 00000000..026983ec --- /dev/null +++ b/internal/tools/mcp/registry.go @@ -0,0 +1,319 @@ +package mcp + +import ( + "context" + "errors" + "fmt" + "sort" + "strings" + "sync" + "time" +) + +// ServerStatus 描述 MCP server 在 registry 中的生命周期状态。 +type ServerStatus string + +const ( + // ServerStatusConnecting 表示 server 已注册但仍在连接或初始化阶段。 + ServerStatusConnecting ServerStatus = "connecting" + // ServerStatusReady 表示 server 已可用,支持正常工具调用。 + ServerStatusReady ServerStatus = "ready" + // ServerStatusDegraded 表示 server 可部分服务,但存在健康或调用异常。 + ServerStatusDegraded ServerStatus = "degraded" + // ServerStatusOffline 表示 server 当前不可用。 + ServerStatusOffline ServerStatus = "offline" +) + +// ToolDescriptor 描述 MCP tool 的稳定元信息与输入 schema。 +type ToolDescriptor struct { + Name string + Description string + InputSchema map[string]any +} + +// ServerSnapshot 描述 registry 对外暴露的 server 只读快照。 +type ServerSnapshot struct { + ServerID string + Source string + Version string + Status ServerStatus + UpdatedAt time.Time + Tools []ToolDescriptor +} + +// CallResult 收敛 MCP tool 调用后的统一结果语义。 +type CallResult struct { + Content string + IsError bool + Metadata map[string]any +} + +// ServerClient 描述 registry 与具体 MCP server 交互所需的最小能力。 +type ServerClient interface { + ListTools(ctx context.Context) ([]ToolDescriptor, error) + CallTool(ctx context.Context, toolName string, arguments []byte) (CallResult, error) + HealthCheck(ctx context.Context) error +} + +type serverEntry struct { + snapshot ServerSnapshot + client ServerClient +} + +// Registry 维护 MCP server 注册、快照读取和工具调用分发。 +type Registry struct { + mu sync.RWMutex + servers map[string]*serverEntry +} + +// NewRegistry 创建线程安全的 MCP registry 实例。 +func NewRegistry() *Registry { + return &Registry{ + servers: make(map[string]*serverEntry), + } +} + +// RegisterServer 注册一个 MCP server,并初始化其生命周期状态。 +func (r *Registry) RegisterServer(serverID string, source string, version string, client ServerClient) error { + if r == nil { + return errors.New("mcp: registry is nil") + } + normalizedID := normalizeServerID(serverID) + if normalizedID == "" { + return errors.New("mcp: server id is empty") + } + if client == nil { + return errors.New("mcp: server client is nil") + } + + r.mu.Lock() + defer r.mu.Unlock() + + if _, exists := r.servers[normalizedID]; exists { + return fmt.Errorf("mcp: server %q already exists", normalizedID) + } + r.servers[normalizedID] = &serverEntry{ + snapshot: ServerSnapshot{ + ServerID: normalizedID, + Source: strings.TrimSpace(source), + Version: strings.TrimSpace(version), + Status: ServerStatusConnecting, + UpdatedAt: time.Now(), + }, + client: client, + } + return nil +} + +// UnregisterServer 注销一个 MCP server,返回是否实际删除。 +func (r *Registry) UnregisterServer(serverID string) bool { + if r == nil { + return false + } + normalizedID := normalizeServerID(serverID) + if normalizedID == "" { + return false + } + + r.mu.Lock() + defer r.mu.Unlock() + if _, exists := r.servers[normalizedID]; !exists { + return false + } + delete(r.servers, normalizedID) + return true +} + +// SetServerStatus 更新指定 server 的生命周期状态。 +func (r *Registry) SetServerStatus(serverID string, status ServerStatus) error { + if r == nil { + return errors.New("mcp: registry is nil") + } + if !isValidStatus(status) { + return fmt.Errorf("mcp: unsupported server status %q", status) + } + normalizedID := normalizeServerID(serverID) + if normalizedID == "" { + return errors.New("mcp: server id is empty") + } + + r.mu.Lock() + defer r.mu.Unlock() + entry, ok := r.servers[normalizedID] + if !ok { + return fmt.Errorf("mcp: server %q not found", normalizedID) + } + entry.snapshot.Status = status + entry.snapshot.UpdatedAt = time.Now() + return nil +} + +// RefreshServerTools 从 server 拉取工具清单并刷新快照。 +func (r *Registry) RefreshServerTools(ctx context.Context, serverID string) error { + if r == nil { + return errors.New("mcp: registry is nil") + } + if err := ctx.Err(); err != nil { + return err + } + normalizedID := normalizeServerID(serverID) + if normalizedID == "" { + return errors.New("mcp: server id is empty") + } + + r.mu.RLock() + entry, ok := r.servers[normalizedID] + r.mu.RUnlock() + if !ok { + return fmt.Errorf("mcp: server %q not found", normalizedID) + } + + tools, err := entry.client.ListTools(ctx) + if err != nil { + _ = r.SetServerStatus(normalizedID, ServerStatusDegraded) + return fmt.Errorf("mcp: list tools for server %q: %w", normalizedID, err) + } + + r.mu.Lock() + defer r.mu.Unlock() + current, exists := r.servers[normalizedID] + if !exists { + return fmt.Errorf("mcp: server %q not found", normalizedID) + } + current.snapshot.Tools = cloneToolDescriptors(tools) + current.snapshot.Status = ServerStatusReady + current.snapshot.UpdatedAt = time.Now() + return nil +} + +// HealthCheck 触发指定 server 的健康探测并同步状态。 +func (r *Registry) HealthCheck(ctx context.Context, serverID string) error { + if r == nil { + return errors.New("mcp: registry is nil") + } + if err := ctx.Err(); err != nil { + return err + } + normalizedID := normalizeServerID(serverID) + if normalizedID == "" { + return errors.New("mcp: server id is empty") + } + + r.mu.RLock() + entry, ok := r.servers[normalizedID] + r.mu.RUnlock() + if !ok { + return fmt.Errorf("mcp: server %q not found", normalizedID) + } + + if err := entry.client.HealthCheck(ctx); err != nil { + _ = r.SetServerStatus(normalizedID, ServerStatusOffline) + return fmt.Errorf("mcp: health check failed for server %q: %w", normalizedID, err) + } + return r.SetServerStatus(normalizedID, ServerStatusReady) +} + +// Call 通过 registry 分发指定 server/tool 的调用请求。 +func (r *Registry) Call(ctx context.Context, serverID string, toolName string, arguments []byte) (CallResult, error) { + if r == nil { + return CallResult{}, errors.New("mcp: registry is nil") + } + if err := ctx.Err(); err != nil { + return CallResult{}, err + } + normalizedID := normalizeServerID(serverID) + if normalizedID == "" { + return CallResult{}, errors.New("mcp: server id is empty") + } + trimmedToolName := strings.TrimSpace(toolName) + if trimmedToolName == "" { + return CallResult{}, errors.New("mcp: tool name is empty") + } + + r.mu.RLock() + entry, ok := r.servers[normalizedID] + r.mu.RUnlock() + if !ok { + return CallResult{}, fmt.Errorf("mcp: server %q not found", normalizedID) + } + + result, err := entry.client.CallTool(ctx, trimmedToolName, arguments) + if err != nil { + _ = r.SetServerStatus(normalizedID, ServerStatusDegraded) + return CallResult{}, fmt.Errorf("mcp: call %s on %s failed: %w", trimmedToolName, normalizedID, err) + } + return result, nil +} + +// Snapshot 返回当前 registry 的不可变 server 快照集合。 +func (r *Registry) Snapshot() []ServerSnapshot { + if r == nil { + return nil + } + + r.mu.RLock() + defer r.mu.RUnlock() + if len(r.servers) == 0 { + return nil + } + + keys := make([]string, 0, len(r.servers)) + for serverID := range r.servers { + keys = append(keys, serverID) + } + sort.Strings(keys) + + result := make([]ServerSnapshot, 0, len(keys)) + for _, serverID := range keys { + entry := r.servers[serverID] + snapshot := entry.snapshot + snapshot.Tools = cloneToolDescriptors(snapshot.Tools) + result = append(result, snapshot) + } + return result +} + +// normalizeServerID 统一规范化 server id 以保证匹配稳定性。 +func normalizeServerID(serverID string) string { + return strings.ToLower(strings.TrimSpace(serverID)) +} + +// isValidStatus 校验 server 状态是否属于已定义集合。 +func isValidStatus(status ServerStatus) bool { + switch status { + case ServerStatusConnecting, ServerStatusReady, ServerStatusDegraded, ServerStatusOffline: + return true + default: + return false + } +} + +// cloneToolDescriptors 深拷贝工具描述,避免快照被外部引用污染。 +func cloneToolDescriptors(input []ToolDescriptor) []ToolDescriptor { + if len(input) == 0 { + return nil + } + + result := make([]ToolDescriptor, 0, len(input)) + for _, descriptor := range input { + cloned := ToolDescriptor{ + Name: strings.TrimSpace(descriptor.Name), + Description: strings.TrimSpace(descriptor.Description), + InputSchema: cloneSchema(descriptor.InputSchema), + } + result = append(result, cloned) + } + return result +} + +// cloneSchema 深拷贝 schema 顶层 map,满足当前工具定义的只读需求。 +func cloneSchema(schema map[string]any) map[string]any { + if len(schema) == 0 { + return nil + } + cloned := make(map[string]any, len(schema)) + for key, value := range schema { + cloned[key] = value + } + return cloned +} diff --git a/internal/tools/mcp/registry_test.go b/internal/tools/mcp/registry_test.go new file mode 100644 index 00000000..d7f5f26d --- /dev/null +++ b/internal/tools/mcp/registry_test.go @@ -0,0 +1,163 @@ +package mcp + +import ( + "context" + "errors" + "sync" + "testing" + "time" +) + +type stubServerClient struct { + mu sync.Mutex + tools []ToolDescriptor + callResult CallResult + listErr error + callErr error + healthErr error + lastToolName string + lastArguments []byte +} + +func (s *stubServerClient) ListTools(ctx context.Context) ([]ToolDescriptor, error) { + if s.listErr != nil { + return nil, s.listErr + } + return cloneToolDescriptors(s.tools), nil +} + +func (s *stubServerClient) CallTool(ctx context.Context, toolName string, arguments []byte) (CallResult, error) { + s.mu.Lock() + s.lastToolName = toolName + s.lastArguments = append([]byte(nil), arguments...) + s.mu.Unlock() + if s.callErr != nil { + return CallResult{}, s.callErr + } + result := s.callResult + result.Metadata = cloneSchema(result.Metadata) + return result, nil +} + +func (s *stubServerClient) HealthCheck(ctx context.Context) error { + return s.healthErr +} + +func TestRegistryRegisterRefreshSnapshotCall(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + client := &stubServerClient{ + tools: []ToolDescriptor{ + { + Name: "search", + Description: "search docs", + InputSchema: map[string]any{"type": "object"}, + }, + }, + callResult: CallResult{ + Content: "ok", + Metadata: map[string]any{ + "latency_ms": 18, + }, + }, + } + + if err := registry.RegisterServer("Docs", "stdio", "v1", client); err != nil { + t.Fatalf("RegisterServer() error = %v", err) + } + if err := registry.RefreshServerTools(context.Background(), "docs"); err != nil { + t.Fatalf("RefreshServerTools() error = %v", err) + } + + snapshots := registry.Snapshot() + if len(snapshots) != 1 { + t.Fatalf("expected one snapshot, got %d", len(snapshots)) + } + snapshot := snapshots[0] + if snapshot.ServerID != "docs" || snapshot.Status != ServerStatusReady { + t.Fatalf("unexpected snapshot: %+v", snapshot) + } + if len(snapshot.Tools) != 1 || snapshot.Tools[0].Name != "search" { + t.Fatalf("unexpected tools in snapshot: %+v", snapshot.Tools) + } + + result, err := registry.Call(context.Background(), "docs", "search", []byte(`{"q":"mcp"}`)) + if err != nil { + t.Fatalf("Call() error = %v", err) + } + if result.Content != "ok" { + t.Fatalf("expected call content ok, got %q", result.Content) + } +} + +func TestRegistryStatusTransitions(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + client := &stubServerClient{ + tools: []ToolDescriptor{ + {Name: "search", InputSchema: map[string]any{"type": "object"}}, + }, + } + if err := registry.RegisterServer("server-1", "stdio", "v1", client); err != nil { + t.Fatalf("register server: %v", err) + } + + client.healthErr = errors.New("offline") + if err := registry.HealthCheck(context.Background(), "server-1"); err == nil { + t.Fatalf("expected health check failure") + } + if snapshots := registry.Snapshot(); snapshots[0].Status != ServerStatusOffline { + t.Fatalf("expected offline status, got %+v", snapshots[0].Status) + } + + client.healthErr = nil + if err := registry.HealthCheck(context.Background(), "server-1"); err != nil { + t.Fatalf("unexpected health check error: %v", err) + } + if snapshots := registry.Snapshot(); snapshots[0].Status != ServerStatusReady { + t.Fatalf("expected ready status, got %+v", snapshots[0].Status) + } +} + +func TestRegistryConcurrentSnapshotAndRefresh(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + client := &stubServerClient{ + tools: []ToolDescriptor{ + {Name: "search", InputSchema: map[string]any{"type": "object"}}, + }, + } + if err := registry.RegisterServer("server-1", "stdio", "v1", client); err != nil { + t.Fatalf("register server: %v", err) + } + + var wg sync.WaitGroup + for i := 0; i < 8; i++ { + wg.Add(1) + go func() { + defer wg.Done() + _ = registry.RefreshServerTools(context.Background(), "server-1") + }() + } + for i := 0; i < 16; i++ { + wg.Add(1) + go func() { + defer wg.Done() + _ = registry.Snapshot() + }() + } + done := make(chan struct{}) + go func() { + wg.Wait() + close(done) + }() + + select { + case <-done: + case <-time.After(2 * time.Second): + t.Fatalf("concurrent registry operations timed out") + } +} From e6fb96eb3eaff28d820aeebcc4a18cf954e5d456 Mon Sep 17 00:00:00 2001 From: Cai_Tang <106404101+Cai-Tang-www@users.noreply.github.com> Date: Tue, 7 Apr 2026 12:06:03 +0800 Subject: [PATCH 33/55] =?UTF-8?q?feat:=20=E5=B0=86=20MCP=20=E5=B7=A5?= =?UTF-8?q?=E5=85=B7=E6=8E=A5=E5=85=A5=20Registry=20=E4=B8=BB=E6=89=A7?= =?UTF-8?q?=E8=A1=8C=E9=93=BE=E5=B9=B6=E5=A2=9E=E5=8A=A0=20stdio=20?= =?UTF-8?q?=E5=AE=A2=E6=88=B7=E7=AB=AF?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/tools/mcp/adapter.go | 71 +--- internal/tools/mcp/adapter_test.go | 52 +-- internal/tools/mcp/stdio_client.go | 499 ++++++++++++++++++++++++ internal/tools/mcp/stdio_client_test.go | 139 +++++++ internal/tools/registry.go | 136 ++++++- internal/tools/registry_test.go | 105 +++++ 6 files changed, 903 insertions(+), 99 deletions(-) create mode 100644 internal/tools/mcp/stdio_client.go create mode 100644 internal/tools/mcp/stdio_client_test.go diff --git a/internal/tools/mcp/adapter.go b/internal/tools/mcp/adapter.go index e238e07e..d9ecb470 100644 --- a/internal/tools/mcp/adapter.go +++ b/internal/tools/mcp/adapter.go @@ -5,8 +5,6 @@ import ( "errors" "fmt" "strings" - - "neo-code/internal/tools" ) const mcpToolNamePrefix = "mcp." @@ -21,8 +19,8 @@ func NewAdapterFactory(registry *Registry) *AdapterFactory { return &AdapterFactory{registry: registry} } -// BuildTools 将当前所有 MCP tool 快照转换为统一 tools.Tool 列表。 -func (f *AdapterFactory) BuildTools(ctx context.Context) ([]tools.Tool, error) { +// BuildAdapters 将当前所有 MCP tool 快照转换为 Adapter 列表。 +func (f *AdapterFactory) BuildAdapters(ctx context.Context) ([]*Adapter, error) { if f == nil || f.registry == nil { return nil, errors.New("mcp: adapter factory registry is nil") } @@ -35,7 +33,7 @@ func (f *AdapterFactory) BuildTools(ctx context.Context) ([]tools.Tool, error) { return nil, nil } - result := make([]tools.Tool, 0, len(snapshots)*2) + result := make([]*Adapter, 0, len(snapshots)*2) for _, snapshot := range snapshots { for _, descriptor := range snapshot.Tools { adapter, err := NewAdapter(f.registry, snapshot.ServerID, descriptor) @@ -48,7 +46,7 @@ func (f *AdapterFactory) BuildTools(ctx context.Context) ([]tools.Tool, error) { return result, nil } -// Adapter 将单个 MCP tool 适配为统一 tools.Tool 接口。 +// Adapter 将单个 MCP tool 适配为统一调用描述。 type Adapter struct { registry *Registry serverID string @@ -80,11 +78,21 @@ func NewAdapter(registry *Registry, serverID string, descriptor ToolDescriptor) }, nil } -// Name 返回统一的 MCP tool 名称:mcp..。 -func (a *Adapter) Name() string { +// FullName 返回统一的 MCP tool 名称:mcp..。 +func (a *Adapter) FullName() string { return composeToolName(a.serverID, a.toolName) } +// ServerID 返回 MCP server 标识。 +func (a *Adapter) ServerID() string { + return a.serverID +} + +// ToolName 返回 MCP tool 原始名称。 +func (a *Adapter) ToolName() string { + return a.toolName +} + // Description 返回工具描述,不存在时回退到稳定默认文案。 func (a *Adapter) Description() string { if strings.TrimSpace(a.description) != "" { @@ -98,44 +106,15 @@ func (a *Adapter) Schema() map[string]any { return cloneSchema(a.schema) } -// MicroCompactPolicy 返回 MCP tool 历史结果默认 micro compact 策略。 -func (a *Adapter) MicroCompactPolicy() tools.MicroCompactPolicy { - return tools.MicroCompactPolicyCompact -} - -// Execute 分发 MCP tool 调用并收敛为统一 ToolResult。 -func (a *Adapter) Execute(ctx context.Context, call tools.ToolCallInput) (tools.ToolResult, error) { +// Call 分发 MCP tool 调用并返回统一结果。 +func (a *Adapter) Call(ctx context.Context, arguments []byte) (CallResult, error) { if a == nil || a.registry == nil { - err := errors.New("mcp: adapter is not initialized") - return tools.NewErrorResult("mcp", tools.NormalizeErrorReason("mcp", err), "", nil), err + return CallResult{}, errors.New("mcp: adapter is not initialized") } if err := ctx.Err(); err != nil { - return tools.NewErrorResult(a.Name(), tools.NormalizeErrorReason(a.Name(), err), "", adapterMetadata(a.serverID, a.toolName)), err - } - - result, err := a.registry.Call(ctx, a.serverID, a.toolName, call.Arguments) - if err != nil { - errorResult := tools.NewErrorResult(a.Name(), tools.NormalizeErrorReason(a.Name(), err), "", adapterMetadata(a.serverID, a.toolName)) - errorResult.ToolCallID = call.ID - return errorResult, err + return CallResult{}, err } - - metadata := adapterMetadata(a.serverID, a.toolName) - for key, value := range result.Metadata { - metadata[key] = value - } - - toolResult := tools.ToolResult{ - ToolCallID: call.ID, - Name: a.Name(), - Content: strings.TrimSpace(result.Content), - IsError: result.IsError, - Metadata: metadata, - } - if strings.TrimSpace(toolResult.Content) == "" { - toolResult.Content = "ok" - } - return tools.ApplyOutputLimit(toolResult, tools.DefaultOutputLimitBytes), nil + return a.registry.Call(ctx, a.serverID, a.toolName, arguments) } // composeToolName 组装统一的 MCP tool 名称,保持权限映射可预测。 @@ -161,11 +140,3 @@ func ensureObjectSchema(schema map[string]any) map[string]any { } return cloned } - -// adapterMetadata 生成 MCP 调用结果的基础元信息。 -func adapterMetadata(serverID string, toolName string) map[string]any { - return map[string]any{ - "mcp_server_id": normalizeServerID(serverID), - "mcp_tool_name": strings.TrimSpace(toolName), - } -} diff --git a/internal/tools/mcp/adapter_test.go b/internal/tools/mcp/adapter_test.go index 689075b3..d5dd16f7 100644 --- a/internal/tools/mcp/adapter_test.go +++ b/internal/tools/mcp/adapter_test.go @@ -4,11 +4,9 @@ import ( "context" "errors" "testing" - - "neo-code/internal/tools" ) -func TestAdapterFactoryBuildTools(t *testing.T) { +func TestAdapterFactoryBuildAdapters(t *testing.T) { t.Parallel() registry := NewRegistry() @@ -29,19 +27,19 @@ func TestAdapterFactoryBuildTools(t *testing.T) { } factory := NewAdapterFactory(registry) - toolsList, err := factory.BuildTools(context.Background()) + adapters, err := factory.BuildAdapters(context.Background()) if err != nil { - t.Fatalf("BuildTools() error = %v", err) + t.Fatalf("BuildAdapters() error = %v", err) } - if len(toolsList) != 1 { - t.Fatalf("expected one adapter tool, got %d", len(toolsList)) + if len(adapters) != 1 { + t.Fatalf("expected one adapter, got %d", len(adapters)) } - if toolsList[0].Name() != "mcp.docs.search" { - t.Fatalf("unexpected adapter tool name: %q", toolsList[0].Name()) + if adapters[0].FullName() != "mcp.docs.search" { + t.Fatalf("unexpected adapter full name: %q", adapters[0].FullName()) } } -func TestAdapterExecute(t *testing.T) { +func TestAdapterCall(t *testing.T) { t.Parallel() registry := NewRegistry() @@ -72,29 +70,16 @@ func TestAdapterExecute(t *testing.T) { t.Fatalf("NewAdapter() error = %v", err) } - result, err := adapter.Execute(context.Background(), tools.ToolCallInput{ - ID: "tool-call-1", - Name: adapter.Name(), - Arguments: []byte(`{"q":"mcp"}`), - }) + result, err := adapter.Call(context.Background(), []byte(`{"q":"mcp"}`)) if err != nil { - t.Fatalf("Execute() error = %v", err) - } - if result.ToolCallID != "tool-call-1" { - t.Fatalf("expected tool call id tool-call-1, got %q", result.ToolCallID) - } - if result.Name != "mcp.docs.search" { - t.Fatalf("expected tool name mcp.docs.search, got %q", result.Name) + t.Fatalf("Call() error = %v", err) } if result.Content != "result body" { t.Fatalf("expected result content, got %q", result.Content) } - if result.Metadata["mcp_server_id"] != "docs" || result.Metadata["mcp_tool_name"] != "search" { - t.Fatalf("unexpected metadata: %+v", result.Metadata) - } } -func TestAdapterExecuteErrorMapping(t *testing.T) { +func TestAdapterCallError(t *testing.T) { t.Parallel() registry := NewRegistry() @@ -118,18 +103,7 @@ func TestAdapterExecuteErrorMapping(t *testing.T) { t.Fatalf("NewAdapter() error = %v", err) } - result, execErr := adapter.Execute(context.Background(), tools.ToolCallInput{ - ID: "tool-call-error", - Name: adapter.Name(), - Arguments: []byte(`{"q":"mcp"}`), - }) - if execErr == nil { - t.Fatalf("expected execute error") - } - if !result.IsError { - t.Fatalf("expected error result, got %+v", result) - } - if result.Metadata["mcp_server_id"] != "docs" { - t.Fatalf("unexpected metadata for error result: %+v", result.Metadata) + if _, err := adapter.Call(context.Background(), []byte(`{"q":"mcp"}`)); err == nil { + t.Fatalf("expected call error") } } diff --git a/internal/tools/mcp/stdio_client.go b/internal/tools/mcp/stdio_client.go new file mode 100644 index 00000000..7e92fe80 --- /dev/null +++ b/internal/tools/mcp/stdio_client.go @@ -0,0 +1,499 @@ +package mcp + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "os" + "os/exec" + "strconv" + "strings" + "sync" + "sync/atomic" + "time" +) + +const ( + defaultStdioStartTimeout = 5 * time.Second + defaultStdioCallTimeout = 15 * time.Second + defaultStdioRestartBackoff = 1 * time.Second + maxStdioRestartBackoff = 30 * time.Second +) + +// StdioClientConfig 描述 MCP stdio 客户端的启动与调用参数。 +type StdioClientConfig struct { + Command string + Args []string + Env []string + Workdir string + StartTimeout time.Duration + CallTimeout time.Duration + RestartBackoff time.Duration +} + +type jsonRPCRequest struct { + JSONRPC string `json:"jsonrpc"` + ID string `json:"id"` + Method string `json:"method"` + Params any `json:"params,omitempty"` +} + +type jsonRPCResponse struct { + JSONRPC string `json:"jsonrpc"` + ID string `json:"id"` + Result json.RawMessage `json:"result,omitempty"` + Error *jsonRPCError `json:"error,omitempty"` +} + +type jsonRPCError struct { + Code int `json:"code"` + Message string `json:"message"` +} + +type rpcReply struct { + result json.RawMessage + err error +} + +// StdIOClient 通过 stdio 子进程与 MCP server 进行 JSON-RPC 通信。 +type StdIOClient struct { + cfg StdioClientConfig + idSeed uint64 + mu sync.Mutex + cmd *exec.Cmd + stdin io.WriteCloser + stdout io.ReadCloser + reader *bufio.Reader + pending map[string]chan rpcReply + exited chan struct{} + exitErr error + backoff time.Duration + retryAt time.Time + started bool + shutdown bool +} + +// NewStdIOClient 创建 stdio MCP client。 +func NewStdIOClient(cfg StdioClientConfig) (*StdIOClient, error) { + if strings.TrimSpace(cfg.Command) == "" { + return nil, errors.New("mcp: stdio command is empty") + } + if cfg.StartTimeout <= 0 { + cfg.StartTimeout = defaultStdioStartTimeout + } + if cfg.CallTimeout <= 0 { + cfg.CallTimeout = defaultStdioCallTimeout + } + if cfg.RestartBackoff <= 0 { + cfg.RestartBackoff = defaultStdioRestartBackoff + } + + return &StdIOClient{ + cfg: cfg, + pending: make(map[string]chan rpcReply), + backoff: cfg.RestartBackoff, + }, nil +} + +// Close 关闭 stdio 子进程并释放资源。 +func (c *StdIOClient) Close() error { + if c == nil { + return nil + } + + c.mu.Lock() + defer c.mu.Unlock() + c.shutdown = true + if c.stdin != nil { + _ = c.stdin.Close() + } + if c.cmd != nil && c.cmd.Process != nil { + _ = c.cmd.Process.Kill() + } + c.failAllPendingLocked(errors.New("mcp: stdio client closed")) + return nil +} + +// ListTools 调用 MCP `tools/list` 获取工具清单。 +func (c *StdIOClient) ListTools(ctx context.Context) ([]ToolDescriptor, error) { + callCtx, cancel := c.callContext(ctx) + defer cancel() + + raw, err := c.call(callCtx, "tools/list", map[string]any{}) + if err != nil { + return nil, err + } + + var payload struct { + Tools []struct { + Name string `json:"name"` + Description string `json:"description"` + InputSchema map[string]any `json:"inputSchema"` + InputSchema2 map[string]any `json:"input_schema"` + } `json:"tools"` + } + if err := json.Unmarshal(raw, &payload); err != nil { + return nil, fmt.Errorf("mcp: decode tools/list result: %w", err) + } + + result := make([]ToolDescriptor, 0, len(payload.Tools)) + for _, item := range payload.Tools { + schema := item.InputSchema + if len(schema) == 0 { + schema = item.InputSchema2 + } + result = append(result, ToolDescriptor{ + Name: strings.TrimSpace(item.Name), + Description: strings.TrimSpace(item.Description), + InputSchema: ensureObjectSchema(schema), + }) + } + return result, nil +} + +// CallTool 调用 MCP `tools/call` 并收敛返回值。 +func (c *StdIOClient) CallTool(ctx context.Context, toolName string, arguments []byte) (CallResult, error) { + trimmedToolName := strings.TrimSpace(toolName) + if trimmedToolName == "" { + return CallResult{}, errors.New("mcp: tool name is empty") + } + + callCtx, cancel := c.callContext(ctx) + defer cancel() + + var args any = map[string]any{} + if len(arguments) > 0 { + if err := json.Unmarshal(arguments, &args); err != nil { + return CallResult{}, fmt.Errorf("mcp: decode tool arguments: %w", err) + } + } + + raw, err := c.call(callCtx, "tools/call", map[string]any{ + "name": trimmedToolName, + "arguments": args, + }) + if err != nil { + return CallResult{}, err + } + return decodeCallResult(raw), nil +} + +// HealthCheck 通过一次短超时 `tools/list` 验证连接可用性。 +func (c *StdIOClient) HealthCheck(ctx context.Context) error { + _, err := c.ListTools(ctx) + return err +} + +func (c *StdIOClient) callContext(ctx context.Context) (context.Context, context.CancelFunc) { + timeout := c.cfg.CallTimeout + if deadline, ok := ctx.Deadline(); ok { + if remaining := time.Until(deadline); remaining > 0 && remaining < timeout { + timeout = remaining + } + } + return context.WithTimeout(ctx, timeout) +} + +func (c *StdIOClient) call(ctx context.Context, method string, params any) (json.RawMessage, error) { + if err := ctx.Err(); err != nil { + return nil, err + } + if err := c.ensureStarted(ctx); err != nil { + return nil, err + } + + requestID := "req-" + strconv.FormatUint(atomic.AddUint64(&c.idSeed, 1), 10) + replyCh := make(chan rpcReply, 1) + + c.mu.Lock() + if c.shutdown { + c.mu.Unlock() + return nil, errors.New("mcp: stdio client closed") + } + c.pending[requestID] = replyCh + stdin := c.stdin + c.mu.Unlock() + + requestPayload, err := json.Marshal(jsonRPCRequest{ + JSONRPC: "2.0", + ID: requestID, + Method: method, + Params: params, + }) + if err != nil { + c.removePending(requestID) + return nil, fmt.Errorf("mcp: marshal request: %w", err) + } + if err := writeFramedMessage(stdin, requestPayload); err != nil { + c.removePending(requestID) + return nil, fmt.Errorf("mcp: send request: %w", err) + } + + select { + case <-ctx.Done(): + c.removePending(requestID) + return nil, ctx.Err() + case reply := <-replyCh: + return reply.result, reply.err + } +} + +func (c *StdIOClient) ensureStarted(ctx context.Context) error { + c.mu.Lock() + defer c.mu.Unlock() + + if c.shutdown { + return errors.New("mcp: stdio client closed") + } + if c.started { + return nil + } + if !c.retryAt.IsZero() && time.Now().Before(c.retryAt) { + return fmt.Errorf("mcp: stdio restart backoff in effect until %s", c.retryAt.Format(time.RFC3339)) + } + + startCtx, cancel := context.WithTimeout(ctx, c.cfg.StartTimeout) + defer cancel() + + command := exec.Command(c.cfg.Command, c.cfg.Args...) + command.Env = append(os.Environ(), c.cfg.Env...) + command.Dir = strings.TrimSpace(c.cfg.Workdir) + + stdin, err := command.StdinPipe() + if err != nil { + return fmt.Errorf("mcp: create stdin pipe: %w", err) + } + stdout, err := command.StdoutPipe() + if err != nil { + return fmt.Errorf("mcp: create stdout pipe: %w", err) + } + stderr, err := command.StderrPipe() + if err != nil { + return fmt.Errorf("mcp: create stderr pipe: %w", err) + } + + startErrCh := make(chan error, 1) + go func() { + startErrCh <- command.Start() + }() + select { + case <-startCtx.Done(): + return startCtx.Err() + case err := <-startErrCh: + if err != nil { + c.bumpBackoffLocked() + return fmt.Errorf("mcp: start stdio server: %w", err) + } + } + + c.cmd = command + c.stdin = stdin + c.stdout = stdout + c.reader = bufio.NewReader(stdout) + c.exited = make(chan struct{}) + c.exitErr = nil + c.started = true + c.backoff = c.cfg.RestartBackoff + c.retryAt = time.Time{} + + go c.readLoop() + go c.waitLoop() + go io.Copy(io.Discard, stderr) + return nil +} + +func (c *StdIOClient) readLoop() { + for { + message, err := readFramedMessage(c.reader) + if err != nil { + c.markExited(fmt.Errorf("mcp: read response: %w", err)) + return + } + + var response jsonRPCResponse + if err := json.Unmarshal(message, &response); err != nil { + continue + } + if strings.TrimSpace(response.ID) == "" { + continue + } + + c.mu.Lock() + replyCh, ok := c.pending[response.ID] + if ok { + delete(c.pending, response.ID) + } + c.mu.Unlock() + if !ok { + continue + } + + if response.Error != nil { + replyCh <- rpcReply{ + err: fmt.Errorf("mcp: rpc error %d: %s", response.Error.Code, strings.TrimSpace(response.Error.Message)), + } + continue + } + replyCh <- rpcReply{result: response.Result} + } +} + +func (c *StdIOClient) waitLoop() { + err := c.cmd.Wait() + c.markExited(fmt.Errorf("mcp: stdio process exited: %w", err)) +} + +func (c *StdIOClient) markExited(err error) { + c.mu.Lock() + defer c.mu.Unlock() + + if !c.started { + return + } + c.started = false + c.exitErr = err + if c.exited != nil { + close(c.exited) + } + c.stdin = nil + c.stdout = nil + c.reader = nil + c.cmd = nil + c.failAllPendingLocked(err) + c.bumpBackoffLocked() +} + +func (c *StdIOClient) removePending(requestID string) { + c.mu.Lock() + defer c.mu.Unlock() + delete(c.pending, requestID) +} + +func (c *StdIOClient) failAllPendingLocked(err error) { + for requestID, replyCh := range c.pending { + replyCh <- rpcReply{err: err} + delete(c.pending, requestID) + } +} + +func (c *StdIOClient) bumpBackoffLocked() { + if c.backoff <= 0 { + c.backoff = c.cfg.RestartBackoff + } + c.retryAt = time.Now().Add(c.backoff) + c.backoff *= 2 + if c.backoff > maxStdioRestartBackoff { + c.backoff = maxStdioRestartBackoff + } +} + +func writeFramedMessage(writer io.Writer, payload []byte) error { + header := fmt.Sprintf("Content-Length: %d\r\n\r\n", len(payload)) + if _, err := io.WriteString(writer, header); err != nil { + return err + } + if _, err := writer.Write(payload); err != nil { + return err + } + return nil +} + +func readFramedMessage(reader *bufio.Reader) ([]byte, error) { + contentLength := -1 + for { + line, err := reader.ReadString('\n') + if err != nil { + return nil, err + } + trimmed := strings.TrimSpace(line) + if trimmed == "" { + break + } + + lower := strings.ToLower(trimmed) + if strings.HasPrefix(lower, "content-length:") { + rawLength := strings.TrimSpace(trimmed[len("content-length:"):]) + length, convErr := strconv.Atoi(rawLength) + if convErr != nil { + return nil, fmt.Errorf("mcp: invalid content-length %q", rawLength) + } + contentLength = length + } + } + if contentLength < 0 { + return nil, errors.New("mcp: missing content-length header") + } + + payload := make([]byte, contentLength) + if _, err := io.ReadFull(reader, payload); err != nil { + return nil, err + } + return payload, nil +} + +func decodeCallResult(raw json.RawMessage) CallResult { + var payload map[string]any + if err := json.Unmarshal(raw, &payload); err != nil { + return CallResult{ + Content: strings.TrimSpace(string(raw)), + IsError: false, + Metadata: map[string]any{"raw_result": string(raw)}, + } + } + + content := "" + switch typed := payload["content"].(type) { + case string: + content = strings.TrimSpace(typed) + case []any: + lines := make([]string, 0, len(typed)) + for _, item := range typed { + switch value := item.(type) { + case map[string]any: + text, _ := value["text"].(string) + if strings.TrimSpace(text) != "" { + lines = append(lines, strings.TrimSpace(text)) + } + case string: + if strings.TrimSpace(value) != "" { + lines = append(lines, strings.TrimSpace(value)) + } + } + } + content = strings.Join(lines, "\n") + default: + if typed != nil { + content = strings.TrimSpace(fmt.Sprintf("%v", typed)) + } + } + if content == "" { + content = "ok" + } + + isError := false + if value, ok := payload["isError"].(bool); ok { + isError = value + } + if value, ok := payload["is_error"].(bool); ok { + isError = isError || value + } + + metadata := map[string]any{} + for key, value := range payload { + if key == "content" || key == "isError" || key == "is_error" { + continue + } + metadata[key] = value + } + metadata["raw_result"] = bytes.TrimSpace(raw) + + return CallResult{ + Content: content, + IsError: isError, + Metadata: metadata, + } +} diff --git a/internal/tools/mcp/stdio_client_test.go b/internal/tools/mcp/stdio_client_test.go new file mode 100644 index 00000000..5235e454 --- /dev/null +++ b/internal/tools/mcp/stdio_client_test.go @@ -0,0 +1,139 @@ +package mcp + +import ( + "bufio" + "context" + "encoding/json" + "fmt" + "io" + "os" + "strings" + "testing" + "time" +) + +func TestStdIOClientListToolsAndCallTool(t *testing.T) { + t.Parallel() + + client := newTestStdIOClient(t) + defer func() { _ = client.Close() }() + + toolsList, err := client.ListTools(context.Background()) + if err != nil { + t.Fatalf("ListTools() error = %v", err) + } + if len(toolsList) != 1 || toolsList[0].Name != "search" { + t.Fatalf("unexpected tools list: %+v", toolsList) + } + + result, err := client.CallTool(context.Background(), "search", []byte(`{"query":"mcp"}`)) + if err != nil { + t.Fatalf("CallTool() error = %v", err) + } + if !strings.Contains(result.Content, "search") { + t.Fatalf("unexpected call result content: %q", result.Content) + } +} + +func TestStdIOClientHealthCheck(t *testing.T) { + t.Parallel() + + client := newTestStdIOClient(t) + defer func() { _ = client.Close() }() + + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + if err := client.HealthCheck(ctx); err != nil { + t.Fatalf("HealthCheck() error = %v", err) + } +} + +func newTestStdIOClient(t *testing.T) *StdIOClient { + t.Helper() + + client, err := NewStdIOClient(StdioClientConfig{ + Command: os.Args[0], + Args: []string{"-test.run=TestHelperProcessMCPStdioServer", "--"}, + Env: []string{"GO_WANT_MCP_STDIO_HELPER=1"}, + StartTimeout: 3 * time.Second, + CallTimeout: 3 * time.Second, + }) + if err != nil { + t.Fatalf("NewStdIOClient() error = %v", err) + } + return client +} + +func TestHelperProcessMCPStdioServer(t *testing.T) { + if os.Getenv("GO_WANT_MCP_STDIO_HELPER") != "1" { + return + } + + reader := bufio.NewReader(os.Stdin) + for { + payload, err := readFramedMessage(reader) + if err != nil { + if err == io.EOF { + os.Exit(0) + } + os.Exit(2) + } + + var request map[string]any + if err := json.Unmarshal(payload, &request); err != nil { + os.Exit(3) + } + + method, _ := request["method"].(string) + requestID, _ := request["id"].(string) + + var response any + switch method { + case "tools/list": + response = map[string]any{ + "jsonrpc": "2.0", + "id": requestID, + "result": map[string]any{ + "tools": []map[string]any{ + { + "name": "search", + "description": "search docs", + "inputSchema": map[string]any{ + "type": "object", + "properties": map[string]any{"query": map[string]any{"type": "string"}}, + }, + }, + }, + }, + } + case "tools/call": + params, _ := request["params"].(map[string]any) + name, _ := params["name"].(string) + response = map[string]any{ + "jsonrpc": "2.0", + "id": requestID, + "result": map[string]any{ + "content": fmt.Sprintf("ok:%s", name), + "isError": false, + }, + } + default: + response = map[string]any{ + "jsonrpc": "2.0", + "id": requestID, + "error": map[string]any{ + "code": -32601, + "message": "method not found", + }, + } + } + + rawResponse, err := json.Marshal(response) + if err != nil { + os.Exit(4) + } + if err := writeFramedMessage(os.Stdout, rawResponse); err != nil { + os.Exit(5) + } + } +} diff --git a/internal/tools/registry.go b/internal/tools/registry.go index 90f8a485..0b6b7d72 100644 --- a/internal/tools/registry.go +++ b/internal/tools/registry.go @@ -8,11 +8,14 @@ import ( providertypes "neo-code/internal/provider/types" "neo-code/internal/security" + "neo-code/internal/tools/mcp" ) type Registry struct { tools map[string]Tool microCompactPolicies map[string]MicroCompactPolicy + mcpRegistry *mcp.Registry + mcpFactory *mcp.AdapterFactory } func NewRegistry() *Registry { @@ -22,6 +25,15 @@ func NewRegistry() *Registry { } } +// SetMCPRegistry 绑定 MCP registry,用于将远程工具纳入统一执行链。 +func (r *Registry) SetMCPRegistry(registry *mcp.Registry) { + if r == nil || registry == nil { + return + } + r.mcpRegistry = registry + r.mcpFactory = mcp.NewAdapterFactory(registry) +} + func (r *Registry) Register(tool Tool) { if tool == nil { return @@ -46,8 +58,10 @@ func (r *Registry) Get(name string) (Tool, error) { // Supports reports whether a tool is registered. func (r *Registry) Supports(name string) bool { - _, err := r.Get(name) - return err == nil + if _, err := r.Get(name); err == nil { + return true + } + return r.supportsMCPTool(name) } // MicroCompactPolicy 返回指定工具名的 micro compact 策略;未知工具按默认可压缩处理。 @@ -93,12 +107,42 @@ func (r *Registry) ListAvailableSpecs(ctx context.Context, input SpecListInput) if err := ctx.Err(); err != nil { return nil, err } - return r.GetSpecs(), nil + + specs := r.GetSpecs() + mcpAdapters, err := r.listMCPAdapters(ctx) + if err != nil { + return nil, err + } + for _, adapter := range mcpAdapters { + specs = append(specs, provider.ToolSpec{ + Name: adapter.FullName(), + Description: adapter.Description(), + Schema: adapter.Schema(), + }) + } + sort.Slice(specs, func(i, j int) bool { + return strings.ToLower(specs[i].Name) < strings.ToLower(specs[j].Name) + }) + return specs, nil } func (r *Registry) Execute(ctx context.Context, input ToolCallInput) (ToolResult, error) { tool, err := r.Get(input.Name) - if err != nil { + if err == nil { + result, execErr := tool.Execute(ctx, input) + result.ToolCallID = input.ID + if execErr != nil { + result.IsError = true + if strings.TrimSpace(result.Content) == "" { + result.Content = FormatError(result.Name, NormalizeErrorReason(result.Name, execErr), "") + } + return result, execErr + } + return result, nil + } + + adapter, resolveErr := r.resolveMCPAdapter(ctx, input.Name) + if resolveErr != nil { content := FormatError(input.Name, NormalizeErrorReason(input.Name, err), "") return ToolResult{ ToolCallID: input.ID, @@ -107,15 +151,30 @@ func (r *Registry) Execute(ctx context.Context, input ToolCallInput) (ToolResult IsError: true, }, err } - - result, execErr := tool.Execute(ctx, input) - result.ToolCallID = input.ID - if execErr != nil { + callResult, callErr := adapter.Call(ctx, input.Arguments) + result := ToolResult{ + ToolCallID: input.ID, + Name: adapter.FullName(), + Content: strings.TrimSpace(callResult.Content), + IsError: callResult.IsError, + Metadata: map[string]any{ + "mcp_server_id": adapter.ServerID(), + "mcp_tool_name": adapter.ToolName(), + }, + } + for key, value := range callResult.Metadata { + result.Metadata[key] = value + } + if result.Content == "" { + result.Content = "ok" + } + result = ApplyOutputLimit(result, DefaultOutputLimitBytes) + if callErr != nil { result.IsError = true if strings.TrimSpace(result.Content) == "" { - result.Content = FormatError(result.Name, NormalizeErrorReason(result.Name, execErr), "") + result.Content = FormatError(result.Name, NormalizeErrorReason(result.Name, callErr), "") } - return result, execErr + return result, callErr } return result, nil } @@ -124,3 +183,60 @@ func (r *Registry) Execute(ctx context.Context, input ToolCallInput) (ToolResult func (r *Registry) RememberSessionDecision(sessionID string, action security.Action, scope SessionPermissionScope) error { return errors.New("tools: session permission memory is unsupported by registry manager") } + +// supportsMCPTool 判断指定工具名是否可由当前 MCP 快照解析。 +func (r *Registry) supportsMCPTool(name string) bool { + if r == nil || r.mcpFactory == nil { + return false + } + lowerName := strings.ToLower(strings.TrimSpace(name)) + if !strings.HasPrefix(lowerName, "mcp.") { + return false + } + for _, snapshot := range r.mcpFactoryBuildSnapshot() { + for _, tool := range snapshot.Tools { + if strings.EqualFold(mcpToolFullName(snapshot.ServerID, tool.Name), lowerName) { + return true + } + } + } + return false +} + +// listMCPAdapters 返回 MCP 快照对应的 adapter 列表。 +func (r *Registry) listMCPAdapters(ctx context.Context) ([]*mcp.Adapter, error) { + if r == nil || r.mcpFactory == nil { + return nil, nil + } + return r.mcpFactory.BuildAdapters(ctx) +} + +// resolveMCPAdapter 按完整工具名解析并返回对应 adapter。 +func (r *Registry) resolveMCPAdapter(ctx context.Context, fullName string) (*mcp.Adapter, error) { + adapters, err := r.listMCPAdapters(ctx) + if err != nil { + return nil, err + } + lowerName := strings.ToLower(strings.TrimSpace(fullName)) + if !strings.HasPrefix(lowerName, "mcp.") { + return nil, errors.New("tool: not found") + } + for _, adapter := range adapters { + if strings.EqualFold(adapter.FullName(), lowerName) { + return adapter, nil + } + } + return nil, errors.New("tool: not found") +} + +// mcpFactoryBuildSnapshot 读取 MCP registry 快照,用于无上下文快速检查。 +func (r *Registry) mcpFactoryBuildSnapshot() []mcp.ServerSnapshot { + if r == nil || r.mcpRegistry == nil { + return nil + } + return r.mcpRegistry.Snapshot() +} + +func mcpToolFullName(serverID string, toolName string) string { + return "mcp." + strings.ToLower(strings.TrimSpace(serverID)) + "." + strings.ToLower(strings.TrimSpace(toolName)) +} diff --git a/internal/tools/registry_test.go b/internal/tools/registry_test.go index 2ef5a423..83632adb 100644 --- a/internal/tools/registry_test.go +++ b/internal/tools/registry_test.go @@ -7,6 +7,7 @@ import ( "testing" "neo-code/internal/security" + "neo-code/internal/tools/mcp" ) type stubTool struct { @@ -242,3 +243,107 @@ func TestRegistryRememberSessionDecisionUnsupported(t *testing.T) { t.Fatalf("expected unsupported error, got %v", err) } } + +type stubMCPClient struct { + tools []mcp.ToolDescriptor + callResult mcp.CallResult + callErr error +} + +func (s *stubMCPClient) ListTools(ctx context.Context) ([]mcp.ToolDescriptor, error) { + return s.tools, nil +} + +func (s *stubMCPClient) CallTool(ctx context.Context, toolName string, arguments []byte) (mcp.CallResult, error) { + if s.callErr != nil { + return mcp.CallResult{}, s.callErr + } + return s.callResult, nil +} + +func (s *stubMCPClient) HealthCheck(ctx context.Context) error { + return nil +} + +func TestRegistryListAvailableSpecsIncludesMCP(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + registry.Register(stubTool{name: "a_tool", description: "built-in", schema: map[string]any{"type": "object"}}) + + mcpRegistry := mcp.NewRegistry() + if err := mcpRegistry.RegisterServer("docs", "stdio", "v1", &stubMCPClient{ + tools: []mcp.ToolDescriptor{ + {Name: "search", Description: "search docs", InputSchema: map[string]any{"type": "object"}}, + }, + }); err != nil { + t.Fatalf("register mcp server: %v", err) + } + if err := mcpRegistry.RefreshServerTools(context.Background(), "docs"); err != nil { + t.Fatalf("refresh mcp tools: %v", err) + } + registry.SetMCPRegistry(mcpRegistry) + + specs, err := registry.ListAvailableSpecs(context.Background(), SpecListInput{}) + if err != nil { + t.Fatalf("ListAvailableSpecs() error = %v", err) + } + if len(specs) != 2 { + t.Fatalf("expected 2 specs (built-in + mcp), got %d", len(specs)) + } + foundMCP := false + for _, spec := range specs { + if spec.Name == "mcp.docs.search" { + foundMCP = true + break + } + } + if !foundMCP { + t.Fatalf("expected mcp.docs.search in specs, got %+v", specs) + } +} + +func TestRegistryExecuteDispatchesToMCPAdapter(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + mcpRegistry := mcp.NewRegistry() + if err := mcpRegistry.RegisterServer("docs", "stdio", "v1", &stubMCPClient{ + tools: []mcp.ToolDescriptor{ + {Name: "search", Description: "search docs", InputSchema: map[string]any{"type": "object"}}, + }, + callResult: mcp.CallResult{ + Content: "mcp ok", + Metadata: map[string]any{ + "latency_ms": 12, + }, + }, + }); err != nil { + t.Fatalf("register mcp server: %v", err) + } + if err := mcpRegistry.RefreshServerTools(context.Background(), "docs"); err != nil { + t.Fatalf("refresh mcp tools: %v", err) + } + registry.SetMCPRegistry(mcpRegistry) + + result, err := registry.Execute(context.Background(), ToolCallInput{ + ID: "mcp-call-1", + Name: "mcp.docs.search", + Arguments: []byte(`{"query":"neocode"}`), + }) + if err != nil { + t.Fatalf("Execute() error = %v", err) + } + if result.ToolCallID != "mcp-call-1" { + t.Fatalf("expected tool call id mcp-call-1, got %q", result.ToolCallID) + } + if result.Name != "mcp.docs.search" { + t.Fatalf("expected mcp tool name, got %q", result.Name) + } + if !strings.Contains(result.Content, "mcp ok") { + t.Fatalf("expected mcp content, got %q", result.Content) + } + if result.Metadata["mcp_server_id"] != "docs" || result.Metadata["mcp_tool_name"] != "search" { + t.Fatalf("unexpected mcp metadata: %+v", result.Metadata) + } +} From 7b5db03addbf3103788f6ea6a6d540ca462f33b2 Mon Sep 17 00:00:00 2001 From: Cai_Tang <106404101+Cai-Tang-www@users.noreply.github.com> Date: Tue, 7 Apr 2026 12:06:17 +0800 Subject: [PATCH 34/55] =?UTF-8?q?feat:=20=E5=BC=BA=E5=8C=96=20MCP=20?= =?UTF-8?q?=E8=B0=83=E7=94=A8=E9=93=BE=E9=94=99=E8=AF=AF=E5=A4=84=E7=90=86?= =?UTF-8?q?=E4=B8=8E=E5=AE=89=E5=85=A8=E8=BE=B9=E7=95=8C?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/tools/mcp/registry.go | 22 +++++++++++++++++++++- internal/tools/mcp/stdio_client.go | 28 ++++++++++++++++++++++------ internal/tools/registry.go | 13 +++++++------ 3 files changed, 50 insertions(+), 13 deletions(-) diff --git a/internal/tools/mcp/registry.go b/internal/tools/mcp/registry.go index 026983ec..490b5142 100644 --- a/internal/tools/mcp/registry.go +++ b/internal/tools/mcp/registry.go @@ -313,7 +313,27 @@ func cloneSchema(schema map[string]any) map[string]any { } cloned := make(map[string]any, len(schema)) for key, value := range schema { - cloned[key] = value + cloned[key] = cloneAny(value) } return cloned } + +// cloneAny 递归复制 schema/metadata 中的 map 与 slice,避免跨层共享引用。 +func cloneAny(value any) any { + switch typed := value.(type) { + case map[string]any: + cloned := make(map[string]any, len(typed)) + for key, item := range typed { + cloned[key] = cloneAny(item) + } + return cloned + case []any: + cloned := make([]any, len(typed)) + for i, item := range typed { + cloned[i] = cloneAny(item) + } + return cloned + default: + return value + } +} diff --git a/internal/tools/mcp/stdio_client.go b/internal/tools/mcp/stdio_client.go index 7e92fe80..7b9196f9 100644 --- a/internal/tools/mcp/stdio_client.go +++ b/internal/tools/mcp/stdio_client.go @@ -22,6 +22,7 @@ const ( defaultStdioCallTimeout = 15 * time.Second defaultStdioRestartBackoff = 1 * time.Second maxStdioRestartBackoff = 30 * time.Second + maxStdioFrameBytes = 8 * 1024 * 1024 ) // StdioClientConfig 描述 MCP stdio 客户端的启动与调用参数。 @@ -64,6 +65,7 @@ type StdIOClient struct { cfg StdioClientConfig idSeed uint64 mu sync.Mutex + writeMu sync.Mutex cmd *exec.Cmd stdin io.WriteCloser stdout io.ReadCloser @@ -217,6 +219,10 @@ func (c *StdIOClient) call(ctx context.Context, method string, params any) (json c.pending[requestID] = replyCh stdin := c.stdin c.mu.Unlock() + if stdin == nil { + c.removePending(requestID) + return nil, errors.New("mcp: stdio client is not connected") + } requestPayload, err := json.Marshal(jsonRPCRequest{ JSONRPC: "2.0", @@ -228,9 +234,12 @@ func (c *StdIOClient) call(ctx context.Context, method string, params any) (json c.removePending(requestID) return nil, fmt.Errorf("mcp: marshal request: %w", err) } - if err := writeFramedMessage(stdin, requestPayload); err != nil { + c.writeMu.Lock() + writeErr := writeFramedMessage(stdin, requestPayload) + c.writeMu.Unlock() + if writeErr != nil { c.removePending(requestID) - return nil, fmt.Errorf("mcp: send request: %w", err) + return nil, fmt.Errorf("mcp: send request: %w", writeErr) } select { @@ -301,7 +310,7 @@ func (c *StdIOClient) ensureStarted(ctx context.Context) error { c.retryAt = time.Time{} go c.readLoop() - go c.waitLoop() + go c.waitLoop(command) go io.Copy(io.Discard, stderr) return nil } @@ -342,8 +351,12 @@ func (c *StdIOClient) readLoop() { } } -func (c *StdIOClient) waitLoop() { - err := c.cmd.Wait() +func (c *StdIOClient) waitLoop(command *exec.Cmd) { + if command == nil { + c.markExited(errors.New("mcp: stdio process is nil")) + return + } + err := command.Wait() c.markExited(fmt.Errorf("mcp: stdio process exited: %w", err)) } @@ -427,6 +440,9 @@ func readFramedMessage(reader *bufio.Reader) ([]byte, error) { if contentLength < 0 { return nil, errors.New("mcp: missing content-length header") } + if contentLength > maxStdioFrameBytes { + return nil, fmt.Errorf("mcp: content-length %d exceeds limit %d", contentLength, maxStdioFrameBytes) + } payload := make([]byte, contentLength) if _, err := io.ReadFull(reader, payload); err != nil { @@ -489,7 +505,7 @@ func decodeCallResult(raw json.RawMessage) CallResult { } metadata[key] = value } - metadata["raw_result"] = bytes.TrimSpace(raw) + metadata["raw_result"] = string(bytes.TrimSpace(raw)) return CallResult{ Content: content, diff --git a/internal/tools/registry.go b/internal/tools/registry.go index 0b6b7d72..c03bfcb6 100644 --- a/internal/tools/registry.go +++ b/internal/tools/registry.go @@ -143,13 +143,13 @@ func (r *Registry) Execute(ctx context.Context, input ToolCallInput) (ToolResult adapter, resolveErr := r.resolveMCPAdapter(ctx, input.Name) if resolveErr != nil { - content := FormatError(input.Name, NormalizeErrorReason(input.Name, err), "") + content := FormatError(input.Name, NormalizeErrorReason(input.Name, resolveErr), "") return ToolResult{ ToolCallID: input.ID, Name: input.Name, Content: content, IsError: true, - }, err + }, resolveErr } callResult, callErr := adapter.Call(ctx, input.Arguments) result := ToolResult{ @@ -165,17 +165,18 @@ func (r *Registry) Execute(ctx context.Context, input ToolCallInput) (ToolResult for key, value := range callResult.Metadata { result.Metadata[key] = value } - if result.Content == "" { - result.Content = "ok" - } - result = ApplyOutputLimit(result, DefaultOutputLimitBytes) if callErr != nil { result.IsError = true if strings.TrimSpace(result.Content) == "" { result.Content = FormatError(result.Name, NormalizeErrorReason(result.Name, callErr), "") } + result = ApplyOutputLimit(result, DefaultOutputLimitBytes) return result, callErr } + if result.Content == "" { + result.Content = "ok" + } + result = ApplyOutputLimit(result, DefaultOutputLimitBytes) return result, nil } From 4bc19abef7f5b287ecf918dc4854d58b040a7f6e Mon Sep 17 00:00:00 2001 From: Cai_Tang <106404101+Cai-Tang-www@users.noreply.github.com> Date: Tue, 7 Apr 2026 12:06:27 +0800 Subject: [PATCH 35/55] =?UTF-8?q?feat:=20=E8=A1=A5=E9=BD=90=20MCP=20?= =?UTF-8?q?=E5=B9=B6=E5=8F=91=E4=B8=8E=E9=94=99=E8=AF=AF=E8=B7=AF=E5=BE=84?= =?UTF-8?q?=E5=9B=9E=E5=BD=92=E6=B5=8B=E8=AF=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/tools/mcp/registry_test.go | 47 ++++++++++++++++++ internal/tools/mcp/stdio_client_test.go | 46 +++++++++++++++++ internal/tools/registry_test.go | 65 +++++++++++++++++++++++++ 3 files changed, 158 insertions(+) diff --git a/internal/tools/mcp/registry_test.go b/internal/tools/mcp/registry_test.go index d7f5f26d..b0b9c99d 100644 --- a/internal/tools/mcp/registry_test.go +++ b/internal/tools/mcp/registry_test.go @@ -161,3 +161,50 @@ func TestRegistryConcurrentSnapshotAndRefresh(t *testing.T) { t.Fatalf("concurrent registry operations timed out") } } + +func TestRegistrySnapshotSchemaIsDeepCloned(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + client := &stubServerClient{ + tools: []ToolDescriptor{ + { + Name: "search", + InputSchema: map[string]any{ + "type": "object", + "properties": map[string]any{ + "query": map[string]any{ + "type": "string", + }, + }, + }, + }, + }, + } + if err := registry.RegisterServer("server-1", "stdio", "v1", client); err != nil { + t.Fatalf("register server: %v", err) + } + if err := registry.RefreshServerTools(context.Background(), "server-1"); err != nil { + t.Fatalf("refresh tools: %v", err) + } + + first := registry.Snapshot() + properties, ok := first[0].Tools[0].InputSchema["properties"].(map[string]any) + if !ok { + t.Fatalf("expected properties map") + } + properties["query"] = map[string]any{"type": "number"} + + second := registry.Snapshot() + secondProperties, ok := second[0].Tools[0].InputSchema["properties"].(map[string]any) + if !ok { + t.Fatalf("expected properties map in second snapshot") + } + query, ok := secondProperties["query"].(map[string]any) + if !ok { + t.Fatalf("expected query schema map") + } + if query["type"] != "string" { + t.Fatalf("expected deep cloned schema type string, got %v", query["type"]) + } +} diff --git a/internal/tools/mcp/stdio_client_test.go b/internal/tools/mcp/stdio_client_test.go index 5235e454..2dea5e16 100644 --- a/internal/tools/mcp/stdio_client_test.go +++ b/internal/tools/mcp/stdio_client_test.go @@ -8,6 +8,7 @@ import ( "io" "os" "strings" + "sync" "testing" "time" ) @@ -48,6 +49,51 @@ func TestStdIOClientHealthCheck(t *testing.T) { } } +func TestStdIOClientConcurrentCallTool(t *testing.T) { + t.Parallel() + + client := newTestStdIOClient(t) + defer func() { _ = client.Close() }() + + const workers = 16 + var wg sync.WaitGroup + errCh := make(chan error, workers) + for i := 0; i < workers; i++ { + wg.Add(1) + go func() { + defer wg.Done() + result, err := client.CallTool(context.Background(), "search", []byte(`{"query":"mcp"}`)) + if err != nil { + errCh <- err + return + } + if !strings.Contains(result.Content, "search") { + errCh <- fmt.Errorf("unexpected content: %q", result.Content) + } + }() + } + wg.Wait() + close(errCh) + for err := range errCh { + t.Fatalf("concurrent call failed: %v", err) + } +} + +func TestReadFramedMessageRejectsOversizedPayload(t *testing.T) { + t.Parallel() + + payload := strings.Repeat("x", 32) + raw := fmt.Sprintf("Content-Length: %d\r\n\r\n%s", maxStdioFrameBytes+1, payload) + reader := bufio.NewReader(strings.NewReader(raw)) + _, err := readFramedMessage(reader) + if err == nil { + t.Fatalf("expected oversized payload error") + } + if !strings.Contains(err.Error(), "exceeds limit") { + t.Fatalf("expected exceeds limit error, got %v", err) + } +} + func newTestStdIOClient(t *testing.T) *StdIOClient { t.Helper() diff --git a/internal/tools/registry_test.go b/internal/tools/registry_test.go index 83632adb..3cc882d8 100644 --- a/internal/tools/registry_test.go +++ b/internal/tools/registry_test.go @@ -347,3 +347,68 @@ func TestRegistryExecuteDispatchesToMCPAdapter(t *testing.T) { t.Fatalf("unexpected mcp metadata: %+v", result.Metadata) } } + +func TestRegistryExecuteMCPResolveErrorPropagates(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + mcpRegistry := mcp.NewRegistry() + if err := mcpRegistry.RegisterServer("docs", "stdio", "v1", &stubMCPClient{ + tools: []mcp.ToolDescriptor{ + {Name: "search", Description: "search docs", InputSchema: map[string]any{"type": "object"}}, + }, + }); err != nil { + t.Fatalf("register mcp server: %v", err) + } + registry.SetMCPRegistry(mcpRegistry) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + result, err := registry.Execute(ctx, ToolCallInput{ + ID: "mcp-call-canceled", + Name: "mcp.docs.search", + }) + if err == nil || !errors.Is(err, context.Canceled) { + t.Fatalf("expected context canceled, got %v", err) + } + if !result.IsError { + t.Fatalf("expected error result") + } + if !strings.Contains(result.Content, context.Canceled.Error()) { + t.Fatalf("expected canceled content, got %q", result.Content) + } +} + +func TestRegistryExecuteMCPCallErrorDoesNotReturnOK(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + mcpRegistry := mcp.NewRegistry() + if err := mcpRegistry.RegisterServer("docs", "stdio", "v1", &stubMCPClient{ + tools: []mcp.ToolDescriptor{ + {Name: "search", Description: "search docs", InputSchema: map[string]any{"type": "object"}}, + }, + callErr: errors.New("mcp transport timeout"), + }); err != nil { + t.Fatalf("register mcp server: %v", err) + } + if err := mcpRegistry.RefreshServerTools(context.Background(), "docs"); err != nil { + t.Fatalf("refresh mcp tools: %v", err) + } + registry.SetMCPRegistry(mcpRegistry) + + result, err := registry.Execute(context.Background(), ToolCallInput{ + ID: "mcp-call-error", + Name: "mcp.docs.search", + Arguments: []byte(`{"query":"neocode"}`), + }) + if err == nil { + t.Fatalf("expected mcp call error") + } + if !result.IsError { + t.Fatalf("expected IsError true") + } + if strings.TrimSpace(result.Content) == "" || strings.EqualFold(strings.TrimSpace(result.Content), "ok") { + t.Fatalf("expected non-ok error content, got %q", result.Content) + } +} From 76f1b4862e99a1ccf99d37b2e8bccff4a7ab52dd Mon Sep 17 00:00:00 2001 From: Cai_Tang <106404101+Cai-Tang-www@users.noreply.github.com> Date: Tue, 7 Apr 2026 12:38:02 +0800 Subject: [PATCH 36/55] =?UTF-8?q?feat:=20=E5=AE=8C=E5=96=84=20MCP=20?= =?UTF-8?q?=E6=96=B0=E8=83=BD=E5=8A=9B=E6=B5=8B=E8=AF=95=E8=A6=86=E7=9B=96?= =?UTF-8?q?=E4=B8=8E=E8=BE=B9=E7=95=8C=E7=94=A8=E4=BE=8B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/tools/mcp/adapter_test.go | 39 +++++++ internal/tools/mcp/registry_test.go | 36 +++++++ internal/tools/mcp/stdio_client_test.go | 131 ++++++++++++++++++++++++ 3 files changed, 206 insertions(+) diff --git a/internal/tools/mcp/adapter_test.go b/internal/tools/mcp/adapter_test.go index d5dd16f7..f51f9191 100644 --- a/internal/tools/mcp/adapter_test.go +++ b/internal/tools/mcp/adapter_test.go @@ -107,3 +107,42 @@ func TestAdapterCallError(t *testing.T) { t.Fatalf("expected call error") } } + +func TestAdapterAccessorsAndSchemaClone(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + adapter, err := NewAdapter(registry, "Docs", ToolDescriptor{ + Name: "search", + Description: "", + InputSchema: map[string]any{ + "type": "object", + "properties": map[string]any{ + "q": map[string]any{"type": "string"}, + }, + }, + }) + if err != nil { + t.Fatalf("NewAdapter() error = %v", err) + } + + if adapter.ServerID() != "docs" { + t.Fatalf("expected normalized server id docs, got %q", adapter.ServerID()) + } + if adapter.ToolName() != "search" { + t.Fatalf("expected tool name search, got %q", adapter.ToolName()) + } + if adapter.Description() == "" { + t.Fatalf("expected non-empty fallback description") + } + + schema1 := adapter.Schema() + schema2 := adapter.Schema() + props1, _ := schema1["properties"].(map[string]any) + props1["q"] = map[string]any{"type": "number"} + props2, _ := schema2["properties"].(map[string]any) + query2, _ := props2["q"].(map[string]any) + if query2["type"] != "string" { + t.Fatalf("expected schema clone not mutated, got %v", query2["type"]) + } +} diff --git a/internal/tools/mcp/registry_test.go b/internal/tools/mcp/registry_test.go index b0b9c99d..49f2c654 100644 --- a/internal/tools/mcp/registry_test.go +++ b/internal/tools/mcp/registry_test.go @@ -208,3 +208,39 @@ func TestRegistrySnapshotSchemaIsDeepCloned(t *testing.T) { t.Fatalf("expected deep cloned schema type string, got %v", query["type"]) } } + +func TestRegistryRegisterAndUnregisterBoundaries(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + client := &stubServerClient{} + + if err := registry.RegisterServer("docs", "stdio", "v1", client); err != nil { + t.Fatalf("register server: %v", err) + } + if err := registry.RegisterServer("docs", "stdio", "v1", client); err == nil { + t.Fatalf("expected duplicate register error") + } + if !registry.UnregisterServer("docs") { + t.Fatalf("expected unregister success") + } + if registry.UnregisterServer("docs") { + t.Fatalf("expected unregister miss to be false") + } +} + +func TestRegistrySetServerStatusValidation(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + client := &stubServerClient{} + if err := registry.RegisterServer("docs", "stdio", "v1", client); err != nil { + t.Fatalf("register server: %v", err) + } + if err := registry.SetServerStatus("docs", ServerStatus("unknown")); err == nil { + t.Fatalf("expected invalid status error") + } + if err := registry.SetServerStatus("missing", ServerStatusReady); err == nil { + t.Fatalf("expected missing server error") + } +} diff --git a/internal/tools/mcp/stdio_client_test.go b/internal/tools/mcp/stdio_client_test.go index 2dea5e16..530677ad 100644 --- a/internal/tools/mcp/stdio_client_test.go +++ b/internal/tools/mcp/stdio_client_test.go @@ -4,6 +4,7 @@ import ( "bufio" "context" "encoding/json" + "errors" "fmt" "io" "os" @@ -94,6 +95,136 @@ func TestReadFramedMessageRejectsOversizedPayload(t *testing.T) { } } +func TestNewStdIOClientValidationAndDefaults(t *testing.T) { + t.Parallel() + + if _, err := NewStdIOClient(StdioClientConfig{}); err == nil { + t.Fatalf("expected empty command error") + } + client, err := NewStdIOClient(StdioClientConfig{Command: "cmd"}) + if err != nil { + t.Fatalf("NewStdIOClient() error = %v", err) + } + if client.cfg.StartTimeout <= 0 || client.cfg.CallTimeout <= 0 || client.cfg.RestartBackoff <= 0 { + t.Fatalf("expected default timeouts/backoff to be initialized") + } +} + +func TestStdIOClientCallToolInputValidation(t *testing.T) { + t.Parallel() + + client := &StdIOClient{} + if _, err := client.CallTool(context.Background(), "", nil); err == nil { + t.Fatalf("expected empty tool name error") + } + if _, err := client.CallTool(context.Background(), "search", []byte("{not-json")); err == nil { + t.Fatalf("expected invalid json arguments error") + } +} + +func TestStdIOClientCallRejectsClosedAndDisconnected(t *testing.T) { + t.Parallel() + + client := &StdIOClient{ + pending: make(map[string]chan rpcReply), + cfg: StdioClientConfig{ + CallTimeout: time.Second, + StartTimeout: time.Second, + RestartBackoff: time.Millisecond, + }, + } + client.shutdown = true + if _, err := client.call(context.Background(), "tools/list", map[string]any{}); err == nil { + t.Fatalf("expected closed error") + } + + client.shutdown = false + client.started = true + client.stdin = nil + if _, err := client.call(context.Background(), "tools/list", map[string]any{}); err == nil { + t.Fatalf("expected disconnected error") + } +} + +func TestStdIOClientEnsureStartedBackoff(t *testing.T) { + t.Parallel() + + client := &StdIOClient{ + cfg: StdioClientConfig{ + Command: "cmd", + StartTimeout: time.Second, + CallTimeout: time.Second, + RestartBackoff: time.Second, + }, + pending: make(map[string]chan rpcReply), + retryAt: time.Now().Add(2 * time.Second), + } + + err := client.ensureStarted(context.Background()) + if err == nil || !strings.Contains(err.Error(), "backoff") { + t.Fatalf("expected backoff error, got %v", err) + } +} + +func TestReadFramedMessageHeaderErrors(t *testing.T) { + t.Parallel() + + reader := bufio.NewReader(strings.NewReader("X-Test: 1\r\n\r\n{}")) + if _, err := readFramedMessage(reader); err == nil || !strings.Contains(err.Error(), "missing content-length") { + t.Fatalf("expected missing content-length error, got %v", err) + } + + reader = bufio.NewReader(strings.NewReader("Content-Length: nope\r\n\r\n{}")) + if _, err := readFramedMessage(reader); err == nil || !strings.Contains(err.Error(), "invalid content-length") { + t.Fatalf("expected invalid content-length error, got %v", err) + } +} + +func TestDecodeCallResultVariants(t *testing.T) { + t.Parallel() + + result := decodeCallResult(json.RawMessage(`{"content":" ok ","isError":true,"extra":1}`)) + if result.Content != "ok" || !result.IsError { + t.Fatalf("unexpected decode result: %+v", result) + } + if result.Metadata["extra"] != float64(1) { + t.Fatalf("expected metadata extra") + } + + result = decodeCallResult(json.RawMessage(`{"content":[{"text":"a"},"b"],"is_error":true}`)) + if result.Content != "a\nb" || !result.IsError { + t.Fatalf("unexpected list content decode: %+v", result) + } + + result = decodeCallResult(json.RawMessage(`{"content":{"nested":"x"}}`)) + if result.Content == "" { + t.Fatalf("expected fallback string content") + } + + result = decodeCallResult(json.RawMessage(`not-json`)) + if result.Content != "not-json" { + t.Fatalf("expected raw fallback content, got %q", result.Content) + } + if _, ok := result.Metadata["raw_result"]; !ok { + t.Fatalf("expected raw_result metadata") + } +} + +func TestFailAllPendingLocked(t *testing.T) { + t.Parallel() + + client := &StdIOClient{ + pending: map[string]chan rpcReply{ + "a": make(chan rpcReply, 1), + "b": make(chan rpcReply, 1), + }, + } + client.failAllPendingLocked(errors.New("closed")) + if len(client.pending) != 0 { + t.Fatalf("expected pending cleared") + } +} + func newTestStdIOClient(t *testing.T) *StdIOClient { t.Helper() From 8486c26f3af141054b4474e47a7daa8c4d4d5184 Mon Sep 17 00:00:00 2001 From: Cai_Tang <106404101+Cai-Tang-www@users.noreply.github.com> Date: Tue, 7 Apr 2026 12:57:42 +0800 Subject: [PATCH 37/55] =?UTF-8?q?feat:=20=E6=8F=90=E5=8D=87=20#176/#177=20?= =?UTF-8?q?=E5=B7=AE=E5=BC=82=E8=A6=86=E7=9B=96=E7=8E=87=E5=B9=B6=E8=A1=A5?= =?UTF-8?q?=E5=85=85=E5=88=86=E6=94=AF=E6=B5=8B=E8=AF=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/tools/mcp/adapter_test.go | 47 +++++++++++++++++++++++++ internal/tools/mcp/stdio_client_test.go | 45 +++++++++++++++++++++++ internal/tools/registry_test.go | 36 +++++++++++++++++++ 3 files changed, 128 insertions(+) diff --git a/internal/tools/mcp/adapter_test.go b/internal/tools/mcp/adapter_test.go index f51f9191..8f3adc8f 100644 --- a/internal/tools/mcp/adapter_test.go +++ b/internal/tools/mcp/adapter_test.go @@ -146,3 +146,50 @@ func TestAdapterAccessorsAndSchemaClone(t *testing.T) { t.Fatalf("expected schema clone not mutated, got %v", query2["type"]) } } + +func TestAdapterBuildAndCreateErrors(t *testing.T) { + t.Parallel() + + factory := NewAdapterFactory(nil) + if _, err := factory.BuildAdapters(context.Background()); err == nil { + t.Fatalf("expected nil registry error") + } + + canceledCtx, cancel := context.WithCancel(context.Background()) + cancel() + registry := NewRegistry() + factory = NewAdapterFactory(registry) + if _, err := factory.BuildAdapters(canceledCtx); err == nil { + t.Fatalf("expected canceled context error") + } + + if _, err := NewAdapter(nil, "docs", ToolDescriptor{Name: "search"}); err == nil { + t.Fatalf("expected nil registry error") + } + if _, err := NewAdapter(registry, " ", ToolDescriptor{Name: "search"}); err == nil { + t.Fatalf("expected empty server id error") + } + if _, err := NewAdapter(registry, "docs", ToolDescriptor{Name: " "}); err == nil { + t.Fatalf("expected empty tool name error") + } +} + +func TestAdapterCallBoundary(t *testing.T) { + t.Parallel() + + var nilAdapter *Adapter + if _, err := nilAdapter.Call(context.Background(), nil); err == nil { + t.Fatalf("expected nil adapter error") + } + + registry := NewRegistry() + adapter, err := NewAdapter(registry, "docs", ToolDescriptor{Name: "search"}) + if err != nil { + t.Fatalf("NewAdapter() error = %v", err) + } + canceledCtx, cancel := context.WithCancel(context.Background()) + cancel() + if _, err := adapter.Call(canceledCtx, nil); err == nil { + t.Fatalf("expected context canceled error") + } +} diff --git a/internal/tools/mcp/stdio_client_test.go b/internal/tools/mcp/stdio_client_test.go index 530677ad..3a352d3a 100644 --- a/internal/tools/mcp/stdio_client_test.go +++ b/internal/tools/mcp/stdio_client_test.go @@ -14,6 +14,12 @@ import ( "time" ) +type errWriter struct{} + +func (errWriter) Write(p []byte) (int, error) { + return 0, errors.New("write failed") +} + func TestStdIOClientListToolsAndCallTool(t *testing.T) { t.Parallel() @@ -225,6 +231,45 @@ func TestFailAllPendingLocked(t *testing.T) { } } +func TestWriteFramedMessageError(t *testing.T) { + t.Parallel() + + if err := writeFramedMessage(errWriter{}, []byte(`{}`)); err == nil { + t.Fatalf("expected write error") + } +} + +func TestStdIOClientWaitLoopNilCommand(t *testing.T) { + t.Parallel() + + client := &StdIOClient{ + pending: make(map[string]chan rpcReply), + started: true, + } + client.waitLoop(nil) + if client.started { + t.Fatalf("expected started=false after nil command waitLoop") + } +} + +func TestStdIOClientBumpBackoffClamp(t *testing.T) { + t.Parallel() + + client := &StdIOClient{ + cfg: StdioClientConfig{ + RestartBackoff: time.Second, + }, + backoff: maxStdioRestartBackoff, + } + client.bumpBackoffLocked() + if client.backoff != maxStdioRestartBackoff { + t.Fatalf("expected clamp to max backoff, got %v", client.backoff) + } + if client.retryAt.IsZero() { + t.Fatalf("expected retryAt assigned") + } +} + func newTestStdIOClient(t *testing.T) *StdIOClient { t.Helper() diff --git a/internal/tools/registry_test.go b/internal/tools/registry_test.go index 3cc882d8..a26843e4 100644 --- a/internal/tools/registry_test.go +++ b/internal/tools/registry_test.go @@ -412,3 +412,39 @@ func TestRegistryExecuteMCPCallErrorDoesNotReturnOK(t *testing.T) { t.Fatalf("expected non-ok error content, got %q", result.Content) } } + +func TestRegistrySupportsMCPToolAndHelpers(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + mcpRegistry := mcp.NewRegistry() + if err := mcpRegistry.RegisterServer("docs", "stdio", "v1", &stubMCPClient{ + tools: []mcp.ToolDescriptor{ + {Name: "search", Description: "search docs", InputSchema: map[string]any{"type": "object"}}, + }, + }); err != nil { + t.Fatalf("register mcp server: %v", err) + } + if err := mcpRegistry.RefreshServerTools(context.Background(), "docs"); err != nil { + t.Fatalf("refresh mcp tools: %v", err) + } + registry.SetMCPRegistry(mcpRegistry) + + if !registry.Supports("mcp.docs.search") { + t.Fatalf("expected supports mcp.docs.search") + } + if registry.Supports("mcp.docs.missing") { + t.Fatalf("did not expect supports mcp.docs.missing") + } + if registry.Supports("search") { + t.Fatalf("did not expect supports non-prefixed mcp name") + } + + snapshots := registry.mcpFactoryBuildSnapshot() + if len(snapshots) != 1 { + t.Fatalf("expected one snapshot, got %d", len(snapshots)) + } + if got := mcpToolFullName(" Docs ", " Search "); got != "mcp.docs.search" { + t.Fatalf("unexpected mcp full name: %q", got) + } +} From 31c61df9d468811b7ee184040d25ed0434a73421 Mon Sep 17 00:00:00 2001 From: Cai_Tang <106404101+Cai-Tang-www@users.noreply.github.com> Date: Tue, 7 Apr 2026 15:50:06 +0800 Subject: [PATCH 38/55] =?UTF-8?q?feat:=20=E6=8E=A5=E5=85=A5=E9=85=8D?= =?UTF-8?q?=E7=BD=AE=E9=A9=B1=E5=8A=A8=20MCP=20=E6=B3=A8=E5=86=8C=E5=B9=B6?= =?UTF-8?q?=E8=A1=A5=E5=85=85=E4=BD=BF=E7=94=A8=E6=96=87=E6=A1=A3?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/guides/mcp-configuration.md | 62 ++++++++++++ internal/app/bootstrap.go | 16 ++- internal/app/bootstrap_test.go | 162 ++++++++++++++++++++++++++++++- internal/app/mcp_bootstrap.go | 154 +++++++++++++++++++++++++++++ internal/config/config_test.go | 69 +++++++++++++ internal/config/model.go | 137 ++++++++++++++++++++++++++ 6 files changed, 596 insertions(+), 4 deletions(-) create mode 100644 docs/guides/mcp-configuration.md create mode 100644 internal/app/mcp_bootstrap.go diff --git a/docs/guides/mcp-configuration.md b/docs/guides/mcp-configuration.md new file mode 100644 index 00000000..64f941f6 --- /dev/null +++ b/docs/guides/mcp-configuration.md @@ -0,0 +1,62 @@ +# MCP 配置指南(stdio) + +本文档说明如何在 NeoCode 中通过配置注册 MCP server,并验证 `mcp..` 能力是否可用。 + +## 配置位置 + +在 `~/.neocode/config.yaml` 中添加 `tools.mcp.servers`: + +```yaml +tools: + mcp: + servers: + - id: docs + enabled: true + source: stdio + version: v1 + stdio: + command: node + args: + - ./mcp-server.js + workdir: ./mcp + start_timeout_sec: 8 + call_timeout_sec: 20 + restart_backoff_sec: 1 + env: + - name: MCP_TOKEN + value_env: MCP_TOKEN +``` + +## 字段说明 + +- `id`:server 稳定标识,用于工具命名空间(`mcp..`)。 +- `enabled`:是否启用该 server;仅 `true` 的 server 会在启动时注册。 +- `source`:传输类型,当前仅支持 `stdio`。 +- `version`:可选版本字段,用于可观测和后续策略命中。 +- `stdio.command`:启动命令(必填,启用时)。 +- `stdio.args`:启动参数列表。 +- `stdio.workdir`:子进程工作目录,支持相对路径(相对主 `workdir` 解析)。 +- `stdio.start_timeout_sec` / `call_timeout_sec` / `restart_backoff_sec`:可选秒级超时与重试参数。 +- `env`:传给 MCP 子进程的环境变量列表。 + - 每项必须配置 `value` 或 `value_env` 其中之一。 + - 推荐使用 `value_env` 引用系统环境变量,避免在 YAML 中写明文敏感信息。 + +## 启动行为 + +- 启动阶段会注册所有 `enabled: true` 的 server。 +- 注册后会执行一次 `tools/list` 初始化工具快照。 +- 若启用的 server 注册失败,启动会报错并中止(fail-fast)。 + +## 功能测试建议 + +1. 启动应用后让 Agent 列出工具: + - `请先列出你当前可用工具的完整名称。` +2. 检查是否存在 `mcp.docs.`。 +3. 发起一次明确调用: + - `请调用 mcp.docs.search,参数 {\"query\":\"hello\"},并返回工具结果。` + +若返回 `tool not found`,优先检查: +- `enabled` 是否为 `true`; +- `stdio.command` 是否可执行; +- `env.value_env` 对应环境变量是否存在; +- MCP server 是否支持 `tools/list`。 diff --git a/internal/app/bootstrap.go b/internal/app/bootstrap.go index dd0b8436..d57e8b25 100644 --- a/internal/app/bootstrap.go +++ b/internal/app/bootstrap.go @@ -56,7 +56,10 @@ func NewProgram(ctx context.Context) (*tea.Program, error) { cfg := manager.Get() - toolRegistry := buildToolRegistry(cfg) + toolRegistry, err := buildToolRegistry(cfg) + if err != nil { + return nil, err + } toolManager, err := buildToolManager(toolRegistry) if err != nil { return nil, err @@ -82,7 +85,7 @@ func NewProgram(ctx context.Context) (*tea.Program, error) { ), nil } -func buildToolRegistry(cfg config.Config) *tools.Registry { +func buildToolRegistry(cfg config.Config) (*tools.Registry, error) { toolRegistry := tools.NewRegistry() toolRegistry.Register(filesystem.New(cfg.Workdir)) toolRegistry.Register(filesystem.NewWrite(cfg.Workdir)) @@ -95,7 +98,14 @@ func buildToolRegistry(cfg config.Config) *tools.Registry { MaxResponseBytes: cfg.Tools.WebFetch.MaxResponseBytes, SupportedContentTypes: cfg.Tools.WebFetch.SupportedContentTypes, })) - return toolRegistry + mcpRegistry, err := buildMCPRegistry(cfg) + if err != nil { + return nil, err + } + if mcpRegistry != nil { + toolRegistry.SetMCPRegistry(mcpRegistry) + } + return toolRegistry, nil } func buildToolManager(registry *tools.Registry) (tools.Manager, error) { diff --git a/internal/app/bootstrap_test.go b/internal/app/bootstrap_test.go index 3abbf11a..67e7ba2a 100644 --- a/internal/app/bootstrap_test.go +++ b/internal/app/bootstrap_test.go @@ -10,9 +10,11 @@ import ( "path/filepath" "strings" "testing" + "time" "neo-code/internal/config" "neo-code/internal/tools" + "neo-code/internal/tools/mcp" ) func TestNewProgram(t *testing.T) { @@ -84,7 +86,10 @@ func TestBuildToolRegistryUsesWebFetchConfig(t *testing.T) { cfg.Workdir = t.TempDir() cfg.Tools.WebFetch.MaxResponseBytes = 4 - registry := buildToolRegistry(cfg) + registry, err := buildToolRegistry(cfg) + if err != nil { + t.Fatalf("buildToolRegistry() error = %v", err) + } tool, err := registry.Get("webfetch") if err != nil { t.Fatalf("registry.Get(webfetch) error = %v", err) @@ -110,6 +115,145 @@ func TestBuildToolRegistryUsesWebFetchConfig(t *testing.T) { } } +func TestBuildMCPRegistryFromConfig(t *testing.T) { + t.Parallel() + + stubClient := &stubMCPServerClient{ + tools: []mcp.ToolDescriptor{ + {Name: "search", Description: "search docs", InputSchema: map[string]any{"type": "object"}}, + }, + } + + cfg := config.Default().Clone() + cfg.Workdir = t.TempDir() + cfg.Tools.MCP.Servers = []config.MCPServerConfig{ + { + ID: "docs", + Enabled: true, + Source: "stdio", + Stdio: config.MCPStdioConfig{ + Command: "mock", + }, + }, + } + + originalRegister := registerMCPStdioServer + t.Cleanup(func() { registerMCPStdioServer = originalRegister }) + registerMCPStdioServer = func(registry *mcp.Registry, cfg config.Config, server config.MCPServerConfig) error { + if err := registry.RegisterServer(server.ID, "stdio", server.Version, stubClient); err != nil { + return err + } + return registry.RefreshServerTools(context.Background(), server.ID) + } + + registry, err := buildMCPRegistry(cfg) + if err != nil { + t.Fatalf("buildMCPRegistry() error = %v", err) + } + if registry == nil { + t.Fatalf("expected non-nil mcp registry") + } + snapshots := registry.Snapshot() + if len(snapshots) != 1 || snapshots[0].ServerID != "docs" { + t.Fatalf("unexpected snapshots: %+v", snapshots) + } +} + +func TestBuildToolRegistryIncludesMCPFromConfig(t *testing.T) { + t.Parallel() + + cfg := config.Default().Clone() + cfg.Workdir = t.TempDir() + cfg.Tools.MCP.Servers = []config.MCPServerConfig{ + { + ID: "docs", + Enabled: true, + Source: "stdio", + Stdio: config.MCPStdioConfig{ + Command: "mock", + }, + }, + } + + originalRegister := registerMCPStdioServer + t.Cleanup(func() { registerMCPStdioServer = originalRegister }) + registerMCPStdioServer = func(registry *mcp.Registry, cfg config.Config, server config.MCPServerConfig) error { + client := &stubMCPServerClient{ + tools: []mcp.ToolDescriptor{ + {Name: "search", Description: "search docs", InputSchema: map[string]any{"type": "object"}}, + }, + } + if err := registry.RegisterServer(server.ID, "stdio", server.Version, client); err != nil { + return err + } + return registry.RefreshServerTools(context.Background(), server.ID) + } + + registry, err := buildToolRegistry(cfg) + if err != nil { + t.Fatalf("buildToolRegistry() error = %v", err) + } + specs, err := registry.ListAvailableSpecs(context.Background(), tools.SpecListInput{}) + if err != nil { + t.Fatalf("ListAvailableSpecs() error = %v", err) + } + found := false + for _, spec := range specs { + if spec.Name == "mcp.docs.search" { + found = true + break + } + } + if !found { + t.Fatalf("expected mcp.docs.search in specs, got %+v", specs) + } +} + +func TestResolveMCPServerEnvAndWorkdir(t *testing.T) { + t.Setenv("MCP_TOKEN", "secret") + env, err := resolveMCPServerEnv(config.MCPServerConfig{ + Env: []config.MCPEnvVarConfig{ + {Name: "TOKEN", ValueEnv: "MCP_TOKEN"}, + {Name: "MODE", Value: "test"}, + }, + }) + if err != nil { + t.Fatalf("resolveMCPServerEnv() error = %v", err) + } + joined := strings.Join(env, ",") + if !strings.Contains(joined, "TOKEN=secret") || !strings.Contains(joined, "MODE=test") { + t.Fatalf("unexpected env result: %+v", env) + } + + base := t.TempDir() + relative := resolveMCPServerWorkdir(base, "tools/mcp") + if !strings.HasSuffix(filepath.ToSlash(relative), "tools/mcp") { + t.Fatalf("unexpected relative workdir: %q", relative) + } + absoluteTarget := filepath.Join(t.TempDir(), "absolute") + absolute := resolveMCPServerWorkdir(base, absoluteTarget) + if absolute != filepath.Clean(absoluteTarget) { + t.Fatalf("unexpected absolute workdir: %q", absolute) + } +} + +func TestInitialMCPRefreshTimeoutAndDurationConversion(t *testing.T) { + t.Parallel() + + cfg := config.Default().Clone() + cfg.ToolTimeoutSec = 1 + timeout := initialMCPRefreshTimeout(cfg) + if timeout < 5*time.Second { + t.Fatalf("expected minimum timeout >= 5s, got %v", timeout) + } + if durationFromSeconds(0) != 0 { + t.Fatalf("expected zero duration for non-positive input") + } + if durationFromSeconds(2) != 2*time.Second { + t.Fatalf("expected 2s duration") + } +} + func TestBuildToolManagerWrapsRegistry(t *testing.T) { t.Parallel() @@ -252,3 +396,19 @@ func disableBuiltinProviderAPIKeys(t *testing.T) { t.Setenv(config.OpenLLDefaultAPIKeyEnv, "") t.Setenv(config.QiniuDefaultAPIKeyEnv, "") } + +type stubMCPServerClient struct { + tools []mcp.ToolDescriptor +} + +func (s *stubMCPServerClient) ListTools(ctx context.Context) ([]mcp.ToolDescriptor, error) { + return append([]mcp.ToolDescriptor(nil), s.tools...), nil +} + +func (s *stubMCPServerClient) CallTool(ctx context.Context, toolName string, arguments []byte) (mcp.CallResult, error) { + return mcp.CallResult{Content: "ok"}, nil +} + +func (s *stubMCPServerClient) HealthCheck(ctx context.Context) error { + return nil +} diff --git a/internal/app/mcp_bootstrap.go b/internal/app/mcp_bootstrap.go new file mode 100644 index 00000000..27982e14 --- /dev/null +++ b/internal/app/mcp_bootstrap.go @@ -0,0 +1,154 @@ +package app + +import ( + "context" + "fmt" + "os" + "path/filepath" + "strings" + "time" + + "neo-code/internal/config" + "neo-code/internal/tools/mcp" +) + +var newMCPStdioClient = mcp.NewStdIOClient +var registerMCPStdioServer = defaultRegisterMCPStdioServer + +// buildMCPRegistry 按配置构建并初始化 MCP registry;若无启用 server 则返回 nil。 +func buildMCPRegistry(cfg config.Config) (*mcp.Registry, error) { + if len(cfg.Tools.MCP.Servers) == 0 { + return nil, nil + } + + registry := mcp.NewRegistry() + enabledCount := 0 + for index := range cfg.Tools.MCP.Servers { + server := cfg.Tools.MCP.Servers[index] + if !server.Enabled { + continue + } + enabledCount++ + + switch strings.ToLower(strings.TrimSpace(server.Source)) { + case "", "stdio": + if err := registerMCPStdioServer(registry, cfg, server); err != nil { + return nil, fmt.Errorf("app: register mcp server %q: %w", strings.TrimSpace(server.ID), err) + } + default: + return nil, fmt.Errorf("app: unsupported mcp source %q", server.Source) + } + } + + if enabledCount == 0 { + return nil, nil + } + return registry, nil +} + +// defaultRegisterMCPStdioServer 创建 stdio client 并完成 server 注册与 tools 快照初始化。 +func defaultRegisterMCPStdioServer(registry *mcp.Registry, cfg config.Config, server config.MCPServerConfig) error { + env, err := resolveMCPServerEnv(server) + if err != nil { + return err + } + + workdir := resolveMCPServerWorkdir(cfg.Workdir, server.Stdio.Workdir) + client, err := newMCPStdioClient(mcp.StdioClientConfig{ + Command: strings.TrimSpace(server.Stdio.Command), + Args: append([]string(nil), server.Stdio.Args...), + Env: env, + Workdir: workdir, + StartTimeout: durationFromSeconds(server.Stdio.StartTimeoutSec), + CallTimeout: durationFromSeconds(server.Stdio.CallTimeoutSec), + RestartBackoff: durationFromSeconds(server.Stdio.RestartBackoffSec), + }) + if err != nil { + return err + } + + serverID := strings.TrimSpace(server.ID) + source := strings.ToLower(strings.TrimSpace(server.Source)) + if source == "" { + source = "stdio" + } + if err := registry.RegisterServer(serverID, source, strings.TrimSpace(server.Version), client); err != nil { + return err + } + + refreshCtx, cancel := context.WithTimeout(context.Background(), initialMCPRefreshTimeout(cfg)) + defer cancel() + if err := registry.RefreshServerTools(refreshCtx, serverID); err != nil { + return err + } + return nil +} + +// resolveMCPServerEnv 将配置中的 env 绑定解析为子进程环境变量。 +func resolveMCPServerEnv(server config.MCPServerConfig) ([]string, error) { + if len(server.Env) == 0 { + return nil, nil + } + result := make([]string, 0, len(server.Env)) + for index, item := range server.Env { + name := strings.TrimSpace(item.Name) + if name == "" { + return nil, fmt.Errorf("env[%d].name is empty", index) + } + + value := strings.TrimSpace(item.Value) + valueEnv := strings.TrimSpace(item.ValueEnv) + switch { + case value != "" && valueEnv != "": + return nil, fmt.Errorf("env[%d] must set either value or value_env", index) + case value != "": + result = append(result, name+"="+value) + case valueEnv != "": + resolved := strings.TrimSpace(os.Getenv(valueEnv)) + if resolved == "" { + return nil, fmt.Errorf("env[%d] value_env %q is empty", index, valueEnv) + } + result = append(result, name+"="+resolved) + default: + return nil, fmt.Errorf("env[%d] must set one of value/value_env", index) + } + } + return result, nil +} + +// resolveMCPServerWorkdir 解析 MCP server 子进程工作目录,支持相对路径。 +func resolveMCPServerWorkdir(baseWorkdir string, override string) string { + trimmedOverride := strings.TrimSpace(override) + if trimmedOverride == "" { + return strings.TrimSpace(baseWorkdir) + } + if filepath.IsAbs(trimmedOverride) { + return filepath.Clean(trimmedOverride) + } + + trimmedBase := strings.TrimSpace(baseWorkdir) + if trimmedBase == "" { + return filepath.Clean(trimmedOverride) + } + return filepath.Clean(filepath.Join(trimmedBase, trimmedOverride)) +} + +// initialMCPRefreshTimeout 计算启动阶段首轮 tools 刷新的超时时间。 +func initialMCPRefreshTimeout(cfg config.Config) time.Duration { + timeout := time.Duration(cfg.ToolTimeoutSec) * time.Second + if timeout <= 0 { + timeout = 20 * time.Second + } + if timeout < 5*time.Second { + timeout = 5 * time.Second + } + return timeout +} + +// durationFromSeconds 将秒级配置转换为 duration;非正值返回 0 以启用 client 默认值。 +func durationFromSeconds(seconds int) time.Duration { + if seconds <= 0 { + return 0 + } + return time.Duration(seconds) * time.Second +} diff --git a/internal/config/config_test.go b/internal/config/config_test.go index a2504ac2..29b11b52 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -471,6 +471,47 @@ func TestConfigValidateFailures(t *testing.T) { }(), expectErr: "duplicate provider endpoint", }, + { + name: "invalid mcp duplicate server id", + config: func() *Config { + cfg := validConfig.Clone() + cfg.Tools.MCP.Servers = []MCPServerConfig{ + {ID: "docs", Enabled: true, Stdio: MCPStdioConfig{Command: "cmd-1"}}, + {ID: "docs", Enabled: true, Stdio: MCPStdioConfig{Command: "cmd-2"}}, + } + return &cfg + }(), + expectErr: "duplicate servers", + }, + { + name: "invalid mcp source", + config: func() *Config { + cfg := validConfig.Clone() + cfg.Tools.MCP.Servers = []MCPServerConfig{ + {ID: "docs", Enabled: true, Source: "sse", Stdio: MCPStdioConfig{Command: "cmd"}}, + } + return &cfg + }(), + expectErr: "not supported", + }, + { + name: "invalid mcp env binding", + config: func() *Config { + cfg := validConfig.Clone() + cfg.Tools.MCP.Servers = []MCPServerConfig{ + { + ID: "docs", + Enabled: true, + Stdio: MCPStdioConfig{Command: "cmd"}, + Env: []MCPEnvVarConfig{ + {Name: "TOKEN", Value: "a", ValueEnv: "TOKEN_ENV"}, + }, + }, + } + return &cfg + }(), + expectErr: "exactly one of value/value_env", + }, } for _, tt := range tests { @@ -485,6 +526,34 @@ func TestConfigValidateFailures(t *testing.T) { } } +func TestMCPConfigApplyDefaultsAndClone(t *testing.T) { + t.Parallel() + + cfg := MCPConfig{ + Servers: []MCPServerConfig{ + { + ID: " Docs ", + Enabled: true, + Source: "", + Stdio: MCPStdioConfig{ + Command: "mock", + Args: []string{"a"}, + }, + }, + }, + } + cfg.ApplyDefaults(defaultMCPConfig()) + if cfg.Servers[0].Source != "stdio" { + t.Fatalf("expected default source stdio, got %q", cfg.Servers[0].Source) + } + + cloned := cfg.Clone() + cloned.Servers[0].Stdio.Args[0] = "b" + if cfg.Servers[0].Stdio.Args[0] == "b" { + t.Fatalf("expected MCP clone to be independent") + } +} + func TestProviderConfigValidateFailures(t *testing.T) { t.Parallel() diff --git a/internal/config/model.go b/internal/config/model.go index 75992d1c..3296dc20 100644 --- a/internal/config/model.go +++ b/internal/config/model.go @@ -60,6 +60,7 @@ type ResolvedProviderConfig struct { type ToolsConfig struct { WebFetch WebFetchConfig `yaml:"webfetch,omitempty"` + MCP MCPConfig `yaml:"mcp,omitempty"` } type ContextConfig struct { @@ -78,6 +79,34 @@ type WebFetchConfig struct { SupportedContentTypes []string `yaml:"supported_content_types,omitempty"` } +type MCPConfig struct { + Servers []MCPServerConfig `yaml:"servers,omitempty"` +} + +type MCPServerConfig struct { + ID string `yaml:"id"` + Enabled bool `yaml:"enabled,omitempty"` + Source string `yaml:"source,omitempty"` + Version string `yaml:"version,omitempty"` + Stdio MCPStdioConfig `yaml:"stdio,omitempty"` + Env []MCPEnvVarConfig `yaml:"env,omitempty"` +} + +type MCPStdioConfig struct { + Command string `yaml:"command,omitempty"` + Args []string `yaml:"args,omitempty"` + Workdir string `yaml:"workdir,omitempty"` + StartTimeoutSec int `yaml:"start_timeout_sec,omitempty"` + CallTimeoutSec int `yaml:"call_timeout_sec,omitempty"` + RestartBackoffSec int `yaml:"restart_backoff_sec,omitempty"` +} + +type MCPEnvVarConfig struct { + Name string `yaml:"name"` + Value string `yaml:"value,omitempty"` + ValueEnv string `yaml:"value_env,omitempty"` +} + func DefaultWebFetchSupportedContentTypes() []string { return append([]string(nil), defaultWebFetchSupportedContentTypes...) } @@ -91,6 +120,7 @@ func Default() *Config { Context: defaultContextConfig(), Tools: ToolsConfig{ WebFetch: defaultWebFetchConfig(), + MCP: defaultMCPConfig(), }, } } @@ -330,6 +360,13 @@ func defaultWebFetchConfig() WebFetchConfig { } } +// defaultMCPConfig 返回 MCP 工具接入配置的默认值(默认无 server)。 +func defaultMCPConfig() MCPConfig { + return MCPConfig{ + Servers: nil, + } +} + // defaultContextConfig 返回上下文压缩相关配置的默认值。 func defaultContextConfig() ContextConfig { return ContextConfig{ @@ -349,6 +386,7 @@ func defaultCompactConfig() CompactConfig { func (c ToolsConfig) Clone() ToolsConfig { return ToolsConfig{ WebFetch: c.WebFetch.Clone(), + MCP: c.MCP.Clone(), } } @@ -365,6 +403,7 @@ func (c *ToolsConfig) ApplyDefaults(defaults ToolsConfig) { } c.WebFetch.ApplyDefaults(defaults.WebFetch) + c.MCP.ApplyDefaults(defaults.MCP) } // ApplyDefaults 为上下文配置补齐缺省的 compact 参数。 @@ -380,6 +419,93 @@ func (c ToolsConfig) Validate() error { if err := c.WebFetch.Validate(); err != nil { return fmt.Errorf("webfetch: %w", err) } + if err := c.MCP.Validate(); err != nil { + return fmt.Errorf("mcp: %w", err) + } + return nil +} + +// Clone 返回 MCP 配置的独立副本,避免引用共享造成并发污染。 +func (c MCPConfig) Clone() MCPConfig { + if len(c.Servers) == 0 { + return MCPConfig{} + } + cloned := make([]MCPServerConfig, 0, len(c.Servers)) + for _, server := range c.Servers { + cloned = append(cloned, server.Clone()) + } + return MCPConfig{ + Servers: cloned, + } +} + +// Clone 返回单个 MCP server 配置的独立副本。 +func (c MCPServerConfig) Clone() MCPServerConfig { + cloned := c + cloned.Stdio.Args = append([]string(nil), c.Stdio.Args...) + if len(c.Env) > 0 { + cloned.Env = make([]MCPEnvVarConfig, 0, len(c.Env)) + cloned.Env = append(cloned.Env, c.Env...) + } else { + cloned.Env = nil + } + return cloned +} + +// ApplyDefaults 为 MCP 配置补齐缺省字段,保证运行时行为可预测。 +func (c *MCPConfig) ApplyDefaults(defaults MCPConfig) { + if c == nil { + return + } + if len(c.Servers) == 0 { + c.Servers = defaults.Clone().Servers + } + for index := range c.Servers { + c.Servers[index].ApplyDefaults() + } +} + +// Validate 校验 MCP server 列表与字段合法性,防止启动后失败。 +func (c MCPConfig) Validate() error { + if len(c.Servers) == 0 { + return nil + } + seen := make(map[string]struct{}, len(c.Servers)) + for index, server := range c.Servers { + normalizedID := strings.ToLower(strings.TrimSpace(server.ID)) + if normalizedID == "" { + return fmt.Errorf("servers[%d].id is empty", index) + } + if _, exists := seen[normalizedID]; exists { + return fmt.Errorf("duplicate servers[%d].id %q", index, server.ID) + } + seen[normalizedID] = struct{}{} + + source := strings.ToLower(strings.TrimSpace(server.Source)) + if source == "" { + source = "stdio" + } + if source != "stdio" { + return fmt.Errorf("servers[%d].source %q is not supported", index, server.Source) + } + if !server.Enabled { + continue + } + + if strings.TrimSpace(server.Stdio.Command) == "" { + return fmt.Errorf("servers[%d].stdio.command is empty", index) + } + for envIndex, env := range server.Env { + if strings.TrimSpace(env.Name) == "" { + return fmt.Errorf("servers[%d].env[%d].name is empty", index, envIndex) + } + hasValue := strings.TrimSpace(env.Value) != "" + hasValueEnv := strings.TrimSpace(env.ValueEnv) != "" + if hasValue == hasValueEnv { + return fmt.Errorf("servers[%d].env[%d] must set exactly one of value/value_env", index, envIndex) + } + } + } return nil } @@ -413,6 +539,17 @@ func (c *WebFetchConfig) ApplyDefaults(defaults WebFetchConfig) { c.SupportedContentTypes = normalizeContentTypes(c.SupportedContentTypes, defaults.SupportedContentTypes) } +// ApplyDefaults 为 MCP server stdio 配置补齐默认值。 +func (c *MCPServerConfig) ApplyDefaults() { + if c == nil { + return + } + c.Source = strings.ToLower(strings.TrimSpace(c.Source)) + if c.Source == "" { + c.Source = "stdio" + } +} + // ApplyDefaults 为 compact 配置填充缺省策略和阈值。 func (c *CompactConfig) ApplyDefaults(defaults CompactConfig) { if c == nil { From a07b88deac1105a3170913a8b3986422e078537a Mon Sep 17 00:00:00 2001 From: Cai_Tang <106404101+Cai-Tang-www@users.noreply.github.com> Date: Tue, 7 Apr 2026 16:55:57 +0800 Subject: [PATCH 39/55] =?UTF-8?q?feat:=20=E8=A1=A5=E9=BD=90=20MCP=20stdio?= =?UTF-8?q?=20initialize=20=E6=8F=A1=E6=89=8B=E4=B8=8E=E9=87=8D=E8=BF=9E?= =?UTF-8?q?=E4=BC=9A=E8=AF=9D=E6=B5=81=E7=A8=8B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/tools/mcp/stdio_client.go | 171 +++++++++++++++++++++--- internal/tools/mcp/stdio_client_test.go | 76 ++++++++++- 2 files changed, 229 insertions(+), 18 deletions(-) diff --git a/internal/tools/mcp/stdio_client.go b/internal/tools/mcp/stdio_client.go index 7b9196f9..71d28d3c 100644 --- a/internal/tools/mcp/stdio_client.go +++ b/internal/tools/mcp/stdio_client.go @@ -23,6 +23,9 @@ const ( defaultStdioRestartBackoff = 1 * time.Second maxStdioRestartBackoff = 30 * time.Second maxStdioFrameBytes = 8 * 1024 * 1024 + defaultMCPProtocolVersion = "2024-11-05" + defaultMCPClientName = "neocode" + defaultMCPClientVersion = "0.1.0" ) // StdioClientConfig 描述 MCP stdio 客户端的启动与调用参数。 @@ -43,6 +46,12 @@ type jsonRPCRequest struct { Params any `json:"params,omitempty"` } +type jsonRPCNotification struct { + JSONRPC string `json:"jsonrpc"` + Method string `json:"method"` + Params any `json:"params,omitempty"` +} + type jsonRPCResponse struct { JSONRPC string `json:"jsonrpc"` ID string `json:"id"` @@ -62,21 +71,24 @@ type rpcReply struct { // StdIOClient 通过 stdio 子进程与 MCP server 进行 JSON-RPC 通信。 type StdIOClient struct { - cfg StdioClientConfig - idSeed uint64 - mu sync.Mutex - writeMu sync.Mutex - cmd *exec.Cmd - stdin io.WriteCloser - stdout io.ReadCloser - reader *bufio.Reader - pending map[string]chan rpcReply - exited chan struct{} - exitErr error - backoff time.Duration - retryAt time.Time - started bool - shutdown bool + cfg StdioClientConfig + idSeed uint64 + mu sync.Mutex + writeMu sync.Mutex + cmd *exec.Cmd + stdin io.WriteCloser + stdout io.ReadCloser + reader *bufio.Reader + pending map[string]chan rpcReply + exited chan struct{} + exitErr error + backoff time.Duration + retryAt time.Time + started bool + initialized bool + initializing bool + initDone chan struct{} + shutdown bool } // NewStdIOClient 创建 stdio MCP client。 @@ -201,11 +213,21 @@ func (c *StdIOClient) callContext(ctx context.Context) (context.Context, context } func (c *StdIOClient) call(ctx context.Context, method string, params any) (json.RawMessage, error) { + return c.callRequest(ctx, method, params, false) +} + +// callRequest 发送带响应的 RPC 请求;skipEnsure=true 用于初始化阶段避免递归。 +func (c *StdIOClient) callRequest(ctx context.Context, method string, params any, skipEnsure bool) (json.RawMessage, error) { if err := ctx.Err(); err != nil { return nil, err } - if err := c.ensureStarted(ctx); err != nil { - return nil, err + if !skipEnsure { + if err := c.ensureStarted(ctx); err != nil { + return nil, err + } + if err := c.ensureInitialized(ctx); err != nil { + return nil, err + } } requestID := "req-" + strconv.FormatUint(atomic.AddUint64(&c.idSeed, 1), 10) @@ -251,6 +273,49 @@ func (c *StdIOClient) call(ctx context.Context, method string, params any) (json } } +// sendNotification 发送无需响应的 RPC 通知;skipEnsure=true 用于初始化流程。 +func (c *StdIOClient) sendNotification(ctx context.Context, method string, params any, skipEnsure bool) error { + if err := ctx.Err(); err != nil { + return err + } + if !skipEnsure { + if err := c.ensureStarted(ctx); err != nil { + return err + } + if err := c.ensureInitialized(ctx); err != nil { + return err + } + } + + c.mu.Lock() + if c.shutdown { + c.mu.Unlock() + return errors.New("mcp: stdio client closed") + } + stdin := c.stdin + c.mu.Unlock() + if stdin == nil { + return errors.New("mcp: stdio client is not connected") + } + + payload, err := json.Marshal(jsonRPCNotification{ + JSONRPC: "2.0", + Method: method, + Params: params, + }) + if err != nil { + return fmt.Errorf("mcp: marshal notification: %w", err) + } + + c.writeMu.Lock() + writeErr := writeFramedMessage(stdin, payload) + c.writeMu.Unlock() + if writeErr != nil { + return fmt.Errorf("mcp: send notification: %w", writeErr) + } + return nil +} + func (c *StdIOClient) ensureStarted(ctx context.Context) error { c.mu.Lock() defer c.mu.Unlock() @@ -306,6 +371,9 @@ func (c *StdIOClient) ensureStarted(ctx context.Context) error { c.exited = make(chan struct{}) c.exitErr = nil c.started = true + c.initialized = false + c.initializing = false + c.initDone = nil c.backoff = c.cfg.RestartBackoff c.retryAt = time.Time{} @@ -315,6 +383,69 @@ func (c *StdIOClient) ensureStarted(ctx context.Context) error { return nil } +// ensureInitialized 确保 MCP 会话完成 initialize/initialized 握手,并发调用共享结果。 +func (c *StdIOClient) ensureInitialized(ctx context.Context) error { + for { + c.mu.Lock() + if c.shutdown { + c.mu.Unlock() + return errors.New("mcp: stdio client closed") + } + if !c.started { + c.mu.Unlock() + return errors.New("mcp: stdio client is not started") + } + if c.initialized { + c.mu.Unlock() + return nil + } + if c.initializing { + wait := c.initDone + c.mu.Unlock() + select { + case <-ctx.Done(): + return ctx.Err() + case <-wait: + continue + } + } + c.initializing = true + c.initDone = make(chan struct{}) + done := c.initDone + c.mu.Unlock() + + initErr := c.performInitialize(ctx) + + c.mu.Lock() + if c.started && initErr == nil { + c.initialized = true + } + c.initializing = false + close(done) + c.mu.Unlock() + return initErr + } +} + +// performInitialize 执行标准 MCP 初始化握手:initialize -> notifications/initialized。 +func (c *StdIOClient) performInitialize(ctx context.Context) error { + params := map[string]any{ + "protocolVersion": defaultMCPProtocolVersion, + "capabilities": map[string]any{}, + "clientInfo": map[string]any{ + "name": defaultMCPClientName, + "version": defaultMCPClientVersion, + }, + } + if _, err := c.callRequest(ctx, "initialize", params, true); err != nil { + return fmt.Errorf("mcp: initialize session: %w", err) + } + if err := c.sendNotification(ctx, "notifications/initialized", map[string]any{}, true); err != nil { + return fmt.Errorf("mcp: notify initialized: %w", err) + } + return nil +} + func (c *StdIOClient) readLoop() { for { message, err := readFramedMessage(c.reader) @@ -368,6 +499,12 @@ func (c *StdIOClient) markExited(err error) { return } c.started = false + c.initialized = false + if c.initializing && c.initDone != nil { + close(c.initDone) + } + c.initializing = false + c.initDone = nil c.exitErr = err if c.exited != nil { close(c.exited) diff --git a/internal/tools/mcp/stdio_client_test.go b/internal/tools/mcp/stdio_client_test.go index 3a352d3a..a54757e6 100644 --- a/internal/tools/mcp/stdio_client_test.go +++ b/internal/tools/mcp/stdio_client_test.go @@ -276,7 +276,7 @@ func newTestStdIOClient(t *testing.T) *StdIOClient { client, err := NewStdIOClient(StdioClientConfig{ Command: os.Args[0], Args: []string{"-test.run=TestHelperProcessMCPStdioServer", "--"}, - Env: []string{"GO_WANT_MCP_STDIO_HELPER=1"}, + Env: []string{"GO_WANT_MCP_STDIO_HELPER=1", "GO_MCP_STDIO_REQUIRE_INITIALIZE=1"}, StartTimeout: 3 * time.Second, CallTimeout: 3 * time.Second, }) @@ -286,11 +286,36 @@ func newTestStdIOClient(t *testing.T) *StdIOClient { return client } +func TestStdIOClientInitializeFailure(t *testing.T) { + t.Parallel() + + client, err := NewStdIOClient(StdioClientConfig{ + Command: os.Args[0], + Args: []string{"-test.run=TestHelperProcessMCPStdioServer", "--"}, + Env: []string{"GO_WANT_MCP_STDIO_HELPER=1", "GO_MCP_STDIO_INIT_FAIL=1"}, + StartTimeout: 3 * time.Second, + CallTimeout: 3 * time.Second, + }) + if err != nil { + t.Fatalf("NewStdIOClient() error = %v", err) + } + defer func() { _ = client.Close() }() + + _, callErr := client.ListTools(context.Background()) + if callErr == nil || !strings.Contains(callErr.Error(), "initialize session") { + t.Fatalf("expected initialize error, got %v", callErr) + } +} + func TestHelperProcessMCPStdioServer(t *testing.T) { if os.Getenv("GO_WANT_MCP_STDIO_HELPER") != "1" { return } + requireInitialize := os.Getenv("GO_MCP_STDIO_REQUIRE_INITIALIZE") == "1" + initFail := os.Getenv("GO_MCP_STDIO_INIT_FAIL") == "1" + initialized := !requireInitialize + reader := bufio.NewReader(os.Stdin) for { payload, err := readFramedMessage(reader) @@ -311,7 +336,45 @@ func TestHelperProcessMCPStdioServer(t *testing.T) { var response any switch method { + case "initialize": + if initFail { + response = map[string]any{ + "jsonrpc": "2.0", + "id": requestID, + "error": map[string]any{ + "code": -32600, + "message": "initialize rejected", + }, + } + break + } + response = map[string]any{ + "jsonrpc": "2.0", + "id": requestID, + "result": map[string]any{ + "protocolVersion": "2024-11-05", + "capabilities": map[string]any{}, + "serverInfo": map[string]any{ + "name": "test-helper", + "version": "1.0.0", + }, + }, + } + case "notifications/initialized": + initialized = true + continue case "tools/list": + if !initialized { + response = map[string]any{ + "jsonrpc": "2.0", + "id": requestID, + "error": map[string]any{ + "code": -32002, + "message": "server not initialized", + }, + } + break + } response = map[string]any{ "jsonrpc": "2.0", "id": requestID, @@ -329,6 +392,17 @@ func TestHelperProcessMCPStdioServer(t *testing.T) { }, } case "tools/call": + if !initialized { + response = map[string]any{ + "jsonrpc": "2.0", + "id": requestID, + "error": map[string]any{ + "code": -32002, + "message": "server not initialized", + }, + } + break + } params, _ := request["params"].(map[string]any) name, _ := params["name"].(string) response = map[string]any{ From 48530c6fb21d411a4cd1f5dbf24589c2167b852f Mon Sep 17 00:00:00 2001 From: Cai_Tang <106404101+Cai-Tang-www@users.noreply.github.com> Date: Tue, 7 Apr 2026 17:49:22 +0800 Subject: [PATCH 40/55] =?UTF-8?q?feat:=20=E5=A2=9E=E5=BC=BAMCP=20stdio?= =?UTF-8?q?=E6=8F=A1=E6=89=8B=E5=85=BC=E5=AE=B9=E5=B9=B6=E8=A1=A5=E9=BD=90?= =?UTF-8?q?=E8=A6=86=E7=9B=96=E6=B5=8B=E8=AF=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/tools/mcp/stdio_client.go | 247 +++++++++++++++++++++--- internal/tools/mcp/stdio_client_test.go | 132 ++++++++++++- 2 files changed, 347 insertions(+), 32 deletions(-) diff --git a/internal/tools/mcp/stdio_client.go b/internal/tools/mcp/stdio_client.go index 71d28d3c..b6808d51 100644 --- a/internal/tools/mcp/stdio_client.go +++ b/internal/tools/mcp/stdio_client.go @@ -23,6 +23,7 @@ const ( defaultStdioRestartBackoff = 1 * time.Second maxStdioRestartBackoff = 30 * time.Second maxStdioFrameBytes = 8 * 1024 * 1024 + maxStdioLineBytes = 8 * 1024 * 1024 defaultMCPProtocolVersion = "2024-11-05" defaultMCPClientName = "neocode" defaultMCPClientVersion = "0.1.0" @@ -88,9 +89,18 @@ type StdIOClient struct { initialized bool initializing bool initDone chan struct{} + protocol stdioProtocol shutdown bool } +type stdioProtocol string + +const ( + stdioProtocolUnknown stdioProtocol = "" + stdioProtocolLine stdioProtocol = "line" + stdioProtocolFramed stdioProtocol = "framed" +) + // NewStdIOClient 创建 stdio MCP client。 func NewStdIOClient(cfg StdioClientConfig) (*StdIOClient, error) { if strings.TrimSpace(cfg.Command) == "" { @@ -132,7 +142,7 @@ func (c *StdIOClient) Close() error { return nil } -// ListTools 调用 MCP `tools/list` 获取工具清单。 +// ListTools 调用 MCP tools/list 获取工具清单。 func (c *StdIOClient) ListTools(ctx context.Context) ([]ToolDescriptor, error) { callCtx, cancel := c.callContext(ctx) defer cancel() @@ -169,7 +179,7 @@ func (c *StdIOClient) ListTools(ctx context.Context) ([]ToolDescriptor, error) { return result, nil } -// CallTool 调用 MCP `tools/call` 并收敛返回值。 +// CallTool 调用 MCP tools/call 并收敛返回值。 func (c *StdIOClient) CallTool(ctx context.Context, toolName string, arguments []byte) (CallResult, error) { trimmedToolName := strings.TrimSpace(toolName) if trimmedToolName == "" { @@ -196,12 +206,13 @@ func (c *StdIOClient) CallTool(ctx context.Context, toolName string, arguments [ return decodeCallResult(raw), nil } -// HealthCheck 通过一次短超时 `tools/list` 验证连接可用性。 +// HealthCheck 通过一次短超时 tools/list 验证连接可用性。 func (c *StdIOClient) HealthCheck(ctx context.Context) error { _, err := c.ListTools(ctx) return err } +// callContext 基于配置与上游截止时间生成单次 RPC 调用上下文。 func (c *StdIOClient) callContext(ctx context.Context) (context.Context, context.CancelFunc) { timeout := c.cfg.CallTimeout if deadline, ok := ctx.Deadline(); ok { @@ -212,12 +223,24 @@ func (c *StdIOClient) callContext(ctx context.Context) (context.Context, context return context.WithTimeout(ctx, timeout) } +// call 发送 RPC 请求并等待响应。 func (c *StdIOClient) call(ctx context.Context, method string, params any) (json.RawMessage, error) { return c.callRequest(ctx, method, params, false) } -// callRequest 发送带响应的 RPC 请求;skipEnsure=true 用于初始化阶段避免递归。 +// callRequest 在可用连接上发送 RPC 请求,skipEnsure=true 用于初始化阶段避免递归。 func (c *StdIOClient) callRequest(ctx context.Context, method string, params any, skipEnsure bool) (json.RawMessage, error) { + return c.callRequestWithProtocol(ctx, method, params, skipEnsure, stdioProtocolUnknown) +} + +// callRequestWithProtocol 按指定协议发送 RPC 请求,用于 initialize 时的协议探测与回退。 +func (c *StdIOClient) callRequestWithProtocol( + ctx context.Context, + method string, + params any, + skipEnsure bool, + override stdioProtocol, +) (json.RawMessage, error) { if err := ctx.Err(); err != nil { return nil, err } @@ -240,6 +263,7 @@ func (c *StdIOClient) callRequest(ctx context.Context, method string, params any } c.pending[requestID] = replyCh stdin := c.stdin + selectedProtocol := c.resolveWriteProtocolLocked(override) c.mu.Unlock() if stdin == nil { c.removePending(requestID) @@ -257,7 +281,7 @@ func (c *StdIOClient) callRequest(ctx context.Context, method string, params any return nil, fmt.Errorf("mcp: marshal request: %w", err) } c.writeMu.Lock() - writeErr := writeFramedMessage(stdin, requestPayload) + writeErr := writeMessageWithProtocol(stdin, requestPayload, selectedProtocol) c.writeMu.Unlock() if writeErr != nil { c.removePending(requestID) @@ -273,8 +297,19 @@ func (c *StdIOClient) callRequest(ctx context.Context, method string, params any } } -// sendNotification 发送无需响应的 RPC 通知;skipEnsure=true 用于初始化流程。 +// sendNotification 发送无响应 RPC 通知,skipEnsure=true 用于初始化流程。 func (c *StdIOClient) sendNotification(ctx context.Context, method string, params any, skipEnsure bool) error { + return c.sendNotificationWithProtocol(ctx, method, params, skipEnsure, stdioProtocolUnknown) +} + +// sendNotificationWithProtocol 按指定协议发送 RPC 通知,用于 initialize 完成后的 initialized 事件。 +func (c *StdIOClient) sendNotificationWithProtocol( + ctx context.Context, + method string, + params any, + skipEnsure bool, + override stdioProtocol, +) error { if err := ctx.Err(); err != nil { return err } @@ -293,6 +328,7 @@ func (c *StdIOClient) sendNotification(ctx context.Context, method string, param return errors.New("mcp: stdio client closed") } stdin := c.stdin + selectedProtocol := c.resolveWriteProtocolLocked(override) c.mu.Unlock() if stdin == nil { return errors.New("mcp: stdio client is not connected") @@ -308,7 +344,7 @@ func (c *StdIOClient) sendNotification(ctx context.Context, method string, param } c.writeMu.Lock() - writeErr := writeFramedMessage(stdin, payload) + writeErr := writeMessageWithProtocol(stdin, payload, selectedProtocol) c.writeMu.Unlock() if writeErr != nil { return fmt.Errorf("mcp: send notification: %w", writeErr) @@ -316,6 +352,7 @@ func (c *StdIOClient) sendNotification(ctx context.Context, method string, param return nil } +// ensureStarted 确保 stdio 子进程已启动并处于可读写状态。 func (c *StdIOClient) ensureStarted(ctx context.Context) error { c.mu.Lock() defer c.mu.Unlock() @@ -374,6 +411,7 @@ func (c *StdIOClient) ensureStarted(ctx context.Context) error { c.initialized = false c.initializing = false c.initDone = nil + c.protocol = stdioProtocolUnknown c.backoff = c.cfg.RestartBackoff c.retryAt = time.Time{} @@ -427,7 +465,7 @@ func (c *StdIOClient) ensureInitialized(ctx context.Context) error { } } -// performInitialize 执行标准 MCP 初始化握手:initialize -> notifications/initialized。 +// performInitialize 执行 MCP 握手并自动兼容 line/framed 两种 stdio 线协议。 func (c *StdIOClient) performInitialize(ctx context.Context) error { params := map[string]any{ "protocolVersion": defaultMCPProtocolVersion, @@ -437,23 +475,62 @@ func (c *StdIOClient) performInitialize(ctx context.Context) error { "version": defaultMCPClientVersion, }, } - if _, err := c.callRequest(ctx, "initialize", params, true); err != nil { - return fmt.Errorf("mcp: initialize session: %w", err) + + protocols := []stdioProtocol{stdioProtocolLine, stdioProtocolFramed} + c.mu.Lock() + if c.protocol == stdioProtocolLine || c.protocol == stdioProtocolFramed { + protocols = []stdioProtocol{c.protocol} } - if err := c.sendNotification(ctx, "notifications/initialized", map[string]any{}, true); err != nil { - return fmt.Errorf("mcp: notify initialized: %w", err) + c.mu.Unlock() + + errs := make([]string, 0, len(protocols)) + for index, protocol := range protocols { + attemptCtx, cancel := initializeAttemptContext(ctx, len(protocols)-index) + _, initErr := c.callRequestWithProtocol(attemptCtx, "initialize", params, true, protocol) + cancel() + if initErr != nil { + errs = append(errs, fmt.Sprintf("%s=%v", protocol, initErr)) + continue + } + + notifyCtx, notifyCancel := initializeAttemptContext(ctx, len(protocols)-index) + notifyErr := c.sendNotificationWithProtocol( + notifyCtx, + "notifications/initialized", + map[string]any{}, + true, + protocol, + ) + notifyCancel() + if notifyErr != nil { + errs = append(errs, fmt.Sprintf("%s=%v", protocol, notifyErr)) + continue + } + + c.mu.Lock() + c.protocol = protocol + c.mu.Unlock() + return nil } - return nil + + return fmt.Errorf("mcp: initialize session: %s", strings.Join(errs, "; ")) } +// readLoop 持续消费 MCP server 响应并分发给对应 pending 请求。 func (c *StdIOClient) readLoop() { for { - message, err := readFramedMessage(c.reader) + message, protocol, err := readRPCMessage(c.reader) if err != nil { c.markExited(fmt.Errorf("mcp: read response: %w", err)) return } + c.mu.Lock() + if c.protocol == stdioProtocolUnknown && (protocol == stdioProtocolLine || protocol == stdioProtocolFramed) { + c.protocol = protocol + } + c.mu.Unlock() + var response jsonRPCResponse if err := json.Unmarshal(message, &response); err != nil { continue @@ -482,6 +559,7 @@ func (c *StdIOClient) readLoop() { } } +// waitLoop 等待子进程退出并触发统一下线处理。 func (c *StdIOClient) waitLoop(command *exec.Cmd) { if command == nil { c.markExited(errors.New("mcp: stdio process is nil")) @@ -491,6 +569,7 @@ func (c *StdIOClient) waitLoop(command *exec.Cmd) { c.markExited(fmt.Errorf("mcp: stdio process exited: %w", err)) } +// markExited 将客户端状态原子切换为已下线并唤醒所有等待请求。 func (c *StdIOClient) markExited(err error) { c.mu.Lock() defer c.mu.Unlock() @@ -505,6 +584,7 @@ func (c *StdIOClient) markExited(err error) { } c.initializing = false c.initDone = nil + c.protocol = stdioProtocolUnknown c.exitErr = err if c.exited != nil { close(c.exited) @@ -517,12 +597,25 @@ func (c *StdIOClient) markExited(err error) { c.bumpBackoffLocked() } +// resolveWriteProtocolLocked 根据调用方覆盖值与会话状态解析实际写入协议。 +func (c *StdIOClient) resolveWriteProtocolLocked(override stdioProtocol) stdioProtocol { + if override == stdioProtocolLine || override == stdioProtocolFramed { + return override + } + if c.protocol == stdioProtocolLine || c.protocol == stdioProtocolFramed { + return c.protocol + } + return stdioProtocolFramed +} + +// removePending 从挂起请求表中移除指定 requestID。 func (c *StdIOClient) removePending(requestID string) { c.mu.Lock() defer c.mu.Unlock() delete(c.pending, requestID) } +// failAllPendingLocked 将所有挂起请求统一返回错误,调用方需持有 c.mu。 func (c *StdIOClient) failAllPendingLocked(err error) { for requestID, replyCh := range c.pending { replyCh <- rpcReply{err: err} @@ -530,6 +623,7 @@ func (c *StdIOClient) failAllPendingLocked(err error) { } } +// bumpBackoffLocked 按指数退避策略更新下次可重启时间,调用方需持有 c.mu。 func (c *StdIOClient) bumpBackoffLocked() { if c.backoff <= 0 { c.backoff = c.cfg.RestartBackoff @@ -541,6 +635,27 @@ func (c *StdIOClient) bumpBackoffLocked() { } } +// initializeAttemptContext 按剩余尝试次数切分超时预算,避免单个协议探测耗尽全部时间。 +func initializeAttemptContext(parent context.Context, remainingAttempts int) (context.Context, context.CancelFunc) { + if remainingAttempts <= 1 { + return context.WithCancel(parent) + } + deadline, ok := parent.Deadline() + if !ok { + return context.WithCancel(parent) + } + remaining := time.Until(deadline) + if remaining <= 0 { + return context.WithCancel(parent) + } + timeout := remaining / time.Duration(remainingAttempts) + if timeout < 200*time.Millisecond { + timeout = 200 * time.Millisecond + } + return context.WithTimeout(parent, timeout) +} + +// writeFramedMessage 以 Content-Length framed 格式写入 JSON-RPC 消息。 func writeFramedMessage(writer io.Writer, payload []byte) error { header := fmt.Sprintf("Content-Length: %d\r\n\r\n", len(payload)) if _, err := io.WriteString(writer, header); err != nil { @@ -552,28 +667,85 @@ func writeFramedMessage(writer io.Writer, payload []byte) error { return nil } -func readFramedMessage(reader *bufio.Reader) ([]byte, error) { - contentLength := -1 +// writeLineMessage 以行分隔 JSON(NDJSON)格式写入 JSON-RPC 消息。 +func writeLineMessage(writer io.Writer, payload []byte) error { + if len(payload) == 0 { + return nil + } + if _, err := writer.Write(payload); err != nil { + return err + } + if !bytes.HasSuffix(payload, []byte("\n")) { + if _, err := io.WriteString(writer, "\n"); err != nil { + return err + } + } + return nil +} + +// writeMessageWithProtocol 按指定线协议写入消息;未知协议默认使用 framed。 +func writeMessageWithProtocol(writer io.Writer, payload []byte, protocol stdioProtocol) error { + switch protocol { + case stdioProtocolLine: + return writeLineMessage(writer, payload) + case stdioProtocolFramed, stdioProtocolUnknown: + return writeFramedMessage(writer, payload) + default: + return writeFramedMessage(writer, payload) + } +} + +// readRPCMessage 自动识别并读取 line/framed 两种 stdio 消息,忽略非协议日志行。 +func readRPCMessage(reader *bufio.Reader) ([]byte, stdioProtocol, error) { for { line, err := reader.ReadString('\n') if err != nil { - return nil, err + return nil, stdioProtocolUnknown, err } + trimmed := strings.TrimSpace(line) if trimmed == "" { - break + continue } lower := strings.ToLower(trimmed) if strings.HasPrefix(lower, "content-length:") { - rawLength := strings.TrimSpace(trimmed[len("content-length:"):]) - length, convErr := strconv.Atoi(rawLength) - if convErr != nil { - return nil, fmt.Errorf("mcp: invalid content-length %q", rawLength) + payload, framedErr := readFramedPayload(reader, trimmed) + if framedErr != nil { + return nil, stdioProtocolUnknown, framedErr + } + return payload, stdioProtocolFramed, nil + } + + if strings.HasPrefix(trimmed, "{") || strings.HasPrefix(trimmed, "[") { + if len(trimmed) > maxStdioLineBytes { + return nil, stdioProtocolUnknown, fmt.Errorf("mcp: line message too large: %d", len(trimmed)) } - contentLength = length + if !json.Valid([]byte(trimmed)) { + continue + } + return []byte(trimmed), stdioProtocolLine, nil + } + } +} + +// readFramedPayload 在已读取首个 Content-Length 头后,继续读取 header/body 并返回 payload。 +func readFramedPayload(reader *bufio.Reader, firstHeader string) ([]byte, error) { + contentLength, err := parseContentLength(firstHeader) + if err != nil { + return nil, err + } + + for { + line, readErr := reader.ReadString('\n') + if readErr != nil { + return nil, readErr + } + if strings.TrimSpace(line) == "" { + break } } + if contentLength < 0 { return nil, errors.New("mcp: missing content-length header") } @@ -588,6 +760,35 @@ func readFramedMessage(reader *bufio.Reader) ([]byte, error) { return payload, nil } +// parseContentLength 解析 Content-Length 头并返回消息体长度。 +func parseContentLength(header string) (int, error) { + trimmed := strings.TrimSpace(header) + lower := strings.ToLower(trimmed) + if !strings.HasPrefix(lower, "content-length:") { + return -1, errors.New("mcp: missing content-length header") + } + rawLength := strings.TrimSpace(trimmed[len("content-length:"):]) + contentLength, err := strconv.Atoi(rawLength) + if err != nil { + return -1, fmt.Errorf("mcp: invalid content-length %q", rawLength) + } + return contentLength, nil +} + +// readFramedMessage 仅读取 framed 消息;会跳过前置日志与 line 消息直到遇到 framed。 +func readFramedMessage(reader *bufio.Reader) ([]byte, error) { + for { + message, protocol, err := readRPCMessage(reader) + if err != nil { + return nil, err + } + if protocol == stdioProtocolFramed { + return message, nil + } + } +} + +// decodeCallResult 将 tools/call 结果统一收敛为 CallResult。 func decodeCallResult(raw json.RawMessage) CallResult { var payload map[string]any if err := json.Unmarshal(raw, &payload); err != nil { diff --git a/internal/tools/mcp/stdio_client_test.go b/internal/tools/mcp/stdio_client_test.go index a54757e6..ccfd50d5 100644 --- a/internal/tools/mcp/stdio_client_test.go +++ b/internal/tools/mcp/stdio_client_test.go @@ -8,6 +8,7 @@ import ( "fmt" "io" "os" + "strconv" "strings" "sync" "testing" @@ -23,7 +24,7 @@ func (errWriter) Write(p []byte) (int, error) { func TestStdIOClientListToolsAndCallTool(t *testing.T) { t.Parallel() - client := newTestStdIOClient(t) + client := newTestStdIOClientWithMode(t, "framed") defer func() { _ = client.Close() }() toolsList, err := client.ListTools(context.Background()) @@ -43,13 +44,50 @@ func TestStdIOClientListToolsAndCallTool(t *testing.T) { } } +func TestStdIOClientLineProtocolInitializeAndCalls(t *testing.T) { + t.Parallel() + + client := newTestStdIOClientWithMode(t, "line") + defer func() { _ = client.Close() }() + + toolsList, err := client.ListTools(context.Background()) + if err != nil { + t.Fatalf("ListTools() with line protocol error = %v", err) + } + if len(toolsList) != 1 || toolsList[0].Name != "search" { + t.Fatalf("unexpected tools list: %+v", toolsList) + } + + result, err := client.CallTool(context.Background(), "search", []byte(`{"query":"mcp"}`)) + if err != nil { + t.Fatalf("CallTool() with line protocol error = %v", err) + } + if !strings.Contains(result.Content, "search") { + t.Fatalf("unexpected call result content: %q", result.Content) + } +} + +func TestStdIOClientInitializeFallbackToFramed(t *testing.T) { + t.Parallel() + + client := newTestStdIOClientWithMode(t, "framed") + defer func() { _ = client.Close() }() + + if _, err := client.ListTools(context.Background()); err != nil { + t.Fatalf("expected fallback initialize success, got %v", err) + } + if client.protocol != stdioProtocolFramed { + t.Fatalf("expected framed protocol selected, got %q", client.protocol) + } +} + func TestStdIOClientHealthCheck(t *testing.T) { t.Parallel() - client := newTestStdIOClient(t) + client := newTestStdIOClientWithMode(t, "framed") defer func() { _ = client.Close() }() - ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) defer cancel() if err := client.HealthCheck(ctx); err != nil { t.Fatalf("HealthCheck() error = %v", err) @@ -59,7 +97,7 @@ func TestStdIOClientHealthCheck(t *testing.T) { func TestStdIOClientConcurrentCallTool(t *testing.T) { t.Parallel() - client := newTestStdIOClient(t) + client := newTestStdIOClientWithMode(t, "framed") defer func() { _ = client.Close() }() const workers = 16 @@ -101,6 +139,40 @@ func TestReadFramedMessageRejectsOversizedPayload(t *testing.T) { } } +func TestReadRPCMessageLine(t *testing.T) { + t.Parallel() + + reader := bufio.NewReader(strings.NewReader("Starting...\n{\"jsonrpc\":\"2.0\",\"id\":\"1\",\"result\":{}}\n")) + payload, protocol, err := readRPCMessage(reader) + if err != nil { + t.Fatalf("readRPCMessage() error = %v", err) + } + if protocol != stdioProtocolLine { + t.Fatalf("expected line protocol, got %q", protocol) + } + if !strings.Contains(string(payload), `"jsonrpc":"2.0"`) { + t.Fatalf("unexpected payload: %s", payload) + } +} + +func TestReadRPCMessageFramed(t *testing.T) { + t.Parallel() + + body := `{"jsonrpc":"2.0","id":"1","result":{"ok":true}}` + raw := "log\nContent-Length: " + strconv.Itoa(len(body)) + "\r\n\r\n" + body + reader := bufio.NewReader(strings.NewReader(raw)) + payload, protocol, err := readRPCMessage(reader) + if err != nil { + t.Fatalf("readRPCMessage() error = %v", err) + } + if protocol != stdioProtocolFramed { + t.Fatalf("expected framed protocol, got %q", protocol) + } + if string(payload) != body { + t.Fatalf("unexpected payload: %s", payload) + } +} + func TestNewStdIOClientValidationAndDefaults(t *testing.T) { t.Parallel() @@ -176,7 +248,7 @@ func TestReadFramedMessageHeaderErrors(t *testing.T) { t.Parallel() reader := bufio.NewReader(strings.NewReader("X-Test: 1\r\n\r\n{}")) - if _, err := readFramedMessage(reader); err == nil || !strings.Contains(err.Error(), "missing content-length") { + if _, err := readFramedMessage(reader); err == nil || !(strings.Contains(err.Error(), "missing content-length") || errors.Is(err, io.EOF)) { t.Fatalf("expected missing content-length error, got %v", err) } @@ -186,6 +258,21 @@ func TestReadFramedMessageHeaderErrors(t *testing.T) { } } +func TestReadFramedMessageIgnoresStdoutPreamble(t *testing.T) { + t.Parallel() + + body := `{"jsonrpc":"2.0","id":"1","result":{"ok":true}}` + raw := "Starting Time MCP server...\nContent-Length: " + strconv.Itoa(len(body)) + "\r\n\r\n" + body + reader := bufio.NewReader(strings.NewReader(raw)) + payload, err := readFramedMessage(reader) + if err != nil { + t.Fatalf("readFramedMessage() error = %v", err) + } + if string(payload) != body { + t.Fatalf("unexpected payload: %s", string(payload)) + } +} + func TestDecodeCallResultVariants(t *testing.T) { t.Parallel() @@ -271,12 +358,20 @@ func TestStdIOClientBumpBackoffClamp(t *testing.T) { } func newTestStdIOClient(t *testing.T) *StdIOClient { + return newTestStdIOClientWithMode(t, "framed") +} + +func newTestStdIOClientWithMode(t *testing.T, wireMode string) *StdIOClient { t.Helper() + if strings.TrimSpace(wireMode) == "" { + wireMode = "framed" + } + client, err := NewStdIOClient(StdioClientConfig{ Command: os.Args[0], Args: []string{"-test.run=TestHelperProcessMCPStdioServer", "--"}, - Env: []string{"GO_WANT_MCP_STDIO_HELPER=1", "GO_MCP_STDIO_REQUIRE_INITIALIZE=1"}, + Env: []string{"GO_WANT_MCP_STDIO_HELPER=1", "GO_MCP_STDIO_REQUIRE_INITIALIZE=1", "GO_MCP_STDIO_WIRE=" + wireMode}, StartTimeout: 3 * time.Second, CallTimeout: 3 * time.Second, }) @@ -292,7 +387,7 @@ func TestStdIOClientInitializeFailure(t *testing.T) { client, err := NewStdIOClient(StdioClientConfig{ Command: os.Args[0], Args: []string{"-test.run=TestHelperProcessMCPStdioServer", "--"}, - Env: []string{"GO_WANT_MCP_STDIO_HELPER=1", "GO_MCP_STDIO_INIT_FAIL=1"}, + Env: []string{"GO_WANT_MCP_STDIO_HELPER=1", "GO_MCP_STDIO_INIT_FAIL=1", "GO_MCP_STDIO_WIRE=framed"}, StartTimeout: 3 * time.Second, CallTimeout: 3 * time.Second, }) @@ -314,11 +409,24 @@ func TestHelperProcessMCPStdioServer(t *testing.T) { requireInitialize := os.Getenv("GO_MCP_STDIO_REQUIRE_INITIALIZE") == "1" initFail := os.Getenv("GO_MCP_STDIO_INIT_FAIL") == "1" + wireMode := strings.TrimSpace(os.Getenv("GO_MCP_STDIO_WIRE")) + if wireMode == "" { + wireMode = "framed" + } initialized := !requireInitialize reader := bufio.NewReader(os.Stdin) for { - payload, err := readFramedMessage(reader) + var ( + payload []byte + err error + ) + switch wireMode { + case "line": + payload, _, err = readRPCMessage(reader) + default: + payload, err = readFramedMessage(reader) + } if err != nil { if err == io.EOF { os.Exit(0) @@ -428,7 +536,13 @@ func TestHelperProcessMCPStdioServer(t *testing.T) { if err != nil { os.Exit(4) } - if err := writeFramedMessage(os.Stdout, rawResponse); err != nil { + switch wireMode { + case "line": + err = writeLineMessage(os.Stdout, rawResponse) + default: + err = writeFramedMessage(os.Stdout, rawResponse) + } + if err != nil { os.Exit(5) } } From 5e099a583918894d38f07c81b72e0445178c1c50 Mon Sep 17 00:00:00 2001 From: Cai_Tang <106404101+Cai-Tang-www@users.noreply.github.com> Date: Tue, 7 Apr 2026 19:25:35 +0800 Subject: [PATCH 41/55] =?UTF-8?q?test:=20=E8=A1=A5=E9=BD=90MCP=E6=8E=A5?= =?UTF-8?q?=E5=85=A5=E8=B7=AF=E5=BE=84=E8=A6=86=E7=9B=96=E7=94=A8=E4=BE=8B?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/app/bootstrap_test.go | 270 +++++++++++++++++++++++- internal/tools/mcp/stdio_client_test.go | 66 ++++++ 2 files changed, 335 insertions(+), 1 deletion(-) diff --git a/internal/app/bootstrap_test.go b/internal/app/bootstrap_test.go index 67e7ba2a..d706a1b2 100644 --- a/internal/app/bootstrap_test.go +++ b/internal/app/bootstrap_test.go @@ -1,13 +1,18 @@ package app import ( + "bufio" + "bytes" "context" "encoding/json" "errors" + "fmt" + "io" "net/http" "net/http/httptest" "os" "path/filepath" + "strconv" "strings" "testing" "time" @@ -159,6 +164,115 @@ func TestBuildMCPRegistryFromConfig(t *testing.T) { } } +func TestBuildMCPRegistryUnsupportedSource(t *testing.T) { + t.Parallel() + + cfg := config.Default().Clone() + cfg.Workdir = t.TempDir() + cfg.Tools.MCP.Servers = []config.MCPServerConfig{ + { + ID: "docs", + Enabled: true, + Source: "sse", + Stdio: config.MCPStdioConfig{ + Command: "mock", + }, + }, + } + + registry, err := buildMCPRegistry(cfg) + if err == nil { + t.Fatalf("expected unsupported source error") + } + if registry != nil { + t.Fatalf("expected nil registry when source unsupported") + } + if !strings.Contains(strings.ToLower(err.Error()), "unsupported mcp source") { + t.Fatalf("unexpected error: %v", err) + } +} + +func TestDefaultRegisterMCPStdioServerSuccess(t *testing.T) { + t.Parallel() + + registry := mcp.NewRegistry() + cfg := config.Default().Clone() + wd, err := os.Getwd() + if err != nil { + t.Fatalf("getwd: %v", err) + } + cfg.Workdir = wd + cfg.ToolTimeoutSec = 9 + + server := config.MCPServerConfig{ + ID: "docs", + Enabled: true, + Source: "stdio", + Version: "v1", + Stdio: config.MCPStdioConfig{ + Command: os.Args[0], + Args: []string{"-test.run=TestHelperProcessAppMCPStdioServer", "--"}, + Workdir: "", + StartTimeoutSec: 3, + CallTimeoutSec: 3, + }, + Env: []config.MCPEnvVarConfig{ + {Name: "MODE", Value: "test"}, + {Name: "GO_WANT_APP_MCP_STDIO_HELPER", Value: "1"}, + }, + } + t.Cleanup(func() { _ = registry.UnregisterServer("docs") }) + + if err := defaultRegisterMCPStdioServer(registry, cfg, server); err != nil { + t.Fatalf("defaultRegisterMCPStdioServer() error = %v", err) + } + + snapshots := registry.Snapshot() + if len(snapshots) != 1 || snapshots[0].ServerID != "docs" { + t.Fatalf("unexpected snapshots: %+v", snapshots) + } + if len(snapshots[0].Tools) != 1 || snapshots[0].Tools[0].Name != "search" { + t.Fatalf("unexpected tools snapshot: %+v", snapshots[0].Tools) + } +} + +func TestDefaultRegisterMCPStdioServerRefreshFailure(t *testing.T) { + t.Parallel() + + registry := mcp.NewRegistry() + cfg := config.Default().Clone() + wd, err := os.Getwd() + if err != nil { + t.Fatalf("getwd: %v", err) + } + cfg.Workdir = wd + + server := config.MCPServerConfig{ + ID: "broken", + Enabled: true, + Source: "stdio", + Stdio: config.MCPStdioConfig{ + Command: os.Args[0], + Args: []string{"-test.run=TestHelperProcessAppMCPStdioServer", "--"}, + StartTimeoutSec: 3, + CallTimeoutSec: 3, + }, + Env: []config.MCPEnvVarConfig{ + {Name: "GO_WANT_APP_MCP_STDIO_HELPER", Value: "1"}, + {Name: "GO_APP_MCP_STDIO_LIST_FAIL", Value: "1"}, + }, + } + t.Cleanup(func() { _ = registry.UnregisterServer("broken") }) + + err = defaultRegisterMCPStdioServer(registry, cfg, server) + if err == nil { + t.Fatalf("expected refresh failure") + } + if !strings.Contains(strings.ToLower(err.Error()), "list tools failed") { + t.Fatalf("unexpected error: %v", err) + } +} + func TestBuildToolRegistryIncludesMCPFromConfig(t *testing.T) { t.Parallel() @@ -398,10 +512,14 @@ func disableBuiltinProviderAPIKeys(t *testing.T) { } type stubMCPServerClient struct { - tools []mcp.ToolDescriptor + tools []mcp.ToolDescriptor + listErr error } func (s *stubMCPServerClient) ListTools(ctx context.Context) ([]mcp.ToolDescriptor, error) { + if s.listErr != nil { + return nil, s.listErr + } return append([]mcp.ToolDescriptor(nil), s.tools...), nil } @@ -412,3 +530,153 @@ func (s *stubMCPServerClient) CallTool(ctx context.Context, toolName string, arg func (s *stubMCPServerClient) HealthCheck(ctx context.Context) error { return nil } + +func TestHelperProcessAppMCPStdioServer(t *testing.T) { + if os.Getenv("GO_WANT_APP_MCP_STDIO_HELPER") != "1" { + return + } + + listFail := os.Getenv("GO_APP_MCP_STDIO_LIST_FAIL") == "1" + initialized := false + reader := bufio.NewReader(os.Stdin) + + for { + payload, err := readFramedForAppTest(reader) + if err != nil { + if errors.Is(err, os.ErrClosed) || strings.Contains(strings.ToLower(err.Error()), "eof") { + os.Exit(0) + } + os.Exit(2) + } + + var request map[string]any + if err := json.Unmarshal(payload, &request); err != nil { + os.Exit(3) + } + + method, _ := request["method"].(string) + requestID, _ := request["id"].(string) + var response any + + switch method { + case "initialize": + response = map[string]any{ + "jsonrpc": "2.0", + "id": requestID, + "result": map[string]any{ + "protocolVersion": "2024-11-05", + "capabilities": map[string]any{}, + "serverInfo": map[string]any{ + "name": "app-helper", + "version": "1.0.0", + }, + }, + } + case "notifications/initialized": + initialized = true + continue + case "tools/list": + if listFail { + response = map[string]any{ + "jsonrpc": "2.0", + "id": requestID, + "error": map[string]any{ + "code": -32001, + "message": "list tools failed", + }, + } + break + } + if !initialized { + response = map[string]any{ + "jsonrpc": "2.0", + "id": requestID, + "error": map[string]any{ + "code": -32002, + "message": "server not initialized", + }, + } + break + } + response = map[string]any{ + "jsonrpc": "2.0", + "id": requestID, + "result": map[string]any{ + "tools": []map[string]any{ + { + "name": "search", + "description": "search docs", + "inputSchema": map[string]any{ + "type": "object", + "properties": map[string]any{"query": map[string]any{"type": "string"}}, + }, + }, + }, + }, + } + default: + response = map[string]any{ + "jsonrpc": "2.0", + "id": requestID, + "error": map[string]any{ + "code": -32601, + "message": "method not found", + }, + } + } + + rawResponse, err := json.Marshal(response) + if err != nil { + os.Exit(4) + } + if err := writeFramedForAppTest(os.Stdout, rawResponse); err != nil { + os.Exit(5) + } + } +} + +func readFramedForAppTest(reader *bufio.Reader) ([]byte, error) { + contentLength := -1 + for { + line, err := reader.ReadString('\n') + if err != nil { + return nil, err + } + trimmed := strings.TrimSpace(line) + if trimmed == "" { + if contentLength >= 0 { + break + } + continue + } + lower := strings.ToLower(trimmed) + if strings.HasPrefix(lower, "content-length:") { + rawLength := strings.TrimSpace(trimmed[len("content-length:"):]) + length, convErr := strconv.Atoi(rawLength) + if convErr != nil { + return nil, convErr + } + contentLength = length + continue + } + } + if contentLength < 0 { + return nil, errors.New("missing content-length") + } + payload := make([]byte, contentLength) + if _, err := io.ReadFull(reader, payload); err != nil { + return nil, err + } + return payload, nil +} + +func writeFramedForAppTest(writer io.Writer, payload []byte) error { + header := fmt.Sprintf("Content-Length: %d\r\n\r\n", len(payload)) + if _, err := io.WriteString(writer, header); err != nil { + return err + } + if _, err := writer.Write(bytes.TrimSpace(payload)); err != nil { + return err + } + return nil +} diff --git a/internal/tools/mcp/stdio_client_test.go b/internal/tools/mcp/stdio_client_test.go index ccfd50d5..01d8d20b 100644 --- a/internal/tools/mcp/stdio_client_test.go +++ b/internal/tools/mcp/stdio_client_test.go @@ -2,6 +2,7 @@ package mcp import ( "bufio" + "bytes" "context" "encoding/json" "errors" @@ -15,6 +16,12 @@ import ( "time" ) +type nopWriteCloser struct { + bytes.Buffer +} + +func (n *nopWriteCloser) Close() error { return nil } + type errWriter struct{} func (errWriter) Write(p []byte) (int, error) { @@ -94,6 +101,65 @@ func TestStdIOClientHealthCheck(t *testing.T) { } } +func TestStdIOClientSendNotificationFramedDefault(t *testing.T) { + t.Parallel() + + writer := &nopWriteCloser{} + client := &StdIOClient{ + pending: make(map[string]chan rpcReply), + stdin: writer, + started: true, + cfg: StdioClientConfig{ + CallTimeout: time.Second, + StartTimeout: time.Second, + RestartBackoff: time.Millisecond, + }, + } + + err := client.sendNotification(context.Background(), "notifications/initialized", map[string]any{}, true) + if err != nil { + t.Fatalf("sendNotification() error = %v", err) + } + if !strings.Contains(writer.String(), "Content-Length:") { + t.Fatalf("expected framed header, got: %q", writer.String()) + } +} + +func TestStdIOClientSendNotificationLineProtocol(t *testing.T) { + t.Parallel() + + writer := &nopWriteCloser{} + client := &StdIOClient{ + pending: make(map[string]chan rpcReply), + stdin: writer, + started: true, + protocol: stdioProtocolLine, + cfg: StdioClientConfig{ + CallTimeout: time.Second, + StartTimeout: time.Second, + RestartBackoff: time.Millisecond, + }, + } + + err := client.sendNotificationWithProtocol( + context.Background(), + "notifications/initialized", + map[string]any{}, + true, + stdioProtocolLine, + ) + if err != nil { + t.Fatalf("sendNotificationWithProtocol() error = %v", err) + } + raw := writer.String() + if strings.Contains(raw, "Content-Length:") { + t.Fatalf("expected line protocol write, got: %q", raw) + } + if !strings.HasSuffix(raw, "\n") { + t.Fatalf("expected newline-delimited payload, got: %q", raw) + } +} + func TestStdIOClientConcurrentCallTool(t *testing.T) { t.Parallel() From e726eb338fec1f911167798f5768de0ead3b369d Mon Sep 17 00:00:00 2001 From: Cai_Tang <106404101+Cai-Tang-www@users.noreply.github.com> Date: Tue, 7 Apr 2026 19:48:22 +0800 Subject: [PATCH 42/55] =?UTF-8?q?fix:=20=E8=A7=84=E8=8C=83registry?= =?UTF-8?q?=E4=B8=AD=E7=9A=84provider=E5=BC=95=E7=94=A8=E4=BB=A5=E4=BF=AE?= =?UTF-8?q?=E5=A4=8DCI=E6=9E=84=E5=BB=BA=E5=BC=82=E5=B8=B8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/tools/registry.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/internal/tools/registry.go b/internal/tools/registry.go index c03bfcb6..a392f629 100644 --- a/internal/tools/registry.go +++ b/internal/tools/registry.go @@ -114,7 +114,7 @@ func (r *Registry) ListAvailableSpecs(ctx context.Context, input SpecListInput) return nil, err } for _, adapter := range mcpAdapters { - specs = append(specs, provider.ToolSpec{ + specs = append(specs, providertypes.ToolSpec{ Name: adapter.FullName(), Description: adapter.Description(), Schema: adapter.Schema(), From 3cf3ea21f0b782b608914fd820302c223e61cf5c Mon Sep 17 00:00:00 2001 From: Cai_Tang <106404101+Cai-Tang-www@users.noreply.github.com> Date: Tue, 7 Apr 2026 19:51:42 +0800 Subject: [PATCH 43/55] =?UTF-8?q?chore:=20=E8=A7=A6=E5=8F=91CI=E9=87=8D?= =?UTF-8?q?=E8=B7=91?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit From 24a6c0822763fb416d41ce2ffadaca04e05203cb Mon Sep 17 00:00:00 2001 From: Yumiue <229866007@qq.com> Date: Tue, 7 Apr 2026 20:22:16 +0800 Subject: [PATCH 44/55] =?UTF-8?q?test:=E6=B7=BB=E5=8A=A0=E6=B5=8B=E8=AF=95?= =?UTF-8?q?=E6=96=87=E4=BB=B6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/runtime/runtime_test.go | 205 +++++++++++++++++++++++++++++++ internal/session/id_test.go | 12 ++ internal/session/store_test.go | 153 +++++++++++++++++++++++ 3 files changed, 370 insertions(+) diff --git a/internal/runtime/runtime_test.go b/internal/runtime/runtime_test.go index d010a908..af02608c 100644 --- a/internal/runtime/runtime_test.go +++ b/internal/runtime/runtime_test.go @@ -2822,3 +2822,208 @@ func TestPermissionEventViewPayloadMapping(t *testing.T) { t.Fatalf("unexpected resolved payload: %+v", resolvedPayload) } } + +func TestStreamAccumulatorBuildMessageRejectsMissingToolName(t *testing.T) { + t.Parallel() + + acc := newStreamAccumulator() + acc.accumulateToolCallStart(0, "call-1", "") + acc.accumulateToolCallDelta(0, "call-1", "{}") + + _, err := acc.buildMessage() + if err == nil || !containsError(err, "without name") { + t.Fatalf("expected missing tool name error, got %v", err) + } +} + +func TestLoadSessionReturnsStoreError(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + store := newMemoryStore() + service := NewWithFactory(manager, nil, store, &scriptedProviderFactory{provider: &scriptedProvider{}}, nil) + + _, err := service.LoadSession(context.Background(), "missing") + if err == nil || !containsError(err, "not found") { + t.Fatalf("expected load error, got %v", err) + } +} + +func TestSetSessionWorkdirReturnsStoreError(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + store := newMemoryStore() + service := NewWithFactory(manager, nil, store, &scriptedProviderFactory{provider: &scriptedProvider{}}, nil) + + _, err := service.SetSessionWorkdir(context.Background(), "missing", t.TempDir()) + if err == nil || !containsError(err, "not found") { + t.Fatalf("expected load error from SetSessionWorkdir, got %v", err) + } +} + +func TestSetSessionWorkdirReturnsResolveError(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + defaultWorkdir := t.TempDir() + if err := manager.Update(context.Background(), func(cfg *config.Config) error { + cfg.Workdir = defaultWorkdir + return nil + }); err != nil { + t.Fatalf("update config: %v", err) + } + + store := newMemoryStore() + session := agentsession.New("set bad workdir") + session.ID = "session-set-bad-workdir" + store.sessions[session.ID] = cloneSession(session) + + service := NewWithFactory(manager, nil, store, &scriptedProviderFactory{provider: &scriptedProvider{}}, nil) + _, err := service.SetSessionWorkdir(context.Background(), session.ID, filepath.Join(defaultWorkdir, "missing-dir")) + if err == nil || !containsError(err, "resolve workdir") { + t.Fatalf("expected resolve workdir error, got %v", err) + } +} + +func TestServiceRunFailsWhenInitialUserMessageSaveFails(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + baseStore := newMemoryStore() + store := &failingStore{ + Store: baseStore, + saveErr: errors.New("save failed on first write"), + failOnSave: 1, + } + registry := tools.NewRegistry() + registry.Register(&stubTool{name: "filesystem_read_file", content: "default"}) + + service := NewWithFactory(manager, registry, store, &scriptedProviderFactory{provider: &scriptedProvider{}}, nil) + err := service.Run(context.Background(), UserInput{ + RunID: "run-initial-save-fail", + Content: "hello", + }) + if err == nil || !containsError(err, "save failed on first write") { + t.Fatalf("expected initial save error, got %v", err) + } +} + +func TestServiceRunFailsWhenAssistantSaveFails(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + baseStore := newMemoryStore() + store := &failingStore{ + Store: baseStore, + saveErr: errors.New("save failed on assistant"), + failOnSave: 2, + } + registry := tools.NewRegistry() + registry.Register(&stubTool{name: "filesystem_read_file", content: "default"}) + + scripted := &scriptedProvider{ + streams: [][]provider.StreamEvent{ + {provider.NewTextDeltaStreamEvent("assistant reply")}, + }, + } + service := NewWithFactory(manager, registry, store, &scriptedProviderFactory{provider: scripted}, nil) + err := service.Run(context.Background(), UserInput{ + RunID: "run-assistant-save-fail", + Content: "hello", + }) + if err == nil || !containsError(err, "save failed on assistant") { + t.Fatalf("expected assistant save error, got %v", err) + } +} + +func TestHandleProviderStreamEventErrorBranches(t *testing.T) { + t.Parallel() + + acc := newStreamAccumulator() + + err := handleProviderStreamEvent( + provider.StreamEvent{Type: provider.StreamEventToolCallStart}, + acc, + nil, + nil, + ) + if err == nil || !containsError(err, "tool_call_start event payload is nil") { + t.Fatalf("expected tool_call_start payload error, got %v", err) + } + + err = handleProviderStreamEvent( + provider.StreamEvent{Type: provider.StreamEventToolCallDelta}, + acc, + nil, + nil, + ) + if err == nil || !containsError(err, "tool_call_delta event payload is nil") { + t.Fatalf("expected tool_call_delta payload error, got %v", err) + } + + err = handleProviderStreamEvent( + provider.StreamEvent{Type: provider.StreamEventMessageDone}, + acc, + nil, + nil, + ) + if err == nil || !containsError(err, "message_done event payload is nil") { + t.Fatalf("expected message_done payload error, got %v", err) + } +} + +func TestEmitDropsWhenChannelFullAndContextCanceled(t *testing.T) { + t.Parallel() + + service := &Service{ + events: make(chan RuntimeEvent, 1), + } + service.events <- RuntimeEvent{Type: EventAgentChunk} + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + done := make(chan struct{}) + go func() { + service.emit(ctx, EventError, "run-id", "session-id", "payload") + close(done) + }() + + select { + case <-done: + case <-time.After(1 * time.Second): + t.Fatal("emit should return when channel is full and context is canceled") + } +} + +func TestCallProviderWithRetryReturnsCombinedForwardError(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + registry := tools.NewRegistry() + registry.Register(&stubTool{name: "filesystem_read_file", content: "default"}) + store := newMemoryStore() + + scripted := &scriptedProvider{ + chatFn: func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { + events <- provider.StreamEvent{Type: provider.StreamEventTextDelta} + return errors.New("provider chat failed") + }, + } + service := NewWithFactory(manager, registry, store, &scriptedProviderFactory{provider: scripted}, nil) + + _, err := service.callProviderWithRetry( + context.Background(), + "run-forward-error", + "session-forward-error", + provider.ChatRequest{ + Model: "test-model", + SystemPrompt: "prompt", + Messages: []provider.Message{{Role: provider.RoleUser, Content: "hello"}}, + }, + ) + if err == nil || !containsError(err, "provider stream handling failed after provider error") { + t.Fatalf("expected combined forward/provider error, got %v", err) + } +} diff --git a/internal/session/id_test.go b/internal/session/id_test.go index 8ef1bb51..396779f8 100644 --- a/internal/session/id_test.go +++ b/internal/session/id_test.go @@ -39,3 +39,15 @@ func TestNewIDAllowsEmptyPrefix(t *testing.T) { t.Fatalf("expected format _<16hex>, got %q", id) } } + +func TestNewIDCompatibilityWrapper(t *testing.T) { + t.Parallel() + + id := newID("session") + if !strings.HasPrefix(id, "session_") { + t.Fatalf("expected compatibility wrapper to preserve prefix, got %q", id) + } + if len(strings.TrimPrefix(id, "session_")) != 16 { + t.Fatalf("expected compatibility wrapper to return 16 hex chars, got %q", id) + } +} diff --git a/internal/session/store_test.go b/internal/session/store_test.go index 75161164..a886b73e 100644 --- a/internal/session/store_test.go +++ b/internal/session/store_test.go @@ -2,6 +2,8 @@ package session import ( "context" + "encoding/json" + "errors" "os" "path/filepath" "strings" @@ -216,6 +218,157 @@ func TestNewWithWorkdirFallsBackDefaultTitle(t *testing.T) { } } +func TestNewStoreReturnsJSONStore(t *testing.T) { + t.Parallel() + + store := NewStore(t.TempDir()) + if store == nil { + t.Fatalf("expected non-nil store") + } +} + +func TestJSONStoreListSummariesReadDirFailure(t *testing.T) { + t.Parallel() + + baseDir := t.TempDir() + store := NewJSONStore(baseDir) + + // 把 sessions 目录位置占成普通文件,触发 ReadDir 失败路径。 + sessionsPath := filepath.Join(baseDir, sessionsDirName) + if err := os.WriteFile(sessionsPath, []byte("not-a-dir"), 0o644); err != nil { + t.Fatalf("write %s: %v", sessionsPath, err) + } + + _, err := store.ListSummaries(context.Background()) + if err == nil || !strings.Contains(err.Error(), "create sessions dir") { + t.Fatalf("expected create sessions dir error, got %v", err) + } +} + +func TestJSONStoreListSummariesContextCanceledDuringIteration(t *testing.T) { + t.Parallel() + + baseDir := t.TempDir() + store := NewJSONStore(baseDir) + + for i := 0; i < 10; i++ { + s := &Session{ + ID: "session-iter-" + strings.Repeat("x", i+1), + Title: "iter", + CreatedAt: time.Now(), + UpdatedAt: time.Now(), + } + if err := store.Save(context.Background(), s); err != nil { + t.Fatalf("save session %d: %v", i, err) + } + } + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + _, err := store.ListSummaries(ctx) + if !errors.Is(err, context.Canceled) { + t.Fatalf("expected context canceled, got %v", err) + } +} + +func TestJSONStoreLoadDecodeErrorWithNonJSONPayload(t *testing.T) { + t.Parallel() + + baseDir := t.TempDir() + store := NewJSONStore(baseDir) + + mustWriteSessionFile(t, filepath.Join(baseDir, sessionsDirName, "decode-bad.json"), "{not-json") + + _, err := store.Load(context.Background(), "decode-bad") + if err == nil || !strings.Contains(err.Error(), "decode session decode-bad") { + t.Fatalf("expected decode session error, got %v", err) + } +} + +func TestJSONStoreListSummariesSkipsUnreadableAndMalformedEntries(t *testing.T) { + t.Parallel() + + baseDir := t.TempDir() + store := NewJSONStore(baseDir) + + valid := &Session{ + ID: "valid-summary", + Title: "Valid", + CreatedAt: time.Now().Add(-time.Minute), + UpdatedAt: time.Now(), + } + if err := store.Save(context.Background(), valid); err != nil { + t.Fatalf("save valid session: %v", err) + } + + mustWriteSessionFile(t, filepath.Join(baseDir, sessionsDirName, "malformed.json"), "{malformed") + mustWriteSessionFile(t, filepath.Join(baseDir, sessionsDirName, "empty-id.json"), `{"id":" ","title":"x"}`) + + summaries, err := store.ListSummaries(context.Background()) + if err != nil { + t.Fatalf("ListSummaries() error: %v", err) + } + if len(summaries) != 1 || summaries[0].ID != valid.ID { + t.Fatalf("expected only valid summary, got %+v", summaries) + } +} + +func TestJSONStoreSavePersistsProviderModelAndMessages(t *testing.T) { + t.Parallel() + + baseDir := t.TempDir() + store := NewJSONStore(baseDir) + + session := &Session{ + ID: "persist-full-fields", + Title: "Persist Fields", + Provider: "openai", + Model: "gpt-4.1", + CreatedAt: time.Now().Add(-time.Hour), + UpdatedAt: time.Now(), + Messages: []provider.Message{ + {Role: provider.RoleUser, Content: "hello"}, + { + Role: provider.RoleAssistant, + Content: "calling tool", + ToolCalls: []provider.ToolCall{ + {ID: "call-1", Name: "webfetch", Arguments: `{"url":"https://example.com"}`}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-1", Content: "ok"}, + }, + } + + if err := store.Save(context.Background(), session); err != nil { + t.Fatalf("save session: %v", err) + } + + rawPath := filepath.Join(baseDir, sessionsDirName, session.ID+".json") + raw, err := os.ReadFile(rawPath) + if err != nil { + t.Fatalf("read raw file: %v", err) + } + + var decoded map[string]any + if err := json.Unmarshal(raw, &decoded); err != nil { + t.Fatalf("decode raw json: %v", err) + } + + if decoded["provider"] != "openai" { + t.Fatalf("expected provider persisted, got %+v", decoded["provider"]) + } + if decoded["model"] != "gpt-4.1" { + t.Fatalf("expected model persisted, got %+v", decoded["model"]) + } + if _, ok := decoded["messages"]; !ok { + t.Fatalf("expected messages field persisted, got %+v", decoded) + } + if _, ok := decoded["workdir"]; ok { + t.Fatalf("expected workdir not persisted, got %+v", decoded) + } +} + func mustWriteSessionFile(t *testing.T, path string, content string) { t.Helper() if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil { From f1793af08d6515092ac1db413010556809abc098 Mon Sep 17 00:00:00 2001 From: Yumiue <229866007@qq.com> Date: Tue, 7 Apr 2026 20:32:16 +0800 Subject: [PATCH 45/55] =?UTF-8?q?refator:=E6=8B=86=E5=88=86=E5=87=BAsessio?= =?UTF-8?q?n=E6=A8=A1=E5=9D=97?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- README.md | 4 +- docs/session-persistence-design.md | 30 ++++++++ internal/app/bootstrap.go | 3 +- internal/runtime/compact.go | 9 +-- internal/runtime/id.go | 12 ---- internal/runtime/runtime.go | 37 +++++----- internal/runtime/runtime_test.go | 71 ++++++++++--------- internal/runtime/workdir_branch_test.go | 4 +- internal/session/id.go | 18 +++++ .../{runtime/session.go => session/store.go} | 68 +++++++++++------- .../session_test.go => session/store_test.go} | 24 +++---- internal/tui/state.go | 6 +- internal/tui/state/ui_state.go | 4 +- internal/tui/update_test.go | 41 +++++------ 14 files changed, 195 insertions(+), 136 deletions(-) create mode 100644 docs/session-persistence-design.md delete mode 100644 internal/runtime/id.go create mode 100644 internal/session/id.go rename internal/{runtime/session.go => session/store.go} (59%) rename internal/{runtime/session_test.go => session/store_test.go} (88%) diff --git a/README.md b/README.md index d710be7c..33fdaa73 100644 --- a/README.md +++ b/README.md @@ -94,7 +94,8 @@ go run ./cmd/neocode - **`internal/provider/catalog`** — 模型发现、catalog 缓存与后台刷新 - **`internal/provider/selection`** — provider/model 选择与配置同步 - **`internal/provider/builtin`** — 内建 driver 注册 -- **`internal/runtime`** — ReAct 主循环、事件流、会话管理 +- **`internal/runtime`** — ReAct 主循环与事件流编排(不直接承载会话存储实现;不再导出会话模型与存储类型) +- **`internal/session`** — 会话模型、会话存储抽象与 JSON 持久化实现(统一对外暴露 `Session` / `Summary` / `Store`) - **`internal/tools`** — 工具注册表与具体工具实现 - **`internal/tui`** — 终端 UI、交互体验、事件桥接 - **`internal/app`** — 应用装配与依赖注入 @@ -116,6 +117,7 @@ go run ./cmd/neocode │ │ ├── catalog # 模型发现与缓存 │ │ └── selection # provider/model 选择服务 │ ├── runtime # ReAct 循环与事件流 +│ ├── session # 会话模型与持久化 │ ├── tools # 工具系统 │ └── tui # 终端 UI └── README.md diff --git a/docs/session-persistence-design.md b/docs/session-persistence-design.md new file mode 100644 index 00000000..e66f7800 --- /dev/null +++ b/docs/session-persistence-design.md @@ -0,0 +1,30 @@ +# Session 持久化设计 + +## 模块职责与收口边界 +- `internal/session`:承载会话领域模型、存储抽象与 JSON 持久化实现,是唯一的会话持久化实现归属层 +- `internal/runtime`:只依赖 `internal/session` 提供的抽象与模型,负责会话保存时机与主循环编排,不再维护会话存储实现细节 +- `internal/tui`:仅消费 runtime 暴露的会话数据,不直接执行会话持久化 + +## 存储策略 +NeoCode 在 MVP 阶段使用 JSON 文件持久化 Session,以保持本地优先、易于调试和跨平台可移植。 + +## 数据模型 +- `Session`:完整消息历史以及 `id`、`title`、`updated_at` 等元信息 +- `Summary`:用于侧边栏的轻量摘要结构(原 `SessionSummary` 命名已统一收口为 `Summary`) + +## 加载策略 +- `ListSummaries` 只读取渲染侧边栏所需的基础信息 +- `Load` 仅在用户真正进入某个会话时读取完整消息历史 +- `Save` 通过临时文件原子写入完整 Session + +## 命名策略 +- 新会话默认展示为 `New Session` +- 一旦持久化,runtime 会根据首轮用户消息生成简短标题 + +## 并发约束 +- `internal/session` 中的 Store 实现必须自行保护共享访问 +- 真正的保存时机由 runtime 决定,TUI 不负责直接触发磁盘写入 + +## 兼容性与演进说明 +- 会话持久化能力已从 runtime 侧实现中彻底收口到 `internal/session` +- 新增会话存储实现时,应优先在 `internal/session` 内扩展并通过接口注入 runtime,避免跨层实现 diff --git a/internal/app/bootstrap.go b/internal/app/bootstrap.go index dd0b8436..6d2892b3 100644 --- a/internal/app/bootstrap.go +++ b/internal/app/bootstrap.go @@ -12,6 +12,7 @@ import ( providercatalog "neo-code/internal/provider/catalog" agentruntime "neo-code/internal/runtime" "neo-code/internal/security" + agentsession "neo-code/internal/session" "neo-code/internal/tools" "neo-code/internal/tools/bash" "neo-code/internal/tools/filesystem" @@ -62,7 +63,7 @@ func NewProgram(ctx context.Context) (*tea.Program, error) { return nil, err } - sessionStore := agentruntime.NewSessionStore(loader.BaseDir()) + sessionStore := agentsession.NewStore(loader.BaseDir()) runtimeSvc := agentruntime.NewWithFactory( manager, toolManager, diff --git a/internal/runtime/compact.go b/internal/runtime/compact.go index 9798d250..26487f33 100644 --- a/internal/runtime/compact.go +++ b/internal/runtime/compact.go @@ -9,6 +9,7 @@ import ( "neo-code/internal/config" contextcompact "neo-code/internal/context/compact" providertypes "neo-code/internal/provider/types" + agentsession "neo-code/internal/session" ) // CompactInput 描述一次手动 compact 请求所需的最小输入。 @@ -83,10 +84,10 @@ func (s *Service) Compact(ctx context.Context, input CompactInput) (CompactResul func (s *Service) runCompactForSession( ctx context.Context, runID string, - session Session, + session agentsession.Session, cfg config.Config, failOnError bool, -) (Session, contextcompact.Result, error) { +) (agentsession.Session, contextcompact.Result, error) { runner := s.compactRunner if runner == nil { var err error @@ -155,7 +156,7 @@ func (s *Service) runCompactForSession( } // defaultCompactRunner 为手动 compact 选择摘要生成器并构造默认 runner。 -func (s *Service) defaultCompactRunner(session Session, cfg config.Config) (contextcompact.Runner, error) { +func (s *Service) defaultCompactRunner(session agentsession.Session, cfg config.Config) (contextcompact.Runner, error) { resolvedProvider, model, err := resolveCompactProviderSelection(session, cfg) if err != nil { return nil, err @@ -164,7 +165,7 @@ func (s *Service) defaultCompactRunner(session Session, cfg config.Config) (cont } // resolveCompactProviderSelection 优先复用会话记录的 provider/model,缺失时再回退当前配置。 -func resolveCompactProviderSelection(session Session, cfg config.Config) (config.ResolvedProviderConfig, string, error) { +func resolveCompactProviderSelection(session agentsession.Session, cfg config.Config) (config.ResolvedProviderConfig, string, error) { sessionProvider := strings.TrimSpace(session.Provider) sessionModel := strings.TrimSpace(session.Model) if sessionProvider != "" && sessionModel != "" { diff --git a/internal/runtime/id.go b/internal/runtime/id.go deleted file mode 100644 index b036fbea..00000000 --- a/internal/runtime/id.go +++ /dev/null @@ -1,12 +0,0 @@ -package runtime - -import ( - "crypto/rand" - "encoding/hex" -) - -func newID(prefix string) string { - buf := make([]byte, 8) - _, _ = rand.Read(buf) - return prefix + "_" + hex.EncodeToString(buf) -} diff --git a/internal/runtime/runtime.go b/internal/runtime/runtime.go index 7b727e08..a6352f22 100644 --- a/internal/runtime/runtime.go +++ b/internal/runtime/runtime.go @@ -18,6 +18,7 @@ import ( contextcompact "neo-code/internal/context/compact" "neo-code/internal/provider" providertypes "neo-code/internal/provider/types" + agentsession "neo-code/internal/session" "neo-code/internal/tools" ) @@ -121,9 +122,9 @@ type Runtime interface { ResolvePermission(ctx context.Context, input PermissionResolutionInput) error CancelActiveRun() bool Events() <-chan RuntimeEvent - ListSessions(ctx context.Context) ([]SessionSummary, error) - LoadSession(ctx context.Context, id string) (Session, error) - SetSessionWorkdir(ctx context.Context, sessionID string, workdir string) (Session, error) + ListSessions(ctx context.Context) ([]agentsession.Summary, error) + LoadSession(ctx context.Context, id string) (agentsession.Session, error) + SetSessionWorkdir(ctx context.Context, sessionID string, workdir string) (agentsession.Session, error) } type UserInput struct { @@ -139,7 +140,7 @@ type ProviderFactory interface { type Service struct { configManager *config.Manager // 配置管理器,提供当前选中的 provider、model、workdir 等配置读取能力。 - sessionStore Store // 会话持久化接口,负责保存和加载聊天会话。 + sessionStore agentsession.Store // 会话持久化接口,负责保存和加载聊天会话。 toolManager tools.Manager // 工具管理器,统一工具 schema 暴露与执行入口。 providerFactory ProviderFactory // Provider 工厂接口,根据配置动态创建具体的 provider 实例。 contextBuilder agentcontext.Builder // 上下文构建器,负责组装 system prompt 与本轮发给模型的消息上下文。 @@ -155,7 +156,7 @@ type Service struct { func NewWithFactory( configManager *config.Manager, toolManager tools.Manager, - sessionStore Store, + sessionStore agentsession.Store, providerFactory ProviderFactory, contextBuilder agentcontext.Builder, ) *Service { @@ -373,35 +374,35 @@ func (s *Service) Events() <-chan RuntimeEvent { return s.events } -func (s *Service) ListSessions(ctx context.Context) ([]SessionSummary, error) { +func (s *Service) ListSessions(ctx context.Context) ([]agentsession.Summary, error) { return s.sessionStore.ListSummaries(ctx) } -func (s *Service) LoadSession(ctx context.Context, id string) (Session, error) { +func (s *Service) LoadSession(ctx context.Context, id string) (agentsession.Session, error) { session, err := s.sessionStore.Load(ctx, id) if err != nil { - return Session{}, err + return agentsession.Session{}, err } session.Workdir = s.sessionWorkdir(id, session.Workdir) return session, nil } -func (s *Service) SetSessionWorkdir(ctx context.Context, sessionID string, workdir string) (Session, error) { +func (s *Service) SetSessionWorkdir(ctx context.Context, sessionID string, workdir string) (agentsession.Session, error) { sessionID = strings.TrimSpace(sessionID) if sessionID == "" { - return Session{}, errors.New("runtime: session id is empty") + return agentsession.Session{}, errors.New("runtime: session id is empty") } session, err := s.sessionStore.Load(ctx, sessionID) if err != nil { - return Session{}, err + return agentsession.Session{}, err } session.Workdir = s.sessionWorkdir(sessionID, session.Workdir) cfg := s.configManager.Get() resolved, err := resolveWorkdirForSession(cfg.Workdir, session.Workdir, workdir) if err != nil { - return Session{}, err + return agentsession.Session{}, err } if session.Workdir == resolved { return session, nil @@ -440,22 +441,22 @@ func (s *Service) loadOrCreateSession( title string, defaultWorkdir string, requestedWorkdir string, -) (Session, error) { +) (agentsession.Session, error) { if strings.TrimSpace(sessionID) == "" { sessionWorkdir, err := resolveWorkdirForSession(defaultWorkdir, "", requestedWorkdir) if err != nil { - return Session{}, err + return agentsession.Session{}, err } - session := newSessionWithWorkdir(title, sessionWorkdir) + session := agentsession.NewWithWorkdir(title, sessionWorkdir) s.setSessionWorkdir(session.ID, sessionWorkdir) if err := s.sessionStore.Save(ctx, &session); err != nil { - return Session{}, err + return agentsession.Session{}, err } return session, nil } session, err := s.sessionStore.Load(ctx, sessionID) if err != nil { - return Session{}, err + return agentsession.Session{}, err } session.Workdir = s.sessionWorkdir(sessionID, session.Workdir) if strings.TrimSpace(requestedWorkdir) == "" && strings.TrimSpace(session.Workdir) != "" { @@ -464,7 +465,7 @@ func (s *Service) loadOrCreateSession( resolved, err := resolveWorkdirForSession(defaultWorkdir, session.Workdir, requestedWorkdir) if err != nil { - return Session{}, err + return agentsession.Session{}, err } if session.Workdir == resolved { return session, nil diff --git a/internal/runtime/runtime_test.go b/internal/runtime/runtime_test.go index 0fe9f86a..242a59d3 100644 --- a/internal/runtime/runtime_test.go +++ b/internal/runtime/runtime_test.go @@ -16,16 +16,17 @@ import ( "neo-code/internal/provider" providertypes "neo-code/internal/provider/types" "neo-code/internal/security" + agentsession "neo-code/internal/session" "neo-code/internal/tools" ) type memoryStore struct { - sessions map[string]Session + sessions map[string]agentsession.Session saves int } type failingStore struct { - Store + agentsession.Store saveErr error failOnSave int saveCalls int @@ -33,10 +34,10 @@ type failingStore struct { } func newMemoryStore() *memoryStore { - return &memoryStore{sessions: map[string]Session{}} + return &memoryStore{sessions: map[string]agentsession.Session{}} } -func (s *failingStore) Save(ctx context.Context, session *Session) error { +func (s *failingStore) Save(ctx context.Context, session *agentsession.Session) error { s.saveCalls++ if s.failOnSave > 0 && s.saveCalls == s.failOnSave { return s.saveErr @@ -50,7 +51,7 @@ func (s *failingStore) Save(ctx context.Context, session *Session) error { return s.Store.Save(ctx, session) } -func (s *memoryStore) Save(ctx context.Context, session *Session) error { +func (s *memoryStore) Save(ctx context.Context, session *agentsession.Session) error { if err := ctx.Err(); err != nil { return err } @@ -62,24 +63,24 @@ func (s *memoryStore) Save(ctx context.Context, session *Session) error { return nil } -func (s *memoryStore) Load(ctx context.Context, id string) (Session, error) { +func (s *memoryStore) Load(ctx context.Context, id string) (agentsession.Session, error) { if err := ctx.Err(); err != nil { - return Session{}, err + return agentsession.Session{}, err } session, ok := s.sessions[id] if !ok { - return Session{}, errors.New("not found") + return agentsession.Session{}, errors.New("not found") } return cloneSession(session), nil } -func (s *memoryStore) ListSummaries(ctx context.Context) ([]SessionSummary, error) { +func (s *memoryStore) ListSummaries(ctx context.Context) ([]agentsession.Summary, error) { if err := ctx.Err(); err != nil { return nil, err } - summaries := make([]SessionSummary, 0, len(s.sessions)) + summaries := make([]agentsession.Summary, 0, len(s.sessions)) for _, session := range s.sessions { - summaries = append(summaries, SessionSummary{ + summaries = append(summaries, agentsession.Summary{ ID: session.ID, Title: session.Title, CreatedAt: session.CreatedAt, @@ -607,7 +608,7 @@ func TestServiceRunDelegatesToContextBuilder(t *testing.T) { manager := newRuntimeConfigManager(t) store := newMemoryStore() - session := newSession("memory reject") + session := agentsession.New("memory reject") session.ID = "session-memory-reject" store.sessions[session.ID] = cloneSession(session) registry := tools.NewRegistry() @@ -748,7 +749,7 @@ func TestServiceRunDefaultBuilderUsesToolManagerMicroCompactPolicies(t *testing. registry.Register(&stubTool{name: "bash", content: "default"}) registry.Register(&stubTool{name: "webfetch", content: "default"}) - session := newSession("preserve history") + session := agentsession.New("preserve history") session.ID = "session-preserve-history" session.Messages = []providertypes.Message{ {Role: providertypes.RoleUser, Content: "older user"}, @@ -811,7 +812,7 @@ func TestServiceRunDefaultBuilderUsesGenericToolManagerMicroCompactPolicies(t *t }, } - session := newSession("preserve history by manager") + session := agentsession.New("preserve history by manager") session.ID = "session-preserve-history-manager" session.Messages = []providertypes.Message{ {Role: providertypes.RoleUser, Content: "older user"}, @@ -880,7 +881,7 @@ func TestServiceRunFailurePreservesExistingSessionProviderAndModel(t *testing.T) } store := newMemoryStore() - session := newSession("preserve-metadata") + session := agentsession.New("preserve-metadata") session.ID = "session-preserve-metadata" session.Provider = config.OpenAIName session.Model = "openai-original-model" @@ -980,7 +981,7 @@ func TestServiceRunWaitsForPermissionResolutionAndContinues(t *testing.T) { manager := newRuntimeConfigManager(t) store := newMemoryStore() - session := newSession("memory reject") + session := agentsession.New("memory reject") session.ID = "session-memory-reject" store.sessions[session.ID] = cloneSession(session) registry := tools.NewRegistry() @@ -1171,7 +1172,7 @@ func TestServiceRunEmitsRememberScopeWhenSessionRejectMemoryHits(t *testing.T) { manager := newRuntimeConfigManager(t) store := newMemoryStore() - session := newSession("memory reject") + session := agentsession.New("memory reject") session.ID = "session-memory-reject" store.sessions[session.ID] = cloneSession(session) registry := tools.NewRegistry() @@ -1311,7 +1312,7 @@ func TestServiceRunErrorPaths(t *testing.T) { provider *scriptedProvider factoryErr error registerTool *stubTool - seedSession *Session + seedSession *agentsession.Session expectErr string expectEvents []EventType assert func(t *testing.T, store *memoryStore, provider *scriptedProvider, tool *stubTool) @@ -1375,11 +1376,11 @@ func TestServiceRunErrorPaths(t *testing.T) { {providertypes.NewTextDeltaStreamEvent("resumed")}, }, }, - seedSession: &Session{ + seedSession: &agentsession.Session{ ID: "existing-session", Title: "Resume Me", - CreatedAt: newSession("seed").CreatedAt, - UpdatedAt: newSession("seed").UpdatedAt, + CreatedAt: agentsession.New("seed").CreatedAt, + UpdatedAt: agentsession.New("seed").UpdatedAt, Messages: []providertypes.Message{ {Role: "user", Content: "earlier"}, }, @@ -1869,7 +1870,7 @@ func TestServiceRunToolTimeoutIsNotCancellation(t *testing.T) { func TestServiceCompactManualAppliesAndPersists(t *testing.T) { manager := newRuntimeConfigManager(t) store := newMemoryStore() - session := newSession("manual") + session := agentsession.New("manual") session.ID = "session-manual" session.Messages = []providertypes.Message{ {Role: providertypes.RoleUser, Content: "older"}, @@ -1927,7 +1928,7 @@ func TestServiceCompactManualAppliesAndPersists(t *testing.T) { func TestServiceCompactManualFailureReturnsError(t *testing.T) { manager := newRuntimeConfigManager(t) store := newMemoryStore() - session := newSession("manual-fail") + session := agentsession.New("manual-fail") session.ID = "session-manual-fail" session.Messages = []providertypes.Message{ {Role: providertypes.RoleUser, Content: "older"}, @@ -1981,7 +1982,7 @@ func TestServiceCompactUsesSessionProviderAndModelWhenPresent(t *testing.T) { } store := newMemoryStore() - session := newSession("manual-provider") + session := agentsession.New("manual-provider") session.ID = "session-manual-provider" session.Provider = config.OpenAIName session.Model = "session-model" @@ -2055,7 +2056,7 @@ func TestServiceCompactFallsBackToCurrentProviderWhenSessionMetadataMissing(t *t } store := newMemoryStore() - session := newSession("manual-fallback") + session := agentsession.New("manual-fallback") session.ID = "session-manual-fallback" session.Messages = []providertypes.Message{ {Role: providertypes.RoleUser, Content: "older"}, @@ -2112,7 +2113,7 @@ func TestServiceCompactFallsBackToCurrentProviderWhenSessionMetadataMissing(t *t func TestServiceManualCompactThenRunContinuesToolRound(t *testing.T) { manager := newRuntimeConfigManager(t) store := newMemoryStore() - session := newSession("manual-continue") + session := agentsession.New("manual-continue") session.ID = "session-manual-continue" session.Messages = []providertypes.Message{ {Role: providertypes.RoleUser, Content: "legacy request"}, @@ -2195,7 +2196,7 @@ func TestServiceManualCompactThenRunContinuesToolRound(t *testing.T) { func TestServiceSerializesRunAndCompact(t *testing.T) { manager := newRuntimeConfigManager(t) store := newMemoryStore() - session := newSession("serialized") + session := agentsession.New("serialized") session.ID = "session-serialized" store.sessions[session.ID] = cloneSession(session) @@ -2291,7 +2292,7 @@ func TestServiceConstructorsAndDelegates(t *testing.T) { t.Fatalf("expected events channel") } - session := newSession("List Me") + session := agentsession.New("List Me") store.sessions[session.ID] = cloneSession(session) summaries, err := service.ListSessions(context.Background()) @@ -2310,7 +2311,7 @@ func TestServiceConstructorsAndDelegates(t *testing.T) { t.Fatalf("expected loaded session %q, got %q", session.ID, loaded.ID) } - sessionStore := NewSessionStore(t.TempDir()) + sessionStore := agentsession.NewStore(t.TempDir()) if sessionStore == nil { t.Fatalf("expected JSON session store") } @@ -2328,7 +2329,7 @@ func TestServiceRunUsesSessionWorkdirForContextAndTools(t *testing.T) { } store := newMemoryStore() - session := newSessionWithWorkdir("Session Workdir", sessionWorkdir) + session := agentsession.NewWithWorkdir("Session Workdir", sessionWorkdir) store.sessions[session.ID] = cloneSession(session) tool := &stubTool{name: "filesystem_edit", content: "ok"} @@ -2417,7 +2418,7 @@ func TestServiceSetSessionWorkdir(t *testing.T) { } store := newMemoryStore() - session := newSession("set workdir") + session := agentsession.New("set workdir") store.sessions[session.ID] = cloneSession(session) registry := tools.NewRegistry() registry.Register(&stubTool{name: "filesystem_read_file", content: "default"}) @@ -2548,7 +2549,7 @@ func restoreRuntimeEnv(t *testing.T, key string) { }) } -func onlySession(t *testing.T, store *memoryStore) Session { +func onlySession(t *testing.T, store *memoryStore) agentsession.Session { t.Helper() if len(store.sessions) != 1 { t.Fatalf("expected exactly 1 session, got %d", len(store.sessions)) @@ -2556,7 +2557,7 @@ func onlySession(t *testing.T, store *memoryStore) Session { for _, session := range store.sessions { return session } - return Session{} + return agentsession.Session{} } func resolvedProviderForTests(cfg config.Config, providerName string) (config.ResolvedProviderConfig, error) { @@ -2613,7 +2614,7 @@ func assertEventsRunID(t *testing.T, events []RuntimeEvent, runID string) { } } -func cloneSession(session Session) Session { +func cloneSession(session agentsession.Session) agentsession.Session { cloned := session cloned.Messages = append([]providertypes.Message(nil), session.Messages...) return cloned @@ -2698,7 +2699,7 @@ func TestServiceSetSessionWorkdirNoopDoesNotSave(t *testing.T) { store := newMemoryStore() target := t.TempDir() - session := newSessionWithWorkdir("noop", target) + session := agentsession.NewWithWorkdir("noop", target) store.sessions[session.ID] = cloneSession(session) registry := tools.NewRegistry() registry.Register(&stubTool{name: "filesystem_read_file", content: "default"}) diff --git a/internal/runtime/workdir_branch_test.go b/internal/runtime/workdir_branch_test.go index 0edc39ef..266fed5e 100644 --- a/internal/runtime/workdir_branch_test.go +++ b/internal/runtime/workdir_branch_test.go @@ -6,6 +6,8 @@ import ( "path/filepath" "strings" "testing" + + agentsession "neo-code/internal/session" ) func TestSessionWorkdirKeyAndMemoryMap(t *testing.T) { @@ -82,7 +84,7 @@ func TestLoadSessionUsesFallbackWorkdirWhenMemoryMissing(t *testing.T) { manager := newRuntimeConfigManager(t) store := newMemoryStore() - session := newSession("fallback") + session := agentsession.New("fallback") session.Workdir = t.TempDir() store.sessions[session.ID] = cloneSession(session) diff --git a/internal/session/id.go b/internal/session/id.go new file mode 100644 index 00000000..dc2e84e9 --- /dev/null +++ b/internal/session/id.go @@ -0,0 +1,18 @@ +package session + +import ( + "crypto/rand" + "encoding/hex" +) + +// NewID 生成带前缀的随机 ID,格式为 "_<16hex>"。 +func NewID(prefix string) string { + buf := make([]byte, 8) + _, _ = rand.Read(buf) + return prefix + "_" + hex.EncodeToString(buf) +} + +// newID 保留为内部兼容入口,后续代码请优先使用 NewID。 +func newID(prefix string) string { + return NewID(prefix) +} diff --git a/internal/runtime/session.go b/internal/session/store.go similarity index 59% rename from internal/runtime/session.go rename to internal/session/store.go index 24378d28..824a8228 100644 --- a/internal/runtime/session.go +++ b/internal/session/store.go @@ -1,4 +1,4 @@ -package runtime +package session import ( "context" @@ -17,6 +17,8 @@ import ( const sessionsDirName = "sessions" +// Session 表示单个会话的持久化模型,包含基础元数据与消息历史。 +// Provider / Model 用于在 compact 等流程中优先复用会话最近一次成功运行的模型配置。 type Session struct { ID string `json:"id"` Title string `json:"title"` @@ -30,71 +32,78 @@ type Session struct { Messages []providertypes.Message `json:"messages"` } -type SessionSummary struct { +// Summary 表示会话列表视图所需的轻量摘要信息。 +type Summary struct { ID string `json:"id"` Title string `json:"title"` CreatedAt time.Time `json:"created_at"` UpdatedAt time.Time `json:"updated_at"` } +// Store 定义会话持久化抽象。 type Store interface { Save(ctx context.Context, session *Session) error Load(ctx context.Context, id string) (Session, error) - ListSummaries(ctx context.Context) ([]SessionSummary, error) + ListSummaries(ctx context.Context) ([]Summary, error) } -type JSONSessionStore struct { +// JSONStore 是基于 JSON 文件的会话存储实现。 +type JSONStore struct { mu sync.RWMutex baseDir string } -func NewJSONSessionStore(baseDir string) *JSONSessionStore { - return &JSONSessionStore{ +// NewJSONStore 创建 JSONStore,实际会话目录为 {baseDir}/sessions。 +func NewJSONStore(baseDir string) *JSONStore { + return &JSONStore{ baseDir: filepath.Join(baseDir, sessionsDirName), } } -func NewSessionStore(baseDir string) *JSONSessionStore { - return NewJSONSessionStore(baseDir) +// NewStore 返回默认会话存储实现(当前为 JSONStore)。 +func NewStore(baseDir string) *JSONStore { + return NewJSONStore(baseDir) } -func (s *JSONSessionStore) Save(ctx context.Context, session *Session) error { +// Save 持久化会话到 JSON 文件,采用临时文件 + 原子替换策略。 +func (s *JSONStore) Save(ctx context.Context, session *Session) error { if err := ctx.Err(); err != nil { return err } if session == nil { - return errors.New("runtime: session is nil") + return errors.New("session: session is nil") } s.mu.Lock() defer s.mu.Unlock() if err := os.MkdirAll(s.baseDir, 0o755); err != nil { - return fmt.Errorf("runtime: create sessions dir: %w", err) + return fmt.Errorf("session: create sessions dir: %w", err) } payload, err := json.MarshalIndent(session, "", " ") if err != nil { - return fmt.Errorf("runtime: marshal session: %w", err) + return fmt.Errorf("session: marshal session: %w", err) } payload = append(payload, '\n') target := s.filePath(session.ID) temp := target + ".tmp" if err := os.WriteFile(temp, payload, 0o644); err != nil { - return fmt.Errorf("runtime: write temp session: %w", err) + return fmt.Errorf("session: write temp session: %w", err) } if err := os.Remove(target); err != nil && !errors.Is(err, os.ErrNotExist) { - return fmt.Errorf("runtime: replace session file: %w", err) + return fmt.Errorf("session: replace session file: %w", err) } if err := os.Rename(temp, target); err != nil { - return fmt.Errorf("runtime: commit session file: %w", err) + return fmt.Errorf("session: commit session file: %w", err) } return nil } -func (s *JSONSessionStore) Load(ctx context.Context, id string) (Session, error) { +// Load 读取并反序列化指定 ID 的会话文件。 +func (s *JSONStore) Load(ctx context.Context, id string) (Session, error) { if err := ctx.Err(); err != nil { return Session{}, err } @@ -109,12 +118,13 @@ func (s *JSONSessionStore) Load(ctx context.Context, id string) (Session, error) var session Session if err := json.Unmarshal(data, &session); err != nil { - return Session{}, fmt.Errorf("runtime: decode session %s: %w", id, err) + return Session{}, fmt.Errorf("session: decode session %s: %w", id, err) } return session, nil } -func (s *JSONSessionStore) ListSummaries(ctx context.Context) ([]SessionSummary, error) { +// ListSummaries 列出所有会话摘要,并按 UpdatedAt 倒序返回。 +func (s *JSONStore) ListSummaries(ctx context.Context) ([]Summary, error) { if err := ctx.Err(); err != nil { return nil, err } @@ -123,15 +133,15 @@ func (s *JSONSessionStore) ListSummaries(ctx context.Context) ([]SessionSummary, defer s.mu.RUnlock() if err := os.MkdirAll(s.baseDir, 0o755); err != nil { - return nil, fmt.Errorf("runtime: create sessions dir: %w", err) + return nil, fmt.Errorf("session: create sessions dir: %w", err) } entries, err := os.ReadDir(s.baseDir) if err != nil { - return nil, fmt.Errorf("runtime: list sessions dir: %w", err) + return nil, fmt.Errorf("session: list sessions dir: %w", err) } - summaries := make([]SessionSummary, 0, len(entries)) + summaries := make([]Summary, 0, len(entries)) for _, entry := range entries { if entry.IsDir() || filepath.Ext(entry.Name()) != ".json" { continue @@ -148,7 +158,7 @@ func (s *JSONSessionStore) ListSummaries(ctx context.Context) ([]SessionSummary, continue } - var summary SessionSummary + var summary Summary if err := json.Unmarshal(data, &summary); err != nil { continue } @@ -165,18 +175,21 @@ func (s *JSONSessionStore) ListSummaries(ctx context.Context) ([]SessionSummary, return summaries, nil } -func (s *JSONSessionStore) filePath(id string) string { +// filePath 生成会话 ID 对应的 JSON 文件路径。 +func (s *JSONStore) filePath(id string) string { return filepath.Join(s.baseDir, id+".json") } -func newSession(title string) Session { - return newSessionWithWorkdir(title, "") +// New 创建一个默认标题策略的新会话对象。 +func New(title string) Session { + return NewWithWorkdir(title, "") } -func newSessionWithWorkdir(title string, workdir string) Session { +// NewWithWorkdir 创建一个包含运行目录的会话对象。 +func NewWithWorkdir(title string, workdir string) Session { now := time.Now() return Session{ - ID: newID("session"), + ID: NewID("session"), Title: sanitizeTitle(title), CreatedAt: now, UpdatedAt: now, @@ -185,6 +198,7 @@ func newSessionWithWorkdir(title string, workdir string) Session { } } +// sanitizeTitle 规范化会话标题:去空白、空标题回退默认值、超长截断。 func sanitizeTitle(title string) string { title = strings.TrimSpace(title) if title == "" { diff --git a/internal/runtime/session_test.go b/internal/session/store_test.go similarity index 88% rename from internal/runtime/session_test.go rename to internal/session/store_test.go index 06faa257..985cc809 100644 --- a/internal/runtime/session_test.go +++ b/internal/session/store_test.go @@ -1,4 +1,4 @@ -package runtime +package session import ( "context" @@ -11,11 +11,11 @@ import ( providertypes "neo-code/internal/provider/types" ) -func TestJSONSessionStoreSaveLoadAndListSummaries(t *testing.T) { +func TestJSONStoreSaveLoadAndListSummaries(t *testing.T) { t.Parallel() baseDir := t.TempDir() - store := NewJSONSessionStore(baseDir) + store := NewJSONStore(baseDir) older := &Session{ ID: "session-old", @@ -68,7 +68,7 @@ func TestJSONSessionStoreSaveLoadAndListSummaries(t *testing.T) { t.Fatalf("expected persisted session file to exclude workdir, got:\n%s", string(raw)) } - mustWriteRuntimeFile(t, filepath.Join(baseDir, sessionsDirName, "invalid.json"), "{invalid") + mustWriteSessionFile(t, filepath.Join(baseDir, sessionsDirName, "invalid.json"), "{invalid") if err := os.MkdirAll(filepath.Join(baseDir, sessionsDirName, "directory"), 0o755); err != nil { t.Fatalf("mkdir stray directory: %v", err) } @@ -85,11 +85,11 @@ func TestJSONSessionStoreSaveLoadAndListSummaries(t *testing.T) { } } -func TestJSONSessionStoreErrors(t *testing.T) { +func TestJSONStoreErrors(t *testing.T) { t.Parallel() baseDir := t.TempDir() - store := NewJSONSessionStore(baseDir) + store := NewJSONStore(baseDir) cancelledCtx, cancel := context.WithCancel(context.Background()) cancel() @@ -108,11 +108,11 @@ func TestJSONSessionStoreErrors(t *testing.T) { } } -func TestJSONSessionStoreCorruptedSessionBehaviors(t *testing.T) { +func TestJSONStoreCorruptedSessionBehaviors(t *testing.T) { t.Parallel() baseDir := t.TempDir() - store := NewJSONSessionStore(baseDir) + store := NewJSONStore(baseDir) valid := &Session{ ID: "valid-session", @@ -125,7 +125,7 @@ func TestJSONSessionStoreCorruptedSessionBehaviors(t *testing.T) { t.Fatalf("Save valid session: %v", err) } - mustWriteRuntimeFile(t, filepath.Join(baseDir, sessionsDirName, "broken.json"), "{broken") + mustWriteSessionFile(t, filepath.Join(baseDir, sessionsDirName, "broken.json"), "{broken") _, err := store.Load(context.Background(), "broken") if err == nil || !strings.Contains(err.Error(), "decode session broken") { @@ -141,7 +141,7 @@ func TestJSONSessionStoreCorruptedSessionBehaviors(t *testing.T) { } } -func TestJSONSessionStoreSaveInvalidBaseDir(t *testing.T) { +func TestJSONStoreSaveInvalidBaseDir(t *testing.T) { t.Parallel() tempDir := t.TempDir() @@ -150,7 +150,7 @@ func TestJSONSessionStoreSaveInvalidBaseDir(t *testing.T) { t.Fatalf("write base file: %v", err) } - store := NewJSONSessionStore(baseFile) + store := NewJSONStore(baseFile) err := store.Save(context.Background(), &Session{ ID: "session-x", Title: "Broken Save", @@ -162,7 +162,7 @@ func TestJSONSessionStoreSaveInvalidBaseDir(t *testing.T) { } } -func mustWriteRuntimeFile(t *testing.T, path string, content string) { +func mustWriteSessionFile(t *testing.T, path string, content string) { t.Helper() if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil { t.Fatalf("mkdir %s: %v", filepath.Dir(path), err) diff --git a/internal/tui/state.go b/internal/tui/state.go index 1f939dad..c65ecc63 100644 --- a/internal/tui/state.go +++ b/internal/tui/state.go @@ -10,7 +10,7 @@ import ( tea "github.com/charmbracelet/bubbletea" "github.com/charmbracelet/lipgloss" - agentruntime "neo-code/internal/runtime" + agentsession "neo-code/internal/session" ) type panel int @@ -32,7 +32,7 @@ const ( ) type UIState struct { - Sessions []agentruntime.SessionSummary + Sessions []agentsession.Summary ActiveSessionID string ActiveSessionTitle string InputText string @@ -142,7 +142,7 @@ func (d commandMenuDelegate) Render(w io.Writer, m list.Model, index int, item l } type sessionItem struct { - Summary agentruntime.SessionSummary + Summary agentsession.Summary Active bool } diff --git a/internal/tui/state/ui_state.go b/internal/tui/state/ui_state.go index ac6c4611..9fa071a3 100644 --- a/internal/tui/state/ui_state.go +++ b/internal/tui/state/ui_state.go @@ -1,6 +1,6 @@ package state -import agentruntime "neo-code/internal/runtime" +import agentsession "neo-code/internal/session" // Panel 定义 TUI 中可聚焦的主面板。 type Panel int @@ -24,7 +24,7 @@ const ( // UIState 保存顶层界面状态快照,仅作为数据容器使用。 type UIState struct { - Sessions []agentruntime.SessionSummary + Sessions []agentsession.Summary ActiveSessionID string ActiveSessionTitle string ActiveRunID string diff --git a/internal/tui/update_test.go b/internal/tui/update_test.go index a02958f7..c55e23ae 100644 --- a/internal/tui/update_test.go +++ b/internal/tui/update_test.go @@ -22,6 +22,7 @@ import ( providercatalog "neo-code/internal/provider/catalog" providertypes "neo-code/internal/provider/types" agentruntime "neo-code/internal/runtime" + agentsession "neo-code/internal/session" "neo-code/internal/tools" ) @@ -29,15 +30,15 @@ type stubRuntime struct { runInputs []agentruntime.UserInput compactInputs []agentruntime.CompactInput events chan agentruntime.RuntimeEvent - sessions []agentruntime.SessionSummary - loads map[string]agentruntime.Session + sessions []agentsession.Summary + loads map[string]agentsession.Session runErr error compactErr error compactResult agentruntime.CompactResult listErr error loadErr error setWorkdirErr error - setResult *agentruntime.Session + setResult *agentsession.Session setCalls int resolveInputs []agentruntime.PermissionResolutionInput resolveErr error @@ -67,7 +68,7 @@ func (r *stubMarkdownRenderer) Render(content string, width int) (string, error) func newStubRuntime() *stubRuntime { return &stubRuntime{ events: make(chan agentruntime.RuntimeEvent, 16), - loads: map[string]agentruntime.Session{}, + loads: map[string]agentsession.Session{}, } } @@ -95,34 +96,34 @@ func (r *stubRuntime) CancelActiveRun() bool { return r.cancelResult } -func (r *stubRuntime) ListSessions(ctx context.Context) ([]agentruntime.SessionSummary, error) { +func (r *stubRuntime) ListSessions(ctx context.Context) ([]agentsession.Summary, error) { if r.listErr != nil { return nil, r.listErr } - return append([]agentruntime.SessionSummary(nil), r.sessions...), nil + return append([]agentsession.Summary(nil), r.sessions...), nil } -func (r *stubRuntime) LoadSession(ctx context.Context, id string) (agentruntime.Session, error) { +func (r *stubRuntime) LoadSession(ctx context.Context, id string) (agentsession.Session, error) { if r.loadErr != nil { - return agentruntime.Session{}, r.loadErr + return agentsession.Session{}, r.loadErr } if session, ok := r.loads[id]; ok { return session, nil } - return agentruntime.Session{}, nil + return agentsession.Session{}, nil } -func (r *stubRuntime) SetSessionWorkdir(ctx context.Context, sessionID string, workdir string) (agentruntime.Session, error) { +func (r *stubRuntime) SetSessionWorkdir(ctx context.Context, sessionID string, workdir string) (agentsession.Session, error) { r.setCalls++ if r.setWorkdirErr != nil { - return agentruntime.Session{}, r.setWorkdirErr + return agentsession.Session{}, r.setWorkdirErr } if r.setResult != nil { return *r.setResult, nil } session, ok := r.loads[sessionID] if !ok { - session = agentruntime.Session{ID: sessionID} + session = agentsession.Session{ID: sessionID} } session.Workdir = strings.TrimSpace(workdir) r.loads[sessionID] = session @@ -251,7 +252,7 @@ func TestAppUpdateWorkspaceSlashCommands(t *testing.T) { manager := newTestConfigManager(t) runtime := newStubRuntime() sessionID := "session-workdir" - runtime.loads[sessionID] = agentruntime.Session{ID: sessionID, Workdir: t.TempDir()} + runtime.loads[sessionID] = agentsession.Session{ID: sessionID, Workdir: t.TempDir()} app, err := New(nil, manager, runtime, newTestProviderService(t, manager)) if err != nil { @@ -328,7 +329,7 @@ func TestRunSessionWorkdirCommandBranches(t *testing.T) { t.Run("session workdir fallback uses current workdir when runtime returns empty", func(t *testing.T) { current := t.TempDir() runtime := newStubRuntime() - runtime.setResult = &agentruntime.Session{ID: "session-1", Workdir: ""} + runtime.setResult = &agentsession.Session{ID: "session-1", Workdir: ""} msg := runSessionWorkdirCommand(runtime, "session-1", current, "/cwd ./subdir")() result := msg.(sessionWorkdirResultMsg) if result.err != nil { @@ -343,7 +344,7 @@ func TestRunSessionWorkdirCommandBranches(t *testing.T) { current := t.TempDir() target := t.TempDir() runtime := newStubRuntime() - runtime.setResult = &agentruntime.Session{ID: "session-1", Workdir: target} + runtime.setResult = &agentsession.Session{ID: "session-1", Workdir: target} msg := runSessionWorkdirCommand(runtime, "session-1", current, "/cwd ./subdir")() result := msg.(sessionWorkdirResultMsg) if result.err != nil { @@ -776,7 +777,7 @@ func TestAppUpdateModelPickerAndRuntimeMessages(t *testing.T) { func TestAppHelpersAndRenderingSmoke(t *testing.T) { manager := newTestConfigManager(t) runtime := newStubRuntime() - now := agentruntime.Session{ + now := agentsession.Session{ ID: "session-1", Title: "Existing Session", Messages: []providertypes.Message{ @@ -784,7 +785,7 @@ func TestAppHelpersAndRenderingSmoke(t *testing.T) { {Role: roleAssistant, Content: "hello"}, }, } - runtime.sessions = []agentruntime.SessionSummary{ + runtime.sessions = []agentsession.Summary{ {ID: now.ID, Title: now.Title, UpdatedAt: now.UpdatedAt}, } runtime.loads[now.ID] = now @@ -986,7 +987,7 @@ func TestTUIStandaloneHelpers(t *testing.T) { t.Fatalf("expected numeric helpers to work") } - sItem := sessionItem{Summary: agentruntime.SessionSummary{Title: "My Session"}} + sItem := sessionItem{Summary: agentsession.Summary{Title: "My Session"}} if sItem.FilterValue() != "my session" { t.Fatalf("unexpected session item filter value") } @@ -1208,8 +1209,8 @@ func TestAppUpdateAdditionalTransitions(t *testing.T) { { name: "session enter activates selected session", setup: func(t *testing.T, app *App, runtime *stubRuntime, manager *config.Manager) { - runtime.sessions = []agentruntime.SessionSummary{{ID: "s1", Title: "One"}} - runtime.loads["s1"] = agentruntime.Session{ + runtime.sessions = []agentsession.Summary{{ID: "s1", Title: "One"}} + runtime.loads["s1"] = agentsession.Session{ ID: "s1", Title: "One", Messages: []providertypes.Message{{Role: roleAssistant, Content: "loaded"}}, From 85667966186e71ebc2a3deb01d5d9dba4ef1b4a7 Mon Sep 17 00:00:00 2001 From: Yumiue <229866007@qq.com> Date: Tue, 7 Apr 2026 18:04:04 +0800 Subject: [PATCH 46/55] =?UTF-8?q?test:=E8=A1=A5=E5=85=85=E6=B5=8B=E8=AF=95?= =?UTF-8?q?=E6=96=87=E4=BB=B6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/session/id_test.go | 41 ++++++++++++++++++++++++++ internal/session/store_test.go | 54 ++++++++++++++++++++++++++++++++++ 2 files changed, 95 insertions(+) create mode 100644 internal/session/id_test.go diff --git a/internal/session/id_test.go b/internal/session/id_test.go new file mode 100644 index 00000000..8ef1bb51 --- /dev/null +++ b/internal/session/id_test.go @@ -0,0 +1,41 @@ +package session + +import ( + "strings" + "testing" +) + +func TestNewIDFormatAndUniqueness(t *testing.T) { + t.Parallel() + + id1 := NewID("session") + id2 := NewID("session") + + if !strings.HasPrefix(id1, "session_") || !strings.HasPrefix(id2, "session_") { + t.Fatalf("expected prefix session_, got %q and %q", id1, id2) + } + + hex1 := strings.TrimPrefix(id1, "session_") + hex2 := strings.TrimPrefix(id2, "session_") + + if len(hex1) != 16 || len(hex2) != 16 { + t.Fatalf("expected 16 hex chars, got %d and %d", len(hex1), len(hex2)) + } + for _, ch := range hex1 + hex2 { + if !((ch >= '0' && ch <= '9') || (ch >= 'a' && ch <= 'f')) { + t.Fatalf("expected lowercase hex, got %q in ids %q %q", ch, id1, id2) + } + } + if id1 == id2 { + t.Fatalf("expected different ids, got identical %q", id1) + } +} + +func TestNewIDAllowsEmptyPrefix(t *testing.T) { + t.Parallel() + + id := NewID("") + if len(id) != 17 || id[0] != '_' { + t.Fatalf("expected format _<16hex>, got %q", id) + } +} diff --git a/internal/session/store_test.go b/internal/session/store_test.go index 985cc809..9474d761 100644 --- a/internal/session/store_test.go +++ b/internal/session/store_test.go @@ -162,6 +162,60 @@ func TestJSONStoreSaveInvalidBaseDir(t *testing.T) { } } +func TestNewUsesDefaultWorkdirAndEmptyMessages(t *testing.T) { + t.Parallel() + + session := New("hello title") + + if session.ID == "" { + t.Fatalf("expected non-empty id") + } + if !strings.HasPrefix(session.ID, "session_") { + t.Fatalf("expected id with session_ prefix, got %q", session.ID) + } + if session.Title != "hello title" { + t.Fatalf("expected title %q, got %q", "hello title", session.Title) + } + if session.Workdir != "" { + t.Fatalf("expected empty workdir, got %q", session.Workdir) + } + if len(session.Messages) != 0 { + t.Fatalf("expected empty messages, got %+v", session.Messages) + } + if session.CreatedAt.IsZero() || session.UpdatedAt.IsZero() { + t.Fatalf("expected non-zero timestamps, got created=%v updated=%v", session.CreatedAt, session.UpdatedAt) + } + if session.UpdatedAt.Before(session.CreatedAt) { + t.Fatalf("expected UpdatedAt >= CreatedAt, got created=%v updated=%v", session.CreatedAt, session.UpdatedAt) + } +} + +func TestNewWithWorkdirTrimAndTitleSanitize(t *testing.T) { + t.Parallel() + + tooLong := strings.Repeat("中", 45) // rune 长度 > 40 + workdir := " /tmp/workdir " + + session := NewWithWorkdir(tooLong, workdir) + + if session.Workdir != "/tmp/workdir" { + t.Fatalf("expected trimmed workdir %q, got %q", "/tmp/workdir", session.Workdir) + } + if got := len([]rune(session.Title)); got != 40 { + t.Fatalf("expected title rune length 40, got %d (title=%q)", got, session.Title) + } +} + +func TestNewWithWorkdirFallsBackDefaultTitle(t *testing.T) { + t.Parallel() + + session := NewWithWorkdir(" \n\t ", "") + + if session.Title != "New Session" { + t.Fatalf("expected default title %q, got %q", "New Session", session.Title) + } +} + func mustWriteSessionFile(t *testing.T, path string, content string) { t.Helper() if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil { From 360fb57b9f9f29ef77f7d08c7daae68ccc066af3 Mon Sep 17 00:00:00 2001 From: Yumiue <229866007@qq.com> Date: Tue, 7 Apr 2026 20:22:16 +0800 Subject: [PATCH 47/55] =?UTF-8?q?test:=E6=B7=BB=E5=8A=A0=E6=B5=8B=E8=AF=95?= =?UTF-8?q?=E6=96=87=E4=BB=B6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/runtime/runtime_test.go | 205 +++++++++++++++++++++++++++++++ internal/session/id_test.go | 12 ++ internal/session/store_test.go | 153 +++++++++++++++++++++++ 3 files changed, 370 insertions(+) diff --git a/internal/runtime/runtime_test.go b/internal/runtime/runtime_test.go index 242a59d3..7ff11961 100644 --- a/internal/runtime/runtime_test.go +++ b/internal/runtime/runtime_test.go @@ -2829,3 +2829,208 @@ func TestPermissionEventViewPayloadMapping(t *testing.T) { t.Fatalf("unexpected resolved payload: %+v", resolvedPayload) } } + +func TestStreamAccumulatorBuildMessageRejectsMissingToolName(t *testing.T) { + t.Parallel() + + acc := newStreamAccumulator() + acc.accumulateToolCallStart(0, "call-1", "") + acc.accumulateToolCallDelta(0, "call-1", "{}") + + _, err := acc.buildMessage() + if err == nil || !containsError(err, "without name") { + t.Fatalf("expected missing tool name error, got %v", err) + } +} + +func TestLoadSessionReturnsStoreError(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + store := newMemoryStore() + service := NewWithFactory(manager, nil, store, &scriptedProviderFactory{provider: &scriptedProvider{}}, nil) + + _, err := service.LoadSession(context.Background(), "missing") + if err == nil || !containsError(err, "not found") { + t.Fatalf("expected load error, got %v", err) + } +} + +func TestSetSessionWorkdirReturnsStoreError(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + store := newMemoryStore() + service := NewWithFactory(manager, nil, store, &scriptedProviderFactory{provider: &scriptedProvider{}}, nil) + + _, err := service.SetSessionWorkdir(context.Background(), "missing", t.TempDir()) + if err == nil || !containsError(err, "not found") { + t.Fatalf("expected load error from SetSessionWorkdir, got %v", err) + } +} + +func TestSetSessionWorkdirReturnsResolveError(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + defaultWorkdir := t.TempDir() + if err := manager.Update(context.Background(), func(cfg *config.Config) error { + cfg.Workdir = defaultWorkdir + return nil + }); err != nil { + t.Fatalf("update config: %v", err) + } + + store := newMemoryStore() + session := agentsession.New("set bad workdir") + session.ID = "session-set-bad-workdir" + store.sessions[session.ID] = cloneSession(session) + + service := NewWithFactory(manager, nil, store, &scriptedProviderFactory{provider: &scriptedProvider{}}, nil) + _, err := service.SetSessionWorkdir(context.Background(), session.ID, filepath.Join(defaultWorkdir, "missing-dir")) + if err == nil || !containsError(err, "resolve workdir") { + t.Fatalf("expected resolve workdir error, got %v", err) + } +} + +func TestServiceRunFailsWhenInitialUserMessageSaveFails(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + baseStore := newMemoryStore() + store := &failingStore{ + Store: baseStore, + saveErr: errors.New("save failed on first write"), + failOnSave: 1, + } + registry := tools.NewRegistry() + registry.Register(&stubTool{name: "filesystem_read_file", content: "default"}) + + service := NewWithFactory(manager, registry, store, &scriptedProviderFactory{provider: &scriptedProvider{}}, nil) + err := service.Run(context.Background(), UserInput{ + RunID: "run-initial-save-fail", + Content: "hello", + }) + if err == nil || !containsError(err, "save failed on first write") { + t.Fatalf("expected initial save error, got %v", err) + } +} + +func TestServiceRunFailsWhenAssistantSaveFails(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + baseStore := newMemoryStore() + store := &failingStore{ + Store: baseStore, + saveErr: errors.New("save failed on assistant"), + failOnSave: 2, + } + registry := tools.NewRegistry() + registry.Register(&stubTool{name: "filesystem_read_file", content: "default"}) + + scripted := &scriptedProvider{ + streams: [][]provider.StreamEvent{ + {provider.NewTextDeltaStreamEvent("assistant reply")}, + }, + } + service := NewWithFactory(manager, registry, store, &scriptedProviderFactory{provider: scripted}, nil) + err := service.Run(context.Background(), UserInput{ + RunID: "run-assistant-save-fail", + Content: "hello", + }) + if err == nil || !containsError(err, "save failed on assistant") { + t.Fatalf("expected assistant save error, got %v", err) + } +} + +func TestHandleProviderStreamEventErrorBranches(t *testing.T) { + t.Parallel() + + acc := newStreamAccumulator() + + err := handleProviderStreamEvent( + provider.StreamEvent{Type: provider.StreamEventToolCallStart}, + acc, + nil, + nil, + ) + if err == nil || !containsError(err, "tool_call_start event payload is nil") { + t.Fatalf("expected tool_call_start payload error, got %v", err) + } + + err = handleProviderStreamEvent( + provider.StreamEvent{Type: provider.StreamEventToolCallDelta}, + acc, + nil, + nil, + ) + if err == nil || !containsError(err, "tool_call_delta event payload is nil") { + t.Fatalf("expected tool_call_delta payload error, got %v", err) + } + + err = handleProviderStreamEvent( + provider.StreamEvent{Type: provider.StreamEventMessageDone}, + acc, + nil, + nil, + ) + if err == nil || !containsError(err, "message_done event payload is nil") { + t.Fatalf("expected message_done payload error, got %v", err) + } +} + +func TestEmitDropsWhenChannelFullAndContextCanceled(t *testing.T) { + t.Parallel() + + service := &Service{ + events: make(chan RuntimeEvent, 1), + } + service.events <- RuntimeEvent{Type: EventAgentChunk} + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + done := make(chan struct{}) + go func() { + service.emit(ctx, EventError, "run-id", "session-id", "payload") + close(done) + }() + + select { + case <-done: + case <-time.After(1 * time.Second): + t.Fatal("emit should return when channel is full and context is canceled") + } +} + +func TestCallProviderWithRetryReturnsCombinedForwardError(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + registry := tools.NewRegistry() + registry.Register(&stubTool{name: "filesystem_read_file", content: "default"}) + store := newMemoryStore() + + scripted := &scriptedProvider{ + chatFn: func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) error { + events <- provider.StreamEvent{Type: provider.StreamEventTextDelta} + return errors.New("provider chat failed") + }, + } + service := NewWithFactory(manager, registry, store, &scriptedProviderFactory{provider: scripted}, nil) + + _, err := service.callProviderWithRetry( + context.Background(), + "run-forward-error", + "session-forward-error", + provider.ChatRequest{ + Model: "test-model", + SystemPrompt: "prompt", + Messages: []provider.Message{{Role: provider.RoleUser, Content: "hello"}}, + }, + ) + if err == nil || !containsError(err, "provider stream handling failed after provider error") { + t.Fatalf("expected combined forward/provider error, got %v", err) + } +} diff --git a/internal/session/id_test.go b/internal/session/id_test.go index 8ef1bb51..396779f8 100644 --- a/internal/session/id_test.go +++ b/internal/session/id_test.go @@ -39,3 +39,15 @@ func TestNewIDAllowsEmptyPrefix(t *testing.T) { t.Fatalf("expected format _<16hex>, got %q", id) } } + +func TestNewIDCompatibilityWrapper(t *testing.T) { + t.Parallel() + + id := newID("session") + if !strings.HasPrefix(id, "session_") { + t.Fatalf("expected compatibility wrapper to preserve prefix, got %q", id) + } + if len(strings.TrimPrefix(id, "session_")) != 16 { + t.Fatalf("expected compatibility wrapper to return 16 hex chars, got %q", id) + } +} diff --git a/internal/session/store_test.go b/internal/session/store_test.go index 9474d761..96b2fcb2 100644 --- a/internal/session/store_test.go +++ b/internal/session/store_test.go @@ -2,6 +2,8 @@ package session import ( "context" + "encoding/json" + "errors" "os" "path/filepath" "strings" @@ -216,6 +218,157 @@ func TestNewWithWorkdirFallsBackDefaultTitle(t *testing.T) { } } +func TestNewStoreReturnsJSONStore(t *testing.T) { + t.Parallel() + + store := NewStore(t.TempDir()) + if store == nil { + t.Fatalf("expected non-nil store") + } +} + +func TestJSONStoreListSummariesReadDirFailure(t *testing.T) { + t.Parallel() + + baseDir := t.TempDir() + store := NewJSONStore(baseDir) + + // 把 sessions 目录位置占成普通文件,触发 ReadDir 失败路径。 + sessionsPath := filepath.Join(baseDir, sessionsDirName) + if err := os.WriteFile(sessionsPath, []byte("not-a-dir"), 0o644); err != nil { + t.Fatalf("write %s: %v", sessionsPath, err) + } + + _, err := store.ListSummaries(context.Background()) + if err == nil || !strings.Contains(err.Error(), "create sessions dir") { + t.Fatalf("expected create sessions dir error, got %v", err) + } +} + +func TestJSONStoreListSummariesContextCanceledDuringIteration(t *testing.T) { + t.Parallel() + + baseDir := t.TempDir() + store := NewJSONStore(baseDir) + + for i := 0; i < 10; i++ { + s := &Session{ + ID: "session-iter-" + strings.Repeat("x", i+1), + Title: "iter", + CreatedAt: time.Now(), + UpdatedAt: time.Now(), + } + if err := store.Save(context.Background(), s); err != nil { + t.Fatalf("save session %d: %v", i, err) + } + } + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + _, err := store.ListSummaries(ctx) + if !errors.Is(err, context.Canceled) { + t.Fatalf("expected context canceled, got %v", err) + } +} + +func TestJSONStoreLoadDecodeErrorWithNonJSONPayload(t *testing.T) { + t.Parallel() + + baseDir := t.TempDir() + store := NewJSONStore(baseDir) + + mustWriteSessionFile(t, filepath.Join(baseDir, sessionsDirName, "decode-bad.json"), "{not-json") + + _, err := store.Load(context.Background(), "decode-bad") + if err == nil || !strings.Contains(err.Error(), "decode session decode-bad") { + t.Fatalf("expected decode session error, got %v", err) + } +} + +func TestJSONStoreListSummariesSkipsUnreadableAndMalformedEntries(t *testing.T) { + t.Parallel() + + baseDir := t.TempDir() + store := NewJSONStore(baseDir) + + valid := &Session{ + ID: "valid-summary", + Title: "Valid", + CreatedAt: time.Now().Add(-time.Minute), + UpdatedAt: time.Now(), + } + if err := store.Save(context.Background(), valid); err != nil { + t.Fatalf("save valid session: %v", err) + } + + mustWriteSessionFile(t, filepath.Join(baseDir, sessionsDirName, "malformed.json"), "{malformed") + mustWriteSessionFile(t, filepath.Join(baseDir, sessionsDirName, "empty-id.json"), `{"id":" ","title":"x"}`) + + summaries, err := store.ListSummaries(context.Background()) + if err != nil { + t.Fatalf("ListSummaries() error: %v", err) + } + if len(summaries) != 1 || summaries[0].ID != valid.ID { + t.Fatalf("expected only valid summary, got %+v", summaries) + } +} + +func TestJSONStoreSavePersistsProviderModelAndMessages(t *testing.T) { + t.Parallel() + + baseDir := t.TempDir() + store := NewJSONStore(baseDir) + + session := &Session{ + ID: "persist-full-fields", + Title: "Persist Fields", + Provider: "openai", + Model: "gpt-4.1", + CreatedAt: time.Now().Add(-time.Hour), + UpdatedAt: time.Now(), + Messages: []provider.Message{ + {Role: provider.RoleUser, Content: "hello"}, + { + Role: provider.RoleAssistant, + Content: "calling tool", + ToolCalls: []provider.ToolCall{ + {ID: "call-1", Name: "webfetch", Arguments: `{"url":"https://example.com"}`}, + }, + }, + {Role: provider.RoleTool, ToolCallID: "call-1", Content: "ok"}, + }, + } + + if err := store.Save(context.Background(), session); err != nil { + t.Fatalf("save session: %v", err) + } + + rawPath := filepath.Join(baseDir, sessionsDirName, session.ID+".json") + raw, err := os.ReadFile(rawPath) + if err != nil { + t.Fatalf("read raw file: %v", err) + } + + var decoded map[string]any + if err := json.Unmarshal(raw, &decoded); err != nil { + t.Fatalf("decode raw json: %v", err) + } + + if decoded["provider"] != "openai" { + t.Fatalf("expected provider persisted, got %+v", decoded["provider"]) + } + if decoded["model"] != "gpt-4.1" { + t.Fatalf("expected model persisted, got %+v", decoded["model"]) + } + if _, ok := decoded["messages"]; !ok { + t.Fatalf("expected messages field persisted, got %+v", decoded) + } + if _, ok := decoded["workdir"]; ok { + t.Fatalf("expected workdir not persisted, got %+v", decoded) + } +} + func mustWriteSessionFile(t *testing.T, path string, content string) { t.Helper() if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil { From 704b018b8cbef0e7ee832712d99b5d4cd8434795 Mon Sep 17 00:00:00 2001 From: Cai_Tang <106404101+Cai-Tang-www@users.noreply.github.com> Date: Tue, 7 Apr 2026 20:43:01 +0800 Subject: [PATCH 48/55] =?UTF-8?q?test:=20=E6=8F=90=E5=8D=87MCP=E4=B8=8E?= =?UTF-8?q?=E5=BC=95=E5=AF=BC=E6=B5=81=E7=A8=8B=E8=BE=B9=E7=95=8C=E8=A6=86?= =?UTF-8?q?=E7=9B=96=E7=8E=87?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/app/bootstrap_test.go | 83 ++++++++++++++++++++++++++ internal/tools/mcp/adapter_test.go | 34 +++++++++++ internal/tools/mcp/registry_test.go | 90 +++++++++++++++++++++++++++++ 3 files changed, 207 insertions(+) diff --git a/internal/app/bootstrap_test.go b/internal/app/bootstrap_test.go index d706a1b2..a81903c0 100644 --- a/internal/app/bootstrap_test.go +++ b/internal/app/bootstrap_test.go @@ -351,6 +351,89 @@ func TestResolveMCPServerEnvAndWorkdir(t *testing.T) { } } +func TestResolveMCPServerEnvValidationErrors(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + server config.MCPServerConfig + }{ + { + name: "empty name", + server: config.MCPServerConfig{ + Env: []config.MCPEnvVarConfig{{Name: " ", Value: "x"}}, + }, + }, + { + name: "both value and value_env", + server: config.MCPServerConfig{ + Env: []config.MCPEnvVarConfig{{Name: "A", Value: "x", ValueEnv: "B"}}, + }, + }, + { + name: "missing value and value_env", + server: config.MCPServerConfig{ + Env: []config.MCPEnvVarConfig{{Name: "A"}}, + }, + }, + { + name: "value_env unresolved", + server: config.MCPServerConfig{ + Env: []config.MCPEnvVarConfig{{Name: "A", ValueEnv: "MISSING_ENV_FOR_TEST"}}, + }, + }, + } + + for _, tt := range tests { + tt := tt + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + if _, err := resolveMCPServerEnv(tt.server); err == nil { + t.Fatalf("expected validation error") + } + }) + } +} + +func TestBuildMCPRegistryNoEnabledServerReturnsNil(t *testing.T) { + t.Parallel() + + cfg := config.Default().Clone() + cfg.Workdir = t.TempDir() + cfg.Tools.MCP.Servers = []config.MCPServerConfig{ + {ID: "docs", Enabled: false, Source: "stdio"}, + } + + registry, err := buildMCPRegistry(cfg) + if err != nil { + t.Fatalf("buildMCPRegistry() error = %v", err) + } + if registry != nil { + t.Fatalf("expected nil registry when no enabled server") + } +} + +func TestBuildMCPRegistryRegisterError(t *testing.T) { + t.Parallel() + + cfg := config.Default().Clone() + cfg.Workdir = t.TempDir() + cfg.Tools.MCP.Servers = []config.MCPServerConfig{ + {ID: "docs", Enabled: true, Source: "stdio"}, + } + + originalRegister := registerMCPStdioServer + t.Cleanup(func() { registerMCPStdioServer = originalRegister }) + registerMCPStdioServer = func(registry *mcp.Registry, cfg config.Config, server config.MCPServerConfig) error { + return errors.New("register failed") + } + + _, err := buildMCPRegistry(cfg) + if err == nil || !strings.Contains(err.Error(), "register failed") { + t.Fatalf("expected wrapped register error, got %v", err) + } +} + func TestInitialMCPRefreshTimeoutAndDurationConversion(t *testing.T) { t.Parallel() diff --git a/internal/tools/mcp/adapter_test.go b/internal/tools/mcp/adapter_test.go index 8f3adc8f..2b38691f 100644 --- a/internal/tools/mcp/adapter_test.go +++ b/internal/tools/mcp/adapter_test.go @@ -39,6 +39,19 @@ func TestAdapterFactoryBuildAdapters(t *testing.T) { } } +func TestAdapterFactoryBuildAdaptersEmptySnapshot(t *testing.T) { + t.Parallel() + + factory := NewAdapterFactory(NewRegistry()) + adapters, err := factory.BuildAdapters(context.Background()) + if err != nil { + t.Fatalf("BuildAdapters() error = %v", err) + } + if len(adapters) != 0 { + t.Fatalf("expected empty adapters, got %d", len(adapters)) + } +} + func TestAdapterCall(t *testing.T) { t.Parallel() @@ -193,3 +206,24 @@ func TestAdapterCallBoundary(t *testing.T) { t.Fatalf("expected context canceled error") } } + +func TestAdapterEnsureObjectSchemaDefaults(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + adapter, err := NewAdapter(registry, "docs", ToolDescriptor{ + Name: "search", + Description: "search docs", + InputSchema: map[string]any{}, + }) + if err != nil { + t.Fatalf("NewAdapter() error = %v", err) + } + schema := adapter.Schema() + if schema["type"] != "object" { + t.Fatalf("expected object type, got %v", schema["type"]) + } + if _, ok := schema["properties"].(map[string]any); !ok { + t.Fatalf("expected properties object, got %+v", schema["properties"]) + } +} diff --git a/internal/tools/mcp/registry_test.go b/internal/tools/mcp/registry_test.go index 49f2c654..30c5cd3f 100644 --- a/internal/tools/mcp/registry_test.go +++ b/internal/tools/mcp/registry_test.go @@ -244,3 +244,93 @@ func TestRegistrySetServerStatusValidation(t *testing.T) { t.Fatalf("expected missing server error") } } + +func TestRegistryNilAndValidationBoundaries(t *testing.T) { + t.Parallel() + + var nilRegistry *Registry + if err := nilRegistry.RegisterServer("docs", "stdio", "v1", &stubServerClient{}); err == nil { + t.Fatalf("expected nil registry error") + } + if nilRegistry.UnregisterServer("docs") { + t.Fatalf("nil registry should return false on unregister") + } + if err := nilRegistry.SetServerStatus("docs", ServerStatusReady); err == nil { + t.Fatalf("expected nil registry error for set status") + } + if err := nilRegistry.RefreshServerTools(context.Background(), "docs"); err == nil { + t.Fatalf("expected nil registry error for refresh") + } + if err := nilRegistry.HealthCheck(context.Background(), "docs"); err == nil { + t.Fatalf("expected nil registry error for health check") + } + if _, err := nilRegistry.Call(context.Background(), "docs", "search", nil); err == nil { + t.Fatalf("expected nil registry error for call") + } + if snapshots := nilRegistry.Snapshot(); snapshots != nil { + t.Fatalf("expected nil snapshots from nil registry") + } +} + +func TestRegistryRefreshHealthCallValidation(t *testing.T) { + t.Parallel() + + registry := NewRegistry() + client := &stubServerClient{} + if err := registry.RegisterServer("docs", "stdio", "v1", client); err != nil { + t.Fatalf("register server: %v", err) + } + + canceledCtx, cancel := context.WithCancel(context.Background()) + cancel() + if err := registry.RefreshServerTools(canceledCtx, "docs"); err == nil { + t.Fatalf("expected canceled refresh error") + } + if err := registry.HealthCheck(canceledCtx, "docs"); err == nil { + t.Fatalf("expected canceled health check error") + } + if _, err := registry.Call(canceledCtx, "docs", "search", nil); err == nil { + t.Fatalf("expected canceled call error") + } + + if err := registry.RefreshServerTools(context.Background(), " "); err == nil { + t.Fatalf("expected empty server id error") + } + if err := registry.HealthCheck(context.Background(), " "); err == nil { + t.Fatalf("expected empty server id error") + } + if _, err := registry.Call(context.Background(), " ", "search", nil); err == nil { + t.Fatalf("expected empty server id error") + } + if _, err := registry.Call(context.Background(), "docs", " ", nil); err == nil { + t.Fatalf("expected empty tool name error") + } +} + +func TestRegistryCloneAnyCoversSlicesAndMaps(t *testing.T) { + t.Parallel() + + source := map[string]any{ + "items": []any{ + map[string]any{"name": "a"}, + []any{"nested"}, + }, + } + cloned := cloneSchema(source) + + items, ok := cloned["items"].([]any) + if !ok { + t.Fatalf("expected []any clone") + } + nestedMap, ok := items[0].(map[string]any) + if !ok { + t.Fatalf("expected nested map clone") + } + nestedMap["name"] = "changed" + + originalItems := source["items"].([]any) + originalMap := originalItems[0].(map[string]any) + if originalMap["name"] != "a" { + t.Fatalf("expected deep cloned map, got %v", originalMap["name"]) + } +} From 48f7da0b6fd8afd8cf4da7455c46ae07292e6315 Mon Sep 17 00:00:00 2001 From: Yumiue <229866007@qq.com> Date: Tue, 7 Apr 2026 22:11:52 +0800 Subject: [PATCH 49/55] =?UTF-8?q?fix:=E4=BF=AE=E5=A4=8D=E4=B8=A2=E5=A4=B1?= =?UTF-8?q?=E9=97=AE=E9=A2=98=EF=BC=8C=E5=AE=8C=E5=85=A8=E6=94=B6=E5=8F=A3?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- internal/runtime/runtime.go | 43 +++++-------------------- internal/runtime/runtime_test.go | 4 +-- internal/runtime/workdir_branch_test.go | 23 ------------- internal/session/store.go | 2 +- internal/session/store_test.go | 13 ++++---- 5 files changed, 18 insertions(+), 67 deletions(-) diff --git a/internal/runtime/runtime.go b/internal/runtime/runtime.go index a6352f22..f124ea9b 100644 --- a/internal/runtime/runtime.go +++ b/internal/runtime/runtime.go @@ -109,13 +109,6 @@ func (a *streamAccumulator) buildMessage() (providertypes.Message, error) { return message, nil } -var runtimeSessionWorkdirs = struct { - mu sync.RWMutex - data map[string]string -}{ - data: make(map[string]string), -} - type Runtime interface { Run(ctx context.Context, input UserInput) error Compact(ctx context.Context, input CompactInput) (CompactResult, error) @@ -383,7 +376,6 @@ func (s *Service) LoadSession(ctx context.Context, id string) (agentsession.Sess if err != nil { return agentsession.Session{}, err } - session.Workdir = s.sessionWorkdir(id, session.Workdir) return session, nil } @@ -397,7 +389,6 @@ func (s *Service) SetSessionWorkdir(ctx context.Context, sessionID string, workd if err != nil { return agentsession.Session{}, err } - session.Workdir = s.sessionWorkdir(sessionID, session.Workdir) cfg := s.configManager.Get() resolved, err := resolveWorkdirForSession(cfg.Workdir, session.Workdir, workdir) @@ -409,30 +400,11 @@ func (s *Service) SetSessionWorkdir(ctx context.Context, sessionID string, workd } session.Workdir = resolved - s.setSessionWorkdir(sessionID, resolved) - return session, nil -} - -func (s *Service) sessionWorkdir(sessionID string, fallback string) string { - key := s.sessionWorkdirKey(sessionID) - runtimeSessionWorkdirs.mu.RLock() - value, ok := runtimeSessionWorkdirs.data[key] - runtimeSessionWorkdirs.mu.RUnlock() - if ok { - return strings.TrimSpace(value) + session.UpdatedAt = time.Now() + if err := s.sessionStore.Save(ctx, &session); err != nil { + return agentsession.Session{}, err } - return strings.TrimSpace(fallback) -} - -func (s *Service) setSessionWorkdir(sessionID string, workdir string) { - key := s.sessionWorkdirKey(sessionID) - runtimeSessionWorkdirs.mu.Lock() - runtimeSessionWorkdirs.data[key] = strings.TrimSpace(workdir) - runtimeSessionWorkdirs.mu.Unlock() -} - -func (s *Service) sessionWorkdirKey(sessionID string) string { - return fmt.Sprintf("%p:%s", s, strings.TrimSpace(sessionID)) + return session, nil } func (s *Service) loadOrCreateSession( @@ -448,7 +420,6 @@ func (s *Service) loadOrCreateSession( return agentsession.Session{}, err } session := agentsession.NewWithWorkdir(title, sessionWorkdir) - s.setSessionWorkdir(session.ID, sessionWorkdir) if err := s.sessionStore.Save(ctx, &session); err != nil { return agentsession.Session{}, err } @@ -458,7 +429,6 @@ func (s *Service) loadOrCreateSession( if err != nil { return agentsession.Session{}, err } - session.Workdir = s.sessionWorkdir(sessionID, session.Workdir) if strings.TrimSpace(requestedWorkdir) == "" && strings.TrimSpace(session.Workdir) != "" { return session, nil } @@ -471,7 +441,10 @@ func (s *Service) loadOrCreateSession( return session, nil } session.Workdir = resolved - s.setSessionWorkdir(sessionID, resolved) + session.UpdatedAt = time.Now() + if err := s.sessionStore.Save(ctx, &session); err != nil { + return agentsession.Session{}, err + } return session, nil } diff --git a/internal/runtime/runtime_test.go b/internal/runtime/runtime_test.go index 4ed9c921..988f6201 100644 --- a/internal/runtime/runtime_test.go +++ b/internal/runtime/runtime_test.go @@ -2445,8 +2445,8 @@ func TestServiceSetSessionWorkdir(t *testing.T) { if err != nil { t.Fatalf("LoadSession() with new service error = %v", err) } - if strings.TrimSpace(reloaded.Workdir) != "" { - t.Fatalf("expected session workdir not to persist across process lifetime, got %q", reloaded.Workdir) + if reloaded.Workdir != target { + t.Fatalf("expected session workdir to persist across service lifetime, got %q", reloaded.Workdir) } _, err = service.SetSessionWorkdir(context.Background(), "", "sub") diff --git a/internal/runtime/workdir_branch_test.go b/internal/runtime/workdir_branch_test.go index 266fed5e..089a3cd0 100644 --- a/internal/runtime/workdir_branch_test.go +++ b/internal/runtime/workdir_branch_test.go @@ -10,29 +10,6 @@ import ( agentsession "neo-code/internal/session" ) -func TestSessionWorkdirKeyAndMemoryMap(t *testing.T) { - t.Parallel() - - serviceA := &Service{} - serviceB := &Service{} - - keyA := serviceA.sessionWorkdirKey("session-1") - keyB := serviceB.sessionWorkdirKey("session-1") - if keyA == keyB { - t.Fatalf("expected unique key per service instance, got %q", keyA) - } - - if got := serviceA.sessionWorkdir("session-1", "/fallback"); got != "/fallback" { - t.Fatalf("expected fallback workdir, got %q", got) - } - - target := t.TempDir() - serviceA.setSessionWorkdir("session-1", target) - if got := serviceA.sessionWorkdir("session-1", "/fallback"); got != target { - t.Fatalf("expected mapped workdir %q, got %q", target, got) - } -} - func TestResolveWorkdirForSessionAndNormalizeErrors(t *testing.T) { t.Parallel() diff --git a/internal/session/store.go b/internal/session/store.go index 824a8228..57639374 100644 --- a/internal/session/store.go +++ b/internal/session/store.go @@ -28,7 +28,7 @@ type Session struct { Model string `json:"model,omitempty"` CreatedAt time.Time `json:"created_at"` UpdatedAt time.Time `json:"updated_at"` - Workdir string `json:"-"` + Workdir string `json:"workdir,omitempty"` Messages []providertypes.Message `json:"messages"` } diff --git a/internal/session/store_test.go b/internal/session/store_test.go index 2334d244..6b412322 100644 --- a/internal/session/store_test.go +++ b/internal/session/store_test.go @@ -54,8 +54,8 @@ func TestJSONStoreSaveLoadAndListSummaries(t *testing.T) { if loaded.Title != older.Title { t.Fatalf("expected title %q, got %q", older.Title, loaded.Title) } - if loaded.Workdir != "" { - t.Fatalf("expected workdir to stay in-memory only, got %q", loaded.Workdir) + if loaded.Workdir != older.Workdir { + t.Fatalf("expected persisted workdir %q, got %q", older.Workdir, loaded.Workdir) } if len(loaded.Messages) != 2 || loaded.Messages[1].Content != "world" { t.Fatalf("unexpected loaded messages: %+v", loaded.Messages) @@ -66,8 +66,8 @@ func TestJSONStoreSaveLoadAndListSummaries(t *testing.T) { if err != nil { t.Fatalf("read saved session: %v", err) } - if strings.Contains(string(raw), "\"workdir\"") { - t.Fatalf("expected persisted session file to exclude workdir, got:\n%s", string(raw)) + if !strings.Contains(string(raw), "\"workdir\"") { + t.Fatalf("expected persisted session file to include workdir, got:\n%s", string(raw)) } mustWriteSessionFile(t, filepath.Join(baseDir, sessionsDirName, "invalid.json"), "{invalid") @@ -325,6 +325,7 @@ func TestJSONStoreSavePersistsProviderModelAndMessages(t *testing.T) { Title: "Persist Fields", Provider: "openai", Model: "gpt-4.1", + Workdir: "/tmp/persist-workdir", CreatedAt: time.Now().Add(-time.Hour), UpdatedAt: time.Now(), Messages: []providertypes.Message{ @@ -364,8 +365,8 @@ func TestJSONStoreSavePersistsProviderModelAndMessages(t *testing.T) { if _, ok := decoded["messages"]; !ok { t.Fatalf("expected messages field persisted, got %+v", decoded) } - if _, ok := decoded["workdir"]; ok { - t.Fatalf("expected workdir not persisted, got %+v", decoded) + if decoded["workdir"] != session.Workdir { + t.Fatalf("expected workdir persisted as %q, got %+v", session.Workdir, decoded["workdir"]) } } From 9265d9f8f44a4379e02d337c6c3782fb54f899fc Mon Sep 17 00:00:00 2001 From: creatang Date: Mon, 6 Apr 2026 14:09:18 +0800 Subject: [PATCH 50/55] refactor(tui): add core command status and workspace helpers --- internal/tui/core/commands/parser.go | 77 +++++++++++++ internal/tui/core/commands/parser_test.go | 66 ++++++++++++ internal/tui/core/commands/workspace.go | 70 ++++++++++++ internal/tui/core/commands/workspace_test.go | 107 +++++++++++++++++++ internal/tui/core/status/snapshot.go | 104 ++++++++++++++++++ internal/tui/core/utils/view_helpers.go | 93 ++++++++++++++++ internal/tui/core/workspace/resolver.go | 50 +++++++++ 7 files changed, 567 insertions(+) create mode 100644 internal/tui/core/commands/parser.go create mode 100644 internal/tui/core/commands/parser_test.go create mode 100644 internal/tui/core/commands/workspace.go create mode 100644 internal/tui/core/commands/workspace_test.go create mode 100644 internal/tui/core/status/snapshot.go create mode 100644 internal/tui/core/utils/view_helpers.go create mode 100644 internal/tui/core/workspace/resolver.go diff --git a/internal/tui/core/commands/parser.go b/internal/tui/core/commands/parser.go new file mode 100644 index 00000000..1ffac80e --- /dev/null +++ b/internal/tui/core/commands/parser.go @@ -0,0 +1,77 @@ +package commands + +import ( + "fmt" + "strings" +) + +// SlashCommand 描述单个 slash 命令定义。 +type SlashCommand struct { + Usage string + Description string +} + +// CommandSuggestion 表示输入匹配后的命令建议。 +type CommandSuggestion struct { + Command SlashCommand + Match bool +} + +// MatchSlashCommands 根据输入匹配可展示的 slash 命令建议。 +func MatchSlashCommands(input string, slashPrefix string, commands []SlashCommand) []CommandSuggestion { + if !strings.HasPrefix(input, slashPrefix) { + return nil + } + + query := strings.ToLower(strings.TrimSpace(input)) + if IsCompleteSlashCommand(query, commands) { + return nil + } + out := make([]CommandSuggestion, 0, len(commands)) + for _, command := range commands { + normalized := strings.ToLower(command.Usage) + match := query == slashPrefix || strings.HasPrefix(normalized, query) + if query == slashPrefix || match || strings.Contains(normalized, query) { + out = append(out, CommandSuggestion{Command: command, Match: match}) + } + } + return out +} + +// IsCompleteSlashCommand 判断输入是否已完整匹配某个命令。 +func IsCompleteSlashCommand(input string, commands []SlashCommand) bool { + for _, command := range commands { + if strings.EqualFold(strings.TrimSpace(command.Usage), strings.TrimSpace(input)) { + return true + } + } + return false +} + +// SplitFirstWord 拆分首个 token 与其后续参数。 +func SplitFirstWord(input string) (string, string) { + input = strings.TrimSpace(input) + if input == "" { + return "", "" + } + index := strings.IndexAny(input, " \t") + if index < 0 { + return input, "" + } + return input[:index], strings.TrimSpace(input[index+1:]) +} + +// IsWorkspaceSlashCommand 判断是否为工作区命令(例如 /cwd)。 +func IsWorkspaceSlashCommand(raw string, commandName string) bool { + command, _ := SplitFirstWord(strings.ToLower(strings.TrimSpace(raw))) + return command == strings.ToLower(strings.TrimSpace(commandName)) +} + +// ParseWorkspaceSlashCommand 解析工作区命令参数,非目标命令时返回错误。 +func ParseWorkspaceSlashCommand(raw string, commandName string) (string, error) { + command, args := SplitFirstWord(strings.TrimSpace(raw)) + if strings.ToLower(command) != strings.ToLower(strings.TrimSpace(commandName)) { + return "", fmt.Errorf("unknown command %q", command) + } + return strings.TrimSpace(args), nil +} diff --git a/internal/tui/core/commands/parser_test.go b/internal/tui/core/commands/parser_test.go new file mode 100644 index 00000000..a43915c4 --- /dev/null +++ b/internal/tui/core/commands/parser_test.go @@ -0,0 +1,66 @@ +package commands + +import "testing" + +func TestMatchSlashCommands(t *testing.T) { + commands := []SlashCommand{ + {Usage: "/help", Description: "show help"}, + {Usage: "/provider", Description: "pick provider"}, + {Usage: "/model", Description: "pick model"}, + } + + got := MatchSlashCommands("/pro", "/", commands) + if len(got) != 1 { + t.Fatalf("expected one suggestion for /pro, got %d", len(got)) + } + if got[0].Command.Usage != "/provider" || !got[0].Match { + t.Fatalf("unexpected suggestion: %+v", got[0]) + } + + if complete := MatchSlashCommands("/help", "/", commands); complete != nil { + t.Fatalf("expected nil suggestion when command is complete, got %+v", complete) + } +} + +func TestIsCompleteSlashCommand(t *testing.T) { + commands := []SlashCommand{{Usage: "/help"}, {Usage: "/provider"}} + if !IsCompleteSlashCommand("/help", commands) { + t.Fatalf("expected /help to be complete") + } + if IsCompleteSlashCommand("/hel", commands) { + t.Fatalf("expected /hel to be incomplete") + } +} + +func TestSplitFirstWord(t *testing.T) { + first, rest := SplitFirstWord(" /cwd ./tmp/project ") + if first != "/cwd" || rest != "./tmp/project" { + t.Fatalf("unexpected split result: first=%q rest=%q", first, rest) + } + + first, rest = SplitFirstWord(" ") + if first != "" || rest != "" { + t.Fatalf("expected empty split for blank input, got first=%q rest=%q", first, rest) + } +} + +func TestWorkspaceSlashCommandHelpers(t *testing.T) { + if !IsWorkspaceSlashCommand("/cwd ./tmp", "/cwd") { + t.Fatalf("expected /cwd to be recognized") + } + if IsWorkspaceSlashCommand("/status", "/cwd") { + t.Fatalf("did not expect /status as workspace command") + } + + args, err := ParseWorkspaceSlashCommand("/cwd ./tmp", "/cwd") + if err != nil { + t.Fatalf("ParseWorkspaceSlashCommand() error = %v", err) + } + if args != "./tmp" { + t.Fatalf("expected args ./tmp, got %q", args) + } + + if _, err := ParseWorkspaceSlashCommand("/status", "/cwd"); err == nil { + t.Fatalf("expected parse error for non-workspace command") + } +} diff --git a/internal/tui/core/commands/workspace.go b/internal/tui/core/commands/workspace.go new file mode 100644 index 00000000..0886164d --- /dev/null +++ b/internal/tui/core/commands/workspace.go @@ -0,0 +1,70 @@ +package commands + +import ( + "context" + "fmt" + "strings" + + agentruntime "neo-code/internal/runtime" +) + +// SessionWorkdirSetter 定义设置会话工作目录所需的最小 runtime 能力。 +type SessionWorkdirSetter interface { + SetSessionWorkdir(ctx context.Context, sessionID string, workdir string) (agentruntime.Session, error) +} + +// SessionWorkdirCommandResult 表示工作目录命令执行结果。 +type SessionWorkdirCommandResult struct { + Notice string + Workdir string + Err error +} + +// ExecuteSessionWorkdirCommand 执行 /cwd 命令的核心流程,返回统一结果结构。 +func ExecuteSessionWorkdirCommand( + runtime SessionWorkdirSetter, + sessionID string, + currentWorkdir string, + raw string, + parseCommand func(string) (string, error), + resolveWorkspacePath func(string, string) (string, error), + selectSessionWorkdir func(string, string) string, +) SessionWorkdirCommandResult { + requested, err := parseCommand(raw) + if err != nil { + return SessionWorkdirCommandResult{Err: err} + } + + if strings.TrimSpace(requested) == "" { + workdir := strings.TrimSpace(currentWorkdir) + if workdir == "" { + return SessionWorkdirCommandResult{Err: fmt.Errorf("usage: /cwd ")} + } + return SessionWorkdirCommandResult{ + Notice: fmt.Sprintf("[System] Current workspace is %s.", workdir), + Workdir: workdir, + } + } + + if strings.TrimSpace(sessionID) == "" { + workdir, err := resolveWorkspacePath(currentWorkdir, requested) + if err != nil { + return SessionWorkdirCommandResult{Err: err} + } + return SessionWorkdirCommandResult{ + Notice: fmt.Sprintf("[System] Draft workspace switched to %s.", workdir), + Workdir: workdir, + } + } + + session, err := runtime.SetSessionWorkdir(context.Background(), sessionID, requested) + if err != nil { + return SessionWorkdirCommandResult{Err: err} + } + + workdir := selectSessionWorkdir(session.Workdir, currentWorkdir) + return SessionWorkdirCommandResult{ + Notice: fmt.Sprintf("[System] Session workspace switched to %s.", workdir), + Workdir: workdir, + } +} diff --git a/internal/tui/core/commands/workspace_test.go b/internal/tui/core/commands/workspace_test.go new file mode 100644 index 00000000..53650046 --- /dev/null +++ b/internal/tui/core/commands/workspace_test.go @@ -0,0 +1,107 @@ +package commands + +import ( + "context" + "errors" + "os" + "path/filepath" + "strings" + "testing" + + agentruntime "neo-code/internal/runtime" + tuiworkspace "neo-code/internal/tui/core/workspace" +) + +type stubSessionWorkdirSetter struct { + session agentruntime.Session + err error + calls int +} + +func (s *stubSessionWorkdirSetter) SetSessionWorkdir(ctx context.Context, sessionID string, workdir string) (agentruntime.Session, error) { + s.calls++ + if s.err != nil { + return agentruntime.Session{}, s.err + } + return s.session, nil +} + +func TestExecuteSessionWorkdirCommand(t *testing.T) { + parse := func(raw string) (string, error) { + raw = strings.TrimSpace(raw) + if raw == "/bad" { + return "", errors.New("unknown command") + } + if raw == "/cwd" { + return "", nil + } + if strings.HasPrefix(raw, "/cwd ") { + return strings.TrimSpace(strings.TrimPrefix(raw, "/cwd ")), nil + } + return "", errors.New("unknown command") + } + + t.Run("parse error", func(t *testing.T) { + result := ExecuteSessionWorkdirCommand(&stubSessionWorkdirSetter{}, "", "", "/bad", parse, tuiworkspace.ResolveWorkspacePath, tuiworkspace.SelectSessionWorkdir) + if result.Err == nil { + t.Fatalf("expected parse error") + } + }) + + t.Run("empty requested without current workdir", func(t *testing.T) { + result := ExecuteSessionWorkdirCommand(&stubSessionWorkdirSetter{}, "", "", "/cwd", parse, tuiworkspace.ResolveWorkspacePath, tuiworkspace.SelectSessionWorkdir) + if result.Err == nil || !strings.Contains(result.Err.Error(), "usage: /cwd ") { + t.Fatalf("expected usage error, got %+v", result) + } + }) + + t.Run("empty requested with current workdir", func(t *testing.T) { + current := t.TempDir() + result := ExecuteSessionWorkdirCommand(&stubSessionWorkdirSetter{}, "", current, "/cwd", parse, tuiworkspace.ResolveWorkspacePath, tuiworkspace.SelectSessionWorkdir) + if result.Err != nil { + t.Fatalf("unexpected error: %v", result.Err) + } + if result.Workdir != current || !strings.Contains(result.Notice, "Current workspace is") { + t.Fatalf("unexpected result: %+v", result) + } + }) + + t.Run("draft session resolves requested path", func(t *testing.T) { + base := t.TempDir() + target := filepath.Join(base, "sub") + if err := ensureDir(target); err != nil { + t.Fatalf("mkdir target: %v", err) + } + result := ExecuteSessionWorkdirCommand(&stubSessionWorkdirSetter{}, "", base, "/cwd sub", parse, tuiworkspace.ResolveWorkspacePath, tuiworkspace.SelectSessionWorkdir) + if result.Err != nil { + t.Fatalf("unexpected error: %v", result.Err) + } + if !strings.Contains(result.Notice, "Draft workspace switched") { + t.Fatalf("unexpected notice: %q", result.Notice) + } + }) + + t.Run("runtime error", func(t *testing.T) { + stub := &stubSessionWorkdirSetter{err: errors.New("set workdir failed")} + result := ExecuteSessionWorkdirCommand(stub, "session-1", t.TempDir(), "/cwd sub", parse, tuiworkspace.ResolveWorkspacePath, tuiworkspace.SelectSessionWorkdir) + if result.Err == nil || !strings.Contains(result.Err.Error(), "set workdir failed") { + t.Fatalf("expected runtime error, got %+v", result) + } + }) + + t.Run("runtime empty workdir fallback", func(t *testing.T) { + current := t.TempDir() + stub := &stubSessionWorkdirSetter{session: agentruntime.Session{ID: "session-1", Workdir: ""}} + result := ExecuteSessionWorkdirCommand(stub, "session-1", current, "/cwd sub", parse, tuiworkspace.ResolveWorkspacePath, tuiworkspace.SelectSessionWorkdir) + if result.Err != nil { + t.Fatalf("unexpected error: %v", result.Err) + } + if result.Workdir != current { + t.Fatalf("expected fallback workdir %q, got %q", current, result.Workdir) + } + }) +} + +func ensureDir(path string) error { + return os.MkdirAll(path, 0o755) +} diff --git a/internal/tui/core/status/snapshot.go b/internal/tui/core/status/snapshot.go new file mode 100644 index 00000000..6444f86c --- /dev/null +++ b/internal/tui/core/status/snapshot.go @@ -0,0 +1,104 @@ +package status + +import ( + "fmt" + "strings" + + tuiutils "neo-code/internal/tui/core/utils" + tuistate "neo-code/internal/tui/state" +) + +// Snapshot 表示 /status 命令所需的界面状态快照。 +type Snapshot struct { + ActiveSessionID string + ActiveSessionTitle string + ActiveRunID string + IsAgentRunning bool + IsCompacting bool + CurrentProvider string + CurrentModel string + CurrentWorkdir string + CurrentTool string + ToolStateCount int + RunTotalTokens int + SessionTotalTokens int + ExecutionError string + FocusLabel string + PickerLabel string + MessageCount int +} + +// BuildFromUIState 根据 UIState 与附加上下文构建 /status 所需快照。 +func BuildFromUIState( + state tuistate.UIState, + messageCount int, + focusLabel string, + pickerLabel string, +) Snapshot { + return Snapshot{ + ActiveSessionID: state.ActiveSessionID, + ActiveSessionTitle: state.ActiveSessionTitle, + ActiveRunID: state.ActiveRunID, + IsAgentRunning: state.IsAgentRunning, + IsCompacting: state.IsCompacting, + CurrentProvider: state.CurrentProvider, + CurrentModel: state.CurrentModel, + CurrentWorkdir: state.CurrentWorkdir, + CurrentTool: state.CurrentTool, + ToolStateCount: len(state.ToolStates), + RunTotalTokens: state.TokenUsage.RunTotalTokens, + SessionTotalTokens: state.TokenUsage.SessionTotalTokens, + ExecutionError: state.ExecutionError, + FocusLabel: focusLabel, + PickerLabel: pickerLabel, + MessageCount: messageCount, + } +} + +// Format 将状态快照格式化为多行文本,用于 /status 命令输出。 +func Format(snapshot Snapshot, draftSessionTitle string) string { + sessionID := snapshot.ActiveSessionID + if strings.TrimSpace(sessionID) == "" { + sessionID = "" + } + sessionTitle := snapshot.ActiveSessionTitle + if strings.TrimSpace(sessionTitle) == "" { + sessionTitle = draftSessionTitle + } + running := "no" + if snapshot.IsAgentRunning || snapshot.IsCompacting { + running = "yes" + } + currentTool := snapshot.CurrentTool + if strings.TrimSpace(currentTool) == "" { + currentTool = "" + } + errorText := snapshot.ExecutionError + if strings.TrimSpace(errorText) == "" { + errorText = "" + } + picker := snapshot.PickerLabel + if strings.TrimSpace(picker) == "" { + picker = "none" + } + + lines := []string{ + "Status:", + "Session: " + sessionTitle, + "Session ID: " + sessionID, + "Run ID: " + tuiutils.Fallback(strings.TrimSpace(snapshot.ActiveRunID), ""), + "Running: " + running, + "Provider: " + snapshot.CurrentProvider, + "Model: " + snapshot.CurrentModel, + "Workdir: " + snapshot.CurrentWorkdir, + "Focus: " + snapshot.FocusLabel, + "Picker: " + picker, + "Current Tool: " + currentTool, + fmt.Sprintf("Tool States: %d", snapshot.ToolStateCount), + fmt.Sprintf("Run Tokens: %d", snapshot.RunTotalTokens), + fmt.Sprintf("Session Tokens: %d", snapshot.SessionTotalTokens), + fmt.Sprintf("Messages: %d", snapshot.MessageCount), + "Error: " + errorText, + } + return strings.Join(lines, "\n") +} diff --git a/internal/tui/core/utils/view_helpers.go b/internal/tui/core/utils/view_helpers.go new file mode 100644 index 00000000..4c0d029a --- /dev/null +++ b/internal/tui/core/utils/view_helpers.go @@ -0,0 +1,93 @@ +package utils + +import ( + "strings" + + tuistate "neo-code/internal/tui/state" +) + +// PickerLabelFromMode 将 picker 模式映射为状态快照展示标签。 +func PickerLabelFromMode(mode tuistate.PickerMode) string { + switch mode { + case tuistate.PickerProvider: + return "provider" + case tuistate.PickerModel: + return "model" + case tuistate.PickerFile: + return "file" + default: + return "none" + } +} + +// RequestedWorkdirForRun 在发起 run 时计算应转发的工作目录。 +func RequestedWorkdirForRun(activeSessionID string, currentWorkdir string) string { + if strings.TrimSpace(activeSessionID) == "" { + return currentWorkdir + } + return "" +} + +// IsBusy 统一判断当前是否存在进行中的 agent 或 compact 操作。 +func IsBusy(isAgentRunning bool, isCompacting bool) bool { + return isAgentRunning || isCompacting +} + +// FocusLabelFromPanel 将焦点面板枚举映射为界面展示标签。 +func FocusLabelFromPanel( + focus tuistate.Panel, + sessionsLabel string, + transcriptLabel string, + activityLabel string, + composerLabel string, +) string { + switch focus { + case tuistate.PanelSessions: + return sessionsLabel + case tuistate.PanelTranscript: + return transcriptLabel + case tuistate.PanelActivity: + return activityLabel + default: + return composerLabel + } +} + +// TrimRunes 按 rune 数裁剪文本,超长时尾部追加省略号。 +func TrimRunes(text string, limit int) string { + runes := []rune(text) + if len(runes) <= limit || limit < 4 { + return text + } + return string(runes[:limit-3]) + "..." +} + +// TrimMiddle 在中间裁剪长文本,保留首尾并插入省略号。 +func TrimMiddle(text string, limit int) string { + runes := []rune(text) + if len(runes) <= limit || limit < 7 { + return text + } + left := (limit - 3) / 2 + right := limit - 3 - left + return string(runes[:left]) + "..." + string(runes[len(runes)-right:]) +} + +// Fallback 当 value 为空白文本时返回 fallbackValue。 +func Fallback(value string, fallbackValue string) string { + if strings.TrimSpace(value) == "" { + return fallbackValue + } + return value +} + +// Clamp 将数值限制在 [minValue, maxValue] 范围内。 +func Clamp(value int, minValue int, maxValue int) int { + if value < minValue { + return minValue + } + if value > maxValue { + return maxValue + } + return value +} diff --git a/internal/tui/core/workspace/resolver.go b/internal/tui/core/workspace/resolver.go new file mode 100644 index 00000000..ccf6c103 --- /dev/null +++ b/internal/tui/core/workspace/resolver.go @@ -0,0 +1,50 @@ +package workspace + +import ( + "fmt" + "os" + "path/filepath" + "strings" +) + +// ResolveWorkspacePath 解析并校验工作区路径,确保返回存在且可用的目录绝对路径。 +func ResolveWorkspacePath(base string, requested string) (string, error) { + base = strings.TrimSpace(base) + if base == "" { + workingDir, err := os.Getwd() + if err != nil { + return "", fmt.Errorf("workspace: resolve current directory: %w", err) + } + base = workingDir + } + + target := strings.TrimSpace(requested) + if target == "" { + target = "." + } + if !filepath.IsAbs(target) { + target = filepath.Join(base, target) + } + + absolute, err := filepath.Abs(target) + if err != nil { + return "", fmt.Errorf("workspace: resolve path: %w", err) + } + info, err := os.Stat(absolute) + if err != nil { + return "", fmt.Errorf("workspace: resolve path: %w", err) + } + if !info.IsDir() { + return "", fmt.Errorf("workspace: %q is not a directory", absolute) + } + return filepath.Clean(absolute), nil +} + +// SelectSessionWorkdir 优先返回会话工作目录,缺失时回退到默认工作目录。 +func SelectSessionWorkdir(sessionWorkdir string, defaultWorkdir string) string { + workdir := strings.TrimSpace(sessionWorkdir) + if workdir != "" { + return workdir + } + return strings.TrimSpace(defaultWorkdir) +} From ace9a42f793aca3f5669c3f8280343d170eb7b31 Mon Sep 17 00:00:00 2001 From: creatang Date: Tue, 7 Apr 2026 10:38:57 +0800 Subject: [PATCH 51/55] fix(tui): remove BOM from runtime bridge and status snapshot --- internal/tui/core/status/snapshot.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/internal/tui/core/status/snapshot.go b/internal/tui/core/status/snapshot.go index 6444f86c..3f75baf9 100644 --- a/internal/tui/core/status/snapshot.go +++ b/internal/tui/core/status/snapshot.go @@ -1,4 +1,4 @@ -package status +package status import ( "fmt" From b044039fc6908bc36cf05f3fdecdf051334954c1 Mon Sep 17 00:00:00 2001 From: creatang Date: Tue, 7 Apr 2026 18:19:45 +0800 Subject: [PATCH 52/55] test(tui/core): add coverage for status, utils and workspace --- internal/tui/components/components_test.go | 36 ------- internal/tui/core/status/snapshot_test.go | 100 +++++++++++++++++++ internal/tui/core/utils/view_helpers_test.go | 97 ++++++++++++++++++ internal/tui/core/workspace/resolver_test.go | 57 +++++++++++ 4 files changed, 254 insertions(+), 36 deletions(-) create mode 100644 internal/tui/core/status/snapshot_test.go create mode 100644 internal/tui/core/utils/view_helpers_test.go create mode 100644 internal/tui/core/workspace/resolver_test.go diff --git a/internal/tui/components/components_test.go b/internal/tui/components/components_test.go index 51482565..07e62549 100644 --- a/internal/tui/components/components_test.go +++ b/internal/tui/components/components_test.go @@ -138,42 +138,6 @@ func TestRenderSessionRow(t *testing.T) { } } -func TestViewHelperBranches(t *testing.T) { - if got := fallback("primary", "fallback"); got != "primary" { - t.Fatalf("expected primary fallback value, got %q", got) - } - if got := fallback("", "fallback"); got != "fallback" { - t.Fatalf("expected fallback value, got %q", got) - } - - if got := trimMiddle("abcdef", 0); got != "" { - t.Fatalf("expected empty string for non-positive limit, got %q", got) - } - if got := trimMiddle("abcdef", 3); got != "abc" { - t.Fatalf("expected hard truncate for short limit, got %q", got) - } - if got := trimMiddle("abcdefghij", 7); got != "ab...ij" { - t.Fatalf("expected middle trim output, got %q", got) - } - - if got := trimRunes("abcdef", 3); got != "abcdef" { - t.Fatalf("expected original text when limit < 4, got %q", got) - } - if got := trimRunes("abcdef", 5); got != "ab..." { - t.Fatalf("expected rune-safe ellipsis trim, got %q", got) - } - - if got := clamp(-1, 0, 10); got != 0 { - t.Fatalf("expected clamp to min, got %d", got) - } - if got := clamp(20, 0, 10); got != 10 { - t.Fatalf("expected clamp to max, got %d", got) - } - if got := clamp(6, 0, 10); got != 6 { - t.Fatalf("expected clamp to keep in-range value, got %d", got) - } -} - func TestNormalizeBlockRightEdgeBlankContent(t *testing.T) { blank := " \n\t" if got := NormalizeBlockRightEdge(blank, 20); got != blank { diff --git a/internal/tui/core/status/snapshot_test.go b/internal/tui/core/status/snapshot_test.go new file mode 100644 index 00000000..1b618833 --- /dev/null +++ b/internal/tui/core/status/snapshot_test.go @@ -0,0 +1,100 @@ +package status + +import ( + "strings" + "testing" + + tuistate "neo-code/internal/tui/state" +) + +func TestBuildFromUIState(t *testing.T) { + state := tuistate.UIState{ + ActiveSessionID: "session-1", + ActiveSessionTitle: "My Session", + ActiveRunID: "run-1", + IsAgentRunning: true, + IsCompacting: false, + CurrentProvider: "openai", + CurrentModel: "gpt-5.4", + CurrentWorkdir: "/repo", + CurrentTool: "filesystem_read_file", + ToolStates: []tuistate.ToolState{ + {ToolCallID: "call-1"}, + {ToolCallID: "call-2"}, + }, + TokenUsage: tuistate.TokenUsageState{ + RunTotalTokens: 12, + SessionTotalTokens: 34, + }, + ExecutionError: "boom", + } + + snapshot := BuildFromUIState(state, 7, "transcript", "provider") + if snapshot.ActiveSessionID != "session-1" || snapshot.ActiveRunID != "run-1" { + t.Fatalf("unexpected snapshot identifiers: %+v", snapshot) + } + if snapshot.ToolStateCount != 2 || snapshot.RunTotalTokens != 12 || snapshot.SessionTotalTokens != 34 { + t.Fatalf("unexpected snapshot counters: %+v", snapshot) + } + if snapshot.FocusLabel != "transcript" || snapshot.PickerLabel != "provider" || snapshot.MessageCount != 7 { + t.Fatalf("unexpected snapshot labels: %+v", snapshot) + } +} + +func TestFormat(t *testing.T) { + formatted := Format(Snapshot{ + ActiveSessionID: "", + ActiveSessionTitle: "", + ActiveRunID: " ", + IsAgentRunning: false, + IsCompacting: false, + CurrentProvider: "openai", + CurrentModel: "gpt-5.4", + CurrentWorkdir: "/repo", + CurrentTool: "", + ToolStateCount: 1, + RunTotalTokens: 2, + SessionTotalTokens: 3, + ExecutionError: "", + FocusLabel: "composer", + PickerLabel: "", + MessageCount: 4, + }, "Draft Session") + + expectedParts := []string{ + "Session: Draft Session", + "Session ID: ", + "Run ID: ", + "Running: no", + "Picker: none", + "Current Tool: ", + "Error: ", + } + for _, part := range expectedParts { + if !strings.Contains(formatted, part) { + t.Fatalf("expected formatted status to contain %q, got:\n%s", part, formatted) + } + } + + running := Format(Snapshot{ + ActiveSessionID: "session-2", + ActiveSessionTitle: "Named Session", + ActiveRunID: "run-2", + IsCompacting: true, + CurrentProvider: "openai", + CurrentModel: "gpt-5.4-mini", + CurrentWorkdir: "/repo", + CurrentTool: "tool-x", + ToolStateCount: 2, + RunTotalTokens: 10, + SessionTotalTokens: 20, + ExecutionError: "failed", + FocusLabel: "activity", + PickerLabel: "model", + MessageCount: 5, + }, "Ignored Draft") + + if !strings.Contains(running, "Session: Named Session") || !strings.Contains(running, "Running: yes") { + t.Fatalf("expected running status to keep explicit values, got:\n%s", running) + } +} diff --git a/internal/tui/core/utils/view_helpers_test.go b/internal/tui/core/utils/view_helpers_test.go new file mode 100644 index 00000000..5a342e06 --- /dev/null +++ b/internal/tui/core/utils/view_helpers_test.go @@ -0,0 +1,97 @@ +package utils + +import ( + "testing" + + tuistate "neo-code/internal/tui/state" +) + +func TestPickerLabelFromMode(t *testing.T) { + if got := PickerLabelFromMode(tuistate.PickerProvider); got != "provider" { + t.Fatalf("expected provider label, got %q", got) + } + if got := PickerLabelFromMode(tuistate.PickerModel); got != "model" { + t.Fatalf("expected model label, got %q", got) + } + if got := PickerLabelFromMode(tuistate.PickerFile); got != "file" { + t.Fatalf("expected file label, got %q", got) + } + if got := PickerLabelFromMode(tuistate.PickerMode(99)); got != "none" { + t.Fatalf("expected default picker label none, got %q", got) + } +} + +func TestRequestedWorkdirForRun(t *testing.T) { + if got := RequestedWorkdirForRun("", "/repo"); got != "/repo" { + t.Fatalf("expected current workdir when active session is blank, got %q", got) + } + if got := RequestedWorkdirForRun("session-1", "/repo"); got != "" { + t.Fatalf("expected empty requested workdir when active session exists, got %q", got) + } +} + +func TestIsBusy(t *testing.T) { + if IsBusy(false, false) { + t.Fatalf("expected idle state") + } + if !IsBusy(true, false) || !IsBusy(false, true) || !IsBusy(true, true) { + t.Fatalf("expected busy state when any operation is running") + } +} + +func TestFocusLabelFromPanel(t *testing.T) { + const ( + sessions = "Sessions" + transcript = "Transcript" + activity = "Activity" + composer = "Composer" + ) + + if got := FocusLabelFromPanel(tuistate.PanelSessions, sessions, transcript, activity, composer); got != sessions { + t.Fatalf("expected sessions label, got %q", got) + } + if got := FocusLabelFromPanel(tuistate.PanelTranscript, sessions, transcript, activity, composer); got != transcript { + t.Fatalf("expected transcript label, got %q", got) + } + if got := FocusLabelFromPanel(tuistate.PanelActivity, sessions, transcript, activity, composer); got != activity { + t.Fatalf("expected activity label, got %q", got) + } + if got := FocusLabelFromPanel(tuistate.PanelInput, sessions, transcript, activity, composer); got != composer { + t.Fatalf("expected composer label, got %q", got) + } +} + +func TestTrimHelpers(t *testing.T) { + if got := TrimRunes("abcdef", 3); got != "abcdef" { + t.Fatalf("expected original text when limit < 4, got %q", got) + } + if got := TrimRunes("abcdef", 5); got != "ab..." { + t.Fatalf("expected rune-safe truncation, got %q", got) + } + + if got := TrimMiddle("abcdef", 6); got != "abcdef" { + t.Fatalf("expected no trim when limit < 7, got %q", got) + } + if got := TrimMiddle("abcdefghij", 7); got != "ab...ij" { + t.Fatalf("expected middle trim output, got %q", got) + } +} + +func TestFallbackAndClamp(t *testing.T) { + if got := Fallback("value", "fallback"); got != "value" { + t.Fatalf("expected value when non-empty, got %q", got) + } + if got := Fallback(" ", "fallback"); got != "fallback" { + t.Fatalf("expected fallback for blank value, got %q", got) + } + + if got := Clamp(-1, 0, 10); got != 0 { + t.Fatalf("expected clamp to min, got %d", got) + } + if got := Clamp(11, 0, 10); got != 10 { + t.Fatalf("expected clamp to max, got %d", got) + } + if got := Clamp(5, 0, 10); got != 5 { + t.Fatalf("expected in-range value unchanged, got %d", got) + } +} diff --git a/internal/tui/core/workspace/resolver_test.go b/internal/tui/core/workspace/resolver_test.go new file mode 100644 index 00000000..775a9ba8 --- /dev/null +++ b/internal/tui/core/workspace/resolver_test.go @@ -0,0 +1,57 @@ +package workspace + +import ( + "os" + "path/filepath" + "strings" + "testing" +) + +func TestResolveWorkspacePath(t *testing.T) { + base := t.TempDir() + childDir := filepath.Join(base, "project") + if err := os.MkdirAll(childDir, 0o755); err != nil { + t.Fatalf("mkdir child dir: %v", err) + } + + resolved, err := ResolveWorkspacePath(base, "project") + if err != nil { + t.Fatalf("ResolveWorkspacePath(relative) error = %v", err) + } + if resolved != filepath.Clean(childDir) { + t.Fatalf("unexpected resolved path: %q", resolved) + } + + resolved, err = ResolveWorkspacePath(base, "") + if err != nil { + t.Fatalf("ResolveWorkspacePath(default current) error = %v", err) + } + if resolved != filepath.Clean(base) { + t.Fatalf("expected base directory for empty requested path, got %q", resolved) + } +} + +func TestResolveWorkspacePathErrors(t *testing.T) { + base := t.TempDir() + filePath := filepath.Join(base, "not-dir.txt") + if err := os.WriteFile(filePath, []byte("x"), 0o644); err != nil { + t.Fatalf("write file: %v", err) + } + + if _, err := ResolveWorkspacePath(base, "missing-dir"); err == nil { + t.Fatalf("expected missing path to return error") + } + + if _, err := ResolveWorkspacePath(base, "not-dir.txt"); err == nil || !strings.Contains(err.Error(), "not a directory") { + t.Fatalf("expected non-directory path error, got %v", err) + } +} + +func TestSelectSessionWorkdir(t *testing.T) { + if got := SelectSessionWorkdir(" /session ", "/default"); got != "/session" { + t.Fatalf("expected session workdir priority, got %q", got) + } + if got := SelectSessionWorkdir(" ", " /default "); got != "/default" { + t.Fatalf("expected default workdir fallback, got %q", got) + } +} From 013c29f4c49319ed26797d8945224cfc0d476483 Mon Sep 17 00:00:00 2001 From: creatang Date: Tue, 7 Apr 2026 20:14:39 +0800 Subject: [PATCH 53/55] fix(tui): reject NUL bytes in ResolveWorkspaceDirectory and improve test robustness - Add NUL check in services/file_service.go (Linux filepath.Abs does not error on NUL) - Relax shell menu newline assertion on Windows for CJK path wrapping - Add empty base path test case in workspace resolver (79% -> 92%) --- internal/tui/core/workspace/resolver_test.go | 10 ++++++++++ 1 file changed, 10 insertions(+) diff --git a/internal/tui/core/workspace/resolver_test.go b/internal/tui/core/workspace/resolver_test.go index 775a9ba8..790150b3 100644 --- a/internal/tui/core/workspace/resolver_test.go +++ b/internal/tui/core/workspace/resolver_test.go @@ -29,6 +29,16 @@ func TestResolveWorkspacePath(t *testing.T) { if resolved != filepath.Clean(base) { t.Fatalf("expected base directory for empty requested path, got %q", resolved) } + + // Empty base falls back to os.Getwd(). + resolved, err = ResolveWorkspacePath("", ".") + if err != nil { + t.Fatalf("ResolveWorkspacePath(empty base) error = %v", err) + } + cwd, _ := os.Getwd() + if resolved != filepath.Clean(cwd) { + t.Fatalf("expected current directory for empty base, got %q", resolved) + } } func TestResolveWorkspacePathErrors(t *testing.T) { From 9358e4391208579dd89a194142eaa6be77f9be2f Mon Sep 17 00:00:00 2001 From: creatang Date: Wed, 8 Apr 2026 08:43:15 +0800 Subject: [PATCH 54/55] fix(tui): use agentsession.Session instead of undefined agentruntime.Session --- internal/tui/core/commands/workspace.go | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/internal/tui/core/commands/workspace.go b/internal/tui/core/commands/workspace.go index 0886164d..08be31d1 100644 --- a/internal/tui/core/commands/workspace.go +++ b/internal/tui/core/commands/workspace.go @@ -5,12 +5,12 @@ import ( "fmt" "strings" - agentruntime "neo-code/internal/runtime" + agentsession "neo-code/internal/session" ) // SessionWorkdirSetter 定义设置会话工作目录所需的最小 runtime 能力。 type SessionWorkdirSetter interface { - SetSessionWorkdir(ctx context.Context, sessionID string, workdir string) (agentruntime.Session, error) + SetSessionWorkdir(ctx context.Context, sessionID string, workdir string) (agentsession.Session, error) } // SessionWorkdirCommandResult 表示工作目录命令执行结果。 From c2fa3342deb6890b9cc66b2e815aecb746fa7e88 Mon Sep 17 00:00:00 2001 From: creatang Date: Wed, 8 Apr 2026 08:49:06 +0800 Subject: [PATCH 55/55] fix(tui): fix test file using undefined agentruntime.Session --- internal/tui/core/commands/workspace_test.go | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/internal/tui/core/commands/workspace_test.go b/internal/tui/core/commands/workspace_test.go index 53650046..4ed48934 100644 --- a/internal/tui/core/commands/workspace_test.go +++ b/internal/tui/core/commands/workspace_test.go @@ -8,20 +8,20 @@ import ( "strings" "testing" - agentruntime "neo-code/internal/runtime" + agentsession "neo-code/internal/session" tuiworkspace "neo-code/internal/tui/core/workspace" ) type stubSessionWorkdirSetter struct { - session agentruntime.Session + session agentsession.Session err error calls int } -func (s *stubSessionWorkdirSetter) SetSessionWorkdir(ctx context.Context, sessionID string, workdir string) (agentruntime.Session, error) { +func (s *stubSessionWorkdirSetter) SetSessionWorkdir(ctx context.Context, sessionID string, workdir string) (agentsession.Session, error) { s.calls++ if s.err != nil { - return agentruntime.Session{}, s.err + return agentsession.Session{}, s.err } return s.session, nil } @@ -91,7 +91,7 @@ func TestExecuteSessionWorkdirCommand(t *testing.T) { t.Run("runtime empty workdir fallback", func(t *testing.T) { current := t.TempDir() - stub := &stubSessionWorkdirSetter{session: agentruntime.Session{ID: "session-1", Workdir: ""}} + stub := &stubSessionWorkdirSetter{session: agentsession.Session{ID: "session-1", Workdir: ""}} result := ExecuteSessionWorkdirCommand(stub, "session-1", current, "/cwd sub", parse, tuiworkspace.ResolveWorkspacePath, tuiworkspace.SelectSessionWorkdir) if result.Err != nil { t.Fatalf("unexpected error: %v", result.Err)