package mcp import ( "bytes" "context" "encoding/json" "net" "net/http" "net/http/httptest" "os" "strings" "testing" "github.com/labstack/echo/v5" sdkmcp "github.com/modelcontextprotocol/go-sdk/mcp" "github.com/stretchr/testify/require" "github.com/usememos/memos/internal/profile" memosproto "github.com/usememos/memos/proto" ) func TestIsAllowedMCPOrigin(t *testing.T) { profile := &profile.Profile{InstanceURL: "https://memos.example.com/app"} tests := []struct { name string host string origin string want bool }{ {name: "empty origin", host: "localhost:5230", origin: "", want: true}, {name: "same http host", host: "localhost:5230", origin: "http://localhost:5230", want: true}, {name: "same https host", host: "memos.example.com", origin: "https://memos.example.com", want: true}, {name: "configured instance URL origin", host: "127.0.0.1:5230", origin: "https://memos.example.com", want: true}, {name: "configured instance URL ignores path", host: "127.0.0.1:5230", origin: "https://memos.example.com", want: true}, {name: "different host", host: "localhost:5230", origin: "https://evil.example.com", want: false}, {name: "invalid origin", host: "localhost:5230", origin: "not a url", want: false}, } for _, test := range tests { t.Run(test.name, func(t *testing.T) { require.Equal(t, test.want, isAllowedMCPOrigin(test.host, test.origin, profile)) }) } } func TestNewMCPServiceRegistersCuratedTools(t *testing.T) { echoServer := echo.New() service, err := NewMCPService(&profile.Profile{Version: "test-version"}, echoServer) require.NoError(t, err) require.NotNil(t, service.handler) require.Len(t, service.operationsByTool, len(curatedOperationIDs)) operation := service.operationsByTool["memo_list_memos"] require.NotNil(t, operation) require.Equal(t, "MemoService_ListMemos", operation.OperationID) require.Equal(t, "GET", operation.Method) require.Equal(t, "/api/v1/memos", operation.Path) } func TestNewMCPServiceUsesEmbeddedOpenAPISpec(t *testing.T) { t.Chdir(t.TempDir()) service, err := NewMCPService(&profile.Profile{Version: "test-version"}, echo.New()) require.NoError(t, err) require.NotNil(t, service.handler) require.Len(t, service.operationsByTool, len(curatedOperationIDs)) } func TestEmbeddedOpenAPISpecMatchesGeneratedFile(t *testing.T) { generated, err := os.ReadFile("../../../proto/gen/openapi.yaml") require.NoError(t, err) require.Equal(t, generated, memosproto.OpenAPIYAML()) } func TestMCPToolHandlerForwardsArgumentsAndAuthorization(t *testing.T) { echoServer := echo.New() echoServer.GET("/api/v1/memos", func(c *echo.Context) error { require.Equal(t, "Bearer token", c.Request().Header.Get("Authorization")) require.Equal(t, "7", c.QueryParam("pageSize")) return c.JSON(http.StatusOK, map[string]any{ "memos": []any{map[string]any{"name": "memos/test"}}, }) }) operation := ®isteredOperation{ Operation: &openAPIOperation{ Method: "GET", Path: "/api/v1/memos", Parameters: []openAPIParameter{{Name: "pageSize", In: "query", Schema: jsonSchema{"type": "integer"}}}, }, } handler := newMCPToolHandler(newAPIAdapter(echoServer), operation) arguments, err := json.Marshal(map[string]any{"pageSize": 7}) require.NoError(t, err) result, err := handler(context.Background(), &sdkmcp.CallToolRequest{ Params: &sdkmcp.CallToolParamsRaw{ Name: "memo_list_memos", Arguments: arguments, }, Extra: &sdkmcp.RequestExtra{ Header: http.Header{"Authorization": []string{"Bearer token"}}, }, }) require.NoError(t, err) require.False(t, result.IsError) require.Equal(t, map[string]any{ "memos": []any{map[string]any{"name": "memos/test"}}, }, result.StructuredContent) } func TestMCPProtocolListsCuratedToolsOnly(t *testing.T) { echoServer := echo.New() service, err := NewMCPService(&profile.Profile{Version: "test-version"}, echoServer) require.NoError(t, err) service.RegisterRoutes(echoServer) initializeMCP(t, echoServer) response := postMCP(t, echoServer, map[string]any{ "jsonrpc": "2.0", "id": 2, "method": "tools/list", }) result, ok := response["result"].(map[string]any) require.True(t, ok) tools, ok := result["tools"].([]any) require.True(t, ok) require.Len(t, tools, len(curatedOperationIDs)) names := map[string]struct{}{} for _, rawTool := range tools { tool, ok := rawTool.(map[string]any) require.True(t, ok) name, ok := tool["name"].(string) require.True(t, ok) names[name] = struct{}{} require.Contains(t, tool, "inputSchema") require.Contains(t, tool, "outputSchema") } require.Contains(t, names, "memo_list_memos") require.Contains(t, names, "memo_create_memo") require.NotContains(t, names, "auth_sign_in") require.NotContains(t, names, "user_create_user") } func TestMCPToolCallReturnsObjectStructuredContent(t *testing.T) { echoServer := echo.New() echoServer.GET("/api/v1/memos", func(c *echo.Context) error { return c.JSON(http.StatusOK, map[string]any{ "memos": []any{map[string]any{"name": "memos/abc123"}}, }) }) service, err := NewMCPService(&profile.Profile{Version: "test-version"}, echoServer) require.NoError(t, err) service.RegisterRoutes(echoServer) initializeMCP(t, echoServer) response := postMCP(t, echoServer, map[string]any{ "jsonrpc": "2.0", "id": 2, "method": "tools/call", "params": map[string]any{ "name": "memo_list_memos", "arguments": map[string]any{ "pageSize": 1, }, }, }) result, ok := response["result"].(map[string]any) require.True(t, ok) require.Equal(t, map[string]any{ "memos": []any{map[string]any{"name": "memos/abc123"}}, }, result["structuredContent"]) } func TestMCPToolCallAllowsGatewayToInferMemoUpdateMask(t *testing.T) { echoServer := echo.New() routeHits := 0 echoServer.PATCH("/api/v1/memos/:memo", func(c *echo.Context) error { routeHits++ require.Equal(t, "abc123", c.Param("memo")) require.Empty(t, c.QueryParam("updateMask")) body := map[string]any{} require.NoError(t, json.NewDecoder(c.Request().Body).Decode(&body)) require.Equal(t, map[string]any{"content": "updated"}, body) return c.JSON(http.StatusOK, map[string]any{ "name": "memos/abc123", "content": "updated", }) }) service, err := NewMCPService(&profile.Profile{Version: "test-version"}, echoServer) require.NoError(t, err) service.RegisterRoutes(echoServer) initializeMCP(t, echoServer) response := postMCP(t, echoServer, map[string]any{ "jsonrpc": "2.0", "id": 2, "method": "tools/call", "params": map[string]any{ "name": "memo_update_memo", "arguments": map[string]any{ "memo": "memos/abc123", "body": map[string]any{"content": "updated"}, }, }, }) result, ok := response["result"].(map[string]any) require.True(t, ok) require.NotEqual(t, true, result["isError"]) require.Equal(t, map[string]any{ "name": "memos/abc123", "content": "updated", }, result["structuredContent"]) require.Equal(t, 1, routeHits) } func TestMCPToolCallBindsMemoFromPathForBodyStarOperations(t *testing.T) { register := func(e *echo.Echo, method, path string, handler func(*echo.Context) error) { switch method { case http.MethodPatch: e.PATCH(path, handler) case http.MethodPost: e.POST(path, handler) default: t.Fatalf("unsupported method %q", method) } } tests := []struct { name string method string path string toolName string body map[string]any response map[string]any }{ { name: "set attachments", method: http.MethodPatch, path: "/api/v1/memos/:memo/attachments", toolName: "memo_set_memo_attachments", body: map[string]any{"attachments": []any{}}, response: map[string]any{}, }, { name: "set relations", method: http.MethodPatch, path: "/api/v1/memos/:memo/relations", toolName: "memo_set_memo_relations", body: map[string]any{"relations": []any{}}, response: map[string]any{}, }, { name: "upsert reaction", method: http.MethodPost, path: "/api/v1/memos/:memo/reactions", toolName: "memo_upsert_memo_reaction", body: map[string]any{ "reaction": map[string]any{ "reactionType": "👍", }, }, response: map[string]any{"reactionType": "👍"}, }, } for _, test := range tests { t.Run(test.name, func(t *testing.T) { echoServer := echo.New() routeHits := 0 register(echoServer, test.method, test.path, func(c *echo.Context) error { routeHits++ // The memo must be bound from the path, and the omitted "name" // property must not reappear in the forwarded body. require.Equal(t, "abc123", c.Param("memo")) body := map[string]any{} require.NoError(t, json.NewDecoder(c.Request().Body).Decode(&body)) require.NotContains(t, body, "name") return c.JSON(http.StatusOK, test.response) }) service, err := NewMCPService(&profile.Profile{Version: "test-version"}, echoServer) require.NoError(t, err) service.RegisterRoutes(echoServer) initializeMCP(t, echoServer) response := postMCP(t, echoServer, map[string]any{ "jsonrpc": "2.0", "id": 2, "method": "tools/call", "params": map[string]any{ "name": test.toolName, "arguments": map[string]any{ "memo": "memos/abc123", "body": test.body, }, }, }) result, ok := response["result"].(map[string]any) require.True(t, ok) require.NotEqual(t, true, result["isError"], result) require.Equal(t, 1, routeHits) }) } } func TestMCPToolCallRejectsInvalidArguments(t *testing.T) { echoServer := echo.New() routeHits := 0 echoServer.GET("/api/v1/memos", func(c *echo.Context) error { routeHits++ return c.JSON(http.StatusOK, map[string]any{"memos": []any{}}) }) echoServer.GET("/api/v1/memos/:memo", func(c *echo.Context) error { routeHits++ return c.JSON(http.StatusOK, map[string]any{"name": c.Param("memo")}) }) service, err := NewMCPService(&profile.Profile{Version: "test-version"}, echoServer) require.NoError(t, err) service.RegisterRoutes(echoServer) initializeMCP(t, echoServer) tests := []struct { name string toolName string arguments map[string]any wantError string }{ { name: "unknown argument", toolName: "memo_list_memos", arguments: map[string]any{"unexpected": true}, wantError: `unknown argument "unexpected"`, }, { name: "missing required argument", toolName: "memo_get_memo", arguments: map[string]any{}, wantError: `missing required argument "memo"`, }, { name: "wrong primitive type", toolName: "memo_list_memos", arguments: map[string]any{"pageSize": "ten"}, wantError: `argument "pageSize" must be integer`, }, } for index, test := range tests { t.Run(test.name, func(t *testing.T) { response := postMCP(t, echoServer, map[string]any{ "jsonrpc": "2.0", "id": index + 2, "method": "tools/call", "params": map[string]any{ "name": test.toolName, "arguments": test.arguments, }, }) result, ok := response["result"].(map[string]any) require.True(t, ok) require.Equal(t, true, result["isError"]) // Error results carry no structuredContent — it would fail // validation against the tool's declared outputSchema in strict // clients. The message travels in the text content instead. _, hasStructured := result["structuredContent"] require.False(t, hasStructured) content, ok := result["content"].([]any) require.True(t, ok) require.NotEmpty(t, content) textBlock, ok := content[0].(map[string]any) require.True(t, ok) require.Contains(t, textBlock["text"], test.wantError) }) } require.Zero(t, routeHits) } // TestMCPLoopbackBehindReverseProxy verifies that a loopback-bound instance // served under a non-loopback Host (the reverse-proxy deployment shape) is no // longer rejected by the SDK's DNS-rebinding guard, while memos' own Origin // allowlist still rejects disallowed origins. func TestMCPLoopbackBehindReverseProxy(t *testing.T) { echoServer := echo.New() service, err := NewMCPService(&profile.Profile{Version: "test-version"}, echoServer) require.NoError(t, err) service.RegisterRoutes(echoServer) initialize, err := json.Marshal(map[string]any{ "jsonrpc": "2.0", "id": 1, "method": "initialize", "params": map[string]any{ "protocolVersion": "2025-06-18", "capabilities": map[string]any{}, "clientInfo": map[string]any{"name": "memos-test", "version": "1.0.0"}, }, }) require.NoError(t, err) // Simulate the proxied deployment: the connection terminates on a loopback // address, but the public Host header is a real domain. newRequest := func() *http.Request { request := httptest.NewRequest(http.MethodPost, "/mcp", bytes.NewReader(initialize)) request.Host = "demo.usememos.com" request.Header.Set("Content-Type", "application/json") request.Header.Set("Accept", "application/json, text/event-stream") loopback := &net.TCPAddr{IP: net.IPv4(127, 0, 0, 1), Port: 5230} ctx := context.WithValue(request.Context(), http.LocalAddrContextKey, loopback) return request.WithContext(ctx) } t.Run("allows non-loopback host with no origin", func(t *testing.T) { recorder := httptest.NewRecorder() echoServer.ServeHTTP(recorder, newRequest()) require.Equal(t, http.StatusOK, recorder.Code, recorder.Body.String()) }) t.Run("still rejects a disallowed origin", func(t *testing.T) { request := newRequest() request.Header.Set("Origin", "https://evil.example.com") recorder := httptest.NewRecorder() echoServer.ServeHTTP(recorder, request) require.Equal(t, http.StatusForbidden, recorder.Code) }) } func initializeMCP(t *testing.T, echoServer *echo.Echo) { t.Helper() response := postMCP(t, echoServer, map[string]any{ "jsonrpc": "2.0", "id": 1, "method": "initialize", "params": map[string]any{ "protocolVersion": "2025-06-18", "capabilities": map[string]any{}, "clientInfo": map[string]any{ "name": "memos-test", "version": "1.0.0", }, }, }) require.NotNil(t, response["result"]) } func postMCP(t *testing.T, echoServer *echo.Echo, payload map[string]any) map[string]any { t.Helper() data, err := json.Marshal(payload) require.NoError(t, err) request := httptest.NewRequest(http.MethodPost, "/mcp", bytes.NewReader(data)) request.Header.Set("Content-Type", "application/json") request.Header.Set("Accept", "application/json, text/event-stream") recorder := httptest.NewRecorder() echoServer.ServeHTTP(recorder, request) require.Equal(t, http.StatusOK, recorder.Code, recorder.Body.String()) var response map[string]any require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &response)) return response } // TestMCPStatelessProtocol2026 drives the endpoint the way a 2026-07-28 client // does: no initialize handshake, per-request _meta, and the standard MCP headers. func TestMCPStatelessProtocol2026(t *testing.T) { echoServer := echo.New() var forwardedAuthorization string echoServer.GET("/api/v1/memos", func(c *echo.Context) error { forwardedAuthorization = c.Request().Header.Get("Authorization") return c.JSON(http.StatusOK, map[string]any{"memos": []any{}}) }) service, err := NewMCPService(&profile.Profile{Version: "test-version"}, echoServer) require.NoError(t, err) service.RegisterRoutes(echoServer) meta := map[string]any{ "io.modelcontextprotocol/protocolVersion": "2026-07-28", "io.modelcontextprotocol/clientCapabilities": map[string]any{}, "io.modelcontextprotocol/clientInfo": map[string]any{"name": "memos-test", "version": "1.0.0"}, } discover := postMCPWithHeaders(t, echoServer, map[string]any{ "jsonrpc": "2.0", "id": 1, "method": "server/discover", "params": map[string]any{"_meta": meta}, }, map[string]string{"Mcp-Method": "server/discover"}) discoverResult, ok := discover["result"].(map[string]any) require.True(t, ok, discover) require.Contains(t, discoverResult["supportedVersions"], "2026-07-28") require.Equal(t, "complete", discoverResult["resultType"]) capabilities, ok := discoverResult["capabilities"].(map[string]any) require.True(t, ok) require.Equal(t, map[string]any{}, capabilities["tools"], "tools capability must not advertise listChanged") require.NotContains(t, capabilities, "logging", "logging is deprecated and unused") require.Positive(t, discoverResult["ttlMs"]) list := postMCPWithHeaders(t, echoServer, map[string]any{ "jsonrpc": "2.0", "id": 2, "method": "tools/list", "params": map[string]any{"_meta": meta}, }, map[string]string{"Mcp-Method": "tools/list"}) listResult, ok := list["result"].(map[string]any) require.True(t, ok, list) require.Len(t, listResult["tools"], len(curatedOperationIDs)) require.Equal(t, float64(toolCatalogTTL.Milliseconds()), listResult["ttlMs"]) require.Equal(t, "public", listResult["cacheScope"]) call := postMCPWithHeaders(t, echoServer, map[string]any{ "jsonrpc": "2.0", "id": 3, "method": "tools/call", "params": map[string]any{ "_meta": meta, "name": "memo_list_memos", "arguments": map[string]any{}, }, }, map[string]string{ "Mcp-Method": "tools/call", "Mcp-Name": "memo_list_memos", "Authorization": "Bearer test-token", }) callResult, ok := call["result"].(map[string]any) require.True(t, ok, call) require.Equal(t, "complete", callResult["resultType"]) require.Equal(t, map[string]any{"memos": []any{}}, callResult["structuredContent"]) require.Equal(t, "Bearer test-token", forwardedAuthorization) } // TestMCPStatelessRejectsSessionMethods confirms the stateless transport shape // required by 2026-07-28: no GET stream, no DELETE session teardown. func TestMCPStatelessRejectsSessionMethods(t *testing.T) { echoServer := echo.New() service, err := NewMCPService(&profile.Profile{Version: "test-version"}, echoServer) require.NoError(t, err) service.RegisterRoutes(echoServer) for _, method := range []string{http.MethodGet, http.MethodDelete} { request := httptest.NewRequest(method, "/mcp", nil) request.Header.Set("Accept", "application/json, text/event-stream") recorder := httptest.NewRecorder() echoServer.ServeHTTP(recorder, request) require.Equal(t, http.StatusMethodNotAllowed, recorder.Code, method) } } // TestMCPRequestBodyLimitMatchesAPI guards against the SDK's 4 MiB default // overriding the API-wide limit, which would break attachment uploads over MCP. func TestMCPRequestBodyLimitMatchesAPI(t *testing.T) { echoServer := echo.New() var receivedBytes int echoServer.POST("/api/v1/attachments", func(c *echo.Context) error { body := map[string]any{} require.NoError(t, json.NewDecoder(c.Request().Body).Decode(&body)) content, _ := body["content"].(string) receivedBytes = len(content) return c.JSON(http.StatusOK, map[string]any{"name": "attachments/1"}) }) service, err := NewMCPService(&profile.Profile{Version: "test-version"}, echoServer) require.NoError(t, err) service.RegisterRoutes(echoServer) content := strings.Repeat("A", 8<<20) response := postMCPWithHeaders(t, echoServer, map[string]any{ "jsonrpc": "2.0", "id": 1, "method": "tools/call", "params": map[string]any{ "name": "attachment_create_attachment", "arguments": map[string]any{ "body": map[string]any{ "filename": "large.bin", "type": "application/octet-stream", "content": content, }, }, }, }, nil) require.Contains(t, response, "result", response) require.Equal(t, len(content), receivedBytes) // Declare an over-limit Content-Length instead of allocating the body. request := httptest.NewRequest(http.MethodPost, "/mcp", strings.NewReader("{}")) request.ContentLength = maxMCPRequestBytes + 1 request.Header.Set("Content-Type", "application/json") request.Header.Set("Accept", "application/json, text/event-stream") recorder := httptest.NewRecorder() echoServer.ServeHTTP(recorder, request) require.Equal(t, http.StatusRequestEntityTooLarge, recorder.Code) } func postMCPWithHeaders(t *testing.T, echoServer *echo.Echo, payload map[string]any, headers map[string]string) map[string]any { t.Helper() data, err := json.Marshal(payload) require.NoError(t, err) request := httptest.NewRequest(http.MethodPost, "/mcp", bytes.NewReader(data)) request.Header.Set("Content-Type", "application/json") request.Header.Set("Accept", "application/json, text/event-stream") if _, hasMeta := payload["params"].(map[string]any)["_meta"]; hasMeta { request.Header.Set("Mcp-Protocol-Version", "2026-07-28") } for key, value := range headers { request.Header.Set(key, value) } recorder := httptest.NewRecorder() echoServer.ServeHTTP(recorder, request) require.Equal(t, http.StatusOK, recorder.Code, recorder.Body.String()) var response map[string]any require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &response)) return response }