package mcp // MCP 端点的判据。 // // # 这些判据真正在钉什么 // // MCP 工具**包装现有 handler**,所以它成立的唯一依据是「与 HTTP 端点 // 同一套语义」。那就不必重新测一遍业务规则(那些在 handler 自己的测试里), // 而要钉三件**只有集成才可能坏**的事: // // 1. **MCP 不能成为越权旁门**。工具参数里没有身份字段;身份只能来自 // AgentAuth 放进的 context。若某处改成从参数取身份,任何人都能冒充 // 别的 Agent —— 这正是 2026-10-02 修掉的 `AgentMayReadSession` // `if scope == nil { return true }` 那种失败模式(默认放行)。 // 2. **参数映射不许偷偷放宽/收紧**。例如 send_mail 的 `attachment_ids` // 若被吞掉,附件会静默不随信发出;workspace 若被丢掉,GetInbox 会 400。 // 3. **协议语义不许退化**。工具失败必须回 result+isError 而不是 // JSON-RPC error(否则模型只看到「协议错误」,拿不到原因就没法改道)。 import ( "context" "encoding/json" "fmt" "log" "net/http" "net/http/httptest" "os" "path/filepath" "strings" "testing" "github.com/agentmail/gateway/internal/db" "github.com/agentmail/gateway/internal/middleware" "github.com/agentmail/gateway/internal/repo" ) func setupDB(t *testing.T) { t.Helper() dir := t.TempDir() if err := db.Connect(context.Background(), filepath.Join(dir, "mcp-test.db")); err != nil { t.Fatalf("connect: %v", err) } if err := db.Migrate(context.Background()); err != nil { t.Fatalf("migrate: %v", err) } t.Cleanup(db.Close) } // newTestServer 起一个挂了全部工具的端点。 func newTestServer() *Server { s := NewServer(log.New(os.Stderr, "", 0)) NewTools().RegisterAll(s) return s } // call 走完整 HTTP 路径调一个工具(这样才是端到端,不是直接调 Run)。 func call(t *testing.T, s *Server, agent, tool string, args map[string]any) (int, *rpcMessage) { t.Helper() if args == nil { args = map[string]any{} } payload, err := json.Marshal(map[string]any{ "jsonrpc": "2.0", "id": 1, "method": "tools/call", "params": map[string]any{"name": tool, "arguments": args}, }) if err != nil { t.Fatalf("marshal: %v", err) } r := httptest.NewRequest(http.MethodPost, "/api/v1/mcp", strings.NewReader(string(payload))) if agent != "" { r = r.WithContext(context.WithValue(r.Context(), middleware.AgentNameKey, agent)) } rec := httptest.NewRecorder() s.HandleHTTP(rec, r) var msg rpcMessage _ = json.Unmarshal(rec.Body.Bytes(), &msg) return rec.Code, &msg } // toolText 取出 tools/call 的文本内容。 func toolText(t *testing.T, msg *rpcMessage) (string, bool) { t.Helper() raw, err := json.Marshal(msg.Result) if err != nil { t.Fatalf("marshal result: %v", err) } var tr toolResult if err := json.Unmarshal(raw, &tr); err != nil { t.Fatalf("result 不是 toolResult:%v", err) } if len(tr.Content) == 0 { return "", tr.IsError } return tr.Content[0].Text, tr.IsError } // seedAgentAndMail 建一个 Agent 与一封寄给它的邮件。 func seedAgentAndMail(t *testing.T, agent, workspace, subject string) (string, string) { t.Helper() ctx := context.Background() if _, err := db.DB.ExecContext(ctx, `INSERT INTO agents (agent_name, secret, platform, default_rounds) VALUES ($1, 'x', 'test', 50)`, agent); err != nil { t.Fatalf("seed agent: %v", err) } fromName, alias := "gui-lab", "s1" mailID := repoInsertMail(t, ctx, agent, fromName, workspace, alias, subject) var sessionID string if err := db.DB.QueryRowContext(ctx, `SELECT session_id FROM mails WHERE mail_id=$1`, mailID).Scan(&sessionID); err != nil { t.Fatalf("取 session_id: %v", err) } _ = fromName return mailID, fromName } // sessionOf 取一封信所属的会话。 func sessionOf(t *testing.T, mailID string) string { t.Helper() var id string if err := db.DB.QueryRowContext(context.Background(), `SELECT session_id FROM mails WHERE mail_id=$1`, mailID).Scan(&id); err != nil { t.Fatalf("sessionOf: %v", err) } return id } // ---- 协议层 ---- func TestInitializeEchoesClientVersion(t *testing.T) { s := newTestServer() payload := `{"jsonrpc":"2.0","id":1,"method":"initialize","params":{"protocolVersion":"2025-06-18"}}` r := httptest.NewRequest(http.MethodPost, "/api/v1/mcp", strings.NewReader(payload)) rec := httptest.NewRecorder() s.HandleHTTP(rec, r) var msg rpcMessage _ = json.Unmarshal(rec.Body.Bytes(), &msg) if msg.Error != nil { t.Fatalf("initialize 不该报错:%v", msg.Error) } res := msg.Result.(map[string]any) if got := res["protocolVersion"]; got != "2025-06-18" { t.Errorf("协议版本应回显客户端给的 2025-06-18,收到 %v", got) } // 反向对照:缺版本时用回退值,不崩 payload = `{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}` rec = httptest.NewRecorder() s.HandleHTTP(rec, httptest.NewRequest(http.MethodPost, "/api/v1/mcp", strings.NewReader(payload))) _ = json.Unmarshal(rec.Body.Bytes(), &msg) if msg.Result.(map[string]any)["protocolVersion"] == nil { t.Error("缺 protocolVersion 时应回退一个值,而不是省略该字段") } } func TestNotificationsGetNoResponse(t *testing.T) { s := newTestServer() // 没有 id = 通知。回了响应,客户端会把响应与请求错配。 for _, m := range []string{ `{"jsonrpc":"2.0","method":"notifications/initialized"}`, `{"jsonrpc":"2.0","method":"tools/list"}`, `{"jsonrpc":"2.0","method":"ping"}`, } { rec := httptest.NewRecorder() s.HandleHTTP(rec, httptest.NewRequest(http.MethodPost, "/api/v1/mcp", strings.NewReader(m))) if rec.Code != http.StatusAccepted { t.Errorf("通知 %s:期望 202,收到 %d(并带响应体 %q)", m, rec.Code, rec.Body.String()) } } } func TestToolsListIsStableAndCarriesAnnotations(t *testing.T) { s := newTestServer() var names []string for i := 0; i < 5; i++ { rec := httptest.NewRecorder() s.HandleHTTP(rec, httptest.NewRequest(http.MethodPost, "/api/v1/mcp", strings.NewReader(`{"jsonrpc":"2.0","id":1,"method":"tools/list"}`))) var msg rpcMessage _ = json.Unmarshal(rec.Body.Bytes(), &msg) list := msg.Result.(map[string]any)["tools"].([]any) names = names[:0] for _, it := range list { names = append(names, it.(map[string]any)["name"].(string)) } if i == 0 { // annotations 必须透传:宿主据此算风险等级并在 plan 档放行。 // 漏传的后果不是「少个提示」,而是工具在该档下全被拒。 found := false for _, it := range list { m := it.(map[string]any) if m["name"] == "read_inbox" { if _, ok := m["annotations"]; !ok { t.Error("read_inbox 缺 annotations") } found = true } } if !found { t.Error("tools/list 里没有 read_inbox") } } } // 顺序必须稳定:map 迭代随机会让客户端每次刷新看到不同排列 if len(names) == 0 { t.Fatal("tools/list 为空") } for j := 1; j < len(names); j++ { if names[j] < names[j-1] { t.Fatalf("工具顺序不稳定:%v", names) } } } func TestUnknownToolIsInvalidParamsNotExecutionFailure(t *testing.T) { s := newTestServer() _, msg := call(t, s, "a", "no_such_tool", nil) if msg.Error == nil { t.Fatal("未知工具名应回 JSON-RPC error") } if msg.Error.Code != rpcInvalidParams { t.Errorf("期望 INVALID_PARAMS(%d),收到 %d", rpcInvalidParams, msg.Error.Code) } // 错误文案要告诉模型可用集合,否则它只能瞎试 if !strings.Contains(msg.Error.Message, "read_inbox") { t.Errorf("未知工具的错误应列出可用工具,收到 %q", msg.Error.Message) } } func TestToolFailureIsResultIsErrorNotRPCError(t *testing.T) { setupDB(t) s := newTestServer() // 不存在的 Agent + 非法 workspace ⇒ handler 报错。 // 关键不是失败本身,而是失败**怎么**回。 _, msg := call(t, s, "ghost", "read_inbox", map[string]any{"workspace": "not-absolute"}) if msg.Error != nil { t.Fatalf("工具失败不该回 JSON-RPC error(模型就看不到原因了),收到 %v", msg.Error) } text, isErr := toolText(t, msg) if !isErr { t.Error("失败时必须带 isError:true") } if text == "" { t.Error("失败时必须给出原因文本") } } func TestMethodNotAllowedOnNonPost(t *testing.T) { s := newTestServer() rec := httptest.NewRecorder() s.HandleHTTP(rec, httptest.NewRequest(http.MethodGet, "/api/v1/mcp", nil)) if rec.Code != http.StatusMethodNotAllowed { t.Errorf("GET /mcp 应 405,收到 %d", rec.Code) } } func TestMalformedJSONGetsParseError(t *testing.T) { s := newTestServer() rec := httptest.NewRecorder() s.HandleHTTP(rec, httptest.NewRequest(http.MethodPost, "/api/v1/mcp", strings.NewReader(`{oops`))) if !strings.Contains(rec.Body.String(), `"id":null`) { t.Errorf("解析错误响应必须带 id:null,收到 %q", rec.Body.String()) } if !strings.Contains(rec.Body.String(), fmt.Sprintf("%d", rpcParse)) { t.Errorf("非法 JSON 应回 PARSE_ERROR(%d),收到 %q", rpcParse, rec.Body.String()) } } // ---- 工具语义与 HTTP 端点同源 ---- func TestReadInboxRequiresWorkspaceSameAsHTTP(t *testing.T) { setupDB(t) seedAgentAndMail(t, "probe", "/tmp/ws", "一封信") s := newTestServer() // 不带 workspace ⇒ 400(这是 GetInbox 的既有语义,包装后必须一致) _, msg := call(t, s, "probe", "read_inbox", map[string]any{}) if msg.Error != nil { t.Fatalf("工具失败不该是协议错误:%v", msg.Error) } text, isErr := toolText(t, msg) if !isErr { t.Errorf("缺 workspace 时 read_inbox 必须失败,收到 %q", text) } if !strings.Contains(text, "workspace") { t.Errorf("错误文案应说清是 workspace 的问题,收到 %q", text) } } func TestReadInboxScopedToOwnWorkspace(t *testing.T) { setupDB(t) ctx := context.Background() // 同一个 Agent 在两个工作区各有信;按工作区收窄后只能看到自己那个的 seedAgentAndMail(t, "probe", "/tmp/wsA", "A 区的信") repoInsertMail(t, ctx, "probe", "gui-lab", "/tmp/wsB", "sB", "B 区的信") s := newTestServer() _, msg := call(t, s, "probe", "read_inbox", map[string]any{"workspace": "/tmp/wsA"}) text, isErr := toolText(t, msg) if isErr { t.Fatalf("read_inbox 失败:%s", text) } if !strings.Contains(text, "A 区的信") { t.Errorf("应看到 A 区的信,收到 %q", text) } if strings.Contains(text, "B 区的信") { t.Errorf("★ 收窄失效:看到了别的工作区的信(B 区的信),输出 %q", text) } } func TestAgentCannotReadOthersMail(t *testing.T) { setupDB(t) ctx := context.Background() // 别人的信 otherMailID, _ := seedAgentAndMail(t, "victim", "/tmp/ws", "别人的信") s := newTestServer() // 攻击者身份,读别人的信 _, msg := call(t, s, "attacker", "read_mail", map[string]any{"mail_id": otherMailID}) text, isErr := toolText(t, msg) if !isErr { t.Errorf("★ 越权:attacker 读到了 victim 的邮件,输出 %q", text) } if strings.Contains(text, "别人的信") { t.Errorf("★ 越权:泄露了别人的正文 %q", text) } _ = ctx } func TestToolArgsCannotForgeIdentity(t *testing.T) { setupDB(t) // 传 identity 相关的参数不应影响权限判定 s := newTestServer() _, msg := call(t, s, "attacker", "read_inbox", map[string]any{ "workspace": "/tmp/ws", "agent_name": "victim", "X-Agent-Name": "victim", "from": "victim", }) // 受害者不在这个工作区/或不存在时应当看到「空」而不是别人的信 text, _ := toolText(t, msg) if strings.Contains(text, "别人的信") { t.Errorf("参数里的身份字段不该影响鉴权,收到 %q", text) } } // ---- 参数映射 ---- func TestAttachmentIDsSurviveMapping(t *testing.T) { setupDB(t) _, _ = seedAgentAndMail(t, "probe", "/tmp/ws", "占位") s := newTestServer() // 用受配额的发信路径发一封带附件的信(附件 id 传下去) _, msg := call(t, s, "probe", "send_mail", map[string]any{ "to": "gui-lab@/tmp.s1", "subject": "带附件", "body": "见附件", "attachment_ids": []any{"att-does-not-exist"}, }) text, isErr := toolText(t, msg) if !isErr { t.Fatalf("传了不存在的附件 id,send_mail 理应失败,收到 %q", text) } // 附件 id 真被记到那封信上 // 发信时被拒的 attachment_id 必须在错误里出现(说明映射到了、校验才起效)。 // 换一个不存在的 id 会被 SendMail 拒绝,而拒绝原因正是参数映射是否生效的直接证据。 if !strings.Contains(text, "attachment") { t.Errorf("★ attachment_ids 似乎没被映射过去(错误里没提附件):%s", text) } } func TestMissingArgumentsTreatedAsEmpty(t *testing.T) { setupDB(t) seedAgentAndMail(t, "probe", "/tmp/ws", "信") s := newTestServer() // suggest_address 无参 ⇒ 返回候选名,不该崩 _, msg := call(t, s, "probe", "suggest_address", nil) if msg.Error != nil { t.Fatalf("无参 suggest_address 不该是协议错误:%v", msg.Error) } _, isErr := toolText(t, msg) if isErr { t.Errorf("无参 suggest_address 不该失败") } } func TestToolCountAndNoExecutionTools(t *testing.T) { s := newTestServer() // 11 个邮件工具,不含任何"会动机器"的工具 if got := s.ToolCount(); got != 11 { t.Errorf("工具数应为 11(只读+收发邮件),收到 %d", got) } for _, name := range []string{"run_command", "write_file", "exec", "shell"} { if _, ok := s.tools[name]; ok { t.Errorf("★ 通用 MCP 面不得提供执行类工具 %q", name) } } } // ---- 辅助 ---- func repoInsertMail(t *testing.T, ctx context.Context, to, from, workspace, alias, subject string) string { t.Helper() sessionID := repoEnsureSession(t, ctx, to, from, workspace, alias) var mailID string err := db.DB.QueryRowContext(ctx, `INSERT INTO mails (session_id, from_name, to_name, subject, body, status) VALUES ($1,$2,$3,$4,'正文',$5) RETURNING mail_id`, sessionID, from, to, subject, "unread").Scan(&mailID) if err != nil { t.Fatalf("insert mail: %v", err) } return mailID } func repoEnsureSession(t *testing.T, ctx context.Context, to, from, workspace, alias string) string { t.Helper() t.Helper() var sessionID string // 收窄走 s.workspace(见 repo.workspaceScope),所以工作区建会话时就得给对 err := db.DB.QueryRowContext(ctx, `INSERT INTO sessions (session_alias, subject, status, workspace, from_agent) VALUES ($1,'t','active',$2,$3) RETURNING session_id`, alias, workspace, to).Scan(&sessionID) if err != nil { t.Fatalf("insert session: %v", err) } _ = from return sessionID } // repoInsertAttachment 造一个附件。upload_attachment 走 blob 存储, // 这里只需要一个**已存在**的 attachment_id 来验证参数映射。 func repoInsertAttachment(t *testing.T, mailID string) string { t.Helper() var id string err := db.DB.QueryRowContext(context.Background(), `INSERT INTO attachments (mail_id, uploader, filename, content_type, size_bytes, sha256) VALUES ($1,'probe','a.txt','text/plain',12,'deadbeef') RETURNING attachment_id`, mailID).Scan(&id) if err != nil { t.Fatalf("insert attachment: %v", err) } return id } // repoAttachmentIDsOfMail 取一封信挂着的附件 id(attachments.mail_id)。 func repoAttachmentIDsOfMail(t *testing.T, mailID string) string { t.Helper() var ids string err := db.DB.QueryRowContext(context.Background(), `SELECT COALESCE(group_concat(attachment_id),'') FROM attachments WHERE mail_id=$1`, mailID).Scan(&ids) if err != nil { return "" } return ids } var _ = repo.ListInboxScoped // 保持 import(工具链会用到 handler,repo 只在测试里用到) // TestEveryErrorResponseCarriesID 钉住 JSON-RPC 的硬要求:**每个**响应都要带 id。 // // 为什么单独写一格:解析错误那格只覆盖了「解析失败」一条路径。曾用 // `ID json.RawMessage` + `omitempty`,nil 时整个 id 字段会从 JSON 里消失 —— // 客户端拿不到配对依据,会一直等这条的响应。这格遍历所有错误出口, // 任何一条漏掉 id 都会红。 func TestEveryErrorResponseCarriesID(t *testing.T) { s := newTestServer() cases := []struct { name string body string want string // 期望的 id(JSON 字面量) }{ {"解析失败", `{oops`, "null"}, {"jsonrpc 版本不对", `{"jsonrpc":"1.0","id":7,"method":"ping"}`, "7"}, {"方法不存在", `{"jsonrpc":"2.0","id":7,"method":"no/such"}`, "7"}, {"工具名未知", `{"jsonrpc":"2.0","id":7,"method":"tools/call","params":{"name":"nope"}}`, "7"}, {"工具名缺失", `{"jsonrpc":"2.0","id":7,"method":"tools/call","params":{}}`, "7"}, } for _, c := range cases { rec := httptest.NewRecorder() s.HandleHTTP(rec, httptest.NewRequest(http.MethodPost, "/api/v1/mcp", strings.NewReader(c.body))) out := rec.Body.String() // 直接查 JSON 里有没有 id 键(而不是靠字符串匹配,那样会漏掉 // `"id":null` 这种情形) var decoded map[string]json.RawMessage if err := json.Unmarshal(rec.Body.Bytes(), &decoded); err != nil { t.Errorf("%s:响应不是合法 JSON:%v(%q)", c.name, err, out) continue } rawID, ok := decoded["id"] if !ok { t.Errorf("★ %s:响应缺少 id 字段 ⇒ 客户端无法配对,会一直等(收到 %q)", c.name, out) continue } if string(rawID) != c.want { t.Errorf("%s:id 应为 %s,收到 %s", c.name, c.want, string(rawID)) } } } // TestPathParamsReachHandlers 钉住「包装 handler 时路径参数必须到位」。 // // ★ 这格是端到端撞出来的:主人读自己的信与越权者读别人的信**都**返回 // 「Invalid id」。只看越权那一次会误判成「收得太紧」,进而把正确的收窄改松; // 做对照才看出是 `chi.URLParam` 拿不到值(httptest 请求没过 chi 的路由)。 // // 判据用「主人能读自己的信」当锚:它同时证明参数到位**且**收窄没有误伤。 func TestPathParamsReachHandlers(t *testing.T) { setupDB(t) mailID, _ := seedAgentAndMail(t, "owner", "/tmp/ws", "主人自己的信") sessionID := sessionOf(t, mailID) s := newTestServer() // agentScope 要求 session_id 与信所属会话一致(这是**正确**的收窄, // 不是路径参数问题)—— 所以这里带上它。 args := map[string]any{"mail_id": mailID, "session_id": sessionID} text, isErr := toolText(t, callMsg(t, s, "owner", "read_mail", args)) if isErr { t.Fatalf("★ 主人读自己的信应该成功,却失败了(路径参数没到位?):%s", text) } if !strings.Contains(text, "主人自己的信") { t.Errorf("应读到自己的信,收到 %q", text) } // read_thread 同样带路径参数,一并钉住 text, isErr = toolText(t, callMsg(t, s, "owner", "read_thread", args)) if isErr { t.Errorf("★ read_thread 也应成功(同一类路径参数):%s", text) } } // callMsg 是 call 的单值版本(拿消息)。 func callMsg(t *testing.T, s *Server, agent, tool string, args map[string]any) *rpcMessage { t.Helper() _, msg := call(t, s, agent, tool, args) return msg }