mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-03 03:51:59 +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:
@@ -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
@@ -8,6 +8,8 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
"reflect"
|
"reflect"
|
||||||
|
"regexp"
|
||||||
|
"runtime/debug"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
|
|
||||||
@@ -40,7 +42,7 @@ func NewHandler(db common.Database, registry common.ModelRegistry, cfg Config) *
|
|||||||
registry: registry,
|
registry: registry,
|
||||||
hooks: NewHookRegistry(),
|
hooks: NewHookRegistry(),
|
||||||
mcpServer: server.NewMCPServer("resolvemcp", "1.0.0"),
|
mcpServer: server.NewMCPServer("resolvemcp", "1.0.0"),
|
||||||
config: cfg,
|
config: cfg.withDefaults(),
|
||||||
name: "resolvemcp",
|
name: "resolvemcp",
|
||||||
version: "1.0.0",
|
version: "1.0.0",
|
||||||
}
|
}
|
||||||
@@ -244,18 +246,28 @@ var errRecordNotFound = errors.New("record not found")
|
|||||||
// Usage: defer recoverPanic(&returnedErr)
|
// Usage: defer recoverPanic(&returnedErr)
|
||||||
func recoverPanic(err *error) {
|
func recoverPanic(err *error) {
|
||||||
if r := recover(); r != nil {
|
if r := recover(); r != nil {
|
||||||
msg := fmt.Sprintf("%v", r)
|
logger.Error("[resolvemcp] panic recovered: %v\n%s", r, debug.Stack())
|
||||||
logger.Error("[resolvemcp] panic recovered: %s", msg)
|
*err = errInternal
|
||||||
*err = fmt.Errorf("internal error: %s", msg)
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// executeRead reads records from the database and returns raw data + metadata.
|
// 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)
|
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)
|
model, err := h.registry.GetModelByEntity(schema, entity)
|
||||||
if err != nil {
|
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)
|
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
|
var metadata *common.Metadata
|
||||||
err = h.runInTx(ctx, hookCtx, func(common.Database) error {
|
err = h.runInTx(ctx, hookCtx, func(common.Database) error {
|
||||||
var err 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
|
return err
|
||||||
})
|
})
|
||||||
if err != nil {
|
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.
|
// 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))
|
sliceType := reflect.SliceOf(reflect.PointerTo(modelType))
|
||||||
modelPtr := reflect.New(sliceType).Interface()
|
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
|
// expandJoins is empty for resolvemcp — no custom SQL join support yet
|
||||||
cursorFilter, err := getCursorFilter(tableName, pkName, modelColumns, options, nil)
|
cursorFilter, err := getCursorFilter(tableName, pkName, modelColumns, options, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, nil, fmt.Errorf("cursor error: %w", err)
|
return nil, nil, invalidArg("invalid cursor")
|
||||||
}
|
}
|
||||||
|
|
||||||
if cursorFilter != "" {
|
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.
|
// Count — must happen before preloads are applied; Bun panics when counting with relations.
|
||||||
total, err := query.Count(ctx)
|
total := 0
|
||||||
if err != nil {
|
if count {
|
||||||
return nil, nil, fmt.Errorf("error counting records: %w", err)
|
var err error
|
||||||
|
total, err = query.Count(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return nil, nil, fmt.Errorf("error counting records: %w", err)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Pagination
|
// 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.
|
// Preloads — applied after count to avoid Bun panic when counting with relations.
|
||||||
if len(options.Preload) > 0 {
|
if len(options.Preload) > 0 {
|
||||||
|
if err := h.validatePreloads(model, options.Preload); err != nil {
|
||||||
|
return nil, nil, err
|
||||||
|
}
|
||||||
var preloadErr error
|
var preloadErr error
|
||||||
query, preloadErr = h.applyPreloads(model, query, options.Preload)
|
query, preloadErr = h.applyPreloads(model, query, options.Preload)
|
||||||
if preloadErr != nil {
|
if preloadErr != nil {
|
||||||
@@ -461,9 +480,14 @@ func (h *Handler) readInTx(ctx context.Context, hookCtx *HookContext, model inte
|
|||||||
// executeCreate inserts one or more records.
|
// executeCreate inserts one or more records.
|
||||||
func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data interface{}) (_ interface{}, retErr error) {
|
func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data interface{}) (_ interface{}, retErr error) {
|
||||||
defer recoverPanic(&retErr)
|
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)
|
model, err := h.registry.GetModelByEntity(schema, entity)
|
||||||
if err != nil {
|
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)
|
result, err := common.ValidateAndUnwrapModel(model)
|
||||||
@@ -515,12 +539,12 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
|
|||||||
for _, item := range v {
|
for _, item := range v {
|
||||||
itemMap, ok := item.(map[string]interface{})
|
itemMap, ok := item.(map[string]interface{})
|
||||||
if !ok {
|
if !ok {
|
||||||
return fmt.Errorf("each item must be an object")
|
return invalidArg("each item must be an object")
|
||||||
}
|
}
|
||||||
originals = append(originals, itemMap)
|
originals = append(originals, itemMap)
|
||||||
}
|
}
|
||||||
default:
|
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))
|
insertedIDs = make([]interface{}, 0, len(originals))
|
||||||
@@ -530,7 +554,7 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if len(cols) == 0 {
|
if len(cols) == 0 {
|
||||||
return fmt.Errorf("no writable fields in data")
|
return invalidArg("no writable fields in data")
|
||||||
}
|
}
|
||||||
q := tx.NewInsert().Table(tableName)
|
q := tx.NewInsert().Table(tableName)
|
||||||
for key, value := range cols {
|
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.
|
// executeUpdate updates a record by ID.
|
||||||
func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string, data interface{}) (_ interface{}, retErr error) {
|
func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string, data interface{}) (_ interface{}, retErr error) {
|
||||||
defer recoverPanic(&retErr)
|
defer recoverPanic(&retErr)
|
||||||
|
ctx, cancel := h.callContext(ctx)
|
||||||
|
defer cancel()
|
||||||
model, err := h.registry.GetModelByEntity(schema, entity)
|
model, err := h.registry.GetModelByEntity(schema, entity)
|
||||||
if err != nil {
|
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)
|
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{})
|
updates, ok := data.(map[string]interface{})
|
||||||
if !ok {
|
if !ok {
|
||||||
return nil, fmt.Errorf("data must be an object")
|
return nil, invalidArg("data must be an object")
|
||||||
}
|
}
|
||||||
|
|
||||||
if id == "" {
|
if id == "" {
|
||||||
@@ -622,7 +648,7 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
if id == "" {
|
if id == "" {
|
||||||
return nil, fmt.Errorf("update requires an ID")
|
return nil, invalidArg("update requires an id")
|
||||||
}
|
}
|
||||||
|
|
||||||
pkName := reflection.GetPrimaryKeyName(model)
|
pkName := reflection.GetPrimaryKeyName(model)
|
||||||
@@ -663,7 +689,7 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
if len(setCols) == 0 {
|
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
|
// 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.
|
// executeDelete deletes a record by ID.
|
||||||
func (h *Handler) executeDelete(ctx context.Context, schema, entity, id string) (_ interface{}, retErr error) {
|
func (h *Handler) executeDelete(ctx context.Context, schema, entity, id string) (_ interface{}, retErr error) {
|
||||||
defer recoverPanic(&retErr)
|
defer recoverPanic(&retErr)
|
||||||
|
ctx, cancel := h.callContext(ctx)
|
||||||
|
defer cancel()
|
||||||
if id == "" {
|
if id == "" {
|
||||||
return nil, fmt.Errorf("delete requires an ID")
|
return nil, invalidArg("delete requires an id")
|
||||||
}
|
}
|
||||||
|
|
||||||
model, err := h.registry.GetModelByEntity(schema, entity)
|
model, err := h.registry.GetModelByEntity(schema, entity)
|
||||||
if err != nil {
|
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)
|
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)
|
return h.hooks.Execute(OnTxBegin, hookCtx)
|
||||||
}, body)
|
}, 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
@@ -3,6 +3,7 @@ package resolvemcp
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"runtime/debug"
|
||||||
"sync"
|
"sync"
|
||||||
|
|
||||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
"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)
|
logger.Debug("Executing %d resolvemcp hook(s) for %s", len(hooks), hookType)
|
||||||
|
|
||||||
for i, hook := range hooks {
|
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)
|
logger.Error("resolvemcp hook %d for %s failed: %v", i+1, hookType, err)
|
||||||
return fmt.Errorf("hook execution failed: %w", err)
|
return fmt.Errorf("hook execution failed: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if ctx.Abort {
|
if ctx.Abort {
|
||||||
logger.Warn("resolvemcp hook %d for %s requested abort: %s", i+1, hookType, ctx.AbortMessage)
|
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
|
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) {
|
func (r *HookRegistry) Clear(hookType HookType) {
|
||||||
r.mu.Lock()
|
r.mu.Lock()
|
||||||
defer r.mu.Unlock()
|
defer r.mu.Unlock()
|
||||||
|
|||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -18,6 +18,7 @@ package resolvemcp
|
|||||||
import (
|
import (
|
||||||
"net/http"
|
"net/http"
|
||||||
"runtime/debug"
|
"runtime/debug"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/gorilla/mux"
|
"github.com/gorilla/mux"
|
||||||
"github.com/uptrace/bun"
|
"github.com/uptrace/bun"
|
||||||
@@ -41,6 +42,25 @@ type Config struct {
|
|||||||
// If empty, the path is detected from each incoming request automatically.
|
// If empty, the path is detected from each incoming request automatically.
|
||||||
BasePath string
|
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
|
// 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
|
// 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.
|
// host, with at most 32 distinct base URLs cached; prefer setting BaseURL.
|
||||||
@@ -53,6 +73,31 @@ type Config struct {
|
|||||||
EnableAnnotations bool
|
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.
|
// NewHandlerWithGORM creates a Handler backed by a GORM database connection.
|
||||||
func NewHandlerWithGORM(db *gorm.DB, cfg Config) *Handler {
|
func NewHandlerWithGORM(db *gorm.DB, cfg Config) *Handler {
|
||||||
return NewHandler(database.NewGormAdapter(db), modelregistry.NewModelRegistry(), cfg)
|
return NewHandler(database.NewGormAdapter(db), modelregistry.NewModelRegistry(), cfg)
|
||||||
|
|||||||
@@ -31,7 +31,7 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList
|
|||||||
hookCtx.Abort = true
|
hookCtx.Abort = true
|
||||||
hookCtx.AbortMessage = err.Error()
|
hookCtx.AbortMessage = err.Error()
|
||||||
hookCtx.AbortCode = http.StatusUnauthorized
|
hookCtx.AbortCode = http.StatusUnauthorized
|
||||||
return err
|
return NewClientError(CodeForbidden, err.Error())
|
||||||
}
|
}
|
||||||
return nil
|
return nil
|
||||||
})
|
})
|
||||||
@@ -83,17 +83,17 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList
|
|||||||
|
|
||||||
// BeforeCreate: enforce CanCreate rule.
|
// BeforeCreate: enforce CanCreate rule.
|
||||||
handler.Hooks().Register(BeforeCreate, func(hookCtx *HookContext) error {
|
handler.Hooks().Register(BeforeCreate, func(hookCtx *HookContext) error {
|
||||||
return security.CheckModelCreateAllowed(newSecurityContext(hookCtx))
|
return forbidden(security.CheckModelCreateAllowed(newSecurityContext(hookCtx)))
|
||||||
})
|
})
|
||||||
|
|
||||||
// BeforeUpdate: enforce CanUpdate rule.
|
// BeforeUpdate: enforce CanUpdate rule.
|
||||||
handler.Hooks().Register(BeforeUpdate, func(hookCtx *HookContext) error {
|
handler.Hooks().Register(BeforeUpdate, func(hookCtx *HookContext) error {
|
||||||
return security.CheckModelUpdateAllowed(newSecurityContext(hookCtx))
|
return forbidden(security.CheckModelUpdateAllowed(newSecurityContext(hookCtx)))
|
||||||
})
|
})
|
||||||
|
|
||||||
// BeforeDelete: enforce CanDelete rule.
|
// BeforeDelete: enforce CanDelete rule.
|
||||||
handler.Hooks().Register(BeforeDelete, func(hookCtx *HookContext) error {
|
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")
|
logger.Info("Security hooks registered for resolvemcp handler")
|
||||||
@@ -167,3 +167,11 @@ func (s *securityContext) GetResult() interface{} {
|
|||||||
func (s *securityContext) SetResult(result interface{}) {
|
func (s *securityContext) SetResult(result interface{}) {
|
||||||
s.ctx.Result = result
|
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())
|
||||||
|
}
|
||||||
|
|||||||
@@ -324,7 +324,7 @@ func registerReadTool(h *Handler, schema, entity string, info modelInfo) {
|
|||||||
|
|
||||||
data, metadata, err := h.executeRead(ctx, schema, entity, id, options)
|
data, metadata, err := h.executeRead(ctx, schema, entity, id, options)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return mcp.NewToolResultError(err.Error()), nil
|
return toolError("tool", err), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
return marshalResult(map[string]interface{}{
|
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)
|
result, err := h.executeCreate(ctx, schema, entity, data)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return mcp.NewToolResultError(err.Error()), nil
|
return toolError("tool", err), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
return marshalResult(map[string]interface{}{
|
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)
|
result, err := h.executeUpdate(ctx, schema, entity, id, dataMap)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return mcp.NewToolResultError(err.Error()), nil
|
return toolError("tool", err), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
return marshalResult(map[string]interface{}{
|
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)
|
result, err := h.executeDelete(ctx, schema, entity, id)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return mcp.NewToolResultError(err.Error()), nil
|
return toolError("tool", err), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
return marshalResult(map[string]interface{}{
|
return marshalResult(map[string]interface{}{
|
||||||
|
|||||||
@@ -1,7 +1,6 @@
|
|||||||
package resolvemcp
|
package resolvemcp
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
|
||||||
"reflect"
|
"reflect"
|
||||||
"sort"
|
"sort"
|
||||||
"strings"
|
"strings"
|
||||||
@@ -23,7 +22,7 @@ func writeColumns(model interface{}, data map[string]interface{}) (map[string]in
|
|||||||
modelType = modelType.Elem()
|
modelType = modelType.Elem()
|
||||||
}
|
}
|
||||||
if modelType == nil || modelType.Kind() != reflect.Struct {
|
if modelType == nil || modelType.Kind() != reflect.Struct {
|
||||||
return nil, fmt.Errorf("invalid model")
|
return nil, errInternal
|
||||||
}
|
}
|
||||||
|
|
||||||
accepted := make(map[string]string)
|
accepted := make(map[string]string)
|
||||||
@@ -44,13 +43,13 @@ func writeColumns(model interface{}, data map[string]interface{}) (map[string]in
|
|||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
if _, dup := out[col]; dup {
|
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
|
out[col] = value
|
||||||
}
|
}
|
||||||
if len(unknown) > 0 {
|
if len(unknown) > 0 {
|
||||||
sort.Strings(unknown)
|
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
|
return out, nil
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user