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
+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
}