memos/server/router/mcp/adapter.go

250 lines
6.9 KiB
Go

package mcp
import (
"bytes"
"context"
"encoding/json"
"io"
"net/http"
"net/http/httptest"
"net/url"
"strconv"
"strings"
"github.com/labstack/echo/v5"
sdkmcp "github.com/modelcontextprotocol/go-sdk/mcp"
"github.com/pkg/errors"
)
type apiAdapter struct {
echoServer *echo.Echo
}
func newAPIAdapter(echoServer *echo.Echo) *apiAdapter {
return &apiAdapter{echoServer: echoServer}
}
func (a *apiAdapter) execute(ctx context.Context, operation *openAPIOperation, arguments map[string]any, authorization string) (*sdkmcp.CallToolResult, error) {
req, err := buildAPIRequest(ctx, operation, arguments, authorization)
if err != nil {
return newToolErrorResult(err.Error()), nil
}
recorder := httptest.NewRecorder()
a.echoServer.ServeHTTP(recorder, req)
value, err := decodeJSONValue(recorder.Body.Bytes())
if err != nil {
return newToolErrorResult(err.Error()), nil
}
if recorder.Code < http.StatusOK || recorder.Code >= http.StatusMultipleChoices {
return newToolErrorResult(apiErrorMessage(recorder.Code, value)), nil
}
return newStructuredToolResult(value)
}
func buildAPIRequest(ctx context.Context, operation *openAPIOperation, arguments map[string]any, authorization string) (*http.Request, error) {
path, err := substitutePathParameters(operation, arguments)
if err != nil {
return nil, err
}
query := url.Values{}
for _, parameter := range operation.Parameters {
if parameter.In != "query" {
continue
}
value, ok := arguments[parameter.Name]
if !ok || value == nil {
continue
}
query.Set(parameter.Name, valueToString(value))
}
if encoded := query.Encode(); encoded != "" {
path += "?" + encoded
}
var body io.Reader
if operation.RequestBody != nil {
bodyValue, ok := arguments["body"]
if !ok || bodyValue == nil {
if operation.RequestBody.Required {
return nil, errors.New(`missing required request body "body"`)
}
bodyValue = map[string]any{}
}
data, err := json.Marshal(bodyValue)
if err != nil {
return nil, errors.Wrap(err, "failed to marshal request body")
}
body = bytes.NewReader(data)
}
req := httptest.NewRequest(operation.Method, path, body).WithContext(ctx)
if body != nil {
req.Header.Set("Content-Type", "application/json")
}
if authorization != "" {
req.Header.Set("Authorization", authorization)
}
return req, nil
}
func substitutePathParameters(operation *openAPIOperation, arguments map[string]any) (string, error) {
for _, parameter := range operation.Parameters {
if parameter.In != "path" {
continue
}
value, ok := arguments[parameter.Name]
if !ok || value == nil || valueToString(value) == "" {
return "", errors.Errorf(`missing required path parameter "%s"`, parameter.Name)
}
}
// Resolve placeholders in the order they appear in the path so a nested
// resource name (e.g. "memos/abc123/reactions/reaction456") can be matched
// against its already-resolved parent segments. Each placeholder is resolved
// exactly once from the argument map and cached in resolved, so a value that
// itself contains a "{" can never be re-expanded into a longer prefix.
path := operation.Path
resolved := map[string]string{}
for _, name := range pathPlaceholderNames(operation.Path) {
value, ok := arguments[name]
if !ok {
continue
}
id := trimResourceNamePrefix(operation.Path, name, valueToString(value), resolved)
resolved[name] = id
path = strings.ReplaceAll(path, "{"+name+"}", url.PathEscape(id))
}
return path, nil
}
// pathPlaceholderNames returns the "{name}" placeholder names in the order they
// appear in path.
func pathPlaceholderNames(path string) []string {
var names []string
for {
start := strings.Index(path, "{")
if start < 0 {
break
}
endOffset := strings.Index(path[start:], "}")
if endOffset < 0 {
break
}
end := start + endOffset
names = append(names, path[start+1:end])
path = path[end+1:]
}
return names
}
// trimResourceNamePrefix accepts canonical resource names for path parameters.
// It uses the already-resolved parent segments so nested names such as
// "memos/abc123/reactions/reaction456" are accepted only when their parent
// segments match the other arguments. Bare IDs pass through unchanged.
func trimResourceNamePrefix(path, parameterName, value string, resolved map[string]string) string {
placeholder := "/{" + parameterName + "}"
index := strings.Index(path, placeholder)
if index < 0 {
return value
}
prefix, ok := resolvedResourceNamePrefix(path[:index], resolved)
if !ok || prefix == "" {
return value
}
id, ok := strings.CutPrefix(value, prefix+"/")
if !ok || id == "" || strings.Contains(id, "/") {
return value
}
return id
}
// resolvedResourceNamePrefix rebuilds the collection prefix preceding a
// placeholder by substituting each earlier placeholder with its already-resolved
// bare id. It reads resolved ids from the map instead of recursing, and rejects
// any id that is not a single bare segment, so every iteration removes one "{"
// and the loop always terminates.
func resolvedResourceNamePrefix(prefix string, resolved map[string]string) (string, bool) {
const apiPrefix = "/api/v1/"
prefix, ok := strings.CutPrefix(prefix, apiPrefix)
if !ok {
return "", false
}
for {
start := strings.Index(prefix, "{")
if start < 0 {
break
}
endOffset := strings.Index(prefix[start:], "}")
if endOffset < 0 {
return "", false
}
end := start + endOffset
parameterName := prefix[start+1 : end]
id, ok := resolved[parameterName]
if !ok || id == "" || strings.ContainsAny(id, "/{}") {
return "", false
}
prefix = prefix[:start] + id + prefix[end+1:]
}
return strings.Trim(prefix, "/"), true
}
func valueToString(value any) string {
switch typed := value.(type) {
case string:
return typed
case bool:
return strconv.FormatBool(typed)
case int:
return strconv.Itoa(typed)
case int8:
return strconv.FormatInt(int64(typed), 10)
case int16:
return strconv.FormatInt(int64(typed), 10)
case int32:
return strconv.FormatInt(int64(typed), 10)
case int64:
return strconv.FormatInt(typed, 10)
case uint:
return strconv.FormatUint(uint64(typed), 10)
case uint8:
return strconv.FormatUint(uint64(typed), 10)
case uint16:
return strconv.FormatUint(uint64(typed), 10)
case uint32:
return strconv.FormatUint(uint64(typed), 10)
case uint64:
return strconv.FormatUint(typed, 10)
case float32:
return strconv.FormatFloat(float64(typed), 'f', -1, 32)
case float64:
return strconv.FormatFloat(typed, 'f', -1, 64)
default:
data, err := json.Marshal(value)
if err != nil {
return ""
}
return strings.Trim(string(data), `"`)
}
}
func apiErrorMessage(statusCode int, value any) string {
status := strings.TrimSpace(strconv.Itoa(statusCode) + " " + http.StatusText(statusCode))
if object, ok := value.(map[string]any); ok {
for _, key := range []string{"message", "error"} {
message, ok := object[key].(string)
if ok && message != "" {
return status + ": " + message
}
}
}
return status
}