memos/server/router/mcp/service_test.go
johnnyjoygh 652957c0f9 fix(mcp): expose standard JSON Schema formats in tool schemas
The gnostic OpenAPI generator hardcodes `format: enum`, `format: bytes`,
and `format: field-mask` on string properties. These are OpenAPI-only
markers; MCP clients treat tool schemas as JSON Schema 2020-12 and warn
or fail on them (zod-based parsers log "unknown format" for every tool
on every load, ajv strict mode rejects the definitions).

Normalize schemas as they leave the resolver and when parameter schemas
are cloned: drop `format: enum` (the `enum` keyword is already present),
rewrite `format: bytes` to `contentEncoding: base64`, and drop
`format: field-mask`. The generated openapi.yaml is unchanged.

A catalog test now walks every curated tool's input and output schema
and fails on any format outside the registered allowlist.

Fixes #6262
2026-09-03 21:53:03 +08:00

630 lines
20 KiB
Go

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 := &registeredOperation{
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
}