feat(resolvemcp): read/write limits, preload validation, query timeout, stable client error codes

Config gains DefaultLimit/MaxLimit/MaxOffset/MaxBatch/MaxPreloadDepth/MaxWriteRows/QueryTimeout/
ConfirmTTL. Reads are capped and the total COUNT is optional. Errors reach clients as
{code,message}; everything else is logged with a reference. Panics (handler and hooks) are
recovered without returning the panic value.
This commit is contained in:
Hein
2026-10-01 13:35:00 +02:00
parent 82f901a49c
commit 276c3814d8
8 changed files with 414 additions and 36 deletions
+81
View File
@@ -0,0 +1,81 @@
package resolvemcp
import (
"crypto/rand"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"github.com/mark3labs/mcp-go/mcp"
"github.com/bitechdev/ResolveSpec/pkg/logger"
)
// Stable error codes returned to MCP clients.
const (
CodeInvalidArgument = "invalid_argument"
CodeNotFound = "not_found"
CodeForbidden = "forbidden"
CodeLimitExceeded = "limit_exceeded"
CodeInternal = "internal"
)
// ClientError is an error whose code and message are safe to show to an MCP client.
// Hooks may return one (NewClientError) to give the client a specific reason; every other error
// is reported as an opaque internal error with a reference that matches the server log.
type ClientError struct {
Code string
Message string
}
func (e *ClientError) Error() string { return e.Message }
// NewClientError returns an error that reaches the client as {code, message}.
func NewClientError(code, message string) error {
return &ClientError{Code: code, Message: message}
}
func invalidArg(format string, a ...any) error {
return NewClientError(CodeInvalidArgument, fmt.Sprintf(format, a...))
}
// errInternal is what recovered panics return: the details are logged, not sent.
var errInternal = NewClientError(CodeInternal, "internal error")
// clientFacing maps err to the code and message the client sees. Anything that is not a
// ClientError or a not-found is logged in full with a short reference and reported as an
// opaque internal error carrying that reference.
func clientFacing(op string, err error) (code, message string) {
var ce *ClientError
switch {
case errors.As(err, &ce):
if ce.Code == CodeInternal {
ref := newRef()
logger.Error("[resolvemcp] %s: internal error ref=%s: %v", op, ref, err)
return CodeInternal, "internal error (ref " + ref + ")"
}
return ce.Code, ce.Message
case errors.Is(err, errRecordNotFound):
return CodeNotFound, "record not found"
}
ref := newRef()
logger.Error("[resolvemcp] %s failed ref=%s: %v", op, ref, err)
return CodeInternal, "internal error (ref " + ref + ")"
}
func newRef() string {
var b [4]byte
_, _ = rand.Read(b[:])
return hex.EncodeToString(b[:])
}
// toolError builds the error result for a tool call.
func toolError(op string, err error) *mcp.CallToolResult {
code, msg := clientFacing(op, err)
b, _ := json.Marshal(map[string]any{
"success": false,
"error": map[string]string{"code": code, "message": msg},
})
return mcp.NewToolResultError(string(b))
}
+103 -22
View File
@@ -8,6 +8,8 @@ import (
"fmt"
"net/http"
"reflect"
"regexp"
"runtime/debug"
"strings"
"sync"
@@ -40,7 +42,7 @@ func NewHandler(db common.Database, registry common.ModelRegistry, cfg Config) *
registry: registry,
hooks: NewHookRegistry(),
mcpServer: server.NewMCPServer("resolvemcp", "1.0.0"),
config: cfg,
config: cfg.withDefaults(),
name: "resolvemcp",
version: "1.0.0",
}
@@ -244,18 +246,28 @@ var errRecordNotFound = errors.New("record not found")
// Usage: defer recoverPanic(&returnedErr)
func recoverPanic(err *error) {
if r := recover(); r != nil {
msg := fmt.Sprintf("%v", r)
logger.Error("[resolvemcp] panic recovered: %s", msg)
*err = fmt.Errorf("internal error: %s", msg)
logger.Error("[resolvemcp] panic recovered: %v\n%s", r, debug.Stack())
*err = errInternal
}
}
// executeRead reads records from the database and returns raw data + metadata.
func (h *Handler) executeRead(ctx context.Context, schema, entity, id string, options common.RequestOptions) (_ interface{}, _ *common.Metadata, retErr error) {
func (h *Handler) executeRead(ctx context.Context, schema, entity, id string, options common.RequestOptions) (interface{}, *common.Metadata, error) {
return h.executeReadCounted(ctx, schema, entity, id, options, true)
}
// executeReadCounted is executeRead with control over the total-row COUNT, which costs a full
// scan of the filtered set and is only run when the caller asks for it.
func (h *Handler) executeReadCounted(ctx context.Context, schema, entity, id string, options common.RequestOptions, count bool) (_ interface{}, _ *common.Metadata, retErr error) {
defer recoverPanic(&retErr)
ctx, cancel := h.callContext(ctx)
defer cancel()
if err := h.checkReadLimits(&options); err != nil {
return nil, nil, err
}
model, err := h.registry.GetModelByEntity(schema, entity)
if err != nil {
return nil, nil, fmt.Errorf("model not found: %w", err)
return nil, nil, invalidArg("model not found: %s", buildModelName(schema, entity))
}
unwrapped, err := common.ValidateAndUnwrapModel(model)
@@ -293,7 +305,7 @@ func (h *Handler) executeRead(ctx context.Context, schema, entity, id string, op
var metadata *common.Metadata
err = h.runInTx(ctx, hookCtx, func(common.Database) error {
var err error
data, metadata, err = h.readInTx(ctx, hookCtx, model, modelType, tableName, id, options)
data, metadata, err = h.readInTx(ctx, hookCtx, model, modelType, tableName, id, options, count)
return err
})
if err != nil {
@@ -303,7 +315,7 @@ func (h *Handler) executeRead(ctx context.Context, schema, entity, id string, op
}
// readInTx runs the read hooks and queries on hookCtx.Tx.
func (h *Handler) readInTx(ctx context.Context, hookCtx *HookContext, model interface{}, modelType reflect.Type, tableName, id string, options common.RequestOptions) (interface{}, *common.Metadata, error) {
func (h *Handler) readInTx(ctx context.Context, hookCtx *HookContext, model interface{}, modelType reflect.Type, tableName, id string, options common.RequestOptions, count bool) (interface{}, *common.Metadata, error) {
sliceType := reflect.SliceOf(reflect.PointerTo(modelType))
modelPtr := reflect.New(sliceType).Interface()
@@ -354,7 +366,7 @@ func (h *Handler) readInTx(ctx context.Context, hookCtx *HookContext, model inte
// expandJoins is empty for resolvemcp — no custom SQL join support yet
cursorFilter, err := getCursorFilter(tableName, pkName, modelColumns, options, nil)
if err != nil {
return nil, nil, fmt.Errorf("cursor error: %w", err)
return nil, nil, invalidArg("invalid cursor")
}
if cursorFilter != "" {
@@ -367,9 +379,13 @@ func (h *Handler) readInTx(ctx context.Context, hookCtx *HookContext, model inte
}
// Count — must happen before preloads are applied; Bun panics when counting with relations.
total, err := query.Count(ctx)
if err != nil {
return nil, nil, fmt.Errorf("error counting records: %w", err)
total := 0
if count {
var err error
total, err = query.Count(ctx)
if err != nil {
return nil, nil, fmt.Errorf("error counting records: %w", err)
}
}
// Pagination
@@ -382,6 +398,9 @@ func (h *Handler) readInTx(ctx context.Context, hookCtx *HookContext, model inte
// Preloads — applied after count to avoid Bun panic when counting with relations.
if len(options.Preload) > 0 {
if err := h.validatePreloads(model, options.Preload); err != nil {
return nil, nil, err
}
var preloadErr error
query, preloadErr = h.applyPreloads(model, query, options.Preload)
if preloadErr != nil {
@@ -461,9 +480,14 @@ func (h *Handler) readInTx(ctx context.Context, hookCtx *HookContext, model inte
// executeCreate inserts one or more records.
func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data interface{}) (_ interface{}, retErr error) {
defer recoverPanic(&retErr)
ctx, cancel := h.callContext(ctx)
defer cancel()
if items, ok := data.([]interface{}); ok && len(items) > h.config.MaxBatch {
return nil, NewClientError(CodeLimitExceeded, fmt.Sprintf("batch of %d exceeds the maximum of %d", len(items), h.config.MaxBatch))
}
model, err := h.registry.GetModelByEntity(schema, entity)
if err != nil {
return nil, fmt.Errorf("model not found: %w", err)
return nil, invalidArg("model not found: %s", buildModelName(schema, entity))
}
result, err := common.ValidateAndUnwrapModel(model)
@@ -515,12 +539,12 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
for _, item := range v {
itemMap, ok := item.(map[string]interface{})
if !ok {
return fmt.Errorf("each item must be an object")
return invalidArg("each item must be an object")
}
originals = append(originals, itemMap)
}
default:
return fmt.Errorf("data must be an object or array of objects")
return invalidArg("data must be an object or array of objects")
}
insertedIDs = make([]interface{}, 0, len(originals))
@@ -530,7 +554,7 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
return err
}
if len(cols) == 0 {
return fmt.Errorf("no writable fields in data")
return invalidArg("no writable fields in data")
}
q := tx.NewInsert().Table(tableName)
for key, value := range cols {
@@ -597,9 +621,11 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
// executeUpdate updates a record by ID.
func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string, data interface{}) (_ interface{}, retErr error) {
defer recoverPanic(&retErr)
ctx, cancel := h.callContext(ctx)
defer cancel()
model, err := h.registry.GetModelByEntity(schema, entity)
if err != nil {
return nil, fmt.Errorf("model not found: %w", err)
return nil, invalidArg("model not found: %s", buildModelName(schema, entity))
}
result, err := common.ValidateAndUnwrapModel(model)
@@ -613,7 +639,7 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
updates, ok := data.(map[string]interface{})
if !ok {
return nil, fmt.Errorf("data must be an object")
return nil, invalidArg("data must be an object")
}
if id == "" {
@@ -622,7 +648,7 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
}
}
if id == "" {
return nil, fmt.Errorf("update requires an ID")
return nil, invalidArg("update requires an id")
}
pkName := reflection.GetPrimaryKeyName(model)
@@ -663,7 +689,7 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
}
}
if len(setCols) == 0 {
return fmt.Errorf("no updatable fields in data")
return invalidArg("no updatable fields in data")
}
// Load the target through the BeforeScan hooks (row security) so a row the caller
@@ -735,13 +761,15 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
// executeDelete deletes a record by ID.
func (h *Handler) executeDelete(ctx context.Context, schema, entity, id string) (_ interface{}, retErr error) {
defer recoverPanic(&retErr)
ctx, cancel := h.callContext(ctx)
defer cancel()
if id == "" {
return nil, fmt.Errorf("delete requires an ID")
return nil, invalidArg("delete requires an id")
}
model, err := h.registry.GetModelByEntity(schema, entity)
if err != nil {
return nil, fmt.Errorf("model not found: %w", err)
return nil, invalidArg("model not found: %s", buildModelName(schema, entity))
}
result, err := common.ValidateAndUnwrapModel(model)
@@ -945,3 +973,56 @@ func (h *Handler) runInTx(ctx context.Context, hookCtx *HookContext, body func(t
return h.hooks.Execute(OnTxBegin, hookCtx)
}, body)
}
// callContext bounds one tool call by Config.QueryTimeout.
func (h *Handler) callContext(ctx context.Context) (context.Context, context.CancelFunc) {
return context.WithTimeout(ctx, h.config.QueryTimeout)
}
// checkReadLimits applies the paging caps to options in place: a missing limit takes
// DefaultLimit, a larger one is clamped to MaxLimit, and an offset above MaxOffset is rejected.
func (h *Handler) checkReadLimits(options *common.RequestOptions) error {
limit := h.config.DefaultLimit
if options.Limit != nil && *options.Limit > 0 {
limit = *options.Limit
}
if limit > h.config.MaxLimit {
limit = h.config.MaxLimit
}
options.Limit = &limit
if options.Offset != nil {
if *options.Offset < 0 {
return invalidArg("offset must not be negative")
}
if *options.Offset > h.config.MaxOffset {
return NewClientError(CodeLimitExceeded, fmt.Sprintf("offset exceeds the maximum of %d; use cursor paging", h.config.MaxOffset))
}
}
return nil
}
var preloadSegmentRe = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*$`)
// validatePreloads checks preload paths against the model: the first segment must be one of
// the model's relations and the path may not be deeper than Config.MaxPreloadDepth.
func (h *Handler) validatePreloads(model interface{}, preloads []common.PreloadOption) error {
relations := map[string]bool{}
for _, name := range buildModelInfo("", "", model).relationNames {
relations[strings.ToLower(name)] = true
}
for _, p := range preloads {
segments := strings.Split(p.Relation, ".")
if len(segments) > h.config.MaxPreloadDepth {
return NewClientError(CodeLimitExceeded, fmt.Sprintf("preload depth exceeds the maximum of %d", h.config.MaxPreloadDepth))
}
for _, seg := range segments {
if !preloadSegmentRe.MatchString(seg) {
return invalidArg("invalid preload relation")
}
}
if !relations[strings.ToLower(segments[0])] {
return invalidArg("unknown relation %q", segments[0])
}
}
return nil
}
+15 -2
View File
@@ -3,6 +3,7 @@ package resolvemcp
import (
"context"
"fmt"
"runtime/debug"
"sync"
"github.com/bitechdev/ResolveSpec/pkg/common"
@@ -107,20 +108,32 @@ func (r *HookRegistry) Execute(hookType HookType, ctx *HookContext) error {
logger.Debug("Executing %d resolvemcp hook(s) for %s", len(hooks), hookType)
for i, hook := range hooks {
if err := hook(ctx); err != nil {
if err := runHook(hook, ctx); err != nil {
logger.Error("resolvemcp hook %d for %s failed: %v", i+1, hookType, err)
return fmt.Errorf("hook execution failed: %w", err)
}
if ctx.Abort {
logger.Warn("resolvemcp hook %d for %s requested abort: %s", i+1, hookType, ctx.AbortMessage)
return fmt.Errorf("operation aborted by hook: %s", ctx.AbortMessage)
return fmt.Errorf("operation aborted by hook: %w", NewClientError(CodeForbidden, ctx.AbortMessage))
}
}
return nil
}
// runHook calls hook and turns a panic into an error, so a faulty hook fails the request
// instead of unwinding through the transaction machinery. The stack is logged, not returned.
func runHook(hook HookFunc, ctx *HookContext) (err error) {
defer func() {
if r := recover(); r != nil {
logger.Error("resolvemcp hook panic: %v\n%s", r, debug.Stack())
err = errInternal
}
}()
return hook(ctx)
}
func (r *HookRegistry) Clear(hookType HookType) {
r.mu.Lock()
defer r.mu.Unlock()
+151
View File
@@ -0,0 +1,151 @@
package resolvemcp
import (
"context"
"errors"
"strings"
"testing"
"time"
"github.com/DATA-DOG/go-sqlmock"
"github.com/bitechdev/ResolveSpec/pkg/common"
)
func ptr[T any](v T) *T { return &v }
func TestConfigDefaults(t *testing.T) {
c := Config{}.withDefaults()
if c.DefaultLimit != 50 || c.MaxLimit != 1000 || c.MaxOffset != 100000 || c.MaxBatch != 100 ||
c.MaxPreloadDepth != 2 || c.MaxWriteRows != 100 || c.QueryTimeout != 30*time.Second || c.ConfirmTTL != 5*time.Minute {
t.Errorf("defaults: %+v", c)
}
if c := (Config{DefaultLimit: 5000, MaxLimit: 200}).withDefaults(); c.DefaultLimit != 200 {
t.Errorf("default limit must not exceed max: %d", c.DefaultLimit)
}
}
func TestCheckReadLimits(t *testing.T) {
h, _, _ := newTxHarness(t)
cases := []struct {
name string
in common.RequestOptions
want int
wantErr string
}{
{"no limit takes default", common.RequestOptions{}, 50, ""},
{"zero takes default", common.RequestOptions{Limit: ptr(0)}, 50, ""},
{"negative takes default", common.RequestOptions{Limit: ptr(-3)}, 50, ""},
{"explicit kept", common.RequestOptions{Limit: ptr(10)}, 10, ""},
{"clamped", common.RequestOptions{Limit: ptr(1 << 30)}, 1000, ""},
{"offset too big", common.RequestOptions{Offset: ptr(100001)}, 0, CodeLimitExceeded},
{"offset negative", common.RequestOptions{Offset: ptr(-1)}, 0, CodeInvalidArgument},
{"offset at max ok", common.RequestOptions{Offset: ptr(100000)}, 50, ""},
}
for _, c := range cases {
opts := c.in
err := h.checkReadLimits(&opts)
if c.wantErr != "" {
var ce *ClientError
if !errors.As(err, &ce) || ce.Code != c.wantErr {
t.Errorf("%s: got %v, want code %s", c.name, err, c.wantErr)
}
continue
}
if err != nil || *opts.Limit != c.want {
t.Errorf("%s: limit %v err %v, want %d", c.name, opts.Limit, err, c.want)
}
}
}
func TestValidatePreloads(t *testing.T) {
type child struct {
ID int `json:"id" bun:"id,pk"`
}
type parent struct {
ID int `json:"id" bun:"id,pk"`
Children []*child `json:"children" bun:"rel:has-many"`
}
h, _, _ := newTxHarness(t)
ok := func(rel string) bool {
return h.validatePreloads(&parent{}, []common.PreloadOption{{Relation: rel}}) == nil
}
if !ok("children") || !ok("Children") || !ok("children.sub") {
t.Error("known relations (and depth 2) must pass")
}
for _, bad := range []string{"nope", "children.a.b", "children; DROP TABLE x", "", "a..b", "id"} {
if ok(bad) {
t.Errorf("preload %q must be rejected", bad)
}
}
}
func TestBatchCap(t *testing.T) {
h, _, ctx := newTxHarness(t)
h.config.MaxBatch = 2
items := []interface{}{map[string]interface{}{}, map[string]interface{}{}, map[string]interface{}{}}
_, err := h.executeCreate(ctx, "public", "items", items)
var ce *ClientError
if !errors.As(err, &ce) || ce.Code != CodeLimitExceeded {
t.Fatalf("want limit_exceeded, got %v", err)
}
}
func TestReadSkipsCountUnlessRequested(t *testing.T) {
h, mock, ctx := newTxHarness(t)
mock.ExpectBegin()
mock.ExpectQuery(`SELECT .* LIMIT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(1, "a"))
mock.ExpectCommit()
if _, _, err := h.executeReadCounted(ctx, "public", "items", "", common.RequestOptions{}, false); err != nil {
t.Fatal(err)
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
}
func TestClientFacingHidesInternals(t *testing.T) {
code, msg := clientFacing("t", errors.New(`pq: relation "secret_table" does not exist`))
if code != CodeInternal || strings.Contains(msg, "secret_table") || !strings.Contains(msg, "ref ") {
t.Errorf("raw error leaked: %s %q", code, msg)
}
if code, msg := clientFacing("t", invalidArg("bad %s", "x")); code != CodeInvalidArgument || msg != "bad x" {
t.Errorf("client error: %s %q", code, msg)
}
if code, _ := clientFacing("t", errRecordNotFound); code != CodeNotFound {
t.Errorf("not found: %s", code)
}
wrapped := errors.Join(errors.New("ctx"), NewClientError(CodeForbidden, "update not allowed for x"))
if code, msg := clientFacing("t", wrapped); code != CodeForbidden || msg != "update not allowed for x" {
t.Errorf("wrapped: %s %q", code, msg)
}
}
func TestHookPanicIsRecovered(t *testing.T) {
h, mock, ctx := newTxHarness(t)
h.Hooks().Register(BeforeDelete, func(*HookContext) error { panic("boom: secret") })
mock.ExpectBegin()
mock.ExpectRollback()
_, err := h.executeDelete(ctx, "public", "items", "7")
if err == nil {
t.Fatal("expected error")
}
if _, msg := clientFacing("t", err); strings.Contains(msg, "boom") || strings.Contains(msg, "secret") {
t.Errorf("panic value leaked: %q", msg)
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
}
func TestToolTimeoutBoundsCall(t *testing.T) {
h, _, _ := newTxHarness(t)
h.config.QueryTimeout = 20 * time.Millisecond
ctx, cancel := h.callContext(context.Background())
defer cancel()
select {
case <-ctx.Done():
case <-time.After(time.Second):
t.Fatal("call context did not time out")
}
}
+45
View File
@@ -18,6 +18,7 @@ package resolvemcp
import (
"net/http"
"runtime/debug"
"time"
"github.com/gorilla/mux"
"github.com/uptrace/bun"
@@ -41,6 +42,25 @@ type Config struct {
// If empty, the path is detected from each incoming request automatically.
BasePath string
// Limits. Zero values take the defaults shown.
// DefaultLimit is the page size when a read gives no limit (50).
DefaultLimit int
// MaxLimit caps a read's limit; larger values are clamped (1000).
MaxLimit int
// MaxOffset rejects a read whose offset is larger (100000).
MaxOffset int
// MaxBatch caps the items in one batch create (100).
MaxBatch int
// MaxPreloadDepth caps the depth of a preload path such as "a.b.c" (2).
MaxPreloadDepth int
// MaxWriteRows caps the rows a filter-based update or delete may touch (100).
MaxWriteRows int
// QueryTimeout bounds one tool call, hooks and queries included (30s).
QueryTimeout time.Duration
// ConfirmTTL is how long a confirmation token for a filter write stays valid (5m).
ConfirmTTL time.Duration
// AllowedHosts restricts the Host header accepted by the SSE transport when BaseURL is
// empty (the message endpoint URL sent to clients is built from it). Empty accepts any
// host, with at most 32 distinct base URLs cached; prefer setting BaseURL.
@@ -53,6 +73,31 @@ type Config struct {
EnableAnnotations bool
}
// withDefaults fills the zero limit fields.
func (c Config) withDefaults() Config {
def := func(v *int, d int) {
if *v <= 0 {
*v = d
}
}
def(&c.DefaultLimit, 50)
def(&c.MaxLimit, 1000)
def(&c.MaxOffset, 100000)
def(&c.MaxBatch, 100)
def(&c.MaxPreloadDepth, 2)
def(&c.MaxWriteRows, 100)
if c.DefaultLimit > c.MaxLimit {
c.DefaultLimit = c.MaxLimit
}
if c.QueryTimeout <= 0 {
c.QueryTimeout = 30 * time.Second
}
if c.ConfirmTTL <= 0 {
c.ConfirmTTL = 5 * time.Minute
}
return c
}
// NewHandlerWithGORM creates a Handler backed by a GORM database connection.
func NewHandlerWithGORM(db *gorm.DB, cfg Config) *Handler {
return NewHandler(database.NewGormAdapter(db), modelregistry.NewModelRegistry(), cfg)
+12 -4
View File
@@ -31,7 +31,7 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList
hookCtx.Abort = true
hookCtx.AbortMessage = err.Error()
hookCtx.AbortCode = http.StatusUnauthorized
return err
return NewClientError(CodeForbidden, err.Error())
}
return nil
})
@@ -83,17 +83,17 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList
// BeforeCreate: enforce CanCreate rule.
handler.Hooks().Register(BeforeCreate, func(hookCtx *HookContext) error {
return security.CheckModelCreateAllowed(newSecurityContext(hookCtx))
return forbidden(security.CheckModelCreateAllowed(newSecurityContext(hookCtx)))
})
// BeforeUpdate: enforce CanUpdate rule.
handler.Hooks().Register(BeforeUpdate, func(hookCtx *HookContext) error {
return security.CheckModelUpdateAllowed(newSecurityContext(hookCtx))
return forbidden(security.CheckModelUpdateAllowed(newSecurityContext(hookCtx)))
})
// BeforeDelete: enforce CanDelete rule.
handler.Hooks().Register(BeforeDelete, func(hookCtx *HookContext) error {
return security.CheckModelDeleteAllowed(newSecurityContext(hookCtx))
return forbidden(security.CheckModelDeleteAllowed(newSecurityContext(hookCtx)))
})
logger.Info("Security hooks registered for resolvemcp handler")
@@ -167,3 +167,11 @@ func (s *securityContext) GetResult() interface{} {
func (s *securityContext) SetResult(result interface{}) {
s.ctx.Result = result
}
// forbidden marks a rule denial as safe to show the client; nil passes through.
func forbidden(err error) error {
if err == nil {
return nil
}
return NewClientError(CodeForbidden, err.Error())
}
+4 -4
View File
@@ -324,7 +324,7 @@ func registerReadTool(h *Handler, schema, entity string, info modelInfo) {
data, metadata, err := h.executeRead(ctx, schema, entity, id, options)
if err != nil {
return mcp.NewToolResultError(err.Error()), nil
return toolError("tool", err), nil
}
return marshalResult(map[string]interface{}{
@@ -394,7 +394,7 @@ func registerCreateTool(h *Handler, schema, entity string, info modelInfo) {
result, err := h.executeCreate(ctx, schema, entity, data)
if err != nil {
return mcp.NewToolResultError(err.Error()), nil
return toolError("tool", err), nil
}
return marshalResult(map[string]interface{}{
@@ -463,7 +463,7 @@ func registerUpdateTool(h *Handler, schema, entity string, info modelInfo) {
result, err := h.executeUpdate(ctx, schema, entity, id, dataMap)
if err != nil {
return mcp.NewToolResultError(err.Error()), nil
return toolError("tool", err), nil
}
return marshalResult(map[string]interface{}{
@@ -504,7 +504,7 @@ func registerDeleteTool(h *Handler, schema, entity string, info modelInfo) {
result, err := h.executeDelete(ctx, schema, entity, id)
if err != nil {
return mcp.NewToolResultError(err.Error()), nil
return toolError("tool", err), nil
}
return marshalResult(map[string]interface{}{
+3 -4
View File
@@ -1,7 +1,6 @@
package resolvemcp
import (
"fmt"
"reflect"
"sort"
"strings"
@@ -23,7 +22,7 @@ func writeColumns(model interface{}, data map[string]interface{}) (map[string]in
modelType = modelType.Elem()
}
if modelType == nil || modelType.Kind() != reflect.Struct {
return nil, fmt.Errorf("invalid model")
return nil, errInternal
}
accepted := make(map[string]string)
@@ -44,13 +43,13 @@ func writeColumns(model interface{}, data map[string]interface{}) (map[string]in
continue
}
if _, dup := out[col]; dup {
return nil, fmt.Errorf("column %q given more than once", col)
return nil, invalidArg("column %q given more than once", col)
}
out[col] = value
}
if len(unknown) > 0 {
sort.Strings(unknown)
return nil, fmt.Errorf("unknown or read-only fields: %s", strings.Join(unknown, ", "))
return nil, invalidArg("unknown or read-only fields: %s", strings.Join(unknown, ", "))
}
return out, nil
}