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/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 fed73849..26a5800e 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,8 @@ 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、不做旧工具结果清理 | + +新增工具默认会参与 micro compact;如果某个工具的历史结果必须保留,需要在 `internal/tools` 的工具实现中显式声明保留策略。 更多行为说明见 [context-compact.md](../context-compact.md)。 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/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,都不需要推翻当前设计。 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/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/app/bootstrap.go b/internal/app/bootstrap.go index d87241ea..ef3a7c9d 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" @@ -56,19 +57,22 @@ 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 } - sessionStore := agentruntime.NewSessionStore(loader.BaseDir()) + sessionStore := agentsession.NewStore(loader.BaseDir()) runtimeSvc := agentruntime.NewWithFactory( manager, toolManager, sessionStore, providerRegistry, - agentcontext.NewBuilder(), + agentcontext.NewBuilderWithToolPolicies(toolRegistry), ) tuiApp, err := tui.New(&cfg, manager, runtimeSvc, providerSelection) @@ -82,7 +86,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,11 +99,18 @@ 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) { - 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..a81903c0 100644 --- a/internal/app/bootstrap_test.go +++ b/internal/app/bootstrap_test.go @@ -1,18 +1,25 @@ package app import ( + "bufio" + "bytes" "context" "encoding/json" "errors" + "fmt" + "io" "net/http" "net/http/httptest" "os" "path/filepath" + "strconv" "strings" "testing" + "time" "neo-code/internal/config" "neo-code/internal/tools" + "neo-code/internal/tools/mcp" ) func TestNewProgram(t *testing.T) { @@ -84,7 +91,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 +120,337 @@ 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 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() + + 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 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() + + 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() @@ -132,16 +473,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 +492,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 @@ -218,6 +579,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 } @@ -229,3 +593,173 @@ func disableBuiltinProviderAPIKeys(t *testing.T) { t.Setenv(config.OpenLLDefaultAPIKeyEnv, "") t.Setenv(config.QiniuDefaultAPIKeyEnv, "") } + +type stubMCPServerClient struct { + 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 +} + +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 +} + +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/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 4516c5e0..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() @@ -876,10 +945,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 +967,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 +984,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..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 { @@ -70,6 +71,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 { @@ -77,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...) } @@ -90,6 +120,7 @@ func Default() *Config { Context: defaultContextConfig(), Tools: ToolsConfig{ WebFetch: defaultWebFetchConfig(), + MCP: defaultMCPConfig(), }, } } @@ -329,6 +360,13 @@ func defaultWebFetchConfig() WebFetchConfig { } } +// defaultMCPConfig 返回 MCP 工具接入配置的默认值(默认无 server)。 +func defaultMCPConfig() MCPConfig { + return MCPConfig{ + Servers: nil, + } +} + // defaultContextConfig 返回上下文压缩相关配置的默认值。 func defaultContextConfig() ContextConfig { return ContextConfig{ @@ -348,6 +386,7 @@ func defaultCompactConfig() CompactConfig { func (c ToolsConfig) Clone() ToolsConfig { return ToolsConfig{ WebFetch: c.WebFetch.Clone(), + MCP: c.MCP.Clone(), } } @@ -364,6 +403,7 @@ func (c *ToolsConfig) ApplyDefaults(defaults ToolsConfig) { } c.WebFetch.ApplyDefaults(defaults.WebFetch) + c.MCP.ApplyDefaults(defaults.MCP) } // ApplyDefaults 为上下文配置补齐缺省的 compact 参数。 @@ -379,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 } @@ -412,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 { diff --git a/internal/context/builder.go b/internal/context/builder.go index ff1487ea..96e3c437 100644 --- a/internal/context/builder.go +++ b/internal/context/builder.go @@ -1,15 +1,25 @@ package context -import "context" +import ( + "context" + + providertypes "neo-code/internal/provider/types" +) // 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{ @@ -17,7 +27,8 @@ func NewBuilder() Builder { &projectRulesSource{}, systemSource, }, - trimPolicy: spanMessageTrimPolicy{}, + trimPolicy: spanMessageTrimPolicy{}, + microCompactPolicies: policies, } } @@ -43,6 +54,14 @@ func (b *DefaultBuilder) Build(ctx context.Context, input BuildInput) (BuildResu return BuildResult{ SystemPrompt: composeSystemPrompt(sections...), - Messages: trimPolicy.Trim(input.Messages), + Messages: applyReadTimeContextProjection(trimPolicy.Trim(input.Messages), input.Compact, b.microCompactPolicies), }, nil } + +// applyReadTimeContextProjection 负责在 provider 请求前按开关应用只读上下文投影,避免改写原始会话消息。 +func applyReadTimeContextProjection(messages []providertypes.Message, options CompactOptions, policies MicroCompactPolicySource) []providertypes.Message { + if options.DisableMicroCompact { + return cloneContextMessages(messages) + } + return microCompactMessagesWithPolicies(messages, policies) +} diff --git a/internal/context/builder_test.go b/internal/context/builder_test.go index 42c002ba..dfe2b838 100644 --- a/internal/context/builder_test.go +++ b/internal/context/builder_test.go @@ -5,11 +5,13 @@ import ( "fmt" "os" "path/filepath" + "reflect" "strings" "testing" "neo-code/internal/context/internalcompact" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" + "neo-code/internal/tools" ) type stubPromptSectionSource struct { @@ -29,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()), @@ -88,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 { @@ -109,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), }) } @@ -150,23 +152,219 @@ func TestDefaultBuilderBuildReturnsPromptSourceError(t *testing.T) { } } +func TestDefaultBuilderBuildAppliesMicroCompactAfterTrim(t *testing.T) { + t.Parallel() + + builder := &DefaultBuilder{ + promptSources: []promptSectionSource{ + stubPromptSectionSource{sections: []promptSection{{title: "Stub", content: "body"}}}, + }, + } + + messages := []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older user"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "old read result"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-2", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-3", Name: "webfetch", Arguments: "{}"}, + }, + }, + {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}) + 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 TestDefaultBuilderBuildSkipsMicroCompactWhenDisabled(t *testing.T) { + t.Parallel() + + builder := &DefaultBuilder{ + promptSources: []promptSectionSource{ + stubPromptSectionSource{sections: []promptSection{{title: "Stub", content: "body"}}}, + }, + } + + messages := []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older user"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "old read result"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-2", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-3", Name: "webfetch", Arguments: "{}"}, + }, + }, + {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, + 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 TestDefaultBuilderBuildHonorsToolMicroCompactPolicies(t *testing.T) { + t.Parallel() + + builder := &DefaultBuilder{ + promptSources: []promptSectionSource{ + stubPromptSectionSource{sections: []promptSection{{title: "Stub", content: "body"}}}, + }, + microCompactPolicies: stubMicroCompactPolicySource{ + "custom_tool": tools.MicroCompactPolicyPreserveHistory, + }, + } + + messages := []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older user"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-1", Name: "custom_tool", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "old custom result"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-2", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-3", Name: "webfetch", Arguments: "{}"}, + }, + }, + {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}) + 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 TestNewBuilderWithToolPoliciesUsesProvidedPolicySource(t *testing.T) { + t.Parallel() + + builder := NewBuilderWithToolPolicies(stubMicroCompactPolicySource{ + "custom_tool": tools.MicroCompactPolicyPreserveHistory, + }) + + messages := []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older user"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-1", Name: "custom_tool", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "old custom result"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-2", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-3", Name: "webfetch", Arguments: "{}"}, + }, + }, + {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}) + 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() - 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) @@ -192,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 { @@ -223,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) @@ -253,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]) } } @@ -263,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") @@ -283,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)) @@ -310,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 new file mode 100644 index 00000000..a881a8b9 --- /dev/null +++ b/internal/context/microcompact.go @@ -0,0 +1,145 @@ +package context + +import ( + "strings" + + "neo-code/internal/context/internalcompact" + providertypes "neo-code/internal/provider/types" + "neo-code/internal/tools" +) + +const ( + // microCompactClearedMessage 是旧工具结果被读时微压缩后的占位符文本。 + microCompactClearedMessage = "[Old tool result content cleared]" + // microCompactRetainedToolSpans 定义默认保留原始内容的最近可压缩工具块数量。 + microCompactRetainedToolSpans = 2 +) + +// microCompactMessages 对裁剪后的消息做只读投影式微压缩,仅清理旧工具结果内容。 +func microCompactMessages(messages []providertypes.Message) []providertypes.Message { + return microCompactMessagesWithPolicies(messages, nil) +} + +// microCompactMessagesWithPolicies 按工具策略对裁剪后的消息做只读投影式微压缩。 +func microCompactMessagesWithPolicies(messages []providertypes.Message, policies MicroCompactPolicySource) []providertypes.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, policies) + if len(compactableIDs) == 0 { + continue + } + if !hasCompactableToolContent(cloned, span, compactableIDs) { + 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 []providertypes.Message) []providertypes.Message { + if len(messages) == 0 { + return nil + } + + cloned := make([]providertypes.Message, 0, len(messages)) + for _, message := range messages { + next := message + next.ToolCalls = append([]providertypes.ToolCall(nil), message.ToolCalls...) + cloned = append(cloned, next) + } + return cloned +} + +// isToolCallSpan 判断当前 span 是否是由 assistant tool call 起始的原子工具块。 +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 == providertypes.RoleAssistant && len(message.ToolCalls) > 0 +} + +// compactableToolCallIDs 返回 assistant tool call 中可参与微压缩的调用 ID 集合。 +func compactableToolCallIDs(calls []providertypes.ToolCall, policies MicroCompactPolicySource) 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 !toolParticipatesInMicroCompact(toolName, policies) { + continue + } + callID := strings.TrimSpace(call.ID) + if callID == "" { + continue + } + ids[callID] = struct{}{} + } + if len(ids) == 0 { + return nil + } + 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 []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 + } + } + return false +} + +// shouldClearToolMessage 判断一条 tool 消息是否满足旧结果清理条件。 +func shouldClearToolMessage(message providertypes.Message, compactableIDs map[string]struct{}) bool { + if message.Role != providertypes.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..24e1e0e4 --- /dev/null +++ b/internal/context/microcompact_test.go @@ -0,0 +1,348 @@ +package context + +import ( + "testing" + + providertypes "neo-code/internal/provider/types" + "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() + + messages := []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older user"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "old read result"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-2", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-3", Name: "webfetch", Arguments: "{}"}, + }, + }, + {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) + 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 TestMicroCompactMessagesHandlesEmptyAndInvalidSpanInputs(t *testing.T) { + t.Parallel() + + if got := microCompactMessages(nil); got != nil { + t.Fatalf("expected nil input to remain nil, got %+v", got) + } + + assistantOnly := []providertypes.Message{ + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.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() + + messages := []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older user"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-0", Name: "filesystem_grep", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-0", Content: "old grep result"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "recent read result"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-2", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + {Role: providertypes.RoleUser, Content: "latest explicit instruction"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-3", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: providertypes.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 TestMicroCompactMessagesKeepsPreservedToolsErrorsAndOrphans(t *testing.T) { + t.Parallel() + + messages := []providertypes.Message{ + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-1", Name: "custom_tool", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "custom result"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-2", Name: "filesystem_edit", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-2", Content: "edit failed", IsError: true}, + {Role: providertypes.RoleTool, ToolCallID: "orphan", Content: "orphan result"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-3", Name: "filesystem_write_file", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-3", Content: microCompactClearedMessage}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-4", Name: "filesystem_grep", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-4", Content: ""}, + } + + got := microCompactMessagesWithPolicies(messages, stubMicroCompactPolicySource{ + "custom_tool": tools.MicroCompactPolicyPreserveHistory, + }) + if got[1].Content != "custom result" { + 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) + } + 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 TestMicroCompactMessagesClearsOnlyNonPreservedResultsInMixedToolSpan(t *testing.T) { + t.Parallel() + + messages := []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older user"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, + {ID: "call-2", Name: "custom_tool", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "read result"}, + {Role: providertypes.RoleTool, ToolCallID: "call-2", Content: "custom result"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-3", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-3", Content: "recent bash result"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-4", Name: "webfetch", Arguments: "{}"}, + }, + }, + {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{ + "custom_tool": tools.MicroCompactPolicyPreserveHistory, + }) + if got[2].Content != microCompactClearedMessage { + t.Fatalf("expected default compactable tool result to be cleared, got %q", got[2].Content) + } + if got[3].Content != "custom result" { + 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 := []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older user"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-1", Name: "repo_search", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "old repo search result"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-2", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-3", Name: "webfetch", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, + {Role: providertypes.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() + + messages := []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older user"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-1", Name: "filesystem_read_file", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "older read result"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-2", Name: "filesystem_grep", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-2", Content: "middle grep result"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-3", Name: "filesystem_edit", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-3", Content: "near edit result"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-4", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-4", Content: "", IsError: true}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-5", Name: "webfetch", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-5", Content: ""}, + {Role: providertypes.RoleUser, Content: "latest explicit instruction"}, + {Role: providertypes.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) + } +} + +func TestMicroCompactMessagesSkipsToolMessagesWhenCompactableIDsMissing(t *testing.T) { + t.Parallel() + + messages := []providertypes.Message{ + {Role: providertypes.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/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 6eb516dd..53af2bfa 100644 --- a/internal/context/types.go +++ b/internal/context/types.go @@ -3,7 +3,8 @@ package context import ( "context" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" + "neo-code/internal/tools" ) // Builder builds the provider-facing context for a single model round. @@ -13,12 +14,23 @@ 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 } // 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 策略的最小依赖。 +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/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..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) { @@ -136,7 +137,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 +235,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 providertypes.ChatRequest, events chan<- providertypes.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/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/errors.go b/internal/provider/errors.go index 49936b9e..daff0da6 100644 --- a/internal/provider/errors.go +++ b/internal/provider/errors.go @@ -10,6 +10,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 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 b1b3599c..00000000 --- a/internal/provider/openai/openai.go +++ /dev/null @@ -1,549 +0,0 @@ -package openai - -import ( - "bufio" - "bytes" - "context" - "encoding/json" - "errors" - "fmt" - "io" - "log" - "net/http" - "sort" - "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) { - provider, err := New(cfg, WithTransport(defaultRetryTransport())) - if err != nil { - return nil, err - } - return provider.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 -} - -func (p *Provider) Chat(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) (provider.ChatResponse, error) { - payload, err := p.buildRequest(req) - if err != nil { - return provider.ChatResponse{}, err - } - - body, err := json.Marshal(payload) - if err != nil { - return provider.ChatResponse{}, 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) - } - 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 provider.ChatResponse{}, 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 provider.ChatResponse{}, 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 -} - -func (p *Provider) consumeStream(ctx context.Context, body io.Reader, events chan<- provider.StreamEvent) (provider.ChatResponse, error) { - reader := bufio.NewReader(body) - - var ( - contentBuilder strings.Builder - finishReason string - usage provider.Usage - done bool - ) - - toolCalls := make(map[int]*provider.ToolCall) - dataLines := make([]string, 0, 4) - - // processChunk 解析单个 SSE data payload,更新累积状态。 - // 返回错误表示应中止流;done 标志通过闭包变量传递。 - 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 != "" { - contentBuilder.WriteString(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() (provider.ChatResponse, error) { - if err := emitMessageDone(ctx, events, finishReason, &usage); err != nil { - return provider.ChatResponse{}, err - } - return finalizeResponse(contentBuilder.String(), toolCalls, finishReason, usage), nil - } - - flushPendingData := func() error { - defer func() { - dataLines = dataLines[:0] - }() - return flushDataLines(dataLines, processChunk) - } - - 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) - } - - line = strings.TrimRight(line, "\r\n") - trimmed := strings.TrimSpace(line) - - switch { - case strings.HasPrefix(trimmed, "data:"): - dataLines = append(dataLines, strings.TrimSpace(strings.TrimPrefix(trimmed, "data:"))) - case trimmed == "": - if flushErr := flushPendingData(); flushErr != nil { - return provider.ChatResponse{}, 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 provider.ChatResponse{}, 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.StreamEvent{ - Type: provider.StreamEventTextDelta, - Text: 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, - }) -} - -// emitToolCallDelta 发送工具调用参数增量事件。 -func emitToolCallDelta(ctx context.Context, events chan<- provider.StreamEvent, index int, argumentsDelta string) error { - if argumentsDelta == "" { - return nil - } - return emitStreamEvent(ctx, events, provider.StreamEvent{ - Type: provider.StreamEventToolCallDelta, - ToolCallIndex: index, - ToolArgumentsDelta: argumentsDelta, - }) -} - -// 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.StreamEvent{ - Type: provider.StreamEventMessageDone, - FinishReason: finishReason, - Usage: 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, args); err != nil { - return err - } - } - 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 - } - - 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 413e5b44..33179a2e 100644 --- a/internal/provider/openai/openai_test.go +++ b/internal/provider/openai/openai_test.go @@ -3,13 +3,17 @@ package openai import ( "context" "encoding/json" + "errors" + "io" "net/http" "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) { @@ -33,7 +37,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) } @@ -43,6 +47,63 @@ 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() + 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 { + return + } + }) +} + +func TestNewDefaultTransportWhenNoOption(t *testing.T) { + t.Parallel() + + cfg := resolvedConfig("", "") + provider, err := New(cfg) + 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() @@ -52,139 +113,761 @@ func TestDefaultRetryTransport(t *testing.T) { } } -func TestDiscoverModels(t *testing.T) { +func TestDiscoverModels(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: %s", r.URL.Path) + } + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(map[string]any{ + "data": []map[string]any{ + {"id": "gpt-4", "name": "GPT-4"}, + {"id": "gpt-3.5-turbo"}, + }, + }) + })) + defer server.Close() + + provider, err := New(resolvedConfig(server.URL, "")) + if err != nil { + t.Fatalf("New() error = %v", err) + } + provider.client = server.Client() + + models, err := provider.DiscoverModels(context.Background()) + if err != nil { + t.Fatalf("DiscoverModels() error = %v", err) + } + if len(models) != 2 { + t.Fatalf("expected 2 models, got %d", len(models)) + } + if models[0].ID != "gpt-4" || models[0].Name != "GPT-4" { + t.Fatalf("unexpected first model: %+v", models[0]) + } +} + +// --- 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) { - if r.URL.Path != "/models" { - t.Fatalf("unexpected path: %s", r.URL.Path) - } - w.Header().Set("Content-Type", "application/json") - _ = json.NewEncoder(w).Encode(map[string]any{ - "data": []map[string]any{ - {"id": "gpt-4", "name": "GPT-4"}, - {"id": "gpt-3.5-turbo"}, - }, - }) + w.WriteHeader(http.StatusInternalServerError) + _, _ = w.Write([]byte("internal error")) })) defer server.Close() - provider, err := New(resolvedConfig(server.URL, "")) + p, err := New(resolvedConfig(server.URL, "")) if err != nil { t.Fatalf("New() error = %v", err) } - provider.client = server.Client() + p.client = server.Client() - models, err := provider.DiscoverModels(context.Background()) - if err != nil { - t.Fatalf("DiscoverModels() error = %v", err) + models, err := p.DiscoverModels(context.Background()) + if err == nil { + t.Fatal("expected error for HTTP 500 response") } - if len(models) != 2 { - t.Fatalf("expected 2 models, got %d", len(models)) + if models != nil { + t.Fatalf("expected nil models on error, got %d models", len(models)) } - if models[0].ID != "gpt-4" || models[0].Name != "GPT-4" { - t.Fatalf("unexpected first model: %+v", models[0]) +} + +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") } } -func TestEmitToolCallDelta(t *testing.T) { +// --- 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, `{"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"}` { - 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 - if got.Type != domain.StreamEventMessageDone || got.FinishReason != "stop" || got.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") @@ -203,155 +886,86 @@ 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) - response, 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) } - 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 + var toolCallStartSeen bool + var toolCallArgs strings.Builder + var messageDone *providertypes.MessageDonePayload + + for _, event := range streamEvents { + switch event.Type { + case providertypes.StreamEventTextDelta: + chunks = append(chunks, requireTextDeltaPayload(t, event).Text) + 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 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 providertypes.StreamEventMessageDone: + p := requireMessageDonePayload(t, event) + messageDone = &p + } } - 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) } } @@ -364,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 { @@ -389,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) } @@ -408,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) @@ -469,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) } @@ -480,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.Text != "chunk" || got.Type != domain.StreamEventTextDelta { + 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") } } @@ -527,20 +1099,86 @@ 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 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: + if !ok { + return drained + } + drained = append(drained, evt) + default: + return drained + } + } +} + +func requireTextDeltaPayload(t *testing.T, event providertypes.StreamEvent) providertypes.TextDeltaPayload { + t.Helper() + payload, err := event.TextDeltaValue() + if err != nil { + t.Fatalf("TextDeltaValue() error = %v", err) + } + return payload +} + +func requireToolCallStartPayload(t *testing.T, event providertypes.StreamEvent) providertypes.ToolCallStartPayload { + t.Helper() + payload, err := event.ToolCallStartValue() + if err != nil { + t.Fatalf("ToolCallStartValue() error = %v", err) + } + return payload +} + +func requireToolCallDeltaPayload(t *testing.T, event providertypes.StreamEvent) providertypes.ToolCallDeltaPayload { + t.Helper() + payload, err := event.ToolCallDeltaValue() + if err != nil { + t.Fatalf("ToolCallDeltaValue() error = %v", err) + } + return payload +} + +func requireMessageDonePayload(t *testing.T, event providertypes.StreamEvent) providertypes.MessageDonePayload { + t.Helper() + payload, err := event.MessageDoneValue() + if err != nil { + t.Fatalf("MessageDoneValue() error = %v", err) + } + 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 { + for _, m := range messages { + if m.Role == "tool" && m.ToolCallID == toolCallID && m.Content == content { return true } } @@ -556,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") + + p, err := New(resolvedConfig("", "")) + if err != nil { + t.Fatalf("New() error = %v", err) + } -type readCloser struct { - *strings.Reader + 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) } @@ -596,66 +1287,55 @@ 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 - if got.Type != domain.StreamEventToolCallStart || got.ToolName != "filesystem_edit" || got.ToolCallID != "call-1" || got.ToolCallIndex != 2 { + payload := requireToolCallStartPayload(t, got) + 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) } - 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 { + if delta.Type != providertypes.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] @@ -672,59 +1352,36 @@ 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) } - close(events) - var foundToolCallStart bool - for evt := range events { - if evt.Type == domain.StreamEventToolCallStart { + for _, evt := range drainStreamEvents(events) { + if evt.Type == providertypes.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) } } } @@ -733,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) } - close(events) + var foundTextDelta, foundToolCallStart, foundToolCallDelta, foundMessageDone bool + var toolCallDeltaContent string + var messageDonePayload *providertypes.MessageDonePayload - var ( - foundTextDelta bool - foundToolCallStart bool - foundToolCallDelta bool - foundMessageDone bool - toolCallDeltaContent string - messageDoneEvt *domain.StreamEvent - ) - - for evt := range events { + for _, evt := range drainStreamEvents(events) { switch evt.Type { - case domain.StreamEventTextDelta: + case providertypes.StreamEventTextDelta: foundTextDelta = true - case domain.StreamEventToolCallStart: + case providertypes.StreamEventToolCallStart: foundToolCallStart = true - if evt.ToolName != "filesystem_edit" { - t.Fatalf("expected ToolName %q, got %q", "filesystem_edit", evt.ToolName) - } - if evt.ToolCallIndex != 0 { - t.Fatalf("expected ToolCallIndex %d for tool_call_start, got %d", 0, evt.ToolCallIndex) + 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 += evt.ToolArgumentsDelta - case domain.StreamEventMessageDone: + toolCallDeltaContent += requireToolCallDeltaPayload(t, evt).ArgumentsDelta + case providertypes.StreamEventMessageDone: foundMessageDone = true - messageDoneEvt = &evt + p := requireMessageDonePayload(t, evt) + messageDonePayload = &p } } @@ -871,24 +1450,19 @@ 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 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/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/sse_reader.go b/internal/provider/openai/sse_reader.go new file mode 100644 index 00000000..3070feec --- /dev/null +++ b/internal/provider/openai/sse_reader.go @@ -0,0 +1,80 @@ +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 行读取器。 +// +// 内部 bufio.Reader 的缓冲区大小设为 maxSSELineSize+1,使得 ReadSlice('\n') +// 在单行超过 L1 上限时立即返回 bufio.ErrBufferFull,避免 ReadString 那样 +// 先把整行全部读进内存才做长度检查。 +func newBoundedSSEReader(r io.Reader) *boundedSSEReader { + return &boundedSSEReader{ + reader: bufio.NewReaderSize(r, maxSSELineSize+1), + } +} + +// ReadLine 读取一行(以 \n 分隔),同时执行 L1 和 L3 检查。 +// 返回去除尾部 \r\n 的行内容;遇到 io.EOF 时返回空字符串和 nil。 +// +// L1 检查通过 bufio.Reader 的缓冲区大小约束实现:如果一行在缓冲区内 +// 未找到 \n,ReadSlice 直接返回 bufio.ErrBufferFull,无需先将整行载入内存。 +func (r *boundedSSEReader) ReadLine() (string, error) { + line, err := r.reader.ReadSlice('\n') + + // L1: 缓冲区溢出 → 单行超过 maxSSELineSize(触发在读取过程中,而非读完后) + if errors.Is(err, bufio.ErrBufferFull) { + return "", provider.ErrLineTooLong + } + + 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 + } + + // 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。 +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 +} diff --git a/internal/provider/openai/sse_reader_test.go b/internal/provider/openai/sse_reader_test.go new file mode 100644 index 00000000..eb7a1803 --- /dev/null +++ b/internal/provider/openai/sse_reader_test.go @@ -0,0 +1,266 @@ +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() + + // 构造输入:每行 1KB(远小于 maxSSELineSize),行数足够多使总量超过 maxStreamTotalSize + line := strings.Repeat("x", 1024) + "\n" // 1025 bytes per line + lineSize := int64(len(line)) + + var sb strings.Builder + linesToWrite := int(maxStreamTotalSize/lineSize) + 1 + for range linesToWrite { + sb.WriteString(line) + } + + r := newBoundedSSEReader(strings.NewReader(sb.String())) + + // 前面的行应能正常读取 + expectedNormal := int(maxStreamTotalSize / lineSize) + for range expectedNormal { + _, err := r.ReadLine() + if err != nil { + t.Fatalf("unexpected error on normal line: %v", err) + } + } + + // 超限的行应返回 ErrStreamTooLarge(而非 ErrLineTooLong) + _, 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) + } + } +} + +// 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 触发后调用方会终止流消费。 +} 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 40a838b9..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) (ChatResponse, 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 f650253d..5ddc5b31 100644 --- a/internal/provider/registry_test.go +++ b/internal/provider/registry_test.go @@ -8,12 +8,13 @@ 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) (provider.ChatResponse, error) { - return provider.ChatResponse{}, nil +func (stubProvider) Chat(ctx context.Context, req providertypes.ChatRequest, events chan<- providertypes.StreamEvent) error { + return nil } func stubDriver(driverType string) provider.DriverDefinition { diff --git a/internal/provider/types.go b/internal/provider/types.go deleted file mode 100644 index eba95f19..00000000 --- a/internal/provider/types.go +++ /dev/null @@ -1,86 +0,0 @@ -package provider - -// Role 常量定义消息角色标识。 -const ( - RoleSystem = "system" - RoleUser = "user" - RoleAssistant = "assistant" - 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"` -} - -// 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"` - OutputTokens int `json:"output_tokens"` - TotalTokens int `json:"total_tokens"` -} - -// StreamEventType 定义流式事件类型。 -type StreamEventType string - -const ( - // StreamEventTextDelta 表示模型输出的文本片段。 - StreamEventTextDelta StreamEventType = "text_delta" - // StreamEventToolCallStart 表示模型开始请求工具调用,TUI 可据此展示过渡提示。 - StreamEventToolCallStart StreamEventType = "tool_call_start" - // StreamEventToolCallDelta 表示工具调用参数的增量片段。 - StreamEventToolCallDelta StreamEventType = "tool_call_delta" - // StreamEventMessageDone 表示本轮消息完成,包含最终统计信息。 - StreamEventMessageDone StreamEventType = "message_done" -) - -// StreamEvent 表示 provider 驱动层向 runtime 推送的流式事件。 -type StreamEvent struct { - Type StreamEventType - - // 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 时有效) -} diff --git a/internal/provider/types/event.go b/internal/provider/types/event.go new file mode 100644 index 00000000..72df4bc1 --- /dev/null +++ b/internal/provider/types/event.go @@ -0,0 +1,127 @@ +package types + +import "fmt" + +// StreamEventType 定义流式事件类型。 +type StreamEventType string + +const ( + // StreamEventTextDelta 表示模型输出的文本片段。 + StreamEventTextDelta StreamEventType = "text_delta" + // StreamEventToolCallStart 表示模型开始请求工具调用。 + StreamEventToolCallStart StreamEventType = "tool_call_start" + // StreamEventToolCallDelta 表示工具调用参数的增量片段。 + StreamEventToolCallDelta StreamEventType = "tool_call_delta" + // StreamEventMessageDone 表示本轮消息完成,并携带最终统计信息。 + StreamEventMessageDone StreamEventType = "message_done" +) + +// StreamEvent 表示 provider 向 runtime 推送的流式事件。 +type StreamEvent struct { + 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 表示文本增量事件的载荷。 +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, + TextDelta: &TextDeltaPayload{Text: text}, + } +} + +// NewToolCallStartStreamEvent 创建工具调用开始流事件。 +func NewToolCallStartStreamEvent(index int, id, name string) StreamEvent { + return StreamEvent{ + Type: StreamEventToolCallStart, + ToolCallStart: &ToolCallStartPayload{Index: index, ID: id, Name: name}, + } +} + +// NewToolCallDeltaStreamEvent 创建工具调用参数增量流事件。 +func NewToolCallDeltaStreamEvent(index int, id, argumentsDelta string) StreamEvent { + return StreamEvent{ + Type: StreamEventToolCallDelta, + ToolCallDelta: &ToolCallDeltaPayload{Index: index, ID: id, ArgumentsDelta: argumentsDelta}, + } +} + +// NewMessageDoneStreamEvent 创建消息完成流事件。 +func NewMessageDoneStreamEvent(finishReason string, usage *Usage) StreamEvent { + return StreamEvent{ + 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/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/types_test.go b/internal/provider/types/types_test.go new file mode 100644 index 00000000..958ebb16 --- /dev/null +++ b/internal/provider/types/types_test.go @@ -0,0 +1,267 @@ +package types + +import ( + "encoding/json" + "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, 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) + } +} + +// --- 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, 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) + } +} + +// --- 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, 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) + } +} + +// --- 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, 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) + } + 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, 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) + } + if payload.Usage != nil { + t.Fatal("expected nil usage") + } + }) + + t.Run("empty finish reason", func(t *testing.T) { + event := NewMessageDoneStreamEvent("", nil) + + 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) + } + }) +} + +// --- 结构体字段覆盖验证 --- + +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, + 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.go b/internal/runtime/compact.go index 4939ca62..26487f33 100644 --- a/internal/runtime/compact.go +++ b/internal/runtime/compact.go @@ -8,7 +8,8 @@ import ( "neo-code/internal/config" contextcompact "neo-code/internal/context/compact" - "neo-code/internal/provider" + 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 @@ -103,7 +104,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 +126,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{ @@ -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/compact_generator.go b/internal/runtime/compact_generator.go index 87d9d658..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 { @@ -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 providertypes.StreamEvent, 32) + streamDone := make(chan error, 1) + acc := newStreamAccumulator() + + go func() { + var streamErr error + defer func() { + streamDone <- streamErr + }() + + for { + select { + case event, ok := <-streamEvents: + if !ok { + return + } + if err := handleProviderStreamEvent(event, acc, nil, nil); err != nil && streamErr == nil { + // 记录首个协议错误后继续排空事件通道,避免 provider 在后续发送时阻塞。 + streamErr = err + } + case <-ctx.Done(): + return + } + } + }() + + 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, }}, - }, nil) + }, streamEvents) + close(streamEvents) + streamErr := <-streamDone + + if err != nil { + return "", err + } + if streamErr != nil { + return "", streamErr + } + + message, err := acc.buildMessage() if err != nil { return "", err } - if len(resp.Message.ToolCalls) > 0 { + 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..00a96ec9 100644 --- a/internal/runtime/compact_generator_test.go +++ b/internal/runtime/compact_generator_test.go @@ -4,10 +4,11 @@ import ( "context" "strings" "testing" + "time" "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) { @@ -20,45 +21,42 @@ 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: [][]providertypes.StreamEvent{ + {providertypes.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") 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, @@ -89,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, "") { @@ -116,21 +114,19 @@ 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: [][]providertypes.StreamEvent{ + { + providertypes.NewToolCallStartStreamEvent(0, "call-1", "filesystem_read_file"), + providertypes.NewToolCallDeltaStreamEvent(0, "call-1", "{}"), }, - }}, + }, } generator := newCompactSummaryGenerator(&scriptedProviderFactory{provider: scripted}, resolvedProvider, "session-model") _, 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, }) @@ -138,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: [][]providertypes.StreamEvent{ + { + {Type: providertypes.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 := []providertypes.StreamEvent{{Type: providertypes.StreamEventTextDelta}} + for i := 0; i < 40; i++ { + stream = append(stream, providertypes.NewTextDeltaStreamEvent("ignored")) + } + scripted := &scriptedProvider{ + streams: [][]providertypes.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/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/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/permission.go b/internal/runtime/permission.go new file mode 100644 index 00000000..0ffe62ee --- /dev/null +++ b/internal/runtime/permission.go @@ -0,0 +1,291 @@ +package runtime + +import ( + "context" + "errors" + "fmt" + "strings" + "sync" + "time" + + providertypes "neo-code/internal/provider/types" + "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 providertypes.ToolCall + Workdir string + ToolTimeout time.Duration +} + +type pendingPermissionRequest struct { + RequestID string + RunID string + SessionID string + Call providertypes.ToolCall + Action security.Action + ResultCh chan PermissionResolutionDecision + Submitted bool +} + +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) + } + if err := ctx.Err(); err != nil { + return err + } + + runtimePendingPermissions.mu.Lock() + pending := runtimePendingPermissions.byRun[s] + 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 resultCh <- decision: + return nil + default: + return nil + } +} + +// 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(), + }) + + 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/permission_test.go b/internal/runtime/permission_test.go new file mode 100644 index 00000000..c82a9526 --- /dev/null +++ b/internal/runtime/permission_test.go @@ -0,0 +1,337 @@ +package runtime + +import ( + "context" + "errors" + "testing" + "time" + + providertypes "neo-code/internal/provider/types" + "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: providertypes.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 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: providertypes.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() + + 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: []scriptedResponse{ + { + Message: providertypes.Message{ + Role: "assistant", + ToolCalls: []providertypes.ToolCall{ + {ID: "call-ask-reject", Name: "webfetch", Arguments: `{"url":"https://example.com/private"}`}, + }, + }, + FinishReason: "tool_calls", + }, + { + Message: providertypes.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: providertypes.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/runtime/runtime.go b/internal/runtime/runtime.go index c95bbcd3..f124ea9b 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" @@ -16,6 +17,8 @@ import ( agentcontext "neo-code/internal/context" 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" ) @@ -29,21 +32,92 @@ const ( providerRetryMaxWait = 5 * time.Second ) -var runtimeSessionWorkdirs = struct { - mu sync.RWMutex - data map[string]string -}{ - data: make(map[string]string), +// streamAccumulator 在流式事件处理过程中累积本轮对话需要持久化的助手消息状态, +// 包括文本内容和工具调用列表。 +type streamAccumulator struct { + content strings.Builder + toolCalls map[int]*providertypes.ToolCall +} + +// newStreamAccumulator 创建并初始化一个空的流式事件累积器。 +func newStreamAccumulator() *streamAccumulator { + return &streamAccumulator{ + toolCalls: make(map[int]*providertypes.ToolCall), + } +} + +// accumulateTextDelta 累积文本增量片段。 +func (a *streamAccumulator) accumulateTextDelta(text string) { + a.content.WriteString(text) +} + +// ensureToolCall 返回指定索引的工具调用条目,不存在时会先创建占位对象。 +func (a *streamAccumulator) ensureToolCall(index int) *providertypes.ToolCall { + call, exists := a.toolCalls[index] + if !exists { + call = &providertypes.ToolCall{} + a.toolCalls[index] = call + } + return call +} + +// accumulateToolCallStart 记录新发现的工具调用(首次出现时创建条目)。 +func (a *streamAccumulator) accumulateToolCallStart(index int, id, name string) { + 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 := a.ensureToolCall(index) + if strings.TrimSpace(id) != "" { + call.ID = id + } + call.Arguments += argumentsDelta +} + +// buildMessage 从累积状态构建最终的 assistant Message 对象,并校验工具调用元数据是否完整。 +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 := providertypes.Message{ + Role: providertypes.RoleAssistant, + Content: a.content.String(), + } + for _, index := range ordered { + call := a.toolCalls[index] + if call == nil { + continue + } + if strings.TrimSpace(call.ID) == "" { + return providertypes.Message{}, fmt.Errorf("runtime: provider emitted tool call %d without id", index) + } + if strings.TrimSpace(call.Name) == "" { + return providertypes.Message{}, fmt.Errorf("runtime: provider emitted tool call %d without name", index) + } + message.ToolCalls = append(message.ToolCalls, *call) + } + return message, nil } 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) - 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 { @@ -59,7 +133,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 与本轮发给模型的消息上下文。 @@ -75,7 +149,7 @@ type Service struct { func NewWithFactory( configManager *config.Manager, toolManager tools.Manager, - sessionStore Store, + sessionStore agentsession.Store, providerFactory ProviderFactory, contextBuilder agentcontext.Builder, ) *Service { @@ -86,7 +160,7 @@ func NewWithFactory( toolManager = tools.NewRegistry() } if contextBuilder == nil { - contextBuilder = agentcontext.NewBuilder() + contextBuilder = agentcontext.NewBuilderWithToolPolicies(toolManager) } return &Service{ @@ -121,8 +195,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) @@ -157,6 +231,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) @@ -169,7 +246,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, providertypes.ChatRequest{ Model: cfg.CurrentModel, SystemPrompt: builtContext.SystemPrompt, Messages: builtContext.Messages, @@ -186,9 +263,12 @@ func (s *Service) Run(ctx context.Context, input UserInput) error { session.Provider = cfg.SelectedProvider session.Model = cfg.CurrentModel - assistant := resp.Message + 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 + assistant.Role = providertypes.RoleAssistant } if strings.TrimSpace(assistant.Content) != "" || len(assistant.ToolCalls) > 0 { @@ -218,18 +298,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,14 +318,11 @@ 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()) } - toolMessage := provider.Message{ - Role: provider.RoleTool, + toolMessage := providertypes.Message{ + Role: providertypes.RoleTool, Content: result.Content, ToolCallID: call.ID, IsError: result.IsError, @@ -295,65 +367,44 @@ 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 } 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( @@ -362,37 +413,38 @@ 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) - s.setSessionWorkdir(session.ID, sessionWorkdir) + session := agentsession.NewWithWorkdir(title, 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) != "" { return session, nil } resolved, err := resolveWorkdirForSession(defaultWorkdir, session.Workdir, requestedWorkdir) if err != nil { - return Session{}, err + return agentsession.Session{}, err } if session.Workdir == resolved { 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 } @@ -417,21 +469,88 @@ func (s *Service) emit(ctx context.Context, kind EventType, runID string, sessio } } -// forwardProviderEvents 将 provider 流式事件转发为 runtime 事件。 +// handleProviderStreamEvent 解析并应用单条 provider 流式事件,缺失载荷或未知类型时返回错误。 +func handleProviderStreamEvent( + event providertypes.StreamEvent, + acc *streamAccumulator, + onTextDelta func(string), + onToolCallStart func(providertypes.ToolCallStartPayload), +) error { + switch event.Type { + case providertypes.StreamEventTextDelta: + payload, err := event.TextDeltaValue() + if err != nil { + return err + } + if onTextDelta != nil { + onTextDelta(payload.Text) + } + if acc != nil { + acc.accumulateTextDelta(payload.Text) + } + case providertypes.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 providertypes.StreamEventToolCallDelta: + payload, err := event.ToolCallDeltaValue() + if err != nil { + return err + } + if acc != nil { + acc.accumulateToolCallDelta(payload.Index, payload.ID, payload.ArgumentsDelta) + } + case providertypes.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{}) { - defer close(done) +func (s *Service) forwardProviderEvents( + ctx context.Context, + runID string, + sessionID string, + input <-chan providertypes.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: - s.emit(ctx, EventAgentChunk, runID, sessionID, event.Text) - case provider.StreamEventToolCallStart: - s.emit(ctx, EventToolCallThinking, runID, sessionID, event.ToolName) + err := handleProviderStreamEvent( + event, + acc, + func(text string) { + s.emit(ctx, EventAgentChunk, runID, sessionID, text) + }, + func(payload providertypes.ToolCallStartPayload) { + s.emit(ctx, EventToolCallThinking, runID, sessionID, payload.Name) + }, + ) + if err != nil && forwardErr == nil { + // 记录首个协议错误后继续排空事件通道,避免 provider 在后续发送时阻塞。 + forwardErr = err } case <-ctx.Done(): return @@ -489,18 +608,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) { + req providertypes.ChatRequest, +) (*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 +631,54 @@ 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) + streamEvents := make(chan providertypes.StreamEvent, 32) + streamDone := make(chan error, 1) + 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 + 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 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 计算指数退避 + 随机抖动的等待时间。 @@ -578,6 +710,7 @@ type permissionEventView struct { decision string reason string ruleID string + scope string resolvedAs string } @@ -609,6 +742,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 +758,7 @@ func (v permissionEventView) toRequestPayload() PermissionRequestPayload { Decision: v.decision, Reason: v.reason, RuleID: v.ruleID, - RememberScope: "", + RememberScope: v.scope, } } @@ -639,7 +773,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..988f6201 100644 --- a/internal/runtime/runtime_test.go +++ b/internal/runtime/runtime_test.go @@ -14,17 +14,19 @@ 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" + 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 +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 @@ -49,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 } @@ -61,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, @@ -90,14 +92,19 @@ func (s *memoryStore) ListSummaries(ctx context.Context) ([]SessionSummary, erro type scriptedProvider struct { name string - responses []provider.ChatResponse - streams [][]provider.StreamEvent - requests []provider.ChatRequest + streams [][]providertypes.StreamEvent + responses []scriptedResponse + requests []providertypes.ChatRequest callCount int - chatFn func(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) (provider.ChatResponse, error) + chatFn func(ctx context.Context, req providertypes.ChatRequest, events chan<- providertypes.StreamEvent) error } -func (p *scriptedProvider) Chat(ctx context.Context, req provider.ChatRequest, events chan<- provider.StreamEvent) (provider.ChatResponse, error) { +type scriptedResponse struct { + Message providertypes.Message + FinishReason string +} + +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 @@ -112,15 +119,39 @@ 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) + if callIndex < len(p.responses) { + response := p.responses[callIndex] + for index, toolCall := range response.Message.ToolCalls { + select { + case events <- providertypes.NewToolCallStartStreamEvent(index, toolCall.ID, toolCall.Name): + case <-ctx.Done(): + return ctx.Err() + } + select { + case events <- providertypes.NewToolCallDeltaStreamEvent(index, toolCall.ID, toolCall.Arguments): + case <-ctx.Done(): + return ctx.Err() + } + } + if response.Message.Content != "" { + select { + case events <- providertypes.NewTextDeltaStreamEvent(response.Message.Content): + case <-ctx.Done(): + return ctx.Err() + } + } + select { + case events <- providertypes.NewMessageDoneStreamEvent(response.FinishReason, nil): + case <-ctx.Done(): + return ctx.Err() + } } - return p.responses[callIndex], nil + + return nil } type scriptedProviderFactory struct { @@ -144,6 +175,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 +193,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 @@ -193,21 +229,28 @@ 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 + policies map[string]tools.MicroCompactPolicy 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) { +func (m *stubToolManager) ListAvailableSpecs(ctx context.Context, input tools.SpecListInput) ([]providertypes.ToolSpec, error) { m.listCalls++ if err := ctx.Err(); err != nil { return nil, err @@ -215,7 +258,14 @@ 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 { + 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) { @@ -228,12 +278,24 @@ 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 input UserInput - providerResponses []provider.ChatResponse - providerStreams [][]provider.StreamEvent + providerStreams [][]providertypes.StreamEvent registerTool tools.Tool contextBuilder agentcontext.Builder expectProviderCalls int @@ -245,26 +307,17 @@ 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{ + providerStreams: [][]providertypes.StreamEvent{ { - {Type: provider.StreamEventTextDelta, Text: "plain "}, - {Type: provider.StreamEventTextDelta, Text: "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 @@ -293,26 +346,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: [][]providertypes.StreamEvent{ { - Message: provider.Message{ - Role: "assistant", - ToolCalls: []provider.ToolCall{ - { - ID: "call-1", - Name: "filesystem_edit", - Arguments: `{"path":"main.go"}`, - }, - }, - }, - FinishReason: "tool_calls", + providertypes.NewToolCallStartStreamEvent(0, "call-1", "filesystem_edit"), + providertypes.NewToolCallDeltaStreamEvent(0, "call-1", `{"path":"main.go"}`), }, { - Message: provider.Message{ - Role: "assistant", - Content: "done", - }, - FinishReason: "stop", + providertypes.NewTextDeltaStreamEvent("done"), }, }, registerTool: &stubTool{ @@ -371,8 +413,7 @@ func TestServiceRun(t *testing.T) { } scripted := &scriptedProvider{ - responses: tt.providerResponses, - streams: tt.providerStreams, + streams: tt.providerStreams, } factory := &scriptedProviderFactory{provider: scripted} @@ -409,6 +450,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: [][]providertypes.StreamEvent{ + { + providertypes.NewToolCallDeltaStreamEvent(0, "", `{"path":"main.go"`), + providertypes.NewToolCallStartStreamEvent(0, "call-late", "filesystem_edit"), + providertypes.NewToolCallDeltaStreamEvent(0, "call-late", `}`), + }, + {providertypes.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: [][]providertypes.StreamEvent{ + { + providertypes.NewToolCallStartStreamEvent(0, "", "filesystem_edit"), + providertypes.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: [][]providertypes.StreamEvent{ + { + {Type: providertypes.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 := []providertypes.StreamEvent{{Type: providertypes.StreamEventTextDelta}} + for i := 0; i < 40; i++ { + stream = append(stream, providertypes.NewTextDeltaStreamEvent("ignored")) + } + scripted := &scriptedProvider{ + streams: [][]providertypes.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 @@ -418,7 +595,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) @@ -431,6 +608,9 @@ func TestServiceRunDelegatesToContextBuilder(t *testing.T) { manager := newRuntimeConfigManager(t) store := newMemoryStore() + session := agentsession.New("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"}) @@ -438,7 +618,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 @@ -446,14 +626,8 @@ func TestServiceRunDelegatesToContextBuilder(t *testing.T) { } scripted := &scriptedProvider{ - responses: []provider.ChatResponse{ - { - Message: provider.Message{ - Role: "assistant", - Content: "done", - }, - FinishReason: "stop", - }, + streams: [][]providertypes.StreamEvent{ + {providertypes.NewTextDeltaStreamEvent("done")}, }, } @@ -478,6 +652,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,21 +669,61 @@ func TestServiceRunDelegatesToContextBuilder(t *testing.T) { } } -func TestServiceRunPersistsSessionProviderAndModel(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([]providertypes.Message(nil), input.Messages...), + }, nil + }, + } + scripted := &scriptedProvider{ - responses: []provider.ChatResponse{{ - Message: provider.Message{Role: provider.RoleAssistant, Content: "done"}, + responses: []scriptedResponse{{ + Message: providertypes.Message{Role: providertypes.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() + + manager := newRuntimeConfigManager(t) + store := newMemoryStore() + registry := tools.NewRegistry() + registry.Register(&stubTool{name: "filesystem_read_file", content: "default"}) + + scripted := &scriptedProvider{ + streams: [][]providertypes.StreamEvent{ + {providertypes.NewTextDeltaStreamEvent("done")}, + }, + } + service := NewWithFactory(manager, registry, store, &scriptedProviderFactory{provider: scripted}, nil) if err := service.Run(context.Background(), UserInput{RunID: "run-session-provider-model", Content: "hello"}); err != nil { t.Fatalf("Run() error = %v", err) @@ -522,6 +739,131 @@ 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 := agentsession.New("preserve history") + session.ID = "session-preserve-history" + session.Messages = []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older user"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-1", Name: "preserve_tool", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "preserved result"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-2", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-3", Name: "webfetch", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, + } + store.sessions[session.ID] = cloneSession(session) + + scripted := &scriptedProvider{ + responses: []scriptedResponse{{ + Message: providertypes.Message{Role: providertypes.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 TestServiceRunDefaultBuilderUsesGenericToolManagerMicroCompactPolicies(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + store := newMemoryStore() + toolManager := &stubToolManager{ + policies: map[string]tools.MicroCompactPolicy{ + "preserve_tool": tools.MicroCompactPolicyPreserveHistory, + }, + } + + session := agentsession.New("preserve history by manager") + session.ID = "session-preserve-history-manager" + session.Messages = []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "older user"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-1", Name: "preserve_tool", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-1", Content: "preserved result"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-2", Name: "bash", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-2", Content: "recent bash result"}, + { + Role: providertypes.RoleAssistant, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-3", Name: "webfetch", Arguments: "{}"}, + }, + }, + {Role: providertypes.RoleTool, ToolCallID: "call-3", Content: "latest webfetch result"}, + } + store.sessions[session.ID] = cloneSession(session) + + scripted := &scriptedProvider{ + responses: []scriptedResponse{{ + Message: providertypes.Message{Role: providertypes.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() @@ -539,12 +881,12 @@ 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" - session.Messages = []provider.Message{ - {Role: provider.RoleUser, Content: "earlier"}, + session.Messages = []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "earlier"}, } store.sessions[session.ID] = cloneSession(session) @@ -584,7 +926,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{ @@ -594,23 +936,12 @@ func TestServiceRunUsesToolManager(t *testing.T) { } scripted := &scriptedProvider{ - responses: []provider.ChatResponse{ - { - Message: provider.Message{ - Role: "assistant", - ToolCalls: []provider.ToolCall{ - {ID: "call-manager", Name: "filesystem_edit", Arguments: `{"path":"main.go"}`}, - }, - }, - FinishReason: "tool_calls", - }, + streams: [][]providertypes.StreamEvent{ { - Message: provider.Message{ - Role: "assistant", - Content: "done", - }, - FinishReason: "stop", + providertypes.NewToolCallStartStreamEvent(0, "call-manager", "filesystem_edit"), + providertypes.NewToolCallDeltaStreamEvent(0, "call-manager", `{"path":"main.go"}`), }, + {providertypes.NewTextDeltaStreamEvent("done")}, }, } @@ -635,7 +966,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 } @@ -645,22 +976,141 @@ func TestServiceRunUsesToolManager(t *testing.T) { } } -func TestServiceRunEmitsPermissionRequestAndResolvedForAsk(t *testing.T) { +func TestServiceRunWaitsForPermissionResolutionAndContinues(t *testing.T) { + t.Parallel() + + manager := newRuntimeConfigManager(t) + store := newMemoryStore() + session := agentsession.New("memory reject") + session.ID = "session-memory-reject" + store.sessions[session.ID] = cloneSession(session) + registry := tools.NewRegistry() + tool := &stubTool{name: "webfetch", content: "fetched"} + 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{ + streams: [][]providertypes.StreamEvent{ + { + providertypes.NewToolCallStartStreamEvent(0, "call-ask", "webfetch"), + providertypes.NewToolCallDeltaStreamEvent(0, "call-ask", `{"url":"https://example.com/private"}`), + }, + {providertypes.NewTextDeltaStreamEvent("done")}, + }, + } + + 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-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 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) + } + + 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 != 1 { + t.Fatalf("expected allowed tool to execute once, got %d", tool.callCount) + } + + events := collectRuntimeEvents(service.Events()) + assertEventSequence(t, events, []EventType{ + EventPermissionResolved, + EventToolResult, + EventAgentDone, + }) + assertNoEventType(t, events, EventError) + + var resolvedPayload PermissionResolvedPayload + for _, event := range events { + switch event.Type { + case EventPermissionResolved: + payload, ok := event.Payload.(PermissionResolvedPayload) + if !ok { + t.Fatalf("expected PermissionResolvedPayload, got %#v", event.Payload) + } + resolvedPayload = payload + } + } + + if resolvedPayload.ToolName != "webfetch" || resolvedPayload.Decision != "allow" { + t.Fatalf("unexpected permission resolved payload: %+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) + } +} + +func TestServiceRunEmitsPermissionResolvedForDeny(t *testing.T) { t.Parallel() manager := newRuntimeConfigManager(t) store := newMemoryStore() registry := tools.NewRegistry() - tool := &stubTool{name: "webfetch", content: "should-not-run"} + tool := &stubTool{name: "bash", 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", + ID: "deny-bash", + Type: security.ActionTypeBash, + Resource: "bash", + Decision: security.DecisionDeny, + Reason: "bash denied", }, }) if err != nil { @@ -672,25 +1122,17 @@ 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: [][]providertypes.StreamEvent{ { - Message: provider.Message{Role: "assistant", Content: "done"}, - FinishReason: "stop", + providertypes.NewToolCallStartStreamEvent(0, "call-deny", "bash"), + providertypes.NewToolCallDeltaStreamEvent(0, "call-deny", `{"command":"echo hi"}`), }, + {providertypes.NewTextDeltaStreamEvent("done")}, }, } 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 { + if err := service.Run(context.Background(), UserInput{RunID: "run-permission-deny", Content: "run bash"}); err != nil { t.Fatalf("Run() error = %v", err) } if tool.callCount != 0 { @@ -699,64 +1141,51 @@ func TestServiceRunEmitsPermissionRequestAndResolvedForAsk(t *testing.T) { events := collectRuntimeEvents(service.Events()) assertEventSequence(t, events, []EventType{ - EventPermissionRequest, EventPermissionResolved, EventToolResult, EventAgentDone, }) + assertNoEventType(t, events, EventPermissionRequest) assertNoEventType(t, events, EventError) - var ( - requestPayload PermissionRequestPayload - 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 { - t.Fatalf("expected PermissionResolvedPayload, got %#v", event.Payload) - } - resolvedPayload = payload + if event.Type != EventPermissionResolved { + continue } + payload, ok := event.Payload.(PermissionResolvedPayload) + if !ok { + t.Fatalf("expected PermissionResolvedPayload, got %#v", event.Payload) + } + if payload.ToolName != "bash" || payload.Decision != "deny" || payload.ResolvedAs != "denied" { + t.Fatalf("unexpected permission resolved payload: %+v", payload) + } + if payload.RuleID != "deny-bash" { + t.Fatalf("expected deny-bash rule id, got %+v", payload) + } + return } - - 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" { - t.Fatalf("unexpected permission resolved payload: %+v", resolvedPayload) - } - if resolvedPayload.ResolvedAs != "rejected" { - t.Fatalf("expected resolved_as rejected, got %+v", resolvedPayload) - } + t.Fatalf("expected permission resolved event payload") } -func TestServiceRunEmitsPermissionResolvedForDeny(t *testing.T) { +func TestServiceRunEmitsRememberScopeWhenSessionRejectMemoryHits(t *testing.T) { t.Parallel() manager := newRuntimeConfigManager(t) store := newMemoryStore() + session := agentsession.New("memory reject") + session.ID = "session-memory-reject" + store.sessions[session.ID] = cloneSession(session) registry := tools.NewRegistry() - tool := &stubTool{name: "bash", content: "should-not-run"} + tool := &stubTool{name: "webfetch", content: "should-not-run"} registry.Register(tool) engine, err := security.NewStaticGateway(security.DecisionAllow, []security.Rule{ { - ID: "deny-bash", - Type: security.ActionTypeBash, - Resource: "bash", - Decision: security.DecisionDeny, - Reason: "bash denied", + ID: "ask-webfetch", + Type: security.ActionTypeRead, + Resource: "webfetch", + Decision: security.DecisionAsk, + Reason: "requires approval", }, }) if err != nil { @@ -766,41 +1195,52 @@ func TestServiceRunEmitsPermissionResolvedForDeny(t *testing.T) { 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{ + responses: []scriptedResponse{ { - Message: provider.Message{ + Message: providertypes.Message{ Role: "assistant", - ToolCalls: []provider.ToolCall{ - {ID: "call-deny", Name: "bash", Arguments: `{"command":"echo hi"}`}, + ToolCalls: []providertypes.ToolCall{ + {ID: "call-memory-reject", Name: "webfetch", Arguments: `{"url":"https://example.com/private"}`}, }, }, FinishReason: "tool_calls", }, { - Message: provider.Message{Role: "assistant", Content: "done"}, + Message: providertypes.Message{Role: "assistant", Content: "done"}, FinishReason: "stop", }, }, } service := NewWithFactory(manager, toolManager, store, &scriptedProviderFactory{provider: scripted}, nil) - if err := service.Run(context.Background(), UserInput{RunID: "run-permission-deny", Content: "run bash"}); err != 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 blocked tool not to execute, got %d", tool.callCount) + t.Fatalf("expected remembered reject to skip tool execution, got %d", tool.callCount) } events := collectRuntimeEvents(service.Events()) - assertEventSequence(t, events, []EventType{ - EventPermissionResolved, - EventToolResult, - EventAgentDone, - }) + assertEventSequence(t, events, []EventType{EventPermissionResolved, EventToolResult, EventAgentDone}) assertNoEventType(t, events, EventPermissionRequest) - assertNoEventType(t, events, EventError) for _, event := range events { if event.Type != EventPermissionResolved { @@ -810,11 +1250,8 @@ func TestServiceRunEmitsPermissionResolvedForDeny(t *testing.T) { if !ok { t.Fatalf("expected PermissionResolvedPayload, got %#v", event.Payload) } - if payload.ToolName != "bash" || payload.Decision != "deny" || payload.ResolvedAs != "denied" { - t.Fatalf("unexpected permission resolved payload: %+v", payload) - } - if payload.RuleID != "deny-bash" { - t.Fatalf("expected deny-bash rule id, got %+v", payload) + if payload.RememberScope != string(tools.SessionPermissionScopeReject) { + t.Fatalf("expected remember_scope reject, got %+v", payload) } return } @@ -843,7 +1280,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) } } @@ -853,14 +1290,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: [][]providertypes.StreamEvent{ + {providertypes.NewTextDeltaStreamEvent("done")}, }, }, }, nil) @@ -881,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) @@ -902,15 +1333,10 @@ func TestServiceRunErrorPaths(t *testing.T) { input: UserInput{RunID: "run-max-loops", Content: "loop"}, maxLoops: 1, provider: &scriptedProvider{ - responses: []provider.ChatResponse{ + streams: [][]providertypes.StreamEvent{ { - Message: provider.Message{ - Role: "assistant", - ToolCalls: []provider.ToolCall{ - {ID: "loop-call", Name: "filesystem_edit", Arguments: `{"path":"x"}`}, - }, - }, - FinishReason: "tool_calls", + providertypes.NewToolCallStartStreamEvent(0, "loop-call", "filesystem_edit"), + providertypes.NewToolCallDeltaStreamEvent(0, "loop-call", `{"path":"x"}`), }, }, }, @@ -946,22 +1372,16 @@ func TestServiceRunErrorPaths(t *testing.T) { Content: "continue", }, provider: &scriptedProvider{ - responses: []provider.ChatResponse{ - { - Message: provider.Message{ - Role: "assistant", - Content: "resumed", - }, - FinishReason: "stop", - }, + streams: [][]providertypes.StreamEvent{ + {providertypes.NewTextDeltaStreamEvent("resumed")}, }, }, - seedSession: &Session{ + seedSession: &agentsession.Session{ ID: "existing-session", Title: "Resume Me", - CreatedAt: newSession("seed").CreatedAt, - UpdatedAt: newSession("seed").UpdatedAt, - Messages: []provider.Message{ + CreatedAt: agentsession.New("seed").CreatedAt, + UpdatedAt: agentsession.New("seed").UpdatedAt, + Messages: []providertypes.Message{ {Role: "user", Content: "earlier"}, }, }, @@ -984,23 +1404,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 providertypes.ChatRequest, events chan<- providertypes.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 <- providertypes.NewTextDeltaStreamEvent("recovered") + return nil }, } }(), @@ -1024,8 +1439,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 providertypes.ChatRequest, events chan<- providertypes.StreamEvent) error { + return &provider.ProviderError{ StatusCode: 401, Code: provider.ErrorCodeAuthFailed, Message: "invalid api key", @@ -1047,8 +1462,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 providertypes.ChatRequest, events chan<- providertypes.StreamEvent) error { + return &provider.ProviderError{ StatusCode: 500, Code: provider.ErrorCodeServer, Message: "internal server error", @@ -1131,10 +1546,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 providertypes.ChatRequest, events chan<- providertypes.StreamEvent) error { close(started) <-ctx.Done() - return provider.ChatResponse{}, ctx.Err() + return ctx.Err() }, } @@ -1174,10 +1589,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 providertypes.ChatRequest, events chan<- providertypes.StreamEvent) error { close(started) <-ctx.Done() - return provider.ChatResponse{}, ctx.Err() + return ctx.Err() }, } @@ -1219,10 +1634,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 providertypes.ChatRequest, events chan<- providertypes.StreamEvent) error { close(started) <-ctx.Done() - return provider.ChatResponse{}, providerErr + return providerErr }, } @@ -1271,15 +1686,10 @@ func TestServiceRunCanceledDuringToolExecution(t *testing.T) { scripted := &scriptedProvider{ name: "tool-cancel-provider", - responses: []provider.ChatResponse{ + streams: [][]providertypes.StreamEvent{ { - Message: provider.Message{ - Role: "assistant", - ToolCalls: []provider.ToolCall{ - {ID: "cancel-call", Name: "filesystem_edit", Arguments: `{"path":"main.go"}`}, - }, - }, - FinishReason: "tool_calls", + providertypes.NewToolCallStartStreamEvent(0, "cancel-call", "filesystem_edit"), + providertypes.NewToolCallDeltaStreamEvent(0, "cancel-call", `{"path":"main.go"}`), }, }, } @@ -1336,15 +1746,10 @@ func TestServiceRunPreservesToolErrorAfterCancel(t *testing.T) { scripted := &scriptedProvider{ name: "tool-error-after-cancel-provider", - responses: []provider.ChatResponse{ + streams: [][]providertypes.StreamEvent{ { - Message: provider.Message{ - Role: "assistant", - ToolCalls: []provider.ToolCall{ - {ID: "tool-error-call", Name: "filesystem_edit", Arguments: `{"path":"main.go"}`}, - }, - }, - FinishReason: "tool_calls", + providertypes.NewToolCallStartStreamEvent(0, "tool-error-call", "filesystem_edit"), + providertypes.NewToolCallDeltaStreamEvent(0, "tool-error-call", `{"path":"main.go"}`), }, }, } @@ -1433,23 +1838,12 @@ func TestServiceRunToolTimeoutIsNotCancellation(t *testing.T) { scripted := &scriptedProvider{ name: "timeout-provider", - responses: []provider.ChatResponse{ - { - Message: provider.Message{ - Role: "assistant", - ToolCalls: []provider.ToolCall{ - {ID: "timeout-call", Name: "filesystem_edit", Arguments: `{"path":"main.go"}`}, - }, - }, - FinishReason: "tool_calls", - }, + streams: [][]providertypes.StreamEvent{ { - Message: provider.Message{ - Role: "assistant", - Content: "done after timeout", - }, - FinishReason: "stop", + providertypes.NewToolCallStartStreamEvent(0, "timeout-call", "filesystem_edit"), + providertypes.NewToolCallDeltaStreamEvent(0, "timeout-call", `{"path":"main.go"}`), }, + {providertypes.NewTextDeltaStreamEvent("done after timeout")}, }, } @@ -1476,12 +1870,12 @@ 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"}, - {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) @@ -1491,9 +1885,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{ @@ -1534,12 +1928,12 @@ 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"}, - {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) @@ -1588,14 +1982,14 @@ 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" - 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) @@ -1603,28 +1997,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: [][]providertypes.StreamEvent{ + {providertypes.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) @@ -1665,12 +2056,12 @@ 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"}, - {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) @@ -1678,28 +2069,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: [][]providertypes.StreamEvent{ + {providertypes.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) @@ -1725,11 +2113,11 @@ 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"}, - {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) @@ -1738,20 +2126,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: [][]providertypes.StreamEvent{ { - Message: provider.Message{Role: "assistant", Content: "done"}, - FinishReason: "stop", + providertypes.NewToolCallStartStreamEvent(0, "call-1", "filesystem_read_file"), + providertypes.NewToolCallDeltaStreamEvent(0, "call-1", `{"path":"main.go"}`), }, + {providertypes.NewTextDeltaStreamEvent("done")}, }, } @@ -1759,9 +2139,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{ @@ -1816,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) @@ -1826,17 +2206,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 providertypes.ChatRequest, events chan<- providertypes.StreamEvent) error { select { case <-providerStarted: default: close(providerStarted) } <-unblockProvider - return provider.ChatResponse{ - Message: provider.Message{Role: provider.RoleAssistant, Content: "done"}, - FinishReason: "stop", - }, nil + events <- providertypes.NewTextDeltaStreamEvent("done") + return nil }, } @@ -1846,7 +2224,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, @@ -1914,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()) @@ -1933,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") } @@ -1951,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"} @@ -1960,20 +2338,12 @@ func TestServiceRunUsesSessionWorkdirForContextAndTools(t *testing.T) { builder := &stubContextBuilder{} scripted := &scriptedProvider{ - responses: []provider.ChatResponse{ - { - Message: provider.Message{ - Role: "assistant", - ToolCalls: []provider.ToolCall{ - {ID: "call-session-workdir", Name: "filesystem_edit", Arguments: `{"path":"main.go"}`}, - }, - }, - FinishReason: "tool_calls", - }, + streams: [][]providertypes.StreamEvent{ { - Message: provider.Message{Role: "assistant", Content: "done"}, - FinishReason: "stop", + providertypes.NewToolCallStartStreamEvent(0, "call-session-workdir", "filesystem_edit"), + providertypes.NewToolCallDeltaStreamEvent(0, "call-session-workdir", `{"path":"main.go"}`), }, + {providertypes.NewTextDeltaStreamEvent("done")}, }, } @@ -2010,11 +2380,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: [][]providertypes.StreamEvent{ + {providertypes.NewTextDeltaStreamEvent("done")}, }, } @@ -2051,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"}) @@ -2078,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") @@ -2182,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)) @@ -2190,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) { @@ -2247,22 +2614,22 @@ 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...) + 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 } @@ -2332,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"}) @@ -2418,3 +2785,252 @@ 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) + } +} + +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: [][]providertypes.StreamEvent{ + {providertypes.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( + providertypes.StreamEvent{Type: providertypes.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( + providertypes.StreamEvent{Type: providertypes.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( + providertypes.StreamEvent{Type: providertypes.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 providertypes.ChatRequest, events chan<- providertypes.StreamEvent) error { + events <- providertypes.StreamEvent{Type: providertypes.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", + providertypes.ChatRequest{ + Model: "test-model", + SystemPrompt: "prompt", + Messages: []providertypes.Message{{Role: providertypes.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/runtime/session_test.go b/internal/runtime/session_test.go deleted file mode 100644 index a82f1215..00000000 --- a/internal/runtime/session_test.go +++ /dev/null @@ -1,173 +0,0 @@ -package runtime - -import ( - "context" - "os" - "path/filepath" - "strings" - "testing" - "time" - - "neo-code/internal/provider" -) - -func TestJSONSessionStoreSaveLoadAndListSummaries(t *testing.T) { - t.Parallel() - - baseDir := t.TempDir() - store := NewJSONSessionStore(baseDir) - - older := &Session{ - ID: "session-old", - Title: "Old Session", - CreatedAt: time.Now().Add(-2 * time.Hour), - UpdatedAt: time.Now().Add(-1 * time.Hour), - Messages: []provider.Message{ - {Role: "user", Content: "hello"}, - {Role: "assistant", Content: "world"}, - }, - } - newer := &Session{ - ID: "session-new", - Title: "New Session", - CreatedAt: time.Now().Add(-30 * time.Minute), - UpdatedAt: time.Now(), - Workdir: t.TempDir(), - Messages: []provider.Message{ - {Role: "user", Content: "new"}, - }, - } - - if err := store.Save(context.Background(), older); err != nil { - t.Fatalf("Save older session: %v", err) - } - if err := store.Save(context.Background(), newer); err != nil { - t.Fatalf("Save newer session: %v", err) - } - - loaded, err := store.Load(context.Background(), older.ID) - if err != nil { - t.Fatalf("Load() error: %v", err) - } - 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 len(loaded.Messages) != 2 || loaded.Messages[1].Content != "world" { - t.Fatalf("unexpected loaded messages: %+v", loaded.Messages) - } - - rawPath := filepath.Join(baseDir, sessionsDirName, newer.ID+".json") - raw, err := os.ReadFile(rawPath) - 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)) - } - - mustWriteRuntimeFile(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) - } - - summaries, err := store.ListSummaries(context.Background()) - if err != nil { - t.Fatalf("ListSummaries() error: %v", err) - } - if len(summaries) != 2 { - t.Fatalf("expected 2 summaries, got %d", len(summaries)) - } - if summaries[0].ID != newer.ID || summaries[1].ID != older.ID { - t.Fatalf("expected summaries sorted by UpdatedAt desc, got %+v", summaries) - } -} - -func TestJSONSessionStoreErrors(t *testing.T) { - t.Parallel() - - baseDir := t.TempDir() - store := NewJSONSessionStore(baseDir) - - cancelledCtx, cancel := context.WithCancel(context.Background()) - cancel() - - if err := store.Save(cancelledCtx, &Session{ID: "x"}); err == nil { - t.Fatalf("expected cancelled save to fail") - } - if err := store.Save(context.Background(), nil); err == nil { - t.Fatalf("expected nil session save to fail") - } - if _, err := store.Load(cancelledCtx, "missing"); err == nil { - t.Fatalf("expected cancelled load to fail") - } - if _, err := store.ListSummaries(cancelledCtx); err == nil { - t.Fatalf("expected cancelled list to fail") - } -} - -func TestJSONSessionStoreCorruptedSessionBehaviors(t *testing.T) { - t.Parallel() - - baseDir := t.TempDir() - store := NewJSONSessionStore(baseDir) - - valid := &Session{ - ID: "valid-session", - Title: "Valid Session", - CreatedAt: time.Now().Add(-time.Minute), - UpdatedAt: time.Now(), - Messages: []provider.Message{{Role: "user", Content: "hello"}}, - } - if err := store.Save(context.Background(), valid); err != nil { - t.Fatalf("Save valid session: %v", err) - } - - mustWriteRuntimeFile(t, filepath.Join(baseDir, sessionsDirName, "broken.json"), "{broken") - - _, err := store.Load(context.Background(), "broken") - if err == nil || !strings.Contains(err.Error(), "decode session broken") { - t.Fatalf("expected corrupted session decode error, got %v", err) - } - - 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 corrupted session file to be skipped, got %+v", summaries) - } -} - -func TestJSONSessionStoreSaveInvalidBaseDir(t *testing.T) { - t.Parallel() - - tempDir := t.TempDir() - baseFile := filepath.Join(tempDir, "not-a-directory") - if err := os.WriteFile(baseFile, []byte("x"), 0o644); err != nil { - t.Fatalf("write base file: %v", err) - } - - store := NewJSONSessionStore(baseFile) - err := store.Save(context.Background(), &Session{ - ID: "session-x", - Title: "Broken Save", - CreatedAt: time.Now(), - UpdatedAt: time.Now(), - }) - if err == nil || !strings.Contains(err.Error(), "create sessions dir") { - t.Fatalf("expected invalid base dir error, got %v", err) - } -} - -func mustWriteRuntimeFile(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) - } - if err := os.WriteFile(path, []byte(content), 0o644); err != nil { - t.Fatalf("write %s: %v", path, err) - } -} diff --git a/internal/runtime/workdir_branch_test.go b/internal/runtime/workdir_branch_test.go index 0edc39ef..089a3cd0 100644 --- a/internal/runtime/workdir_branch_test.go +++ b/internal/runtime/workdir_branch_test.go @@ -6,30 +6,9 @@ import ( "path/filepath" "strings" "testing" -) - -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) - } -} + agentsession "neo-code/internal/session" +) func TestResolveWorkdirForSessionAndNormalizeErrors(t *testing.T) { t.Parallel() @@ -82,7 +61,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/security/policy.go b/internal/security/policy.go new file mode 100644 index 00000000..f414655c --- /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"}, + 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"}, + 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..ab12d2ee --- /dev/null +++ b/internal/security/policy_test.go @@ -0,0 +1,189 @@ +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", + }, + { + 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 { + 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/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/session/id_test.go b/internal/session/id_test.go new file mode 100644 index 00000000..396779f8 --- /dev/null +++ b/internal/session/id_test.go @@ -0,0 +1,53 @@ +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) + } +} + +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/runtime/session.go b/internal/session/store.go similarity index 52% rename from internal/runtime/session.go rename to internal/session/store.go index af64317d..57639374 100644 --- a/internal/runtime/session.go +++ b/internal/session/store.go @@ -1,4 +1,4 @@ -package runtime +package session import ( "context" @@ -12,89 +12,98 @@ import ( "sync" "time" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" ) const sessionsDirName = "sessions" +// Session 表示单个会话的持久化模型,包含基础元数据与消息历史。 +// Provider / Model 用于在 compact 等流程中优先复用会话最近一次成功运行的模型配置。 type Session struct { ID string `json:"id"` Title string `json:"title"` // 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:"workdir,omitempty"` + 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,26 +175,30 @@ 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, Workdir: strings.TrimSpace(workdir), - Messages: []provider.Message{}, + Messages: []providertypes.Message{}, } } +// sanitizeTitle 规范化会话标题:去空白、空标题回退默认值、超长截断。 func sanitizeTitle(title string) string { title = strings.TrimSpace(title) if title == "" { diff --git a/internal/session/store_test.go b/internal/session/store_test.go new file mode 100644 index 00000000..6b412322 --- /dev/null +++ b/internal/session/store_test.go @@ -0,0 +1,381 @@ +package session + +import ( + "context" + "encoding/json" + "errors" + "os" + "path/filepath" + "strings" + "testing" + "time" + + providertypes "neo-code/internal/provider/types" +) + +func TestJSONStoreSaveLoadAndListSummaries(t *testing.T) { + t.Parallel() + + baseDir := t.TempDir() + store := NewJSONStore(baseDir) + + older := &Session{ + ID: "session-old", + Title: "Old Session", + CreatedAt: time.Now().Add(-2 * time.Hour), + UpdatedAt: time.Now().Add(-1 * time.Hour), + Messages: []providertypes.Message{ + {Role: "user", Content: "hello"}, + {Role: "assistant", Content: "world"}, + }, + } + newer := &Session{ + ID: "session-new", + Title: "New Session", + CreatedAt: time.Now().Add(-30 * time.Minute), + UpdatedAt: time.Now(), + Workdir: t.TempDir(), + Messages: []providertypes.Message{ + {Role: "user", Content: "new"}, + }, + } + + if err := store.Save(context.Background(), older); err != nil { + t.Fatalf("Save older session: %v", err) + } + if err := store.Save(context.Background(), newer); err != nil { + t.Fatalf("Save newer session: %v", err) + } + + loaded, err := store.Load(context.Background(), older.ID) + if err != nil { + t.Fatalf("Load() error: %v", err) + } + if loaded.Title != older.Title { + t.Fatalf("expected title %q, got %q", older.Title, loaded.Title) + } + 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) + } + + rawPath := filepath.Join(baseDir, sessionsDirName, newer.ID+".json") + raw, err := os.ReadFile(rawPath) + if err != nil { + t.Fatalf("read saved session: %v", err) + } + 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") + if err := os.MkdirAll(filepath.Join(baseDir, sessionsDirName, "directory"), 0o755); err != nil { + t.Fatalf("mkdir stray directory: %v", err) + } + + summaries, err := store.ListSummaries(context.Background()) + if err != nil { + t.Fatalf("ListSummaries() error: %v", err) + } + if len(summaries) != 2 { + t.Fatalf("expected 2 summaries, got %d", len(summaries)) + } + if summaries[0].ID != newer.ID || summaries[1].ID != older.ID { + t.Fatalf("expected summaries sorted by UpdatedAt desc, got %+v", summaries) + } +} + +func TestJSONStoreErrors(t *testing.T) { + t.Parallel() + + baseDir := t.TempDir() + store := NewJSONStore(baseDir) + + cancelledCtx, cancel := context.WithCancel(context.Background()) + cancel() + + if err := store.Save(cancelledCtx, &Session{ID: "x"}); err == nil { + t.Fatalf("expected cancelled save to fail") + } + if err := store.Save(context.Background(), nil); err == nil { + t.Fatalf("expected nil session save to fail") + } + if _, err := store.Load(cancelledCtx, "missing"); err == nil { + t.Fatalf("expected cancelled load to fail") + } + if _, err := store.ListSummaries(cancelledCtx); err == nil { + t.Fatalf("expected cancelled list to fail") + } +} + +func TestJSONStoreCorruptedSessionBehaviors(t *testing.T) { + t.Parallel() + + baseDir := t.TempDir() + store := NewJSONStore(baseDir) + + valid := &Session{ + ID: "valid-session", + Title: "Valid Session", + CreatedAt: time.Now().Add(-time.Minute), + UpdatedAt: time.Now(), + Messages: []providertypes.Message{{Role: "user", Content: "hello"}}, + } + if err := store.Save(context.Background(), valid); err != nil { + t.Fatalf("Save valid session: %v", err) + } + + 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") { + t.Fatalf("expected corrupted session decode error, got %v", err) + } + + 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 corrupted session file to be skipped, got %+v", summaries) + } +} + +func TestJSONStoreSaveInvalidBaseDir(t *testing.T) { + t.Parallel() + + tempDir := t.TempDir() + baseFile := filepath.Join(tempDir, "not-a-directory") + if err := os.WriteFile(baseFile, []byte("x"), 0o644); err != nil { + t.Fatalf("write base file: %v", err) + } + + store := NewJSONStore(baseFile) + err := store.Save(context.Background(), &Session{ + ID: "session-x", + Title: "Broken Save", + CreatedAt: time.Now(), + UpdatedAt: time.Now(), + }) + if err == nil || !strings.Contains(err.Error(), "create sessions dir") { + t.Fatalf("expected invalid base dir error, got %v", err) + } +} + +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 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", + Workdir: "/tmp/persist-workdir", + CreatedAt: time.Now().Add(-time.Hour), + UpdatedAt: time.Now(), + Messages: []providertypes.Message{ + {Role: providertypes.RoleUser, Content: "hello"}, + { + Role: providertypes.RoleAssistant, + Content: "calling tool", + ToolCalls: []providertypes.ToolCall{ + {ID: "call-1", Name: "webfetch", Arguments: `{"url":"https://example.com"}`}, + }, + }, + {Role: providertypes.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 decoded["workdir"] != session.Workdir { + t.Fatalf("expected workdir persisted as %q, got %+v", session.Workdir, decoded["workdir"]) + } +} + +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) + } + if err := os.WriteFile(path, []byte(content), 0o644); err != nil { + t.Fatalf("write %s: %v", path, err) + } +} diff --git a/internal/tools/bash/tool.go b/internal/tools/bash/tool.go index f2559d97..e02bce21 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 { @@ -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/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/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..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,17 +18,23 @@ 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 } // 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 } +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) @@ -50,6 +56,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 +119,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,20 +153,32 @@ 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 } // 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") } 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) { @@ -175,6 +203,28 @@ func (m *DefaultManager) Execute(ctx context.Context, input ToolCallInput) (Tool result.ToolCallID = input.ID return result, err } + // 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, + 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 +244,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 +280,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..45e7de66 100644 --- a/internal/tools/manager_test.go +++ b/internal/tools/manager_test.go @@ -8,6 +8,7 @@ import ( "strings" "testing" + providertypes "neo-code/internal/provider/types" "neo-code/internal/security" ) @@ -15,6 +16,7 @@ type managerStubTool struct { name string content string err error + policy MicroCompactPolicy callCount int lastCall ToolCallInput } @@ -25,6 +27,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 t.policy } + func (t *managerStubTool) Execute(ctx context.Context, call ToolCallInput) (ToolResult, error) { t.callCount++ t.lastCall = call @@ -36,17 +40,33 @@ func (t *managerStubTool) Execute(ctx context.Context, call ToolCallInput) (Tool type stubSandbox struct { err error + plan *security.WorkspaceExecutionPlan callCount int lastAction security.Action } +type executorWithoutMicroCompactPolicy struct{} + +func (executorWithoutMicroCompactPolicy) ListAvailableSpecs(ctx context.Context, input SpecListInput) ([]providertypes.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 if err := ctx.Err(); err != nil { return nil, err } - return nil, s.err + return s.plan, s.err } func TestDefaultManagerListAvailableSpecs(t *testing.T) { @@ -68,6 +88,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() @@ -343,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() @@ -377,6 +474,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,11 +491,364 @@ 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") } + 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 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("category matching shares decision across same tool category", 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) + } + + 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":"internal/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 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) + } + }) + + 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/mcp/adapter.go b/internal/tools/mcp/adapter.go new file mode 100644 index 00000000..d9ecb470 --- /dev/null +++ b/internal/tools/mcp/adapter.go @@ -0,0 +1,142 @@ +package mcp + +import ( + "context" + "errors" + "fmt" + "strings" +) + +const mcpToolNamePrefix = "mcp." + +// AdapterFactory 基于 registry 快照构造 MCP tool 适配器集合。 +type AdapterFactory struct { + registry *Registry +} + +// NewAdapterFactory 创建 MCP adapter 工厂。 +func NewAdapterFactory(registry *Registry) *AdapterFactory { + return &AdapterFactory{registry: registry} +} + +// 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") + } + if err := ctx.Err(); err != nil { + return nil, err + } + + snapshots := f.registry.Snapshot() + if len(snapshots) == 0 { + return nil, nil + } + + result := make([]*Adapter, 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 适配为统一调用描述。 +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 +} + +// 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) != "" { + 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) +} + +// Call 分发 MCP tool 调用并返回统一结果。 +func (a *Adapter) Call(ctx context.Context, arguments []byte) (CallResult, error) { + if a == nil || a.registry == nil { + return CallResult{}, errors.New("mcp: adapter is not initialized") + } + if err := ctx.Err(); err != nil { + return CallResult{}, err + } + return a.registry.Call(ctx, a.serverID, a.toolName, arguments) +} + +// 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 +} diff --git a/internal/tools/mcp/adapter_test.go b/internal/tools/mcp/adapter_test.go new file mode 100644 index 00000000..2b38691f --- /dev/null +++ b/internal/tools/mcp/adapter_test.go @@ -0,0 +1,229 @@ +package mcp + +import ( + "context" + "errors" + "testing" +) + +func TestAdapterFactoryBuildAdapters(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) + adapters, err := factory.BuildAdapters(context.Background()) + if err != nil { + t.Fatalf("BuildAdapters() error = %v", err) + } + if len(adapters) != 1 { + t.Fatalf("expected one adapter, got %d", len(adapters)) + } + if adapters[0].FullName() != "mcp.docs.search" { + t.Fatalf("unexpected adapter full name: %q", adapters[0].FullName()) + } +} + +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() + + 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.Call(context.Background(), []byte(`{"q":"mcp"}`)) + if err != nil { + t.Fatalf("Call() error = %v", err) + } + if result.Content != "result body" { + t.Fatalf("expected result content, got %q", result.Content) + } +} + +func TestAdapterCallError(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) + } + + if _, err := adapter.Call(context.Background(), []byte(`{"q":"mcp"}`)); err == nil { + 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"]) + } +} + +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") + } +} + +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.go b/internal/tools/mcp/registry.go new file mode 100644 index 00000000..490b5142 --- /dev/null +++ b/internal/tools/mcp/registry.go @@ -0,0 +1,339 @@ +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] = 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/registry_test.go b/internal/tools/mcp/registry_test.go new file mode 100644 index 00000000..30c5cd3f --- /dev/null +++ b/internal/tools/mcp/registry_test.go @@ -0,0 +1,336 @@ +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") + } +} + +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"]) + } +} + +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") + } +} + +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"]) + } +} diff --git a/internal/tools/mcp/stdio_client.go b/internal/tools/mcp/stdio_client.go new file mode 100644 index 00000000..b6808d51 --- /dev/null +++ b/internal/tools/mcp/stdio_client.go @@ -0,0 +1,853 @@ +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 + maxStdioFrameBytes = 8 * 1024 * 1024 + maxStdioLineBytes = 8 * 1024 * 1024 + defaultMCPProtocolVersion = "2024-11-05" + defaultMCPClientName = "neocode" + defaultMCPClientVersion = "0.1.0" +) + +// 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 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"` + 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 + 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{} + 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) == "" { + 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 +} + +// callContext 基于配置与上游截止时间生成单次 RPC 调用上下文。 +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) +} + +// 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 用于初始化阶段避免递归。 +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 + } + 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) + 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 + selectedProtocol := c.resolveWriteProtocolLocked(override) + 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", + ID: requestID, + Method: method, + Params: params, + }) + if err != nil { + c.removePending(requestID) + return nil, fmt.Errorf("mcp: marshal request: %w", err) + } + c.writeMu.Lock() + writeErr := writeMessageWithProtocol(stdin, requestPayload, selectedProtocol) + c.writeMu.Unlock() + if writeErr != nil { + c.removePending(requestID) + return nil, fmt.Errorf("mcp: send request: %w", writeErr) + } + + select { + case <-ctx.Done(): + c.removePending(requestID) + return nil, ctx.Err() + case reply := <-replyCh: + return reply.result, reply.err + } +} + +// 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 + } + 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 + selectedProtocol := c.resolveWriteProtocolLocked(override) + 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 := writeMessageWithProtocol(stdin, payload, selectedProtocol) + c.writeMu.Unlock() + if writeErr != nil { + return fmt.Errorf("mcp: send notification: %w", writeErr) + } + return nil +} + +// ensureStarted 确保 stdio 子进程已启动并处于可读写状态。 +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.initialized = false + c.initializing = false + c.initDone = nil + c.protocol = stdioProtocolUnknown + c.backoff = c.cfg.RestartBackoff + c.retryAt = time.Time{} + + go c.readLoop() + go c.waitLoop(command) + go io.Copy(io.Discard, stderr) + 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 握手并自动兼容 line/framed 两种 stdio 线协议。 +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, + }, + } + + protocols := []stdioProtocol{stdioProtocolLine, stdioProtocolFramed} + c.mu.Lock() + if c.protocol == stdioProtocolLine || c.protocol == stdioProtocolFramed { + protocols = []stdioProtocol{c.protocol} + } + 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 fmt.Errorf("mcp: initialize session: %s", strings.Join(errs, "; ")) +} + +// readLoop 持续消费 MCP server 响应并分发给对应 pending 请求。 +func (c *StdIOClient) readLoop() { + for { + 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 + } + 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} + } +} + +// waitLoop 等待子进程退出并触发统一下线处理。 +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)) +} + +// markExited 将客户端状态原子切换为已下线并唤醒所有等待请求。 +func (c *StdIOClient) markExited(err error) { + c.mu.Lock() + defer c.mu.Unlock() + + if !c.started { + return + } + c.started = false + c.initialized = false + if c.initializing && c.initDone != nil { + close(c.initDone) + } + c.initializing = false + c.initDone = nil + c.protocol = stdioProtocolUnknown + 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() +} + +// 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} + delete(c.pending, requestID) + } +} + +// bumpBackoffLocked 按指数退避策略更新下次可重启时间,调用方需持有 c.mu。 +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 + } +} + +// 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 { + return err + } + if _, err := writer.Write(payload); err != nil { + return err + } + return nil +} + +// 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, stdioProtocolUnknown, err + } + + trimmed := strings.TrimSpace(line) + if trimmed == "" { + continue + } + + lower := strings.ToLower(trimmed) + if strings.HasPrefix(lower, "content-length:") { + 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)) + } + 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") + } + 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 { + return nil, err + } + 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 { + 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"] = string(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..01d8d20b --- /dev/null +++ b/internal/tools/mcp/stdio_client_test.go @@ -0,0 +1,615 @@ +package mcp + +import ( + "bufio" + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "os" + "strconv" + "strings" + "sync" + "testing" + "time" +) + +type nopWriteCloser struct { + bytes.Buffer +} + +func (n *nopWriteCloser) Close() error { return nil } + +type errWriter struct{} + +func (errWriter) Write(p []byte) (int, error) { + return 0, errors.New("write failed") +} + +func TestStdIOClientListToolsAndCallTool(t *testing.T) { + t.Parallel() + + client := newTestStdIOClientWithMode(t, "framed") + 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 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 := newTestStdIOClientWithMode(t, "framed") + defer func() { _ = client.Close() }() + + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + if err := client.HealthCheck(ctx); err != nil { + t.Fatalf("HealthCheck() error = %v", err) + } +} + +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() + + client := newTestStdIOClientWithMode(t, "framed") + 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 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() + + 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") || errors.Is(err, io.EOF)) { + 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 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() + + 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 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 { + 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", "GO_MCP_STDIO_WIRE=" + wireMode}, + StartTimeout: 3 * time.Second, + CallTimeout: 3 * time.Second, + }) + if err != nil { + t.Fatalf("NewStdIOClient() error = %v", err) + } + 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", "GO_MCP_STDIO_WIRE=framed"}, + 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" + wireMode := strings.TrimSpace(os.Getenv("GO_MCP_STDIO_WIRE")) + if wireMode == "" { + wireMode = "framed" + } + initialized := !requireInitialize + + reader := bufio.NewReader(os.Stdin) + for { + 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) + } + 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": + 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, + "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": + 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{ + "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) + } + switch wireMode { + case "line": + err = writeLineMessage(os.Stdout, rawResponse) + default: + err = writeFramedMessage(os.Stdout, rawResponse) + } + if err != nil { + os.Exit(5) + } + } +} 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/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/registry.go b/internal/tools/registry.go index c057ee86..a392f629 100644 --- a/internal/tools/registry.go +++ b/internal/tools/registry.go @@ -6,24 +6,46 @@ import ( "sort" "strings" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" + "neo-code/internal/security" + "neo-code/internal/tools/mcp" ) type Registry struct { - tools map[string]Tool + tools map[string]Tool + microCompactPolicies map[string]MicroCompactPolicy + mcpRegistry *mcp.Registry + mcpFactory *mcp.AdapterFactory } func NewRegistry() *Registry { return &Registry{ - tools: map[string]Tool{}, + tools: map[string]Tool{}, + microCompactPolicies: map[string]MicroCompactPolicy{}, } } +// 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 } - 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) { @@ -36,21 +58,38 @@ 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) } -func (r *Registry) GetSpecs() []provider.ToolSpec { +// 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() []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(), @@ -59,38 +98,146 @@ 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 } - 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, providertypes.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 { - content := FormatError(input.Name, NormalizeErrorReason(input.Name, err), "") + 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, resolveErr), "") return ToolResult{ ToolCallID: input.ID, Name: input.Name, Content: content, IsError: true, - }, err + }, resolveErr } - - 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 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 + result = ApplyOutputLimit(result, DefaultOutputLimitBytes) + return result, callErr + } + if result.Content == "" { + result.Content = "ok" } + result = ApplyOutputLimit(result, DefaultOutputLimitBytes) 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") +} + +// 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 e8d2db35..a26843e4 100644 --- a/internal/tools/registry_test.go +++ b/internal/tools/registry_test.go @@ -5,12 +5,16 @@ import ( "errors" "strings" "testing" + + "neo-code/internal/security" + "neo-code/internal/tools/mcp" ) type stubTool struct { name string description string schema map[string]any + policy MicroCompactPolicy result ToolResult err error } @@ -20,6 +24,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 +166,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 +193,258 @@ 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) + } +} + +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) + } +} + +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) + } +} + +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) + } +} + +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) + } +} + +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) + } +} diff --git a/internal/tools/session_memory.go b/internal/tools/session_memory.go new file mode 100644 index 00000000..91006a14 --- /dev/null +++ b/internal/tools/session_memory.go @@ -0,0 +1,209 @@ +package tools + +import ( + "errors" + "fmt" + "net/url" + "path/filepath" + "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 { + return strings.Join([]string{ + string(action.Type), + sessionPermissionCategory(action), + sessionPermissionTargetScope(action), + }, "|") +} + +// sessionPermissionCategory 将安全动作归一为稳定的工具类别。 +// 类别用于聚合同类工具,再配合 target scope 控制最小授权范围。 +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 +} + +// 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 new file mode 100644 index 00000000..194b316e --- /dev/null +++ b/internal/tools/session_memory_test.go @@ -0,0 +1,283 @@ +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.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/tools/types.go b/internal/tools/types.go index 39147ee9..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" ) @@ -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) } @@ -34,4 +35,4 @@ type ToolResult struct { Metadata map[string]any } -type ToolSpec = provider.ToolSpec +type ToolSpec = providertypes.ToolSpec diff --git a/internal/tools/webfetch/tool.go b/internal/tools/webfetch/tool.go index efa5474f..83daa3f5 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" @@ -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 { diff --git a/internal/tui/app.go b/internal/tui/app.go index 2fdfaf81..19816f03 100644 --- a/internal/tui/app.go +++ b/internal/tui/app.go @@ -15,48 +15,49 @@ import ( "github.com/charmbracelet/lipgloss" "neo-code/internal/config" - "neo-code/internal/provider" + providertypes "neo-code/internal/provider/types" agentruntime "neo-code/internal/runtime" ) 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 []providertypes.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/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/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/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..08be31d1 --- /dev/null +++ b/internal/tui/core/commands/workspace.go @@ -0,0 +1,70 @@ +package commands + +import ( + "context" + "fmt" + "strings" + + agentsession "neo-code/internal/session" +) + +// SessionWorkdirSetter 定义设置会话工作目录所需的最小 runtime 能力。 +type SessionWorkdirSetter interface { + SetSessionWorkdir(ctx context.Context, sessionID string, workdir string) (agentsession.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..4ed48934 --- /dev/null +++ b/internal/tui/core/commands/workspace_test.go @@ -0,0 +1,107 @@ +package commands + +import ( + "context" + "errors" + "os" + "path/filepath" + "strings" + "testing" + + agentsession "neo-code/internal/session" + tuiworkspace "neo-code/internal/tui/core/workspace" +) + +type stubSessionWorkdirSetter struct { + session agentsession.Session + err error + calls int +} + +func (s *stubSessionWorkdirSetter) SetSessionWorkdir(ctx context.Context, sessionID string, workdir string) (agentsession.Session, error) { + s.calls++ + if s.err != nil { + return agentsession.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: 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) + } + 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..3f75baf9 --- /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/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.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/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.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) +} diff --git a/internal/tui/core/workspace/resolver_test.go b/internal/tui/core/workspace/resolver_test.go new file mode 100644 index 00000000..790150b3 --- /dev/null +++ b/internal/tui/core/workspace/resolver_test.go @@ -0,0 +1,67 @@ +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) + } + + // 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) { + 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) + } +} diff --git a/internal/tui/state.go b/internal/tui/state.go index a4200e1c..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 @@ -58,6 +58,15 @@ type activityEntry struct { IsError bool } +type pendingPermissionPrompt struct { + RequestID string + ToolCallID string + ToolName string + ToolCategory string + Target string + Submitted bool +} + type commandMenuMeta struct { Title string } @@ -133,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.go b/internal/tui/update.go index db804480..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" ) @@ -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,20 @@ 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 { + 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) + 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 +256,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 @@ -398,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) == "" { @@ -683,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) @@ -693,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, @@ -725,18 +753,20 @@ 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 } - 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: 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 +776,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 +789,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 { @@ -799,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 } @@ -813,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) { @@ -1396,6 +1468,36 @@ 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 + } + 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 +} + 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 +1509,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 32c932e8..c04a5e73 100644 --- a/internal/tui/update_test.go +++ b/internal/tui/update_test.go @@ -21,7 +21,9 @@ 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" + agentsession "neo-code/internal/session" "neo-code/internal/tools" ) @@ -29,16 +31,18 @@ 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 cancelCalls int cancelResult bool } @@ -65,7 +69,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{}, } } @@ -79,6 +83,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 } @@ -88,34 +97,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 @@ -244,7 +253,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 { @@ -321,7 +330,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 { @@ -336,7 +345,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 { @@ -490,6 +499,131 @@ 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) { + 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) + } + 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 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 @@ -544,7 +678,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", }, @@ -644,15 +778,15 @@ 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{ + Messages: []providertypes.Message{ {Role: roleUser, Content: "hi"}, {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 @@ -801,12 +935,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 == "" { @@ -854,7 +988,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") } @@ -1063,7 +1197,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) { @@ -1076,11 +1210,11 @@ 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"}}, + Messages: []providertypes.Message{{Role: roleAssistant, Content: "loaded"}}, } if err := app.refreshSessions(); err != nil { t.Fatalf("refresh sessions: %v", err) @@ -1655,7 +1789,7 @@ func TestAppHandleRuntimeEventAdditionalBranches(t *testing.T) { event: agentruntime.RuntimeEvent{ Type: agentruntime.EventToolStart, SessionID: "s1", - Payload: provider.ToolCall{ + Payload: providertypes.ToolCall{ Name: "filesystem_edit", }, }, @@ -2315,7 +2449,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") @@ -2611,8 +2745,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 providertypes.ChatRequest, events chan<- providertypes.StreamEvent) error { + return nil } type tUItestCatalogStore struct { 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