mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-01 19:20:31 +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"
|
||||
"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
@@ -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()
|
||||
|
||||
@@ -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 (
|
||||
"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)
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
|
||||
@@ -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{}{
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user