162 lines
5.6 KiB
Go
162 lines
5.6 KiB
Go
package mcp
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"net/http"
|
|
"time"
|
|
|
|
"github.com/labstack/echo/v5"
|
|
"github.com/labstack/echo/v5/middleware"
|
|
sdkmcp "github.com/modelcontextprotocol/go-sdk/mcp"
|
|
"github.com/pkg/errors"
|
|
"gopkg.in/yaml.v3"
|
|
|
|
"github.com/usememos/memos/internal/profile"
|
|
memosproto "github.com/usememos/memos/proto"
|
|
apiv1 "github.com/usememos/memos/server/router/api/v1"
|
|
)
|
|
|
|
// maxMCPRequestBytes caps the /mcp request body. It tracks the API limit because
|
|
// every tool call is forwarded in-process through the API routes. The SDK
|
|
// enforces its own 4 MiB default, so the value is passed to the transport as
|
|
// well; otherwise attachment uploads over MCP fail with 413 before Echo sees them.
|
|
const maxMCPRequestBytes int64 = apiv1.MaxAPIRequestBytes
|
|
|
|
// toolCatalogTTL is the freshness hint advertised on tools/list and
|
|
// server/discover results. The catalog is fixed at startup, so clients may cache
|
|
// it for a long time instead of re-listing on every session.
|
|
const toolCatalogTTL = 24 * time.Hour
|
|
|
|
// MCPService serves the OpenAPI-driven MCP endpoint.
|
|
type MCPService struct {
|
|
profile *profile.Profile
|
|
|
|
operationsByTool map[string]*registeredOperation
|
|
handler http.Handler
|
|
}
|
|
|
|
// NewMCPService creates an MCP service backed by the in-process API routes.
|
|
func NewMCPService(profile *profile.Profile, echoServer *echo.Echo) (*MCPService, error) {
|
|
spec, err := loadMCPServiceOpenAPISpec()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
registry, err := buildOperationRegistry(spec)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
tools, operationsByTool, err := buildCuratedTools(registry)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
version := "dev"
|
|
if profile != nil && profile.Version != "" {
|
|
version = profile.Version
|
|
}
|
|
server := sdkmcp.NewServer(&sdkmcp.Implementation{
|
|
Name: "memos",
|
|
Version: version,
|
|
}, &sdkmcp.ServerOptions{
|
|
// The tool catalog never changes after startup, so do not advertise
|
|
// listChanged: new-protocol clients would otherwise hold a
|
|
// subscriptions/listen stream open for notifications that never come.
|
|
// Leaving Logging unset also drops the deprecated logging capability.
|
|
Capabilities: &sdkmcp.ServerCapabilities{
|
|
Tools: &sdkmcp.ToolCapabilities{},
|
|
},
|
|
})
|
|
server.AddReceivingMiddleware(catalogCacheMiddleware)
|
|
|
|
adapter := newAPIAdapter(echoServer)
|
|
for _, tool := range tools {
|
|
operation := operationsByTool[tool.Name]
|
|
server.AddTool(tool, newMCPToolHandler(adapter, operation))
|
|
}
|
|
|
|
streamableHandler := sdkmcp.NewStreamableHTTPHandler(func(*http.Request) *sdkmcp.Server {
|
|
return server
|
|
}, &sdkmcp.StreamableHTTPOptions{
|
|
// Stateless is also what lets the SDK serve protocol version
|
|
// 2026-07-28; stateful handlers negotiate down to 2025-11-25.
|
|
Stateless: true,
|
|
JSONResponse: true,
|
|
MaxRequestBodyBytes: maxMCPRequestBytes,
|
|
// memos is typically served behind a reverse proxy with the app bound to a
|
|
// loopback address while the public Host header is a real domain. The SDK's
|
|
// DNS-rebinding guard treats that shape as an attack and rejects every
|
|
// request with 403 ("invalid Host header"). Disable it and rely on memos'
|
|
// own Origin/Host allowlist (see RegisterRoutes -> isAllowedMCPOrigin) for
|
|
// CSRF / DNS-rebinding protection instead.
|
|
DisableLocalhostProtection: true,
|
|
})
|
|
return &MCPService{
|
|
profile: profile,
|
|
operationsByTool: operationsByTool,
|
|
handler: streamableHandler,
|
|
}, nil
|
|
}
|
|
|
|
// catalogCacheMiddleware stamps the static-catalog TTL onto results the SDK
|
|
// would otherwise mark immediately stale (ttlMs: 0).
|
|
func catalogCacheMiddleware(next sdkmcp.MethodHandler) sdkmcp.MethodHandler {
|
|
return func(ctx context.Context, method string, request sdkmcp.Request) (sdkmcp.Result, error) {
|
|
result, err := next(ctx, method, request)
|
|
if err != nil {
|
|
return result, err
|
|
}
|
|
switch typed := result.(type) {
|
|
case *sdkmcp.ListToolsResult:
|
|
typed.TTLMs = int(toolCatalogTTL.Milliseconds())
|
|
case *sdkmcp.DiscoverResult:
|
|
typed.TTLMs = int(toolCatalogTTL.Milliseconds())
|
|
default:
|
|
// Other results carry no catalog TTL.
|
|
}
|
|
return result, nil
|
|
}
|
|
}
|
|
|
|
func loadMCPServiceOpenAPISpec() (*openAPISpec, error) {
|
|
spec := &openAPISpec{}
|
|
if err := yaml.Unmarshal(memosproto.OpenAPIYAML(), spec); err != nil {
|
|
return nil, errors.Wrap(err, "failed to parse embedded OpenAPI spec")
|
|
}
|
|
if spec.Paths == nil {
|
|
return nil, errors.New("embedded OpenAPI spec has no paths")
|
|
}
|
|
return spec, nil
|
|
}
|
|
|
|
func newMCPToolHandler(adapter *apiAdapter, operation *registeredOperation) sdkmcp.ToolHandler {
|
|
return func(ctx context.Context, request *sdkmcp.CallToolRequest) (*sdkmcp.CallToolResult, error) {
|
|
arguments := map[string]any{}
|
|
if request.Params != nil && len(request.Params.Arguments) > 0 {
|
|
if err := json.Unmarshal(request.Params.Arguments, &arguments); err != nil {
|
|
return newToolErrorResult(errors.Wrap(err, "failed to decode MCP tool arguments").Error()), nil
|
|
}
|
|
}
|
|
if err := validateToolArguments(operation.InputSchema, arguments); err != nil {
|
|
return newToolErrorResult(err.Error()), nil
|
|
}
|
|
|
|
authorization := ""
|
|
if request.Extra != nil {
|
|
authorization = request.Extra.Header.Get("Authorization")
|
|
}
|
|
return adapter.execute(ctx, operation.Operation, arguments, authorization)
|
|
}
|
|
}
|
|
|
|
// RegisterRoutes registers the streamable HTTP MCP endpoint.
|
|
func (s *MCPService) RegisterRoutes(echoServer *echo.Echo) {
|
|
echoServer.Any("/mcp", func(c *echo.Context) error {
|
|
request := c.Request()
|
|
if !isAllowedMCPOrigin(request.Host, request.Header.Get("Origin"), s.profile) {
|
|
return c.NoContent(http.StatusForbidden)
|
|
}
|
|
s.handler.ServeHTTP(c.Response(), request)
|
|
return nil
|
|
}, middleware.BodyLimit(maxMCPRequestBytes))
|
|
}
|