mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-02 11:31:57 +00:00
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:
+103
-22
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user