From 276c3814d8f3585525ea0c79965db9cc85e7def4 Mon Sep 17 00:00:00 2001 From: Hein Date: Thu, 1 Oct 2026 13:35:00 +0200 Subject: [PATCH] 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. --- pkg/resolvemcp/errors.go | 81 +++++++++++++++++ pkg/resolvemcp/handler.go | 125 ++++++++++++++++++++----- pkg/resolvemcp/hooks.go | 17 +++- pkg/resolvemcp/limits_test.go | 151 +++++++++++++++++++++++++++++++ pkg/resolvemcp/resolvemcp.go | 45 +++++++++ pkg/resolvemcp/security_hooks.go | 16 +++- pkg/resolvemcp/tools.go | 8 +- pkg/resolvemcp/writecols.go | 7 +- 8 files changed, 414 insertions(+), 36 deletions(-) create mode 100644 pkg/resolvemcp/errors.go create mode 100644 pkg/resolvemcp/limits_test.go diff --git a/pkg/resolvemcp/errors.go b/pkg/resolvemcp/errors.go new file mode 100644 index 0000000..826b40d --- /dev/null +++ b/pkg/resolvemcp/errors.go @@ -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)) +} diff --git a/pkg/resolvemcp/handler.go b/pkg/resolvemcp/handler.go index a38cc66..584eebc 100644 --- a/pkg/resolvemcp/handler.go +++ b/pkg/resolvemcp/handler.go @@ -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 +} diff --git a/pkg/resolvemcp/hooks.go b/pkg/resolvemcp/hooks.go index 5187d81..7d59d25 100644 --- a/pkg/resolvemcp/hooks.go +++ b/pkg/resolvemcp/hooks.go @@ -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() diff --git a/pkg/resolvemcp/limits_test.go b/pkg/resolvemcp/limits_test.go new file mode 100644 index 0000000..3c81f6c --- /dev/null +++ b/pkg/resolvemcp/limits_test.go @@ -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") + } +} diff --git a/pkg/resolvemcp/resolvemcp.go b/pkg/resolvemcp/resolvemcp.go index fc7a996..5ab3620 100644 --- a/pkg/resolvemcp/resolvemcp.go +++ b/pkg/resolvemcp/resolvemcp.go @@ -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) diff --git a/pkg/resolvemcp/security_hooks.go b/pkg/resolvemcp/security_hooks.go index 8af2406..c17fd82 100644 --- a/pkg/resolvemcp/security_hooks.go +++ b/pkg/resolvemcp/security_hooks.go @@ -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()) +} diff --git a/pkg/resolvemcp/tools.go b/pkg/resolvemcp/tools.go index 1e7a44e..f8901bf 100644 --- a/pkg/resolvemcp/tools.go +++ b/pkg/resolvemcp/tools.go @@ -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{}{ diff --git a/pkg/resolvemcp/writecols.go b/pkg/resolvemcp/writecols.go index f6fcf16..fb8e7b7 100644 --- a/pkg/resolvemcp/writecols.go +++ b/pkg/resolvemcp/writecols.go @@ -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 }