0.29.1原版
This commit is contained in:
@@ -0,0 +1,167 @@
|
||||
# MCP Server
|
||||
|
||||
This package implements a [Model Context Protocol (MCP)](https://modelcontextprotocol.io) server embedded in the Memos HTTP process. It exposes memo operations as MCP tools, making Memos accessible to any MCP-compatible AI client (Claude Desktop, Cursor, Zed, etc.).
|
||||
|
||||
## Endpoint
|
||||
|
||||
```
|
||||
POST /mcp (tool calls, initialize)
|
||||
GET /mcp (optional SSE stream for server-to-client messages)
|
||||
DELETE /mcp (optional session termination)
|
||||
```
|
||||
|
||||
Transport: [Streamable HTTP](https://modelcontextprotocol.io/specification/2025-03-26/basic/transports) (single endpoint, MCP spec 2025-03-26).
|
||||
|
||||
### Tool Filtering
|
||||
|
||||
The default `/mcp` endpoint exposes all tools. Clients can opt into a smaller
|
||||
tool surface with GitHub-style headers or route aliases:
|
||||
|
||||
| Control | Description |
|
||||
|---|---|
|
||||
| `X-MCP-Readonly: true` | Hide and block mutating tools |
|
||||
| `X-MCP-Toolsets: memos,tags,attachments,relations,reactions` | Limit the default tool list to selected toolsets |
|
||||
| `X-MCP-Tools: list_tags,get_memo` | Add specific tools to the selected toolset list |
|
||||
| `X-MCP-Exclude-Tools: delete_memo` | Remove specific tools |
|
||||
|
||||
Equivalent aliases:
|
||||
|
||||
```text
|
||||
/mcp/readonly
|
||||
/mcp/x/{toolsets}
|
||||
/mcp/x/{toolsets}/readonly
|
||||
```
|
||||
|
||||
Examples:
|
||||
|
||||
```text
|
||||
/mcp/x/memos,tags/readonly
|
||||
X-MCP-Toolsets: memos
|
||||
X-MCP-Tools: list_tags
|
||||
X-MCP-Exclude-Tools: delete_memo
|
||||
```
|
||||
|
||||
## Capabilities
|
||||
|
||||
The server advertises the following MCP capabilities:
|
||||
|
||||
| Capability | Enabled | Details |
|
||||
|---|---|---|
|
||||
| Tools | Yes | List changed notifications supported |
|
||||
| Resources | Yes | Subscribe + list changed supported |
|
||||
| Prompts | Yes | List changed notifications supported |
|
||||
| Logging | Yes | Structured log events |
|
||||
|
||||
## Authentication
|
||||
|
||||
Public reads can be used without authentication. Personal Access Tokens (PATs) or short-lived JWT session tokens are required for:
|
||||
|
||||
- Reading non-public memos or attachments
|
||||
- Any tool that mutates data
|
||||
|
||||
When authenticating, send a Bearer token:
|
||||
|
||||
```
|
||||
Authorization: Bearer <your-PAT>
|
||||
```
|
||||
|
||||
PATs are long-lived tokens created in Settings → My Account → Access Tokens. Short-lived JWT session tokens are also accepted. Requests with an invalid token receive `HTTP 401`.
|
||||
|
||||
## Origin Validation
|
||||
|
||||
For Streamable HTTP safety, requests with an `Origin` header must be same-origin with the current request host or match the configured `instance-url`. Requests without an `Origin` header, such as desktop MCP clients and CLI tools, are allowed.
|
||||
|
||||
## Tools
|
||||
|
||||
### Memo Tools
|
||||
|
||||
| Tool | Description | Required params | Optional params |
|
||||
|---|---|---|---|
|
||||
| `list_memos` | List memos | — | `page_size`, `page`, `state`, `order_by_pinned`, `filter` (supported subset of standard CEL syntax) |
|
||||
| `get_memo` | Get a single memo | `name` | — |
|
||||
| `search_memos` | Full-text search | `query` | — |
|
||||
| `create_memo` | Create a memo | `content` | `visibility` |
|
||||
| `update_memo` | Update a memo | `name` | `content`, `visibility`, `pinned`, `state` |
|
||||
| `delete_memo` | Delete a memo | `name` | — |
|
||||
| `list_memo_comments` | List comments | `name` | — |
|
||||
| `create_memo_comment` | Add a comment | `name`, `content` | — |
|
||||
|
||||
### Tag Tools
|
||||
|
||||
| Tool | Description | Required params |
|
||||
|---|---|---|
|
||||
| `list_tags` | List all tags with counts | — |
|
||||
|
||||
### Attachment Tools
|
||||
|
||||
| Tool | Description | Required params | Optional params |
|
||||
|---|---|---|---|
|
||||
| `list_attachments` | List user's attachments | — | `page_size`, `page`, `memo` |
|
||||
| `get_attachment` | Get attachment metadata | `name` | — |
|
||||
| `delete_attachment` | Delete an attachment | `name` | — |
|
||||
| `link_attachment_to_memo` | Link attachment to a memo you own | `name`, `memo` | — |
|
||||
|
||||
### Relation Tools
|
||||
|
||||
| Tool | Description | Required params | Optional params |
|
||||
|---|---|---|---|
|
||||
| `list_memo_relations` | List relations (refs + comments) | `name` | `type` |
|
||||
| `create_memo_relation` | Create a reference relation from a memo you own to a memo you can read | `name`, `related_memo` | — |
|
||||
| `delete_memo_relation` | Delete a reference relation from a memo you own | `name`, `related_memo` | — |
|
||||
|
||||
### Reaction Tools
|
||||
|
||||
| Tool | Description | Required params |
|
||||
|---|---|---|
|
||||
| `list_reactions` | List reactions on a memo | `name` |
|
||||
| `upsert_reaction` | Add a reaction emoji | `name`, `reaction_type` |
|
||||
| `delete_reaction` | Remove a reaction | `id` |
|
||||
|
||||
## Resources
|
||||
|
||||
| URI Template | Description | MIME Type |
|
||||
|---|---|---|
|
||||
| `memo://memos/{uid}` | Memo content with YAML frontmatter | `text/markdown` |
|
||||
|
||||
## Prompts
|
||||
|
||||
| Prompt | Description | Arguments |
|
||||
|---|---|---|
|
||||
| `capture` | Quick-save a thought as a memo | `content` (required), `tags`, `visibility` |
|
||||
| `review` | Search and summarize memos on a topic | `topic` (required) |
|
||||
| `daily_digest` | Summarize recent memo activity | `days` |
|
||||
| `organize` | Suggest tags/relations for unorganized memos | `scope` |
|
||||
|
||||
## Resource Names
|
||||
|
||||
- Memos: `memos/<uid>` (e.g. `memos/abc123`)
|
||||
- Attachments: `attachments/<uid>` (e.g. `attachments/def456`)
|
||||
|
||||
## Connecting Claude Code
|
||||
|
||||
```bash
|
||||
claude mcp add --transport http memos http://localhost:5230/mcp \
|
||||
--header "Authorization: Bearer <your-PAT>"
|
||||
```
|
||||
|
||||
Use `--scope user` to make it available across all projects:
|
||||
|
||||
```bash
|
||||
claude mcp add --scope user --transport http memos http://localhost:5230/mcp \
|
||||
--header "Authorization: Bearer <your-PAT>"
|
||||
```
|
||||
|
||||
## Package Structure
|
||||
|
||||
| File | Responsibility |
|
||||
|---|---|
|
||||
| `mcp.go` | `MCPService` struct, constructor, route registration, auth middleware, tool filtering |
|
||||
| `tool_metadata.go` | Toolsets, read-only metadata, annotations, structured result helpers |
|
||||
| `api_helpers.go` | Conversion helpers for calling API service methods from MCP tools |
|
||||
| `tools_memo.go` | Memo CRUD tools + helpers (JSON types, visibility/access checks) |
|
||||
| `tools_tag.go` | Tag listing tool |
|
||||
| `tools_attachment.go` | Attachment listing, metadata, deletion, linking tools |
|
||||
| `tools_relation.go` | Memo relation (reference) tools |
|
||||
| `tools_reaction.go` | Reaction (emoji) tools |
|
||||
| `resources_memo.go` | Memo resource template handler |
|
||||
| `prompts.go` | Prompt handlers (capture, review, daily_digest, organize) |
|
||||
@@ -0,0 +1,117 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
// checkMemoAccess returns an error if the caller cannot read the memo.
|
||||
// userID == 0 means anonymous.
|
||||
func checkMemoAccess(memo *store.Memo, userID int32) error {
|
||||
if memo.RowStatus == store.Archived && memo.CreatorID != userID {
|
||||
return errors.New("permission denied")
|
||||
}
|
||||
|
||||
switch memo.Visibility {
|
||||
case store.Protected:
|
||||
if userID == 0 {
|
||||
return errors.New("permission denied")
|
||||
}
|
||||
case store.Private:
|
||||
if memo.CreatorID != userID {
|
||||
return errors.New("permission denied")
|
||||
}
|
||||
default:
|
||||
// store.Public and any unknown visibility: allow.
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func checkMemoOwnership(memo *store.Memo, userID int32) error {
|
||||
if memo.CreatorID != userID {
|
||||
return errors.New("permission denied")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func hasMemoOwnership(memo *store.Memo, userID int32) bool {
|
||||
return memo.CreatorID == userID
|
||||
}
|
||||
|
||||
// applyVisibilityFilter restricts find to memos the caller may see.
|
||||
func applyVisibilityFilter(find *store.FindMemo, userID int32, rowStatus *store.RowStatus) {
|
||||
if rowStatus != nil && *rowStatus == store.Archived {
|
||||
if userID == 0 {
|
||||
impossibleCreatorID := int32(-1)
|
||||
find.CreatorID = &impossibleCreatorID
|
||||
return
|
||||
}
|
||||
find.CreatorID = &userID
|
||||
return
|
||||
}
|
||||
if userID == 0 {
|
||||
find.VisibilityList = []store.Visibility{store.Public}
|
||||
return
|
||||
}
|
||||
find.Filters = append(find.Filters, "creator_id == "+itoa32(userID)+` || visibility in ["PUBLIC", "PROTECTED"]`)
|
||||
}
|
||||
|
||||
func (s *MCPService) checkAttachmentAccess(ctx context.Context, attachment *store.Attachment, userID int32) error {
|
||||
if attachment.CreatorID == userID {
|
||||
return nil
|
||||
}
|
||||
if attachment.MemoID == nil {
|
||||
return errors.New("permission denied")
|
||||
}
|
||||
|
||||
memo, err := s.store.GetMemo(ctx, &store.FindMemo{ID: attachment.MemoID})
|
||||
if err != nil {
|
||||
return errors.Wrap(err, "failed to get linked memo")
|
||||
}
|
||||
if memo == nil {
|
||||
return errors.New("linked memo not found")
|
||||
}
|
||||
return checkMemoAccess(memo, userID)
|
||||
}
|
||||
|
||||
func (s *MCPService) isAllowedOrigin(r *http.Request) bool {
|
||||
origin := r.Header.Get("Origin")
|
||||
if origin == "" {
|
||||
return true
|
||||
}
|
||||
|
||||
originURL, err := url.Parse(origin)
|
||||
if err != nil || originURL.Scheme == "" || originURL.Host == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
if sameOriginHost(originURL.Host, r.Host) {
|
||||
return true
|
||||
}
|
||||
|
||||
if s.profile.InstanceURL == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
instanceURL, err := url.Parse(s.profile.InstanceURL)
|
||||
if err != nil || instanceURL.Scheme == "" || instanceURL.Host == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
return strings.EqualFold(originURL.Scheme, instanceURL.Scheme) && sameOriginHost(originURL.Host, instanceURL.Host)
|
||||
}
|
||||
|
||||
func sameOriginHost(a, b string) bool {
|
||||
return strings.EqualFold(a, b)
|
||||
}
|
||||
|
||||
func itoa32(v int32) string {
|
||||
return strconv.FormatInt(int64(v), 10)
|
||||
}
|
||||
@@ -0,0 +1,74 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
apiv1 "github.com/usememos/memos/server/router/api/v1"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func visibilityToProto(visibility store.Visibility) v1pb.Visibility {
|
||||
switch visibility {
|
||||
case store.Protected:
|
||||
return v1pb.Visibility_PROTECTED
|
||||
case store.Public:
|
||||
return v1pb.Visibility_PUBLIC
|
||||
default:
|
||||
return v1pb.Visibility_PRIVATE
|
||||
}
|
||||
}
|
||||
|
||||
func rowStatusToProto(rowStatus store.RowStatus) v1pb.State {
|
||||
switch rowStatus {
|
||||
case store.Archived:
|
||||
return v1pb.State_ARCHIVED
|
||||
default:
|
||||
return v1pb.State_NORMAL
|
||||
}
|
||||
}
|
||||
|
||||
func (s *MCPService) loadMemoJSONByName(ctx context.Context, name string) (memoJSON, error) {
|
||||
uid, err := parseMemoUID(name)
|
||||
if err != nil {
|
||||
return memoJSON{}, err
|
||||
}
|
||||
memo, err := s.store.GetMemo(ctx, &store.FindMemo{UID: &uid})
|
||||
if err != nil {
|
||||
return memoJSON{}, errors.Wrap(err, "failed to get memo")
|
||||
}
|
||||
if memo == nil {
|
||||
return memoJSON{}, errors.New("memo not found")
|
||||
}
|
||||
return storeMemoToJSONWithStore(ctx, s.store, memo)
|
||||
}
|
||||
|
||||
func (s *MCPService) loadReactionJSONByID(ctx context.Context, reactionID int32) (reactionJSON, error) {
|
||||
reaction, err := s.store.GetReaction(ctx, &store.FindReaction{ID: &reactionID})
|
||||
if err != nil {
|
||||
return reactionJSON{}, errors.Wrap(err, "failed to get reaction")
|
||||
}
|
||||
if reaction == nil {
|
||||
return reactionJSON{}, errors.New("reaction not found")
|
||||
}
|
||||
creator, err := lookupUsername(ctx, s.store, reaction.CreatorID)
|
||||
if err != nil {
|
||||
return reactionJSON{}, errors.Wrap(err, "failed to resolve reaction creator")
|
||||
}
|
||||
return reactionJSON{
|
||||
ID: reaction.ID,
|
||||
Creator: creator,
|
||||
ReactionType: reaction.ReactionType,
|
||||
CreateTime: reaction.CreatedTs,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *MCPService) loadReactionJSONByName(ctx context.Context, name string) (reactionJSON, error) {
|
||||
_, reactionID, err := apiv1.ExtractMemoReactionIDFromName(name)
|
||||
if err != nil {
|
||||
return reactionJSON{}, err
|
||||
}
|
||||
return s.loadReactionJSONByID(ctx, reactionID)
|
||||
}
|
||||
@@ -0,0 +1,251 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/labstack/echo/v5"
|
||||
"github.com/mark3labs/mcp-go/mcp"
|
||||
mcpserver "github.com/mark3labs/mcp-go/server"
|
||||
|
||||
"github.com/usememos/memos/internal/profile"
|
||||
"github.com/usememos/memos/server/auth"
|
||||
apiv1 "github.com/usememos/memos/server/router/api/v1"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
const (
|
||||
headerMCPReadonly = "X-MCP-Readonly"
|
||||
headerMCPToolsets = "X-MCP-Toolsets"
|
||||
headerMCPTools = "X-MCP-Tools"
|
||||
headerMCPExcludeTools = "X-MCP-Exclude-Tools"
|
||||
)
|
||||
|
||||
type mcpRequestConfigContextKey struct{}
|
||||
|
||||
type MCPService struct {
|
||||
profile *profile.Profile
|
||||
store *store.Store
|
||||
apiV1Service *apiv1.APIV1Service
|
||||
authenticator *auth.Authenticator
|
||||
}
|
||||
|
||||
func NewMCPService(profile *profile.Profile, store *store.Store, secret string, apiV1Service *apiv1.APIV1Service) *MCPService {
|
||||
return &MCPService{
|
||||
profile: profile,
|
||||
store: store,
|
||||
apiV1Service: apiV1Service,
|
||||
authenticator: auth.NewAuthenticator(store, secret),
|
||||
}
|
||||
}
|
||||
|
||||
func (s *MCPService) RegisterRoutes(echoServer *echo.Echo) {
|
||||
mcpSrv := mcpserver.NewMCPServer("Memos", "1.0.0",
|
||||
mcpserver.WithToolCapabilities(true),
|
||||
mcpserver.WithResourceCapabilities(true, true),
|
||||
mcpserver.WithPromptCapabilities(true),
|
||||
mcpserver.WithLogging(),
|
||||
mcpserver.WithToolFilter(s.filterTools),
|
||||
mcpserver.WithToolHandlerMiddleware(s.enforceToolAccess),
|
||||
mcpserver.WithRecovery(),
|
||||
mcpserver.WithResourceRecovery(),
|
||||
)
|
||||
s.registerMemoTools(mcpSrv)
|
||||
s.registerTagTools(mcpSrv)
|
||||
s.registerAttachmentTools(mcpSrv)
|
||||
s.registerRelationTools(mcpSrv)
|
||||
s.registerReactionTools(mcpSrv)
|
||||
s.registerMemoResources(mcpSrv)
|
||||
s.registerPrompts(mcpSrv)
|
||||
|
||||
httpHandler := mcpserver.NewStreamableHTTPServer(mcpSrv,
|
||||
mcpserver.WithHTTPContextFunc(s.withRequestConfig),
|
||||
)
|
||||
|
||||
mcpGroup := echoServer.Group("")
|
||||
mcpGroup.Use(func(next echo.HandlerFunc) echo.HandlerFunc {
|
||||
return func(c *echo.Context) error {
|
||||
if !s.isAllowedOrigin(c.Request()) {
|
||||
return c.JSON(http.StatusForbidden, map[string]string{"message": "invalid origin"})
|
||||
}
|
||||
if origin := c.Request().Header.Get("Origin"); origin != "" {
|
||||
headers := c.Response().Header()
|
||||
headers.Set("Vary", "Origin")
|
||||
headers.Set("Access-Control-Allow-Origin", origin)
|
||||
headers.Set("Access-Control-Allow-Headers", strings.Join([]string{
|
||||
"Authorization",
|
||||
"Content-Type",
|
||||
"Accept",
|
||||
"Mcp-Session-Id",
|
||||
"MCP-Protocol-Version",
|
||||
"Last-Event-ID",
|
||||
headerMCPReadonly,
|
||||
headerMCPToolsets,
|
||||
headerMCPTools,
|
||||
headerMCPExcludeTools,
|
||||
}, ", "))
|
||||
headers.Set("Access-Control-Allow-Methods", "GET, POST, DELETE, OPTIONS")
|
||||
if c.Request().Method == http.MethodOptions {
|
||||
return c.NoContent(http.StatusNoContent)
|
||||
}
|
||||
}
|
||||
|
||||
authHeader := c.Request().Header.Get("Authorization")
|
||||
if authHeader != "" {
|
||||
result := s.authenticator.Authenticate(c.Request().Context(), authHeader)
|
||||
if result == nil {
|
||||
return c.JSON(http.StatusUnauthorized, map[string]string{"message": "invalid or expired token"})
|
||||
}
|
||||
ctx := auth.ApplyToContext(c.Request().Context(), result)
|
||||
c.SetRequest(c.Request().WithContext(ctx))
|
||||
}
|
||||
return next(c)
|
||||
}
|
||||
})
|
||||
mcpGroup.Any("/mcp", echo.WrapHandler(httpHandler))
|
||||
mcpGroup.Any("/mcp/readonly", echo.WrapHandler(httpHandler))
|
||||
mcpGroup.Any("/mcp/x/:toolsets", echo.WrapHandler(httpHandler))
|
||||
mcpGroup.Any("/mcp/x/:toolsets/readonly", echo.WrapHandler(httpHandler))
|
||||
}
|
||||
|
||||
func (*MCPService) withRequestConfig(ctx context.Context, r *http.Request) context.Context {
|
||||
return context.WithValue(ctx, mcpRequestConfigContextKey{}, parseMCPRequestConfig(r))
|
||||
}
|
||||
|
||||
func (*MCPService) filterTools(ctx context.Context, tools []mcp.Tool) []mcp.Tool {
|
||||
cfg := mcpRequestConfigFromContext(ctx)
|
||||
filtered := make([]mcp.Tool, 0, len(tools))
|
||||
for _, tool := range tools {
|
||||
if cfg.allowsTool(tool.Name) {
|
||||
filtered = append(filtered, tool)
|
||||
}
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
func (*MCPService) enforceToolAccess(next mcpserver.ToolHandlerFunc) mcpserver.ToolHandlerFunc {
|
||||
return func(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
cfg := mcpRequestConfigFromContext(ctx)
|
||||
if !cfg.allowsTool(req.Params.Name) {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("tool %q is not enabled by MCP configuration", req.Params.Name)), nil
|
||||
}
|
||||
return next(ctx, req)
|
||||
}
|
||||
}
|
||||
|
||||
type mcpRequestConfig struct {
|
||||
readOnly bool
|
||||
toolsets map[string]struct{}
|
||||
includeTools map[string]struct{}
|
||||
excludeTools map[string]struct{}
|
||||
}
|
||||
|
||||
func mcpRequestConfigFromContext(ctx context.Context) mcpRequestConfig {
|
||||
if cfg, ok := ctx.Value(mcpRequestConfigContextKey{}).(mcpRequestConfig); ok {
|
||||
return cfg
|
||||
}
|
||||
return mcpRequestConfig{}
|
||||
}
|
||||
|
||||
func parseMCPRequestConfig(r *http.Request) mcpRequestConfig {
|
||||
cfg := mcpRequestConfig{}
|
||||
|
||||
pathToolsets, pathReadonly := parseMCPPathConfig(r.URL.Path)
|
||||
cfg.readOnly = pathReadonly || parseBoolHeader(r.Header.Get(headerMCPReadonly))
|
||||
cfg.toolsets = mergeStringSets(cfg.toolsets, pathToolsets)
|
||||
cfg.toolsets = mergeStringSets(cfg.toolsets, parseCommaSet(r.Header.Get(headerMCPToolsets), strings.ToLower))
|
||||
cfg.includeTools = parseCommaSet(r.Header.Get(headerMCPTools), keepString)
|
||||
cfg.excludeTools = parseCommaSet(r.Header.Get(headerMCPExcludeTools), keepString)
|
||||
return cfg
|
||||
}
|
||||
|
||||
func parseMCPPathConfig(path string) (map[string]struct{}, bool) {
|
||||
trimmed := strings.Trim(path, "/")
|
||||
if trimmed == "mcp/readonly" {
|
||||
return nil, true
|
||||
}
|
||||
const prefix = "mcp/x/"
|
||||
if !strings.HasPrefix(trimmed, prefix) {
|
||||
return nil, false
|
||||
}
|
||||
|
||||
rest := strings.TrimPrefix(trimmed, prefix)
|
||||
readOnly := false
|
||||
if strings.HasSuffix(rest, "/readonly") {
|
||||
readOnly = true
|
||||
rest = strings.TrimSuffix(rest, "/readonly")
|
||||
}
|
||||
return parseCommaSet(rest, strings.ToLower), readOnly
|
||||
}
|
||||
|
||||
func (cfg mcpRequestConfig) allowsTool(name string) bool {
|
||||
if _, known := allMCPToolNames[name]; !known {
|
||||
return false
|
||||
}
|
||||
if cfg.readOnly {
|
||||
if _, mutates := mcpMutationTools[name]; mutates {
|
||||
return false
|
||||
}
|
||||
}
|
||||
if _, excluded := cfg.excludeTools[name]; excluded {
|
||||
return false
|
||||
}
|
||||
if _, included := cfg.includeTools[name]; included {
|
||||
return true
|
||||
}
|
||||
if len(cfg.toolsets) == 0 {
|
||||
return true
|
||||
}
|
||||
for toolset := range cfg.toolsets {
|
||||
if _, ok := mcpToolsByToolset[toolset][name]; ok {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func parseBoolHeader(value string) bool {
|
||||
switch strings.ToLower(strings.TrimSpace(value)) {
|
||||
case "1", "t", "true", "y", "yes", "on":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func parseCommaSet(value string, normalize func(string) string) map[string]struct{} {
|
||||
if value == "" {
|
||||
return nil
|
||||
}
|
||||
result := map[string]struct{}{}
|
||||
for _, item := range strings.Split(value, ",") {
|
||||
item = strings.TrimSpace(item)
|
||||
if item == "" {
|
||||
continue
|
||||
}
|
||||
result[normalize(item)] = struct{}{}
|
||||
}
|
||||
if len(result) == 0 {
|
||||
return nil
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func mergeStringSets(dst map[string]struct{}, src map[string]struct{}) map[string]struct{} {
|
||||
if len(src) == 0 {
|
||||
return dst
|
||||
}
|
||||
if dst == nil {
|
||||
dst = map[string]struct{}{}
|
||||
}
|
||||
for item := range src {
|
||||
dst[item] = struct{}{}
|
||||
}
|
||||
return dst
|
||||
}
|
||||
|
||||
func keepString(s string) string {
|
||||
return s
|
||||
}
|
||||
@@ -0,0 +1,605 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"reflect"
|
||||
"testing"
|
||||
"time"
|
||||
"unsafe"
|
||||
|
||||
"github.com/labstack/echo/v5"
|
||||
"github.com/lithammer/shortuuid/v4"
|
||||
"github.com/mark3labs/mcp-go/mcp"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/usememos/memos/internal/profile"
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
"github.com/usememos/memos/server/auth"
|
||||
apiv1service "github.com/usememos/memos/server/router/api/v1"
|
||||
"github.com/usememos/memos/store"
|
||||
teststore "github.com/usememos/memos/store/test"
|
||||
)
|
||||
|
||||
type testMCPService struct {
|
||||
service *MCPService
|
||||
store *store.Store
|
||||
}
|
||||
|
||||
func newTestMCPService(t *testing.T) *testMCPService {
|
||||
t.Helper()
|
||||
|
||||
ctx := context.Background()
|
||||
stores := teststore.NewTestingStore(ctx, t)
|
||||
t.Cleanup(func() {
|
||||
require.NoError(t, stores.Close())
|
||||
})
|
||||
|
||||
profile := &profile.Profile{
|
||||
Driver: "sqlite",
|
||||
InstanceURL: "https://notes.example.com",
|
||||
}
|
||||
apiV1Service := apiv1service.NewAPIV1Service("test-secret", profile, stores)
|
||||
svc := NewMCPService(profile, stores, "test-secret", apiV1Service)
|
||||
return &testMCPService{
|
||||
service: svc,
|
||||
store: stores,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *testMCPService) createUser(t *testing.T, username string) *store.User {
|
||||
t.Helper()
|
||||
|
||||
user, err := s.store.CreateUser(context.Background(), &store.User{
|
||||
Username: username,
|
||||
Role: store.RoleUser,
|
||||
Email: username + "@example.com",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
return user
|
||||
}
|
||||
|
||||
func (s *testMCPService) createMemo(t *testing.T, creatorID int32, visibility store.Visibility, content string) *store.Memo {
|
||||
t.Helper()
|
||||
|
||||
memo, err := s.store.CreateMemo(context.Background(), &store.Memo{
|
||||
UID: shortuuid.New(),
|
||||
CreatorID: creatorID,
|
||||
RowStatus: store.Normal,
|
||||
Visibility: visibility,
|
||||
Content: content,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
return memo
|
||||
}
|
||||
|
||||
func (s *testMCPService) archiveMemo(t *testing.T, memoID int32) {
|
||||
t.Helper()
|
||||
|
||||
rowStatus := store.Archived
|
||||
require.NoError(t, s.store.UpdateMemo(context.Background(), &store.UpdateMemo{
|
||||
ID: memoID,
|
||||
RowStatus: &rowStatus,
|
||||
}))
|
||||
}
|
||||
|
||||
func (s *testMCPService) createAttachment(t *testing.T, creatorID int32, memoID *int32) *store.Attachment {
|
||||
t.Helper()
|
||||
|
||||
attachment, err := s.store.CreateAttachment(context.Background(), &store.Attachment{
|
||||
UID: shortuuid.New(),
|
||||
CreatorID: creatorID,
|
||||
Filename: "note.txt",
|
||||
Type: "text/plain",
|
||||
Size: 4,
|
||||
StorageType: storepb.AttachmentStorageType_ATTACHMENT_STORAGE_TYPE_UNSPECIFIED,
|
||||
Reference: "db://attachment/note.txt",
|
||||
MemoID: memoID,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
return attachment
|
||||
}
|
||||
|
||||
func withUser(ctx context.Context, userID int32) context.Context {
|
||||
return context.WithValue(ctx, auth.UserIDContextKey, userID)
|
||||
}
|
||||
|
||||
func toolRequest(name string, arguments map[string]any) mcp.CallToolRequest {
|
||||
return mcp.CallToolRequest{
|
||||
Params: mcp.CallToolParams{
|
||||
Name: name,
|
||||
Arguments: arguments,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func firstText(t *testing.T, result *mcp.CallToolResult) string {
|
||||
t.Helper()
|
||||
require.NotEmpty(t, result.Content)
|
||||
text, ok := result.Content[0].(mcp.TextContent)
|
||||
require.True(t, ok)
|
||||
return text.Text
|
||||
}
|
||||
|
||||
func initializeMCPHTTP(t *testing.T, e *echo.Echo, path string, headers map[string]string) string {
|
||||
t.Helper()
|
||||
payload := 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": "mcp-test",
|
||||
"version": "1.0.0",
|
||||
},
|
||||
},
|
||||
}
|
||||
resp := postMCPHTTP(t, e, path, "", headers, payload)
|
||||
require.Equal(t, http.StatusOK, resp.Code, resp.Body.String())
|
||||
sessionID := resp.Header().Get("Mcp-Session-Id")
|
||||
require.NotEmpty(t, sessionID)
|
||||
return sessionID
|
||||
}
|
||||
|
||||
func callMCPHTTP(t *testing.T, e *echo.Echo, path string, sessionID string, headers map[string]string, method string, params any) map[string]any {
|
||||
t.Helper()
|
||||
payload := map[string]any{
|
||||
"jsonrpc": "2.0",
|
||||
"id": 2,
|
||||
"method": method,
|
||||
}
|
||||
if params != nil {
|
||||
payload["params"] = params
|
||||
}
|
||||
resp := postMCPHTTP(t, e, path, sessionID, headers, payload)
|
||||
require.Equal(t, http.StatusOK, resp.Code, resp.Body.String())
|
||||
var decoded map[string]any
|
||||
require.NoError(t, json.Unmarshal(resp.Body.Bytes(), &decoded))
|
||||
return decoded
|
||||
}
|
||||
|
||||
func postMCPHTTP(t *testing.T, e *echo.Echo, path string, sessionID string, headers map[string]string, payload map[string]any) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
body, err := json.Marshal(payload)
|
||||
require.NoError(t, err)
|
||||
req := httptest.NewRequest(http.MethodPost, path, bytes.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
if sessionID != "" {
|
||||
req.Header.Set("Mcp-Session-Id", sessionID)
|
||||
}
|
||||
for key, value := range headers {
|
||||
req.Header.Set(key, value)
|
||||
}
|
||||
resp := httptest.NewRecorder()
|
||||
e.ServeHTTP(resp, req)
|
||||
return resp
|
||||
}
|
||||
|
||||
func toolNamesFromListResponse(t *testing.T, response map[string]any) map[string]struct{} {
|
||||
t.Helper()
|
||||
result, ok := response["result"].(map[string]any)
|
||||
require.True(t, ok, "missing result: %#v", response)
|
||||
rawTools, ok := result["tools"].([]any)
|
||||
require.True(t, ok, "missing tools: %#v", result)
|
||||
names := map[string]struct{}{}
|
||||
for _, rawTool := range rawTools {
|
||||
tool, ok := rawTool.(map[string]any)
|
||||
require.True(t, ok)
|
||||
name, ok := tool["name"].(string)
|
||||
require.True(t, ok)
|
||||
names[name] = struct{}{}
|
||||
}
|
||||
return names
|
||||
}
|
||||
|
||||
func requireToolPresent(t *testing.T, names map[string]struct{}, name string) {
|
||||
t.Helper()
|
||||
_, ok := names[name]
|
||||
require.True(t, ok, "expected tool %q to be present in %#v", name, names)
|
||||
}
|
||||
|
||||
func requireToolAbsent(t *testing.T, names map[string]struct{}, name string) {
|
||||
t.Helper()
|
||||
_, ok := names[name]
|
||||
require.False(t, ok, "expected tool %q to be absent in %#v", name, names)
|
||||
}
|
||||
|
||||
func nextSSEEvent(t *testing.T, client *apiv1service.SSEClient) *apiv1service.SSEEvent {
|
||||
t.Helper()
|
||||
events := sseClientEvents(t, client)
|
||||
var data []byte
|
||||
select {
|
||||
case eventData, ok := <-events:
|
||||
require.True(t, ok, "SSE client channel closed")
|
||||
data = eventData
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("timed out waiting for SSE event")
|
||||
}
|
||||
var event apiv1service.SSEEvent
|
||||
require.NoError(t, json.Unmarshal(data, &event))
|
||||
return &event
|
||||
}
|
||||
|
||||
func requireNoSSEEvent(t *testing.T, client *apiv1service.SSEClient) {
|
||||
t.Helper()
|
||||
select {
|
||||
case eventData, ok := <-sseClientEvents(t, client):
|
||||
require.True(t, ok, "SSE client channel closed")
|
||||
t.Fatalf("unexpected SSE event received: %s", string(eventData))
|
||||
case <-time.After(150 * time.Millisecond):
|
||||
}
|
||||
}
|
||||
|
||||
func sseClientEvents(t *testing.T, client *apiv1service.SSEClient) <-chan []byte {
|
||||
t.Helper()
|
||||
field := reflect.ValueOf(client).Elem().FieldByName("events")
|
||||
events, ok := reflect.NewAt(field.Type(), unsafe.Pointer(field.UnsafeAddr())).Elem().Interface().(chan []byte)
|
||||
require.True(t, ok)
|
||||
return events
|
||||
}
|
||||
|
||||
func TestHandleGetMemoAndReadResourceDenyArchivedMemoToNonCreator(t *testing.T) {
|
||||
ts := newTestMCPService(t)
|
||||
owner := ts.createUser(t, "owner")
|
||||
other := ts.createUser(t, "other")
|
||||
|
||||
memo := ts.createMemo(t, owner.ID, store.Public, "archived")
|
||||
ts.archiveMemo(t, memo.ID)
|
||||
|
||||
ctx := withUser(context.Background(), other.ID)
|
||||
result, err := ts.service.handleGetMemo(ctx, toolRequest("get_memo", map[string]any{
|
||||
"name": "memos/" + memo.UID,
|
||||
}))
|
||||
require.NoError(t, err)
|
||||
require.True(t, result.IsError)
|
||||
require.Contains(t, firstText(t, result), "permission denied")
|
||||
|
||||
_, err = ts.service.handleReadMemoResource(ctx, mcp.ReadResourceRequest{
|
||||
Params: mcp.ReadResourceParams{
|
||||
URI: "memo://memos/" + memo.UID,
|
||||
},
|
||||
})
|
||||
require.ErrorContains(t, err, "permission denied")
|
||||
}
|
||||
|
||||
func TestHandleListMemosArchivedOnlyReturnsCreatorMemos(t *testing.T) {
|
||||
ts := newTestMCPService(t)
|
||||
owner := ts.createUser(t, "owner")
|
||||
other := ts.createUser(t, "other")
|
||||
|
||||
ownerMemo := ts.createMemo(t, owner.ID, store.Public, "owner archived")
|
||||
ts.archiveMemo(t, ownerMemo.ID)
|
||||
otherMemo := ts.createMemo(t, other.ID, store.Public, "other archived")
|
||||
ts.archiveMemo(t, otherMemo.ID)
|
||||
|
||||
result, err := ts.service.handleListMemos(withUser(context.Background(), owner.ID), toolRequest("list_memos", map[string]any{
|
||||
"state": "ARCHIVED",
|
||||
}))
|
||||
require.NoError(t, err)
|
||||
require.False(t, result.IsError)
|
||||
|
||||
var payload struct {
|
||||
Memos []memoJSON `json:"memos"`
|
||||
}
|
||||
require.NoError(t, json.Unmarshal([]byte(firstText(t, result)), &payload))
|
||||
require.Len(t, payload.Memos, 1)
|
||||
require.Equal(t, "memos/"+ownerMemo.UID, payload.Memos[0].Name)
|
||||
|
||||
anonResult, err := ts.service.handleListMemos(context.Background(), toolRequest("list_memos", map[string]any{
|
||||
"state": "ARCHIVED",
|
||||
}))
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, json.Unmarshal([]byte(firstText(t, anonResult)), &payload))
|
||||
require.Empty(t, payload.Memos)
|
||||
}
|
||||
|
||||
func TestHandleListMemoRelationsFiltersUnreadableTargets(t *testing.T) {
|
||||
ts := newTestMCPService(t)
|
||||
owner := ts.createUser(t, "owner")
|
||||
privateUser := ts.createUser(t, "private-user")
|
||||
publicUser := ts.createUser(t, "public-user")
|
||||
|
||||
source := ts.createMemo(t, owner.ID, store.Public, "source")
|
||||
privateTarget := ts.createMemo(t, privateUser.ID, store.Private, "private")
|
||||
publicTarget := ts.createMemo(t, publicUser.ID, store.Public, "public")
|
||||
|
||||
_, err := ts.store.UpsertMemoRelation(context.Background(), &store.MemoRelation{
|
||||
MemoID: source.ID,
|
||||
RelatedMemoID: privateTarget.ID,
|
||||
Type: store.MemoRelationReference,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
_, err = ts.store.UpsertMemoRelation(context.Background(), &store.MemoRelation{
|
||||
MemoID: source.ID,
|
||||
RelatedMemoID: publicTarget.ID,
|
||||
Type: store.MemoRelationReference,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
result, err := ts.service.handleListMemoRelations(context.Background(), toolRequest("list_memo_relations", map[string]any{
|
||||
"name": "memos/" + source.UID,
|
||||
}))
|
||||
require.NoError(t, err)
|
||||
require.False(t, result.IsError)
|
||||
|
||||
var relations []relationJSON
|
||||
require.NoError(t, json.Unmarshal([]byte(firstText(t, result)), &relations))
|
||||
require.Len(t, relations, 1)
|
||||
require.Equal(t, "memos/"+publicTarget.UID, relations[0].RelatedMemo)
|
||||
|
||||
denied, err := ts.service.handleListMemoRelations(context.Background(), toolRequest("list_memo_relations", map[string]any{
|
||||
"name": "memos/" + privateTarget.UID,
|
||||
}))
|
||||
require.NoError(t, err)
|
||||
require.True(t, denied.IsError)
|
||||
require.Contains(t, firstText(t, denied), "permission denied")
|
||||
}
|
||||
|
||||
func TestHandleLinkAttachmentToMemoRequiresMemoOwnership(t *testing.T) {
|
||||
ts := newTestMCPService(t)
|
||||
attachmentOwner := ts.createUser(t, "attachment-owner")
|
||||
memoOwner := ts.createUser(t, "memo-owner")
|
||||
|
||||
attachment := ts.createAttachment(t, attachmentOwner.ID, nil)
|
||||
memo := ts.createMemo(t, memoOwner.ID, store.Public, "target")
|
||||
|
||||
result, err := ts.service.handleLinkAttachmentToMemo(withUser(context.Background(), attachmentOwner.ID), toolRequest("link_attachment_to_memo", map[string]any{
|
||||
"name": "attachments/" + attachment.UID,
|
||||
"memo": "memos/" + memo.UID,
|
||||
}))
|
||||
require.NoError(t, err)
|
||||
require.True(t, result.IsError)
|
||||
require.Contains(t, firstText(t, result), "permission denied")
|
||||
}
|
||||
|
||||
func TestHandleGetAttachmentDeniesArchivedLinkedMemoToNonCreator(t *testing.T) {
|
||||
ts := newTestMCPService(t)
|
||||
owner := ts.createUser(t, "owner")
|
||||
other := ts.createUser(t, "other")
|
||||
|
||||
memo := ts.createMemo(t, owner.ID, store.Public, "memo")
|
||||
ts.archiveMemo(t, memo.ID)
|
||||
attachment := ts.createAttachment(t, owner.ID, &memo.ID)
|
||||
|
||||
result, err := ts.service.handleGetAttachment(withUser(context.Background(), other.ID), toolRequest("get_attachment", map[string]any{
|
||||
"name": "attachments/" + attachment.UID,
|
||||
}))
|
||||
require.NoError(t, err)
|
||||
require.True(t, result.IsError)
|
||||
require.Contains(t, firstText(t, result), "permission denied")
|
||||
}
|
||||
|
||||
func TestIsAllowedOrigin(t *testing.T) {
|
||||
ts := newTestMCPService(t)
|
||||
|
||||
t.Run("allow missing origin", func(t *testing.T) {
|
||||
req := httptest.NewRequest("POST", "http://localhost:5230/mcp", nil)
|
||||
require.True(t, ts.service.isAllowedOrigin(req))
|
||||
})
|
||||
|
||||
t.Run("allow same origin as request host", func(t *testing.T) {
|
||||
req := httptest.NewRequest("POST", "http://localhost:5230/mcp", nil)
|
||||
req.Header.Set("Origin", "http://localhost:5230")
|
||||
require.True(t, ts.service.isAllowedOrigin(req))
|
||||
})
|
||||
|
||||
t.Run("allow configured instance origin", func(t *testing.T) {
|
||||
req := httptest.NewRequest("POST", "http://127.0.0.1:5230/mcp", nil)
|
||||
req.Host = "127.0.0.1:5230"
|
||||
req.Header.Set("Origin", "https://notes.example.com")
|
||||
require.True(t, ts.service.isAllowedOrigin(req))
|
||||
})
|
||||
|
||||
t.Run("reject cross origin", func(t *testing.T) {
|
||||
req := httptest.NewRequest("POST", "http://localhost:5230/mcp", nil)
|
||||
req.Header.Set("Origin", "https://evil.example.com")
|
||||
require.False(t, ts.service.isAllowedOrigin(req))
|
||||
})
|
||||
}
|
||||
|
||||
func TestMCPToolFilteringRoutesAndHeaders(t *testing.T) {
|
||||
ts := newTestMCPService(t)
|
||||
e := echo.New()
|
||||
ts.service.RegisterRoutes(e)
|
||||
|
||||
t.Run("default endpoint lists all tools", func(t *testing.T) {
|
||||
sessionID := initializeMCPHTTP(t, e, "/mcp", nil)
|
||||
response := callMCPHTTP(t, e, "/mcp", sessionID, nil, "tools/list", map[string]any{})
|
||||
names := toolNamesFromListResponse(t, response)
|
||||
require.Len(t, names, len(allMCPToolNames))
|
||||
requireToolPresent(t, names, "create_memo")
|
||||
requireToolPresent(t, names, "list_tags")
|
||||
requireToolPresent(t, names, "upsert_reaction")
|
||||
})
|
||||
|
||||
t.Run("readonly header hides and blocks mutation tools", func(t *testing.T) {
|
||||
headers := map[string]string{headerMCPReadonly: "true"}
|
||||
sessionID := initializeMCPHTTP(t, e, "/mcp", nil)
|
||||
response := callMCPHTTP(t, e, "/mcp", sessionID, headers, "tools/list", map[string]any{})
|
||||
names := toolNamesFromListResponse(t, response)
|
||||
requireToolPresent(t, names, "list_memos")
|
||||
requireToolPresent(t, names, "list_tags")
|
||||
requireToolAbsent(t, names, "create_memo")
|
||||
requireToolAbsent(t, names, "delete_memo")
|
||||
|
||||
callResponse := callMCPHTTP(t, e, "/mcp", sessionID, headers, "tools/call", map[string]any{
|
||||
"name": "create_memo",
|
||||
"arguments": map[string]any{"content": "blocked"},
|
||||
})
|
||||
result, ok := callResponse["result"].(map[string]any)
|
||||
require.True(t, ok)
|
||||
require.Equal(t, true, result["isError"])
|
||||
rawContent, ok := result["content"].([]any)
|
||||
require.True(t, ok)
|
||||
content, ok := rawContent[0].(map[string]any)
|
||||
require.True(t, ok)
|
||||
require.Contains(t, content["text"], "not enabled")
|
||||
})
|
||||
|
||||
t.Run("readonly alias applies path config", func(t *testing.T) {
|
||||
sessionID := initializeMCPHTTP(t, e, "/mcp/readonly", nil)
|
||||
response := callMCPHTTP(t, e, "/mcp/readonly", sessionID, nil, "tools/list", map[string]any{})
|
||||
names := toolNamesFromListResponse(t, response)
|
||||
requireToolPresent(t, names, "get_memo")
|
||||
requireToolAbsent(t, names, "create_memo")
|
||||
requireToolAbsent(t, names, "upsert_reaction")
|
||||
})
|
||||
|
||||
t.Run("toolsets include and exclude compose", func(t *testing.T) {
|
||||
headers := map[string]string{
|
||||
headerMCPToolsets: "memos",
|
||||
headerMCPTools: "list_tags",
|
||||
headerMCPExcludeTools: "get_memo",
|
||||
}
|
||||
sessionID := initializeMCPHTTP(t, e, "/mcp", nil)
|
||||
response := callMCPHTTP(t, e, "/mcp", sessionID, headers, "tools/list", map[string]any{})
|
||||
names := toolNamesFromListResponse(t, response)
|
||||
requireToolPresent(t, names, "list_memos")
|
||||
requireToolPresent(t, names, "list_tags")
|
||||
requireToolAbsent(t, names, "get_memo")
|
||||
requireToolAbsent(t, names, "list_attachments")
|
||||
})
|
||||
|
||||
t.Run("path toolsets and readonly compose", func(t *testing.T) {
|
||||
sessionID := initializeMCPHTTP(t, e, "/mcp/x/memos,tags/readonly", nil)
|
||||
response := callMCPHTTP(t, e, "/mcp/x/memos,tags/readonly", sessionID, nil, "tools/list", map[string]any{})
|
||||
names := toolNamesFromListResponse(t, response)
|
||||
requireToolPresent(t, names, "list_memos")
|
||||
requireToolPresent(t, names, "list_tags")
|
||||
requireToolAbsent(t, names, "create_memo")
|
||||
requireToolAbsent(t, names, "list_attachments")
|
||||
})
|
||||
|
||||
t.Run("unknown toolset returns empty tool list", func(t *testing.T) {
|
||||
sessionID := initializeMCPHTTP(t, e, "/mcp/x/notreal", nil)
|
||||
response := callMCPHTTP(t, e, "/mcp/x/notreal", sessionID, nil, "tools/list", map[string]any{})
|
||||
names := toolNamesFromListResponse(t, response)
|
||||
require.Empty(t, names)
|
||||
})
|
||||
}
|
||||
|
||||
func TestMCPMemoAndReactionMutationsEmitSSEEvents(t *testing.T) {
|
||||
ts := newTestMCPService(t)
|
||||
user := ts.createUser(t, "author")
|
||||
ctx := withUser(context.Background(), user.ID)
|
||||
client := ts.service.apiV1Service.SSEHub.Subscribe(user.ID, store.RoleUser)
|
||||
defer ts.service.apiV1Service.SSEHub.Unsubscribe(client)
|
||||
|
||||
createResult, err := ts.service.handleCreateMemo(ctx, toolRequest("create_memo", map[string]any{
|
||||
"content": "created from MCP",
|
||||
"visibility": "PRIVATE",
|
||||
}))
|
||||
require.NoError(t, err)
|
||||
require.False(t, createResult.IsError)
|
||||
createEvent := nextSSEEvent(t, client)
|
||||
require.Equal(t, apiv1service.SSEEventMemoCreated, createEvent.Type)
|
||||
|
||||
var created memoJSON
|
||||
require.NoError(t, json.Unmarshal([]byte(firstText(t, createResult)), &created))
|
||||
|
||||
updateResult, err := ts.service.handleUpdateMemo(ctx, toolRequest("update_memo", map[string]any{
|
||||
"name": created.Name,
|
||||
"content": "updated from MCP",
|
||||
}))
|
||||
require.NoError(t, err)
|
||||
require.False(t, updateResult.IsError)
|
||||
updateEvent := nextSSEEvent(t, client)
|
||||
require.Equal(t, apiv1service.SSEEventMemoUpdated, updateEvent.Type)
|
||||
|
||||
commentResult, err := ts.service.handleCreateMemoComment(ctx, toolRequest("create_memo_comment", map[string]any{
|
||||
"name": created.Name,
|
||||
"content": "comment from MCP",
|
||||
}))
|
||||
require.NoError(t, err)
|
||||
require.False(t, commentResult.IsError)
|
||||
commentEvent := nextSSEEvent(t, client)
|
||||
require.Equal(t, apiv1service.SSEEventMemoCommentCreated, commentEvent.Type)
|
||||
require.Equal(t, created.Name, commentEvent.Name)
|
||||
|
||||
upsertReactionResult, err := ts.service.handleUpsertReaction(ctx, toolRequest("upsert_reaction", map[string]any{
|
||||
"name": created.Name,
|
||||
"reaction_type": "👍",
|
||||
}))
|
||||
require.NoError(t, err)
|
||||
require.False(t, upsertReactionResult.IsError)
|
||||
reactionEvent := nextSSEEvent(t, client)
|
||||
require.Equal(t, apiv1service.SSEEventReactionUpserted, reactionEvent.Type)
|
||||
|
||||
var reaction reactionJSON
|
||||
require.NoError(t, json.Unmarshal([]byte(firstText(t, upsertReactionResult)), &reaction))
|
||||
deleteReactionResult, err := ts.service.handleDeleteReaction(ctx, toolRequest("delete_reaction", map[string]any{
|
||||
"id": float64(reaction.ID),
|
||||
}))
|
||||
require.NoError(t, err)
|
||||
require.False(t, deleteReactionResult.IsError)
|
||||
deleteReactionEvent := nextSSEEvent(t, client)
|
||||
require.Equal(t, apiv1service.SSEEventReactionDeleted, deleteReactionEvent.Type)
|
||||
|
||||
deleteResult, err := ts.service.handleDeleteMemo(ctx, toolRequest("delete_memo", map[string]any{
|
||||
"name": created.Name,
|
||||
}))
|
||||
require.NoError(t, err)
|
||||
require.False(t, deleteResult.IsError)
|
||||
deleteEvent := nextSSEEvent(t, client)
|
||||
require.Equal(t, apiv1service.SSEEventMemoDeleted, deleteEvent.Type)
|
||||
}
|
||||
|
||||
func TestMCPRelationAndAttachmentMutationsEmitMemoUpdated(t *testing.T) {
|
||||
ts := newTestMCPService(t)
|
||||
user := ts.createUser(t, "owner")
|
||||
ctx := withUser(context.Background(), user.ID)
|
||||
source := ts.createMemo(t, user.ID, store.Private, "source")
|
||||
target := ts.createMemo(t, user.ID, store.Private, "target")
|
||||
attachment := ts.createAttachment(t, user.ID, nil)
|
||||
client := ts.service.apiV1Service.SSEHub.Subscribe(user.ID, store.RoleUser)
|
||||
defer ts.service.apiV1Service.SSEHub.Unsubscribe(client)
|
||||
|
||||
relationResult, err := ts.service.handleCreateMemoRelation(ctx, toolRequest("create_memo_relation", map[string]any{
|
||||
"name": "memos/" + source.UID,
|
||||
"related_memo": "memos/" + target.UID,
|
||||
}))
|
||||
require.NoError(t, err)
|
||||
require.False(t, relationResult.IsError)
|
||||
relationEvent := nextSSEEvent(t, client)
|
||||
require.Equal(t, apiv1service.SSEEventMemoUpdated, relationEvent.Type)
|
||||
require.Equal(t, "memos/"+source.UID, relationEvent.Name)
|
||||
|
||||
duplicateRelationResult, err := ts.service.handleCreateMemoRelation(ctx, toolRequest("create_memo_relation", map[string]any{
|
||||
"name": "memos/" + source.UID,
|
||||
"related_memo": "memos/" + target.UID,
|
||||
}))
|
||||
require.NoError(t, err)
|
||||
require.False(t, duplicateRelationResult.IsError)
|
||||
requireNoSSEEvent(t, client)
|
||||
|
||||
selfRelationResult, err := ts.service.handleCreateMemoRelation(ctx, toolRequest("create_memo_relation", map[string]any{
|
||||
"name": "memos/" + source.UID,
|
||||
"related_memo": "memos/" + source.UID,
|
||||
}))
|
||||
require.NoError(t, err)
|
||||
require.True(t, selfRelationResult.IsError)
|
||||
require.Contains(t, firstText(t, selfRelationResult), "itself")
|
||||
|
||||
linkResult, err := ts.service.handleLinkAttachmentToMemo(ctx, toolRequest("link_attachment_to_memo", map[string]any{
|
||||
"name": "attachments/" + attachment.UID,
|
||||
"memo": "memos/" + source.UID,
|
||||
}))
|
||||
require.NoError(t, err)
|
||||
require.False(t, linkResult.IsError)
|
||||
attachmentEvent := nextSSEEvent(t, client)
|
||||
require.Equal(t, apiv1service.SSEEventMemoUpdated, attachmentEvent.Type)
|
||||
require.Equal(t, "memos/"+source.UID, attachmentEvent.Name)
|
||||
|
||||
relinkResult, err := ts.service.handleLinkAttachmentToMemo(ctx, toolRequest("link_attachment_to_memo", map[string]any{
|
||||
"name": "attachments/" + attachment.UID,
|
||||
"memo": "memos/" + source.UID,
|
||||
}))
|
||||
require.NoError(t, err)
|
||||
require.False(t, relinkResult.IsError)
|
||||
requireNoSSEEvent(t, client)
|
||||
}
|
||||
@@ -0,0 +1,192 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/mark3labs/mcp-go/mcp"
|
||||
mcpserver "github.com/mark3labs/mcp-go/server"
|
||||
)
|
||||
|
||||
func (s *MCPService) registerPrompts(mcpSrv *mcpserver.MCPServer) {
|
||||
// capture — turns free-form user input into a structured create_memo call.
|
||||
mcpSrv.AddPrompt(
|
||||
mcp.NewPrompt("capture",
|
||||
mcp.WithPromptDescription("Capture a thought, idea, or note as a new memo. "+
|
||||
"Use this prompt when the user wants to quickly save something. "+
|
||||
"The assistant will call create_memo with the provided content."),
|
||||
mcp.WithArgument("content",
|
||||
mcp.ArgumentDescription("The text to save as a memo"),
|
||||
mcp.RequiredArgument(),
|
||||
),
|
||||
mcp.WithArgument("tags",
|
||||
mcp.ArgumentDescription("Comma-separated tags to apply, e.g. \"work,project\""),
|
||||
),
|
||||
mcp.WithArgument("visibility",
|
||||
mcp.ArgumentDescription("Memo visibility: PRIVATE (default), PROTECTED, or PUBLIC"),
|
||||
),
|
||||
),
|
||||
s.handleCapturePrompt,
|
||||
)
|
||||
|
||||
// review — surfaces existing memos on a topic for summarisation.
|
||||
mcpSrv.AddPrompt(
|
||||
mcp.NewPrompt("review",
|
||||
mcp.WithPromptDescription("Search and review memos on a given topic. "+
|
||||
"The assistant will call search_memos and summarise the results, "+
|
||||
"including memo resource URIs for easy reference."),
|
||||
mcp.WithArgument("topic",
|
||||
mcp.ArgumentDescription("Topic or keyword to search for"),
|
||||
mcp.RequiredArgument(),
|
||||
),
|
||||
),
|
||||
s.handleReviewPrompt,
|
||||
)
|
||||
|
||||
// daily_digest — summarise recent activity.
|
||||
mcpSrv.AddPrompt(
|
||||
mcp.NewPrompt("daily_digest",
|
||||
mcp.WithPromptDescription("Get a summary of recent memo activity. "+
|
||||
"The assistant will list recent memos, group them by tags, and highlight "+
|
||||
"any incomplete tasks or pinned items."),
|
||||
mcp.WithArgument("days",
|
||||
mcp.ArgumentDescription("Number of days to look back (default: 1)"),
|
||||
),
|
||||
),
|
||||
s.handleDailyDigestPrompt,
|
||||
)
|
||||
|
||||
// organize — suggest tags and relations for untagged memos.
|
||||
mcpSrv.AddPrompt(
|
||||
mcp.NewPrompt("organize",
|
||||
mcp.WithPromptDescription("Analyze untagged or loosely organized memos and suggest "+
|
||||
"tags, relations, and groupings to improve discoverability."),
|
||||
mcp.WithArgument("scope",
|
||||
mcp.ArgumentDescription("Scope of analysis: \"untagged\" (default) for memos without tags, \"all\" for all recent memos"),
|
||||
),
|
||||
),
|
||||
s.handleOrganizePrompt,
|
||||
)
|
||||
}
|
||||
|
||||
func (*MCPService) handleCapturePrompt(_ context.Context, req mcp.GetPromptRequest) (*mcp.GetPromptResult, error) {
|
||||
content := req.Params.Arguments["content"]
|
||||
if content == "" {
|
||||
return nil, errors.New("content argument is required")
|
||||
}
|
||||
|
||||
tags := req.Params.Arguments["tags"]
|
||||
visibility := req.Params.Arguments["visibility"]
|
||||
if visibility == "" {
|
||||
visibility = "PRIVATE"
|
||||
}
|
||||
|
||||
var sb strings.Builder
|
||||
sb.WriteString("Save the following as a new memo using the create_memo tool.\n\n")
|
||||
fmt.Fprintf(&sb, "Visibility: %s\n\n", visibility)
|
||||
sb.WriteString("Content:\n")
|
||||
sb.WriteString(content)
|
||||
if tags != "" {
|
||||
fmt.Fprintf(&sb, "\n\nAppend these tags inline using #tag syntax: %s", tags)
|
||||
}
|
||||
sb.WriteString("\n\nAfter creating the memo, confirm by showing the memo resource name (e.g. memo://memos/<uid>) so it can be referenced later.")
|
||||
|
||||
return &mcp.GetPromptResult{
|
||||
Description: "Capture a memo",
|
||||
Messages: []mcp.PromptMessage{
|
||||
mcp.NewPromptMessage(mcp.RoleUser, mcp.NewTextContent(sb.String())),
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (*MCPService) handleReviewPrompt(_ context.Context, req mcp.GetPromptRequest) (*mcp.GetPromptResult, error) {
|
||||
topic := req.Params.Arguments["topic"]
|
||||
if topic == "" {
|
||||
return nil, errors.New("topic argument is required")
|
||||
}
|
||||
|
||||
instruction := fmt.Sprintf(
|
||||
`Use the search_memos tool to find memos about %q, then:
|
||||
|
||||
1. Group results by theme or tag
|
||||
2. For each memo, include its resource reference (memo://memos/<uid>) so the user can access it directly
|
||||
3. Provide a concise summary of what has been written on this topic
|
||||
4. Highlight any memos with incomplete tasks (has_incomplete_tasks)
|
||||
5. Note the most recent update times to show currency of the information`,
|
||||
topic,
|
||||
)
|
||||
|
||||
return &mcp.GetPromptResult{
|
||||
Description: fmt.Sprintf("Review memos about %q", topic),
|
||||
Messages: []mcp.PromptMessage{
|
||||
mcp.NewPromptMessage(mcp.RoleUser, mcp.NewTextContent(instruction)),
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (*MCPService) handleDailyDigestPrompt(_ context.Context, req mcp.GetPromptRequest) (*mcp.GetPromptResult, error) {
|
||||
days := req.Params.Arguments["days"]
|
||||
if days == "" {
|
||||
days = "1"
|
||||
}
|
||||
|
||||
instruction := fmt.Sprintf(
|
||||
`Generate a daily digest of memo activity from the last %s day(s):
|
||||
|
||||
1. Use list_memos to fetch recent memos (order by update time, check multiple pages if needed)
|
||||
2. Use list_tags to get the current tag landscape
|
||||
3. Group memos by tags and summarize each group
|
||||
4. Highlight:
|
||||
- Pinned memos (important items)
|
||||
- Memos with incomplete tasks (action items)
|
||||
- New memos created vs. memos updated
|
||||
5. Include memo resource references (memo://memos/<uid>) for each item
|
||||
6. End with a brief "action items" section listing incomplete tasks across all memos`, days,
|
||||
)
|
||||
|
||||
return &mcp.GetPromptResult{
|
||||
Description: "Daily memo digest",
|
||||
Messages: []mcp.PromptMessage{
|
||||
mcp.NewPromptMessage(mcp.RoleUser, mcp.NewTextContent(instruction)),
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (*MCPService) handleOrganizePrompt(_ context.Context, req mcp.GetPromptRequest) (*mcp.GetPromptResult, error) {
|
||||
scope := req.Params.Arguments["scope"]
|
||||
if scope == "" {
|
||||
scope = "untagged"
|
||||
}
|
||||
|
||||
var filter string
|
||||
if scope == "untagged" {
|
||||
filter = `Focus on memos that have no tags. Use list_memos and identify those with empty tag arrays.`
|
||||
} else {
|
||||
filter = `Analyze all recent memos regardless of tagging status.`
|
||||
}
|
||||
|
||||
instruction := fmt.Sprintf(
|
||||
`Analyze memos and suggest organizational improvements:
|
||||
|
||||
1. %s
|
||||
2. Use list_tags to understand the existing tag taxonomy
|
||||
3. For each unorganized memo, suggest:
|
||||
- Appropriate tags from the existing taxonomy, or new tags if needed
|
||||
- Potential relations (references) to other memos on similar topics
|
||||
4. Present suggestions as a structured list:
|
||||
- Memo: memo://memos/<uid> (first line of content as preview)
|
||||
- Suggested tags: #tag1, #tag2
|
||||
- Related to: memo://memos/<other-uid> (brief reason)
|
||||
5. After presenting suggestions, ask the user which changes to apply
|
||||
6. Apply approved changes using update_memo (for tags in content) and create_memo_relation (for references)`, filter,
|
||||
)
|
||||
|
||||
return &mcp.GetPromptResult{
|
||||
Description: fmt.Sprintf("Organize memos (scope: %s)", scope),
|
||||
Messages: []mcp.PromptMessage{
|
||||
mcp.NewPromptMessage(mcp.RoleUser, mcp.NewTextContent(instruction)),
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
@@ -0,0 +1,88 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/mark3labs/mcp-go/mcp"
|
||||
mcpserver "github.com/mark3labs/mcp-go/server"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/usememos/memos/server/auth"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
// Memo resource URI scheme: memo://memos/{uid}
|
||||
// Clients can read any memo they have access to by URI without calling a tool.
|
||||
|
||||
func (s *MCPService) registerMemoResources(mcpSrv *mcpserver.MCPServer) {
|
||||
mcpSrv.AddResourceTemplate(
|
||||
mcp.NewResourceTemplate(
|
||||
"memo://memos/{uid}",
|
||||
"Memo",
|
||||
mcp.WithTemplateDescription("A single Memos note identified by its UID. Returns the memo content as Markdown with a YAML frontmatter header containing metadata."),
|
||||
mcp.WithTemplateMIMEType("text/markdown"),
|
||||
),
|
||||
s.handleReadMemoResource,
|
||||
)
|
||||
}
|
||||
|
||||
func (s *MCPService) handleReadMemoResource(ctx context.Context, req mcp.ReadResourceRequest) ([]mcp.ResourceContents, error) {
|
||||
userID := auth.GetUserID(ctx)
|
||||
|
||||
// URI format: memo://memos/{uid}
|
||||
uid := strings.TrimPrefix(req.Params.URI, "memo://memos/")
|
||||
if uid == req.Params.URI || uid == "" {
|
||||
return nil, errors.Errorf("invalid memo URI %q: expected memo://memos/<uid>", req.Params.URI)
|
||||
}
|
||||
|
||||
memo, err := s.store.GetMemo(ctx, &store.FindMemo{UID: &uid})
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to get memo")
|
||||
}
|
||||
if memo == nil {
|
||||
return nil, errors.Errorf("memo not found: %s", uid)
|
||||
}
|
||||
if err := checkMemoAccess(memo, userID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
j, err := storeMemoToJSONWithStore(ctx, s.store, memo)
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to resolve memo creator")
|
||||
}
|
||||
text := formatMemoMarkdown(j)
|
||||
|
||||
return []mcp.ResourceContents{
|
||||
mcp.TextResourceContents{
|
||||
URI: req.Params.URI,
|
||||
MIMEType: "text/markdown",
|
||||
Text: text,
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
// formatMemoMarkdown renders a memo as Markdown with a YAML frontmatter header.
|
||||
func formatMemoMarkdown(j memoJSON) string {
|
||||
var sb strings.Builder
|
||||
|
||||
sb.WriteString("---\n")
|
||||
fmt.Fprintf(&sb, "name: %s\n", j.Name)
|
||||
fmt.Fprintf(&sb, "creator: %s\n", j.Creator)
|
||||
fmt.Fprintf(&sb, "visibility: %s\n", j.Visibility)
|
||||
fmt.Fprintf(&sb, "state: %s\n", j.State)
|
||||
fmt.Fprintf(&sb, "pinned: %v\n", j.Pinned)
|
||||
if len(j.Tags) > 0 {
|
||||
fmt.Fprintf(&sb, "tags: [%s]\n", strings.Join(j.Tags, ", "))
|
||||
}
|
||||
fmt.Fprintf(&sb, "create_time: %d\n", j.CreateTime)
|
||||
fmt.Fprintf(&sb, "update_time: %d\n", j.UpdateTime)
|
||||
if j.Parent != "" {
|
||||
fmt.Fprintf(&sb, "parent: %s\n", j.Parent)
|
||||
}
|
||||
sb.WriteString("---\n\n")
|
||||
sb.WriteString(j.Content)
|
||||
|
||||
return sb.String()
|
||||
}
|
||||
@@ -0,0 +1,102 @@
|
||||
package mcp
|
||||
|
||||
import "github.com/mark3labs/mcp-go/mcp"
|
||||
|
||||
var mcpToolsByToolset = map[string]map[string]struct{}{
|
||||
"memos": stringSet(
|
||||
"list_memos",
|
||||
"get_memo",
|
||||
"create_memo",
|
||||
"update_memo",
|
||||
"delete_memo",
|
||||
"search_memos",
|
||||
"list_memo_comments",
|
||||
"create_memo_comment",
|
||||
),
|
||||
"tags": stringSet(
|
||||
"list_tags",
|
||||
),
|
||||
"attachments": stringSet(
|
||||
"list_attachments",
|
||||
"get_attachment",
|
||||
"delete_attachment",
|
||||
"link_attachment_to_memo",
|
||||
),
|
||||
"relations": stringSet(
|
||||
"list_memo_relations",
|
||||
"create_memo_relation",
|
||||
"delete_memo_relation",
|
||||
),
|
||||
"reactions": stringSet(
|
||||
"list_reactions",
|
||||
"upsert_reaction",
|
||||
"delete_reaction",
|
||||
),
|
||||
}
|
||||
|
||||
var allMCPToolNames = func() map[string]struct{} {
|
||||
names := map[string]struct{}{}
|
||||
for _, tools := range mcpToolsByToolset {
|
||||
for name := range tools {
|
||||
names[name] = struct{}{}
|
||||
}
|
||||
}
|
||||
return names
|
||||
}()
|
||||
|
||||
var mcpMutationTools = stringSet(
|
||||
"create_memo",
|
||||
"update_memo",
|
||||
"delete_memo",
|
||||
"create_memo_comment",
|
||||
"delete_attachment",
|
||||
"link_attachment_to_memo",
|
||||
"create_memo_relation",
|
||||
"delete_memo_relation",
|
||||
"upsert_reaction",
|
||||
"delete_reaction",
|
||||
)
|
||||
|
||||
type deletedJSON struct {
|
||||
Deleted bool `json:"deleted"`
|
||||
}
|
||||
|
||||
func stringSet(values ...string) map[string]struct{} {
|
||||
result := make(map[string]struct{}, len(values))
|
||||
for _, value := range values {
|
||||
result[value] = struct{}{}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func readOnlyToolOptions(title string, description string, opts ...mcp.ToolOption) []mcp.ToolOption {
|
||||
return annotatedToolOptions(title, description, true, false, true, false, opts...)
|
||||
}
|
||||
|
||||
func createToolOptions(title string, description string, idempotent bool, opts ...mcp.ToolOption) []mcp.ToolOption {
|
||||
return annotatedToolOptions(title, description, false, false, idempotent, false, opts...)
|
||||
}
|
||||
|
||||
func updateToolOptions(title string, description string, opts ...mcp.ToolOption) []mcp.ToolOption {
|
||||
return annotatedToolOptions(title, description, false, true, false, false, opts...)
|
||||
}
|
||||
|
||||
func annotatedToolOptions(title string, description string, readOnly bool, destructive bool, idempotent bool, openWorld bool, opts ...mcp.ToolOption) []mcp.ToolOption {
|
||||
base := []mcp.ToolOption{
|
||||
mcp.WithTitleAnnotation(title),
|
||||
mcp.WithDescription(description),
|
||||
mcp.WithReadOnlyHintAnnotation(readOnly),
|
||||
mcp.WithDestructiveHintAnnotation(destructive),
|
||||
mcp.WithIdempotentHintAnnotation(idempotent),
|
||||
mcp.WithOpenWorldHintAnnotation(openWorld),
|
||||
}
|
||||
return append(base, opts...)
|
||||
}
|
||||
|
||||
func newToolResultJSON(v any) (*mcp.CallToolResult, error) {
|
||||
return mcp.NewToolResultJSON(v)
|
||||
}
|
||||
|
||||
func newDeletedToolResult() (*mcp.CallToolResult, error) {
|
||||
return newToolResultJSON(deletedJSON{Deleted: true})
|
||||
}
|
||||
@@ -0,0 +1,327 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/mark3labs/mcp-go/mcp"
|
||||
mcpserver "github.com/mark3labs/mcp-go/server"
|
||||
"github.com/pkg/errors"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
storepb "github.com/usememos/memos/proto/gen/store"
|
||||
"github.com/usememos/memos/server/auth"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
type attachmentJSON struct {
|
||||
Name string `json:"name"`
|
||||
Creator string `json:"creator"`
|
||||
CreateTime int64 `json:"create_time"`
|
||||
Filename string `json:"filename"`
|
||||
Type string `json:"type"`
|
||||
Size int64 `json:"size"`
|
||||
StorageType string `json:"storage_type"`
|
||||
ExternalLink string `json:"external_link,omitempty"`
|
||||
Memo string `json:"memo,omitempty"`
|
||||
}
|
||||
|
||||
type attachmentListJSON struct {
|
||||
Attachments []attachmentJSON `json:"attachments"`
|
||||
HasMore bool `json:"has_more"`
|
||||
}
|
||||
|
||||
func storeAttachmentToJSON(ctx context.Context, stores *store.Store, a *store.Attachment) (attachmentJSON, error) {
|
||||
creator, err := lookupUsername(ctx, stores, a.CreatorID)
|
||||
if err != nil {
|
||||
return attachmentJSON{}, errors.Wrap(err, "lookup attachment creator username")
|
||||
}
|
||||
j := attachmentJSON{
|
||||
Name: "attachments/" + a.UID,
|
||||
Creator: creator,
|
||||
CreateTime: a.CreatedTs,
|
||||
Filename: a.Filename,
|
||||
Type: a.Type,
|
||||
Size: a.Size,
|
||||
}
|
||||
switch a.StorageType {
|
||||
case storepb.AttachmentStorageType_LOCAL:
|
||||
j.StorageType = "LOCAL"
|
||||
case storepb.AttachmentStorageType_S3:
|
||||
j.StorageType = "S3"
|
||||
j.ExternalLink = a.Reference
|
||||
case storepb.AttachmentStorageType_EXTERNAL:
|
||||
j.StorageType = "EXTERNAL"
|
||||
j.ExternalLink = a.Reference
|
||||
default:
|
||||
j.StorageType = "DATABASE"
|
||||
}
|
||||
if a.MemoUID != nil && *a.MemoUID != "" {
|
||||
j.Memo = "memos/" + *a.MemoUID
|
||||
}
|
||||
return j, nil
|
||||
}
|
||||
|
||||
func storeAttachmentToJSONWithUsernames(a *store.Attachment, usernamesByID map[int32]string) (attachmentJSON, error) {
|
||||
creator, err := lookupUsernameFromCache(usernamesByID, a.CreatorID)
|
||||
if err != nil {
|
||||
return attachmentJSON{}, errors.Wrap(err, "lookup attachment creator username from cache")
|
||||
}
|
||||
j := attachmentJSON{
|
||||
Name: "attachments/" + a.UID,
|
||||
Creator: creator,
|
||||
CreateTime: a.CreatedTs,
|
||||
Filename: a.Filename,
|
||||
Type: a.Type,
|
||||
Size: a.Size,
|
||||
}
|
||||
switch a.StorageType {
|
||||
case storepb.AttachmentStorageType_LOCAL:
|
||||
j.StorageType = "LOCAL"
|
||||
case storepb.AttachmentStorageType_S3:
|
||||
j.StorageType = "S3"
|
||||
j.ExternalLink = a.Reference
|
||||
case storepb.AttachmentStorageType_EXTERNAL:
|
||||
j.StorageType = "EXTERNAL"
|
||||
j.ExternalLink = a.Reference
|
||||
default:
|
||||
j.StorageType = "DATABASE"
|
||||
}
|
||||
if a.MemoUID != nil && *a.MemoUID != "" {
|
||||
j.Memo = "memos/" + *a.MemoUID
|
||||
}
|
||||
return j, nil
|
||||
}
|
||||
|
||||
func parseAttachmentUID(name string) (string, error) {
|
||||
uid, ok := strings.CutPrefix(name, "attachments/")
|
||||
if !ok || uid == "" {
|
||||
return "", errors.Errorf(`attachment name must be "attachments/<uid>", got %q`, name)
|
||||
}
|
||||
return uid, nil
|
||||
}
|
||||
|
||||
func (s *MCPService) registerAttachmentTools(mcpSrv *mcpserver.MCPServer) {
|
||||
mcpSrv.AddTool(mcp.NewTool("list_attachments",
|
||||
readOnlyToolOptions("List attachments", "List attachments owned by the authenticated user. Supports pagination and optional filtering by linked memo.",
|
||||
mcp.WithNumber("page_size", mcp.Description("Maximum attachments to return (1–100, default 20)")),
|
||||
mcp.WithNumber("page", mcp.Description("Zero-based page index (default 0)")),
|
||||
mcp.WithString("memo", mcp.Description(`Filter by linked memo resource name, e.g. "memos/abc123"`)),
|
||||
mcp.WithOutputSchema[attachmentListJSON](),
|
||||
)...,
|
||||
), s.handleListAttachments)
|
||||
|
||||
mcpSrv.AddTool(mcp.NewTool("get_attachment",
|
||||
readOnlyToolOptions("Get attachment", "Get a single attachment's metadata by resource name. Requires authentication.",
|
||||
mcp.WithString("name", mcp.Required(), mcp.Description(`Attachment resource name, e.g. "attachments/abc123"`)),
|
||||
mcp.WithOutputSchema[attachmentJSON](),
|
||||
)...,
|
||||
), s.handleGetAttachment)
|
||||
|
||||
mcpSrv.AddTool(mcp.NewTool("delete_attachment",
|
||||
updateToolOptions("Delete attachment", "Permanently delete an attachment and its stored file. Requires authentication and ownership.",
|
||||
mcp.WithString("name", mcp.Required(), mcp.Description(`Attachment resource name, e.g. "attachments/abc123"`)),
|
||||
mcp.WithOutputSchema[deletedJSON](),
|
||||
)...,
|
||||
), s.handleDeleteAttachment)
|
||||
|
||||
mcpSrv.AddTool(mcp.NewTool("link_attachment_to_memo",
|
||||
createToolOptions("Link attachment to memo", "Link an existing attachment to a memo. Requires authentication and ownership of the attachment.", true,
|
||||
mcp.WithString("name", mcp.Required(), mcp.Description(`Attachment resource name, e.g. "attachments/abc123"`)),
|
||||
mcp.WithString("memo", mcp.Required(), mcp.Description(`Memo resource name, e.g. "memos/abc123"`)),
|
||||
mcp.WithOutputSchema[attachmentJSON](),
|
||||
)...,
|
||||
), s.handleLinkAttachmentToMemo)
|
||||
}
|
||||
|
||||
func (s *MCPService) handleListAttachments(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
userID, err := extractUserID(ctx)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
pageSize := req.GetInt("page_size", 20)
|
||||
if pageSize <= 0 {
|
||||
pageSize = 20
|
||||
}
|
||||
if pageSize > 100 {
|
||||
pageSize = 100
|
||||
}
|
||||
page := req.GetInt("page", 0)
|
||||
if page < 0 {
|
||||
page = 0
|
||||
}
|
||||
|
||||
limit := pageSize + 1
|
||||
offset := page * pageSize
|
||||
find := &store.FindAttachment{
|
||||
CreatorID: &userID,
|
||||
Limit: &limit,
|
||||
Offset: &offset,
|
||||
}
|
||||
|
||||
if memoName := req.GetString("memo", ""); memoName != "" {
|
||||
memoUID, err := parseMemoUID(memoName)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
memo, err := s.store.GetMemo(ctx, &store.FindMemo{UID: &memoUID})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to find memo: %v", err)), nil
|
||||
}
|
||||
if memo == nil {
|
||||
return mcp.NewToolResultError("memo not found"), nil
|
||||
}
|
||||
find.MemoID = &memo.ID
|
||||
}
|
||||
|
||||
attachments, err := s.store.ListAttachments(ctx, find)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to list attachments: %v", err)), nil
|
||||
}
|
||||
|
||||
hasMore := len(attachments) > pageSize
|
||||
if hasMore {
|
||||
attachments = attachments[:pageSize]
|
||||
}
|
||||
creatorIDs := make([]int32, 0, len(attachments))
|
||||
for _, attachment := range attachments {
|
||||
creatorIDs = append(creatorIDs, attachment.CreatorID)
|
||||
}
|
||||
usernamesByID, err := preloadUsernames(ctx, s.store, creatorIDs)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to preload attachment creators: %v", err)), nil
|
||||
}
|
||||
|
||||
results := make([]attachmentJSON, len(attachments))
|
||||
for i, a := range attachments {
|
||||
result, err := storeAttachmentToJSONWithUsernames(a, usernamesByID)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to resolve attachment creator: %v", err)), nil
|
||||
}
|
||||
results[i] = result
|
||||
}
|
||||
|
||||
return newToolResultJSON(attachmentListJSON{Attachments: results, HasMore: hasMore})
|
||||
}
|
||||
|
||||
func (s *MCPService) handleGetAttachment(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
userID := auth.GetUserID(ctx)
|
||||
|
||||
uid, err := parseAttachmentUID(req.GetString("name", ""))
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
attachment, err := s.store.GetAttachment(ctx, &store.FindAttachment{UID: &uid})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to get attachment: %v", err)), nil
|
||||
}
|
||||
if attachment == nil {
|
||||
return mcp.NewToolResultError("attachment not found"), nil
|
||||
}
|
||||
|
||||
if err := s.checkAttachmentAccess(ctx, attachment, userID); err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
result, err := storeAttachmentToJSON(ctx, s.store, attachment)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to resolve attachment creator: %v", err)), nil
|
||||
}
|
||||
return newToolResultJSON(result)
|
||||
}
|
||||
|
||||
func (s *MCPService) handleDeleteAttachment(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
if _, err := extractUserID(ctx); err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
uid, err := parseAttachmentUID(req.GetString("name", ""))
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
if _, err := s.apiV1Service.DeleteAttachment(ctx, &v1pb.DeleteAttachmentRequest{Name: "attachments/" + uid}); err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to delete attachment: %v", err)), nil
|
||||
}
|
||||
return newDeletedToolResult()
|
||||
}
|
||||
|
||||
func (s *MCPService) handleLinkAttachmentToMemo(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
userID, err := extractUserID(ctx)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
uid, err := parseAttachmentUID(req.GetString("name", ""))
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
attachment, err := s.store.GetAttachment(ctx, &store.FindAttachment{UID: &uid})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to get attachment: %v", err)), nil
|
||||
}
|
||||
if attachment == nil {
|
||||
return mcp.NewToolResultError("attachment not found"), nil
|
||||
}
|
||||
if attachment.CreatorID != userID {
|
||||
return mcp.NewToolResultError("permission denied"), nil
|
||||
}
|
||||
|
||||
memoUID, err := parseMemoUID(req.GetString("memo", ""))
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
memo, err := s.store.GetMemo(ctx, &store.FindMemo{UID: &memoUID})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to get memo: %v", err)), nil
|
||||
}
|
||||
if memo == nil {
|
||||
return mcp.NewToolResultError("memo not found"), nil
|
||||
}
|
||||
if err := checkMemoOwnership(memo, userID); err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
currentAttachments, err := s.store.ListAttachments(ctx, &store.FindAttachment{MemoID: &memo.ID})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to list memo attachments: %v", err)), nil
|
||||
}
|
||||
requestAttachments := make([]*v1pb.Attachment, 0, len(currentAttachments)+1)
|
||||
var currentTarget *store.Attachment
|
||||
for _, current := range currentAttachments {
|
||||
requestAttachments = append(requestAttachments, &v1pb.Attachment{Name: "attachments/" + current.UID})
|
||||
if current.ID == attachment.ID {
|
||||
currentTarget = current
|
||||
}
|
||||
}
|
||||
if currentTarget != nil {
|
||||
result, err := storeAttachmentToJSON(ctx, s.store, currentTarget)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to resolve attachment creator: %v", err)), nil
|
||||
}
|
||||
return newToolResultJSON(result)
|
||||
}
|
||||
requestAttachments = append(requestAttachments, &v1pb.Attachment{Name: "attachments/" + uid})
|
||||
|
||||
if _, err := s.apiV1Service.SetMemoAttachments(ctx, &v1pb.SetMemoAttachmentsRequest{
|
||||
Name: "memos/" + memoUID,
|
||||
Attachments: requestAttachments,
|
||||
}); err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to link attachment: %v", err)), nil
|
||||
}
|
||||
|
||||
// Re-fetch to get updated memo UID.
|
||||
updated, err := s.store.GetAttachment(ctx, &store.FindAttachment{ID: &attachment.ID})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to fetch updated attachment: %v", err)), nil
|
||||
}
|
||||
result, err := storeAttachmentToJSON(ctx, s.store, updated)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to resolve attachment creator: %v", err)), nil
|
||||
}
|
||||
return newToolResultJSON(result)
|
||||
}
|
||||
@@ -0,0 +1,607 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/mark3labs/mcp-go/mcp"
|
||||
mcpserver "github.com/mark3labs/mcp-go/server"
|
||||
"github.com/pkg/errors"
|
||||
"google.golang.org/protobuf/types/known/fieldmaskpb"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
"github.com/usememos/memos/server/auth"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
// propertyJSON is the serialisable form of MemoPayload.Property.
|
||||
type propertyJSON struct {
|
||||
HasLink bool `json:"has_link"`
|
||||
HasTaskList bool `json:"has_task_list"`
|
||||
HasCode bool `json:"has_code"`
|
||||
HasIncompleteTasks bool `json:"has_incomplete_tasks"`
|
||||
}
|
||||
|
||||
// memoJSON is the canonical response shape for all MCP memo results.
|
||||
// It serialises correctly with standard encoding/json (no proto marshalling needed).
|
||||
type memoJSON struct {
|
||||
Name string `json:"name"`
|
||||
Creator string `json:"creator"`
|
||||
CreateTime int64 `json:"create_time"`
|
||||
UpdateTime int64 `json:"update_time"`
|
||||
Content string `json:"content,omitempty"`
|
||||
Visibility string `json:"visibility"`
|
||||
Tags []string `json:"tags"`
|
||||
Pinned bool `json:"pinned"`
|
||||
State string `json:"state"`
|
||||
Property *propertyJSON `json:"property,omitempty"`
|
||||
Parent string `json:"parent,omitempty"`
|
||||
}
|
||||
|
||||
type memoListJSON struct {
|
||||
Memos []memoJSON `json:"memos"`
|
||||
HasMore bool `json:"has_more"`
|
||||
}
|
||||
|
||||
func storeMemoToJSON(m *store.Memo) memoJSON {
|
||||
j := memoJSON{
|
||||
Name: "memos/" + m.UID,
|
||||
CreateTime: m.CreatedTs,
|
||||
UpdateTime: m.UpdatedTs,
|
||||
Content: m.Content,
|
||||
Visibility: string(m.Visibility),
|
||||
Pinned: m.Pinned,
|
||||
State: string(m.RowStatus),
|
||||
Tags: []string{},
|
||||
}
|
||||
if m.Payload != nil {
|
||||
if len(m.Payload.Tags) > 0 {
|
||||
j.Tags = m.Payload.Tags
|
||||
}
|
||||
if p := m.Payload.Property; p != nil && (p.HasLink || p.HasTaskList || p.HasCode || p.HasIncompleteTasks) {
|
||||
j.Property = &propertyJSON{
|
||||
HasLink: p.HasLink,
|
||||
HasTaskList: p.HasTaskList,
|
||||
HasCode: p.HasCode,
|
||||
HasIncompleteTasks: p.HasIncompleteTasks,
|
||||
}
|
||||
}
|
||||
}
|
||||
if m.ParentUID != nil {
|
||||
j.Parent = "memos/" + *m.ParentUID
|
||||
}
|
||||
return j
|
||||
}
|
||||
|
||||
func lookupUsername(ctx context.Context, stores *store.Store, userID int32) (string, error) {
|
||||
user, err := stores.GetUser(ctx, &store.FindUser{ID: &userID})
|
||||
if err != nil {
|
||||
return "", errors.Wrapf(err, "failed to get creator user %d", userID)
|
||||
}
|
||||
if user == nil {
|
||||
return "", errors.Errorf("creator user %d not found", userID)
|
||||
}
|
||||
return "users/" + user.Username, nil
|
||||
}
|
||||
|
||||
func preloadUsernames(ctx context.Context, stores *store.Store, userIDs []int32) (map[int32]string, error) {
|
||||
if len(userIDs) == 0 {
|
||||
return map[int32]string{}, nil
|
||||
}
|
||||
|
||||
uniqueUserIDs := make([]int32, 0, len(userIDs))
|
||||
seenUserIDs := make(map[int32]struct{}, len(userIDs))
|
||||
for _, userID := range userIDs {
|
||||
if _, seen := seenUserIDs[userID]; seen {
|
||||
continue
|
||||
}
|
||||
seenUserIDs[userID] = struct{}{}
|
||||
uniqueUserIDs = append(uniqueUserIDs, userID)
|
||||
}
|
||||
|
||||
users, err := stores.ListUsers(ctx, &store.FindUser{IDList: uniqueUserIDs})
|
||||
if err != nil {
|
||||
return nil, errors.Wrap(err, "failed to list creator users")
|
||||
}
|
||||
|
||||
usernamesByID := make(map[int32]string, len(users))
|
||||
for _, user := range users {
|
||||
usernamesByID[user.ID] = "users/" + user.Username
|
||||
}
|
||||
return usernamesByID, nil
|
||||
}
|
||||
|
||||
func lookupUsernameFromCache(usernamesByID map[int32]string, userID int32) (string, error) {
|
||||
username, ok := usernamesByID[userID]
|
||||
if !ok {
|
||||
return "", errors.Errorf("creator user %d not found", userID)
|
||||
}
|
||||
return username, nil
|
||||
}
|
||||
|
||||
func storeMemoToJSONWithStore(ctx context.Context, stores *store.Store, m *store.Memo) (memoJSON, error) {
|
||||
j := storeMemoToJSON(m)
|
||||
creator, err := lookupUsername(ctx, stores, m.CreatorID)
|
||||
if err != nil {
|
||||
return memoJSON{}, err
|
||||
}
|
||||
j.Creator = creator
|
||||
return j, nil
|
||||
}
|
||||
|
||||
func storeMemoToJSONWithUsernames(m *store.Memo, usernamesByID map[int32]string) (memoJSON, error) {
|
||||
j := storeMemoToJSON(m)
|
||||
creator, err := lookupUsernameFromCache(usernamesByID, m.CreatorID)
|
||||
if err != nil {
|
||||
return memoJSON{}, err
|
||||
}
|
||||
j.Creator = creator
|
||||
return j, nil
|
||||
}
|
||||
|
||||
// parseMemoUID extracts the UID from a "memos/<uid>" resource name.
|
||||
func parseMemoUID(name string) (string, error) {
|
||||
uid, ok := strings.CutPrefix(name, "memos/")
|
||||
if !ok || uid == "" {
|
||||
return "", errors.Errorf(`memo name must be in the format "memos/<uid>", got %q`, name)
|
||||
}
|
||||
return uid, nil
|
||||
}
|
||||
|
||||
// parseVisibility validates a visibility string and returns the store constant.
|
||||
func parseVisibility(s string) (store.Visibility, error) {
|
||||
switch v := store.Visibility(s); v {
|
||||
case store.Public, store.Protected, store.Private:
|
||||
return v, nil
|
||||
default:
|
||||
return "", errors.Errorf("visibility must be PRIVATE, PROTECTED, or PUBLIC; got %q", s)
|
||||
}
|
||||
}
|
||||
|
||||
// parseRowStatus validates a state string and returns the store constant.
|
||||
func parseRowStatus(s string) (store.RowStatus, error) {
|
||||
switch rs := store.RowStatus(s); rs {
|
||||
case store.Normal, store.Archived:
|
||||
return rs, nil
|
||||
default:
|
||||
return "", errors.Errorf("state must be NORMAL or ARCHIVED; got %q", s)
|
||||
}
|
||||
}
|
||||
|
||||
func extractUserID(ctx context.Context) (int32, error) {
|
||||
id := auth.GetUserID(ctx)
|
||||
if id == 0 {
|
||||
return 0, errors.New("unauthenticated: a personal access token is required")
|
||||
}
|
||||
return id, nil
|
||||
}
|
||||
|
||||
func (s *MCPService) registerMemoTools(mcpSrv *mcpserver.MCPServer) {
|
||||
mcpSrv.AddTool(mcp.NewTool("list_memos",
|
||||
readOnlyToolOptions("List memos", "List memos visible to the caller. Authenticated users see their own memos plus public and protected memos; unauthenticated callers see only public memos.",
|
||||
mcp.WithNumber("page_size", mcp.Description("Maximum memos to return (1–100, default 20)")),
|
||||
mcp.WithNumber("page", mcp.Description("Zero-based page index for pagination (default 0)")),
|
||||
mcp.WithString("state",
|
||||
mcp.Enum("NORMAL", "ARCHIVED"),
|
||||
mcp.Description("Filter by state: NORMAL (default) or ARCHIVED"),
|
||||
),
|
||||
mcp.WithBoolean("order_by_pinned", mcp.Description("When true, pinned memos appear first (default false)")),
|
||||
mcp.WithString("filter", mcp.Description(`Optional CEL filter (supported subset of standard CEL syntax), e.g. content.contains("keyword") or tags.exists(t, t == "work")`)),
|
||||
mcp.WithOutputSchema[memoListJSON](),
|
||||
)...,
|
||||
), s.handleListMemos)
|
||||
|
||||
mcpSrv.AddTool(mcp.NewTool("get_memo",
|
||||
readOnlyToolOptions("Get memo", "Get a single memo by resource name. Public memos are accessible without authentication.",
|
||||
mcp.WithString("name", mcp.Required(), mcp.Description(`Memo resource name, e.g. "memos/abc123"`)),
|
||||
mcp.WithOutputSchema[memoJSON](),
|
||||
)...,
|
||||
), s.handleGetMemo)
|
||||
|
||||
mcpSrv.AddTool(mcp.NewTool("create_memo",
|
||||
createToolOptions("Create memo", "Create a new memo. Requires authentication.", false,
|
||||
mcp.WithString("content", mcp.Required(), mcp.Description("Memo content in Markdown. Use #tag syntax for tagging.")),
|
||||
mcp.WithString("visibility",
|
||||
mcp.Enum("PRIVATE", "PROTECTED", "PUBLIC"),
|
||||
mcp.Description("Visibility (default: PRIVATE)"),
|
||||
),
|
||||
mcp.WithOutputSchema[memoJSON](),
|
||||
)...,
|
||||
), s.handleCreateMemo)
|
||||
|
||||
mcpSrv.AddTool(mcp.NewTool("update_memo",
|
||||
updateToolOptions("Update memo", "Update a memo's content, visibility, pin state, or archive state. Requires authentication and ownership. Omit any field to leave it unchanged.",
|
||||
mcp.WithString("name", mcp.Required(), mcp.Description(`Memo resource name, e.g. "memos/abc123"`)),
|
||||
mcp.WithString("content", mcp.Description("New Markdown content")),
|
||||
mcp.WithString("visibility",
|
||||
mcp.Enum("PRIVATE", "PROTECTED", "PUBLIC"),
|
||||
mcp.Description("New visibility"),
|
||||
),
|
||||
mcp.WithBoolean("pinned", mcp.Description("Pin or unpin the memo")),
|
||||
mcp.WithString("state",
|
||||
mcp.Enum("NORMAL", "ARCHIVED"),
|
||||
mcp.Description("Set to ARCHIVED to archive, NORMAL to restore"),
|
||||
),
|
||||
mcp.WithOutputSchema[memoJSON](),
|
||||
)...,
|
||||
), s.handleUpdateMemo)
|
||||
|
||||
mcpSrv.AddTool(mcp.NewTool("delete_memo",
|
||||
updateToolOptions("Delete memo", "Permanently delete a memo. Requires authentication and ownership.",
|
||||
mcp.WithString("name", mcp.Required(), mcp.Description(`Memo resource name, e.g. "memos/abc123"`)),
|
||||
mcp.WithOutputSchema[deletedJSON](),
|
||||
)...,
|
||||
), s.handleDeleteMemo)
|
||||
|
||||
mcpSrv.AddTool(mcp.NewTool("search_memos",
|
||||
readOnlyToolOptions("Search memos", "Search memo content. Authenticated users search their own and visible memos; unauthenticated callers search public memos only.",
|
||||
mcp.WithString("query", mcp.Required(), mcp.Description("Text to search for in memo content")),
|
||||
)...,
|
||||
), s.handleSearchMemos)
|
||||
|
||||
mcpSrv.AddTool(mcp.NewTool("list_memo_comments",
|
||||
readOnlyToolOptions("List memo comments", "List comments on a memo. Visibility rules for comments match those of the parent memo.",
|
||||
mcp.WithString("name", mcp.Required(), mcp.Description(`Memo resource name, e.g. "memos/abc123"`)),
|
||||
)...,
|
||||
), s.handleListMemoComments)
|
||||
|
||||
mcpSrv.AddTool(mcp.NewTool("create_memo_comment",
|
||||
createToolOptions("Create memo comment", "Add a comment to a memo. The comment inherits the parent memo's visibility. Requires authentication.", false,
|
||||
mcp.WithString("name", mcp.Required(), mcp.Description(`Memo resource name to comment on, e.g. "memos/abc123"`)),
|
||||
mcp.WithString("content", mcp.Required(), mcp.Description("Comment content in Markdown")),
|
||||
mcp.WithOutputSchema[memoJSON](),
|
||||
)...,
|
||||
), s.handleCreateMemoComment)
|
||||
}
|
||||
|
||||
func (s *MCPService) handleListMemos(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
userID := auth.GetUserID(ctx)
|
||||
|
||||
pageSize := req.GetInt("page_size", 20)
|
||||
if pageSize <= 0 {
|
||||
pageSize = 20
|
||||
}
|
||||
if pageSize > 100 {
|
||||
pageSize = 100
|
||||
}
|
||||
page := req.GetInt("page", 0)
|
||||
if page < 0 {
|
||||
page = 0
|
||||
}
|
||||
|
||||
var rowStatus *store.RowStatus
|
||||
if state := req.GetString("state", "NORMAL"); state != "" {
|
||||
rs, err := parseRowStatus(state)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
rowStatus = &rs
|
||||
}
|
||||
|
||||
limit := pageSize + 1
|
||||
offset := page * pageSize
|
||||
find := &store.FindMemo{
|
||||
ExcludeComments: true,
|
||||
RowStatus: rowStatus,
|
||||
Limit: &limit,
|
||||
Offset: &offset,
|
||||
OrderByPinned: req.GetBool("order_by_pinned", false),
|
||||
}
|
||||
applyVisibilityFilter(find, userID, rowStatus)
|
||||
if filter := req.GetString("filter", ""); filter != "" {
|
||||
find.Filters = append(find.Filters, filter)
|
||||
}
|
||||
|
||||
memos, err := s.store.ListMemos(ctx, find)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to list memos: %v", err)), nil
|
||||
}
|
||||
|
||||
hasMore := len(memos) > pageSize
|
||||
if hasMore {
|
||||
memos = memos[:pageSize]
|
||||
}
|
||||
creatorIDs := make([]int32, 0, len(memos))
|
||||
for _, memo := range memos {
|
||||
creatorIDs = append(creatorIDs, memo.CreatorID)
|
||||
}
|
||||
usernamesByID, err := preloadUsernames(ctx, s.store, creatorIDs)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to preload memo creators: %v", err)), nil
|
||||
}
|
||||
|
||||
results := make([]memoJSON, len(memos))
|
||||
for i, m := range memos {
|
||||
result, err := storeMemoToJSONWithUsernames(m, usernamesByID)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to resolve memo creator: %v", err)), nil
|
||||
}
|
||||
results[i] = result
|
||||
}
|
||||
|
||||
return newToolResultJSON(memoListJSON{Memos: results, HasMore: hasMore})
|
||||
}
|
||||
|
||||
func (s *MCPService) handleGetMemo(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
userID := auth.GetUserID(ctx)
|
||||
|
||||
uid, err := parseMemoUID(req.GetString("name", ""))
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
memo, err := s.store.GetMemo(ctx, &store.FindMemo{UID: &uid})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to get memo: %v", err)), nil
|
||||
}
|
||||
if memo == nil {
|
||||
return mcp.NewToolResultError("memo not found"), nil
|
||||
}
|
||||
if err := checkMemoAccess(memo, userID); err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
result, err := storeMemoToJSONWithStore(ctx, s.store, memo)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to resolve memo creator: %v", err)), nil
|
||||
}
|
||||
return newToolResultJSON(result)
|
||||
}
|
||||
|
||||
func (s *MCPService) handleCreateMemo(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
if _, err := extractUserID(ctx); err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
content := req.GetString("content", "")
|
||||
if content == "" {
|
||||
return mcp.NewToolResultError("content is required"), nil
|
||||
}
|
||||
visibility, err := parseVisibility(req.GetString("visibility", "PRIVATE"))
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
created, err := s.apiV1Service.CreateMemo(ctx, &v1pb.CreateMemoRequest{
|
||||
Memo: &v1pb.Memo{
|
||||
Content: content,
|
||||
Visibility: visibilityToProto(visibility),
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to create memo: %v", err)), nil
|
||||
}
|
||||
|
||||
result, err := s.loadMemoJSONByName(ctx, created.Name)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
return newToolResultJSON(result)
|
||||
}
|
||||
|
||||
func (s *MCPService) handleUpdateMemo(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
if _, err := extractUserID(ctx); err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
uid, err := parseMemoUID(req.GetString("name", ""))
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
update := &v1pb.Memo{Name: "memos/" + uid}
|
||||
updateMask := &fieldmaskpb.FieldMask{}
|
||||
args := req.GetArguments()
|
||||
|
||||
if v := req.GetString("content", ""); v != "" {
|
||||
update.Content = v
|
||||
updateMask.Paths = append(updateMask.Paths, "content")
|
||||
}
|
||||
if v := req.GetString("visibility", ""); v != "" {
|
||||
vis, err := parseVisibility(v)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
update.Visibility = visibilityToProto(vis)
|
||||
updateMask.Paths = append(updateMask.Paths, "visibility")
|
||||
}
|
||||
if v := req.GetString("state", ""); v != "" {
|
||||
rs, err := parseRowStatus(v)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
update.State = rowStatusToProto(rs)
|
||||
updateMask.Paths = append(updateMask.Paths, "state")
|
||||
}
|
||||
if _, ok := args["pinned"]; ok {
|
||||
update.Pinned = req.GetBool("pinned", false)
|
||||
updateMask.Paths = append(updateMask.Paths, "pinned")
|
||||
}
|
||||
|
||||
if len(updateMask.Paths) == 0 {
|
||||
return mcp.NewToolResultError("at least one field must be provided to update"), nil
|
||||
}
|
||||
|
||||
updated, err := s.apiV1Service.UpdateMemo(ctx, &v1pb.UpdateMemoRequest{
|
||||
Memo: update,
|
||||
UpdateMask: updateMask,
|
||||
})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to update memo: %v", err)), nil
|
||||
}
|
||||
|
||||
result, err := s.loadMemoJSONByName(ctx, updated.Name)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
return newToolResultJSON(result)
|
||||
}
|
||||
|
||||
func (s *MCPService) handleDeleteMemo(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
if _, err := extractUserID(ctx); err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
uid, err := parseMemoUID(req.GetString("name", ""))
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
if _, err := s.apiV1Service.DeleteMemo(ctx, &v1pb.DeleteMemoRequest{Name: "memos/" + uid}); err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to delete memo: %v", err)), nil
|
||||
}
|
||||
return newDeletedToolResult()
|
||||
}
|
||||
|
||||
func (s *MCPService) handleSearchMemos(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
userID := auth.GetUserID(ctx)
|
||||
|
||||
query := req.GetString("query", "")
|
||||
if query == "" {
|
||||
return mcp.NewToolResultError("query is required"), nil
|
||||
}
|
||||
|
||||
limit := 50
|
||||
zero := 0
|
||||
rowStatus := store.Normal
|
||||
find := &store.FindMemo{
|
||||
ExcludeComments: true,
|
||||
RowStatus: &rowStatus,
|
||||
Limit: &limit,
|
||||
Offset: &zero,
|
||||
Filters: []string{fmt.Sprintf(`content.contains(%q)`, query)},
|
||||
}
|
||||
applyVisibilityFilter(find, userID, find.RowStatus)
|
||||
|
||||
memos, err := s.store.ListMemos(ctx, find)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to search memos: %v", err)), nil
|
||||
}
|
||||
creatorIDs := make([]int32, 0, len(memos))
|
||||
for _, memo := range memos {
|
||||
creatorIDs = append(creatorIDs, memo.CreatorID)
|
||||
}
|
||||
usernamesByID, err := preloadUsernames(ctx, s.store, creatorIDs)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to preload memo creators: %v", err)), nil
|
||||
}
|
||||
|
||||
results := make([]memoJSON, len(memos))
|
||||
for i, m := range memos {
|
||||
result, err := storeMemoToJSONWithUsernames(m, usernamesByID)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to resolve memo creator: %v", err)), nil
|
||||
}
|
||||
results[i] = result
|
||||
}
|
||||
return newToolResultJSON(results)
|
||||
}
|
||||
|
||||
func (s *MCPService) handleListMemoComments(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
userID := auth.GetUserID(ctx)
|
||||
|
||||
uid, err := parseMemoUID(req.GetString("name", ""))
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
parent, err := s.store.GetMemo(ctx, &store.FindMemo{UID: &uid})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to get memo: %v", err)), nil
|
||||
}
|
||||
if parent == nil {
|
||||
return mcp.NewToolResultError("memo not found"), nil
|
||||
}
|
||||
if err := checkMemoAccess(parent, userID); err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
relationType := store.MemoRelationComment
|
||||
relations, err := s.store.ListMemoRelations(ctx, &store.FindMemoRelation{
|
||||
RelatedMemoID: &parent.ID,
|
||||
Type: &relationType,
|
||||
})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to list relations: %v", err)), nil
|
||||
}
|
||||
if len(relations) == 0 {
|
||||
return newToolResultJSON([]memoJSON{})
|
||||
}
|
||||
|
||||
commentIDs := make([]int32, len(relations))
|
||||
for i, r := range relations {
|
||||
commentIDs[i] = r.MemoID
|
||||
}
|
||||
|
||||
memos, err := s.store.ListMemos(ctx, &store.FindMemo{IDList: commentIDs})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to list comments: %v", err)), nil
|
||||
}
|
||||
creatorIDs := make([]int32, 0, len(memos))
|
||||
for _, memo := range memos {
|
||||
if checkMemoAccess(memo, userID) == nil {
|
||||
creatorIDs = append(creatorIDs, memo.CreatorID)
|
||||
}
|
||||
}
|
||||
usernamesByID, err := preloadUsernames(ctx, s.store, creatorIDs)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to preload memo creators: %v", err)), nil
|
||||
}
|
||||
|
||||
results := make([]memoJSON, 0, len(memos))
|
||||
for _, m := range memos {
|
||||
if checkMemoAccess(m, userID) == nil {
|
||||
result, err := storeMemoToJSONWithUsernames(m, usernamesByID)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to resolve memo creator: %v", err)), nil
|
||||
}
|
||||
results = append(results, result)
|
||||
}
|
||||
}
|
||||
return newToolResultJSON(results)
|
||||
}
|
||||
|
||||
func (s *MCPService) handleCreateMemoComment(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
userID, err := extractUserID(ctx)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
uid, err := parseMemoUID(req.GetString("name", ""))
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
content := req.GetString("content", "")
|
||||
if content == "" {
|
||||
return mcp.NewToolResultError("content is required"), nil
|
||||
}
|
||||
|
||||
parent, err := s.store.GetMemo(ctx, &store.FindMemo{UID: &uid})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to get memo: %v", err)), nil
|
||||
}
|
||||
if parent == nil {
|
||||
return mcp.NewToolResultError("memo not found"), nil
|
||||
}
|
||||
if err := checkMemoAccess(parent, userID); err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
comment, err := s.apiV1Service.CreateMemoComment(ctx, &v1pb.CreateMemoCommentRequest{
|
||||
Name: "memos/" + uid,
|
||||
Comment: &v1pb.Memo{
|
||||
Content: content,
|
||||
Visibility: visibilityToProto(parent.Visibility),
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to create comment: %v", err)), nil
|
||||
}
|
||||
|
||||
result, err := s.loadMemoJSONByName(ctx, comment.Name)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
return newToolResultJSON(result)
|
||||
}
|
||||
@@ -0,0 +1,180 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"github.com/mark3labs/mcp-go/mcp"
|
||||
mcpserver "github.com/mark3labs/mcp-go/server"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
"github.com/usememos/memos/server/auth"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
type reactionJSON struct {
|
||||
ID int32 `json:"id"`
|
||||
Creator string `json:"creator"`
|
||||
ReactionType string `json:"reaction_type"`
|
||||
CreateTime int64 `json:"create_time"`
|
||||
}
|
||||
|
||||
func (s *MCPService) registerReactionTools(mcpSrv *mcpserver.MCPServer) {
|
||||
mcpSrv.AddTool(mcp.NewTool("list_reactions",
|
||||
readOnlyToolOptions("List reactions", "List all reactions on a memo. Returns reaction type and creator for each reaction.",
|
||||
mcp.WithString("name", mcp.Required(), mcp.Description(`Memo resource name, e.g. "memos/abc123"`)),
|
||||
)...,
|
||||
), s.handleListReactions)
|
||||
|
||||
mcpSrv.AddTool(mcp.NewTool("upsert_reaction",
|
||||
createToolOptions("Upsert reaction", "Add a reaction (emoji) to a memo. If the same reaction already exists from the same user, this is a no-op. Requires authentication.", true,
|
||||
mcp.WithString("name", mcp.Required(), mcp.Description(`Memo resource name, e.g. "memos/abc123"`)),
|
||||
mcp.WithString("reaction_type", mcp.Required(), mcp.Description(`Reaction emoji, e.g. "👍", "❤️", "🎉"`)),
|
||||
mcp.WithOutputSchema[reactionJSON](),
|
||||
)...,
|
||||
), s.handleUpsertReaction)
|
||||
|
||||
mcpSrv.AddTool(mcp.NewTool("delete_reaction",
|
||||
updateToolOptions("Delete reaction", "Remove a reaction by its ID. Requires authentication and ownership of the reaction.",
|
||||
mcp.WithNumber("id", mcp.Required(), mcp.Description("Reaction ID to delete")),
|
||||
mcp.WithOutputSchema[deletedJSON](),
|
||||
)...,
|
||||
), s.handleDeleteReaction)
|
||||
}
|
||||
|
||||
func (s *MCPService) handleListReactions(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
userID := auth.GetUserID(ctx)
|
||||
|
||||
uid, err := parseMemoUID(req.GetString("name", ""))
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
memo, err := s.store.GetMemo(ctx, &store.FindMemo{UID: &uid})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to get memo: %v", err)), nil
|
||||
}
|
||||
if memo == nil {
|
||||
return mcp.NewToolResultError("memo not found"), nil
|
||||
}
|
||||
if err := checkMemoAccess(memo, userID); err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
contentID := "memos/" + uid
|
||||
reactions, err := s.store.ListReactions(ctx, &store.FindReaction{ContentID: &contentID})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to list reactions: %v", err)), nil
|
||||
}
|
||||
creatorIDs := make([]int32, 0, len(reactions))
|
||||
for _, reaction := range reactions {
|
||||
creatorIDs = append(creatorIDs, reaction.CreatorID)
|
||||
}
|
||||
usernamesByID, err := preloadUsernames(ctx, s.store, creatorIDs)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to preload reaction creators: %v", err)), nil
|
||||
}
|
||||
|
||||
results := make([]reactionJSON, len(reactions))
|
||||
for i, r := range reactions {
|
||||
creator, err := lookupUsernameFromCache(usernamesByID, r.CreatorID)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to resolve reaction creator: %v", err)), nil
|
||||
}
|
||||
results[i] = reactionJSON{
|
||||
ID: r.ID,
|
||||
Creator: creator,
|
||||
ReactionType: r.ReactionType,
|
||||
CreateTime: r.CreatedTs,
|
||||
}
|
||||
}
|
||||
|
||||
return newToolResultJSON(results)
|
||||
}
|
||||
|
||||
func (s *MCPService) handleUpsertReaction(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
userID, err := extractUserID(ctx)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
uid, err := parseMemoUID(req.GetString("name", ""))
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
reactionType := req.GetString("reaction_type", "")
|
||||
if reactionType == "" {
|
||||
return mcp.NewToolResultError("reaction_type is required"), nil
|
||||
}
|
||||
|
||||
memo, err := s.store.GetMemo(ctx, &store.FindMemo{UID: &uid})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to get memo: %v", err)), nil
|
||||
}
|
||||
if memo == nil {
|
||||
return mcp.NewToolResultError("memo not found"), nil
|
||||
}
|
||||
if err := checkMemoAccess(memo, userID); err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
// Validate reaction type against allowed reactions.
|
||||
memoRelatedSetting, err := s.store.GetInstanceMemoRelatedSetting(ctx)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to get reaction settings: %v", err)), nil
|
||||
}
|
||||
allowed := false
|
||||
for _, r := range memoRelatedSetting.Reactions {
|
||||
if r == reactionType {
|
||||
allowed = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !allowed {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("reaction %q is not in the allowed reaction list", reactionType)), nil
|
||||
}
|
||||
|
||||
contentID := "memos/" + uid
|
||||
reaction, err := s.apiV1Service.UpsertMemoReaction(ctx, &v1pb.UpsertMemoReactionRequest{
|
||||
Name: contentID,
|
||||
Reaction: &v1pb.Reaction{
|
||||
ContentId: contentID,
|
||||
ReactionType: reactionType,
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to upsert reaction: %v", err)), nil
|
||||
}
|
||||
|
||||
result, err := s.loadReactionJSONByName(ctx, reaction.Name)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
return newToolResultJSON(result)
|
||||
}
|
||||
|
||||
func (s *MCPService) handleDeleteReaction(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
if _, err := extractUserID(ctx); err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
reactionID := int32(req.GetInt("id", 0))
|
||||
if reactionID == 0 {
|
||||
return mcp.NewToolResultError("id is required"), nil
|
||||
}
|
||||
|
||||
reaction, err := s.store.GetReaction(ctx, &store.FindReaction{ID: &reactionID})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to get reaction: %v", err)), nil
|
||||
}
|
||||
if reaction == nil {
|
||||
return mcp.NewToolResultError("reaction not found"), nil
|
||||
}
|
||||
|
||||
if _, err := s.apiV1Service.DeleteMemoReaction(ctx, &v1pb.DeleteMemoReactionRequest{
|
||||
Name: fmt.Sprintf("%s/reactions/%d", reaction.ContentID, reactionID),
|
||||
}); err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to delete reaction: %v", err)), nil
|
||||
}
|
||||
return newDeletedToolResult()
|
||||
}
|
||||
@@ -0,0 +1,297 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"github.com/mark3labs/mcp-go/mcp"
|
||||
mcpserver "github.com/mark3labs/mcp-go/server"
|
||||
|
||||
v1pb "github.com/usememos/memos/proto/gen/api/v1"
|
||||
"github.com/usememos/memos/server/auth"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
type relationJSON struct {
|
||||
Memo string `json:"memo"`
|
||||
RelatedMemo string `json:"related_memo"`
|
||||
Type string `json:"type"`
|
||||
}
|
||||
|
||||
func (s *MCPService) registerRelationTools(mcpSrv *mcpserver.MCPServer) {
|
||||
mcpSrv.AddTool(mcp.NewTool("list_memo_relations",
|
||||
readOnlyToolOptions("List memo relations", "List all relations (references and comments) for a memo. Requires read access to the memo.",
|
||||
mcp.WithString("name", mcp.Required(), mcp.Description(`Memo resource name, e.g. "memos/abc123"`)),
|
||||
mcp.WithString("type",
|
||||
mcp.Enum("REFERENCE", "COMMENT"),
|
||||
mcp.Description("Filter by relation type (optional)"),
|
||||
),
|
||||
)...,
|
||||
), s.handleListMemoRelations)
|
||||
|
||||
mcpSrv.AddTool(mcp.NewTool("create_memo_relation",
|
||||
createToolOptions("Create memo relation", "Create a reference relation between two memos. Requires authentication. For comments, use create_memo_comment instead.", true,
|
||||
mcp.WithString("name", mcp.Required(), mcp.Description(`Source memo resource name, e.g. "memos/abc123"`)),
|
||||
mcp.WithString("related_memo", mcp.Required(), mcp.Description(`Target memo resource name, e.g. "memos/def456"`)),
|
||||
mcp.WithOutputSchema[relationJSON](),
|
||||
)...,
|
||||
), s.handleCreateMemoRelation)
|
||||
|
||||
mcpSrv.AddTool(mcp.NewTool("delete_memo_relation",
|
||||
updateToolOptions("Delete memo relation", "Delete a reference relation between two memos. Requires authentication and ownership of the source memo.",
|
||||
mcp.WithString("name", mcp.Required(), mcp.Description(`Source memo resource name, e.g. "memos/abc123"`)),
|
||||
mcp.WithString("related_memo", mcp.Required(), mcp.Description(`Target memo resource name, e.g. "memos/def456"`)),
|
||||
mcp.WithOutputSchema[deletedJSON](),
|
||||
)...,
|
||||
), s.handleDeleteMemoRelation)
|
||||
}
|
||||
|
||||
func (s *MCPService) handleListMemoRelations(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
userID := auth.GetUserID(ctx)
|
||||
|
||||
uid, err := parseMemoUID(req.GetString("name", ""))
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
memo, err := s.store.GetMemo(ctx, &store.FindMemo{UID: &uid})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to get memo: %v", err)), nil
|
||||
}
|
||||
if memo == nil {
|
||||
return mcp.NewToolResultError("memo not found"), nil
|
||||
}
|
||||
if err := checkMemoAccess(memo, userID); err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
find := &store.FindMemoRelation{
|
||||
MemoIDList: []int32{memo.ID},
|
||||
}
|
||||
if typeStr := req.GetString("type", ""); typeStr != "" {
|
||||
switch store.MemoRelationType(typeStr) {
|
||||
case store.MemoRelationReference, store.MemoRelationComment:
|
||||
t := store.MemoRelationType(typeStr)
|
||||
find.Type = &t
|
||||
default:
|
||||
return mcp.NewToolResultError(fmt.Sprintf("type must be REFERENCE or COMMENT, got %q", typeStr)), nil
|
||||
}
|
||||
}
|
||||
|
||||
relations, err := s.store.ListMemoRelations(ctx, find)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to list relations: %v", err)), nil
|
||||
}
|
||||
|
||||
// Resolve memo IDs to UIDs.
|
||||
idSet := make(map[int32]struct{})
|
||||
for _, r := range relations {
|
||||
idSet[r.MemoID] = struct{}{}
|
||||
idSet[r.RelatedMemoID] = struct{}{}
|
||||
}
|
||||
ids := make([]int32, 0, len(idSet))
|
||||
for id := range idSet {
|
||||
ids = append(ids, id)
|
||||
}
|
||||
memos, err := s.store.ListMemos(ctx, &store.FindMemo{IDList: ids, ExcludeContent: true})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to resolve memos: %v", err)), nil
|
||||
}
|
||||
memoByID := make(map[int32]*store.Memo, len(memos))
|
||||
for _, m := range memos {
|
||||
memoByID[m.ID] = m
|
||||
}
|
||||
|
||||
results := make([]relationJSON, 0, len(relations))
|
||||
for _, r := range relations {
|
||||
srcMemo, ok1 := memoByID[r.MemoID]
|
||||
relatedMemo, ok2 := memoByID[r.RelatedMemoID]
|
||||
if !ok1 || !ok2 {
|
||||
continue
|
||||
}
|
||||
if checkMemoAccess(srcMemo, userID) != nil || checkMemoAccess(relatedMemo, userID) != nil {
|
||||
continue
|
||||
}
|
||||
results = append(results, relationJSON{
|
||||
Memo: "memos/" + srcMemo.UID,
|
||||
RelatedMemo: "memos/" + relatedMemo.UID,
|
||||
Type: string(r.Type),
|
||||
})
|
||||
}
|
||||
|
||||
return newToolResultJSON(results)
|
||||
}
|
||||
|
||||
func (s *MCPService) handleCreateMemoRelation(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
userID, err := extractUserID(ctx)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
srcUID, err := parseMemoUID(req.GetString("name", ""))
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
dstUID, err := parseMemoUID(req.GetString("related_memo", ""))
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
if srcUID == dstUID {
|
||||
return mcp.NewToolResultError("cannot create a relation from a memo to itself"), nil
|
||||
}
|
||||
|
||||
srcMemo, err := s.store.GetMemo(ctx, &store.FindMemo{UID: &srcUID})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to get source memo: %v", err)), nil
|
||||
}
|
||||
if srcMemo == nil {
|
||||
return mcp.NewToolResultError("source memo not found"), nil
|
||||
}
|
||||
if !hasMemoOwnership(srcMemo, userID) {
|
||||
return mcp.NewToolResultError("permission denied: must own the source memo"), nil
|
||||
}
|
||||
|
||||
dstMemo, err := s.store.GetMemo(ctx, &store.FindMemo{UID: &dstUID})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to get related memo: %v", err)), nil
|
||||
}
|
||||
if dstMemo == nil {
|
||||
return mcp.NewToolResultError("related memo not found"), nil
|
||||
}
|
||||
if err := checkMemoAccess(dstMemo, userID); err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
relations, changed, err := s.buildReferenceRelationSet(ctx, srcMemo, &dstMemo.UID, nil)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to build relation set: %v", err)), nil
|
||||
}
|
||||
if changed {
|
||||
if _, err := s.apiV1Service.SetMemoRelations(ctx, &v1pb.SetMemoRelationsRequest{
|
||||
Name: "memos/" + srcUID,
|
||||
Relations: relations,
|
||||
}); err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to create relation: %v", err)), nil
|
||||
}
|
||||
}
|
||||
|
||||
return newToolResultJSON(relationJSON{
|
||||
Memo: "memos/" + srcUID,
|
||||
RelatedMemo: "memos/" + dstUID,
|
||||
Type: string(store.MemoRelationReference),
|
||||
})
|
||||
}
|
||||
|
||||
func (s *MCPService) handleDeleteMemoRelation(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
userID, err := extractUserID(ctx)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
srcUID, err := parseMemoUID(req.GetString("name", ""))
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
dstUID, err := parseMemoUID(req.GetString("related_memo", ""))
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
srcMemo, err := s.store.GetMemo(ctx, &store.FindMemo{UID: &srcUID})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to get source memo: %v", err)), nil
|
||||
}
|
||||
if srcMemo == nil {
|
||||
return mcp.NewToolResultError("source memo not found"), nil
|
||||
}
|
||||
if !hasMemoOwnership(srcMemo, userID) {
|
||||
return mcp.NewToolResultError("permission denied: must own the source memo"), nil
|
||||
}
|
||||
|
||||
dstMemo, err := s.store.GetMemo(ctx, &store.FindMemo{UID: &dstUID})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to get related memo: %v", err)), nil
|
||||
}
|
||||
if dstMemo == nil {
|
||||
return mcp.NewToolResultError("related memo not found"), nil
|
||||
}
|
||||
if err := checkMemoAccess(dstMemo, userID); err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
|
||||
relations, changed, err := s.buildReferenceRelationSet(ctx, srcMemo, nil, &dstMemo.UID)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to build relation set: %v", err)), nil
|
||||
}
|
||||
if changed {
|
||||
if _, err := s.apiV1Service.SetMemoRelations(ctx, &v1pb.SetMemoRelationsRequest{
|
||||
Name: "memos/" + srcUID,
|
||||
Relations: relations,
|
||||
}); err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to delete relation: %v", err)), nil
|
||||
}
|
||||
}
|
||||
return newDeletedToolResult()
|
||||
}
|
||||
|
||||
func (s *MCPService) buildReferenceRelationSet(ctx context.Context, source *store.Memo, includeUID *string, excludeUID *string) ([]*v1pb.MemoRelation, bool, error) {
|
||||
referenceType := store.MemoRelationReference
|
||||
relations, err := s.store.ListMemoRelations(ctx, &store.FindMemoRelation{
|
||||
MemoIDList: []int32{source.ID},
|
||||
Type: &referenceType,
|
||||
})
|
||||
if err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
|
||||
idSet := make(map[int32]struct{}, len(relations))
|
||||
for _, relation := range relations {
|
||||
idSet[relation.RelatedMemoID] = struct{}{}
|
||||
}
|
||||
ids := make([]int32, 0, len(idSet))
|
||||
for id := range idSet {
|
||||
ids = append(ids, id)
|
||||
}
|
||||
|
||||
memosByID := map[int32]*store.Memo{}
|
||||
if len(ids) > 0 {
|
||||
memos, err := s.store.ListMemos(ctx, &store.FindMemo{IDList: ids, ExcludeContent: true})
|
||||
if err != nil {
|
||||
return nil, false, err
|
||||
}
|
||||
for _, memo := range memos {
|
||||
memosByID[memo.ID] = memo
|
||||
}
|
||||
}
|
||||
|
||||
result := make([]*v1pb.MemoRelation, 0, len(relations)+1)
|
||||
seenUIDs := map[string]struct{}{}
|
||||
changed := false
|
||||
for _, relation := range relations {
|
||||
relatedMemo := memosByID[relation.RelatedMemoID]
|
||||
if relatedMemo == nil {
|
||||
continue
|
||||
}
|
||||
if excludeUID != nil && relatedMemo.UID == *excludeUID {
|
||||
changed = true
|
||||
continue
|
||||
}
|
||||
result = append(result, newReferenceRelation(source.UID, relatedMemo.UID))
|
||||
seenUIDs[relatedMemo.UID] = struct{}{}
|
||||
}
|
||||
if includeUID != nil {
|
||||
if _, seen := seenUIDs[*includeUID]; !seen && source.UID != *includeUID {
|
||||
result = append(result, newReferenceRelation(source.UID, *includeUID))
|
||||
changed = true
|
||||
}
|
||||
}
|
||||
return result, changed, nil
|
||||
}
|
||||
|
||||
func newReferenceRelation(sourceUID string, relatedUID string) *v1pb.MemoRelation {
|
||||
return &v1pb.MemoRelation{
|
||||
Memo: &v1pb.MemoRelation_Memo{Name: "memos/" + sourceUID},
|
||||
RelatedMemo: &v1pb.MemoRelation_Memo{Name: "memos/" + relatedUID},
|
||||
Type: v1pb.MemoRelation_REFERENCE,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,74 @@
|
||||
package mcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"slices"
|
||||
|
||||
"github.com/mark3labs/mcp-go/mcp"
|
||||
mcpserver "github.com/mark3labs/mcp-go/server"
|
||||
|
||||
"github.com/usememos/memos/server/auth"
|
||||
"github.com/usememos/memos/store"
|
||||
)
|
||||
|
||||
func (s *MCPService) registerTagTools(mcpSrv *mcpserver.MCPServer) {
|
||||
mcpSrv.AddTool(mcp.NewTool("list_tags",
|
||||
readOnlyToolOptions("List tags", "List all tags with their memo counts. Authenticated users see tags from their own and visible memos; unauthenticated callers see tags from public memos only. Results are sorted by count descending, then alphabetically.")...,
|
||||
), s.handleListTags)
|
||||
}
|
||||
|
||||
type tagEntry struct {
|
||||
Tag string `json:"tag"`
|
||||
Count int `json:"count"`
|
||||
}
|
||||
|
||||
func (s *MCPService) handleListTags(ctx context.Context, _ mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
userID := auth.GetUserID(ctx)
|
||||
|
||||
rowStatus := store.Normal
|
||||
find := &store.FindMemo{
|
||||
ExcludeComments: true,
|
||||
ExcludeContent: true,
|
||||
RowStatus: &rowStatus,
|
||||
}
|
||||
applyVisibilityFilter(find, userID, find.RowStatus)
|
||||
|
||||
memos, err := s.store.ListMemos(ctx, find)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to list memos: %v", err)), nil
|
||||
}
|
||||
|
||||
counts := make(map[string]int)
|
||||
for _, m := range memos {
|
||||
if m.Payload == nil {
|
||||
continue
|
||||
}
|
||||
for _, tag := range m.Payload.Tags {
|
||||
counts[tag]++
|
||||
}
|
||||
}
|
||||
|
||||
entries := make([]tagEntry, 0, len(counts))
|
||||
for tag, count := range counts {
|
||||
entries = append(entries, tagEntry{Tag: tag, Count: count})
|
||||
}
|
||||
slices.SortFunc(entries, func(a, b tagEntry) int {
|
||||
if a.Count != b.Count {
|
||||
if a.Count > b.Count {
|
||||
return -1
|
||||
}
|
||||
return 1
|
||||
}
|
||||
switch {
|
||||
case a.Tag < b.Tag:
|
||||
return -1
|
||||
case a.Tag > b.Tag:
|
||||
return 1
|
||||
default:
|
||||
return 0
|
||||
}
|
||||
})
|
||||
|
||||
return newToolResultJSON(entries)
|
||||
}
|
||||
Reference in New Issue
Block a user