feat(resolvemcp): require authentication on MCP endpoints and enforce model rules on writes

Guard() rejects unauthenticated callers (no guest/optional mode); Setup*/New* helpers take a
SecurityList and have explicit *Unauthenticated variants. Model rules now reach the security
hooks, create checks CanCreate (security.CheckModelCreateAllowed), create/update validate keys
against the model's writable columns, update sets only given keys (NULL allowed), update and
delete go through row security via a new BeforeScan hook, and the annotation tool is opt-in
(Config.EnableAnnotations) and runs BeforeHandle.
This commit is contained in:
Hein
2026-10-01 13:31:13 +02:00
parent 7662d5055c
commit ad2f54693f
13 changed files with 643 additions and 104 deletions
+52 -31
View File
@@ -43,7 +43,9 @@ func NewHandler(db common.Database, registry common.ModelRegistry, cfg Config) *
name: "resolvemcp",
version: "1.0.0",
}
registerAnnotationTool(h)
if cfg.EnableAnnotations {
registerAnnotationTool(h)
}
return h
}
@@ -226,7 +228,7 @@ func (h *Handler) executeRead(ctx context.Context, schema, entity, id string, op
model = unwrapped.Model
modelType := unwrapped.ModelType
tableName := h.getTableName(schema, entity, model)
ctx = withRequestData(ctx, schema, entity, tableName, model, unwrapped.ModelPtr)
ctx = withRequestData(h.withModelRules(ctx, schema, entity), schema, entity, tableName, model, unwrapped.ModelPtr)
validator := common.NewColumnValidator(model)
options = validator.FilterRequestOptions(options)
@@ -433,7 +435,7 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
model = result.Model
tableName := h.getTableName(schema, entity, model)
ctx = withRequestData(ctx, schema, entity, tableName, model, result.ModelPtr)
ctx = withRequestData(h.withModelRules(ctx, schema, entity), schema, entity, tableName, model, result.ModelPtr)
hookCtx := &HookContext{
Context: ctx,
@@ -485,8 +487,15 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
insertedIDs = make([]interface{}, 0, len(originals))
for _, itemMap := range originals {
cols, err := writeColumns(model, itemMap)
if err != nil {
return err
}
if len(cols) == 0 {
return fmt.Errorf("no writable fields in data")
}
q := tx.NewInsert().Table(tableName)
for key, value := range itemMap {
for key, value := range cols {
q = q.Value(key, value)
}
if pkName == "" {
@@ -567,7 +576,7 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
model = result.Model
tableName := h.getTableName(schema, entity, model)
ctx = withRequestData(ctx, schema, entity, tableName, model, result.ModelPtr)
ctx = withRequestData(h.withModelRules(ctx, schema, entity), schema, entity, tableName, model, result.ModelPtr)
updates, ok := data.(map[string]interface{})
if !ok {
@@ -602,23 +611,47 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
var updateResult interface{}
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
// Read existing record
if err := h.hooks.Execute(BeforeUpdate, hookCtx); err != nil {
return err
}
if modifiedData, ok := hookCtx.Data.(map[string]interface{}); ok {
updates = modifiedData
}
// SET only the validated incoming keys; the primary key addresses the row, it is not
// rewritten. nil and "" are real values (NULL / empty string).
setCols, err := writeColumns(model, updates)
if err != nil {
return err
}
for col := range setCols {
if strings.EqualFold(col, pkName) {
delete(setCols, col)
}
}
if len(setCols) == 0 {
return fmt.Errorf("no updatable fields in data")
}
// Load the target through the BeforeScan hooks (row security) so a row the caller
// cannot see is reported as not found and never written.
modelType := reflect.TypeOf(model)
if modelType.Kind() == reflect.Pointer {
modelType = modelType.Elem()
}
existingRecord := reflect.New(modelType).Interface()
selectQuery := tx.NewSelect().Model(existingRecord).Column("*").
hookCtx.Query = tx.NewSelect().Model(existingRecord).Column("*").
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id)
if err := selectQuery.ScanModel(ctx); err != nil {
if err := h.hooks.Execute(BeforeScan, hookCtx); err != nil {
return err
}
if err := hookCtx.Query.ScanModel(ctx); err != nil {
if err == sql.ErrNoRows {
return fmt.Errorf("no records found to update")
}
return fmt.Errorf("error fetching existing record: %w", err)
}
// Convert to map
existingMap := make(map[string]interface{})
jsonData, err := json.Marshal(existingRecord)
if err != nil {
@@ -627,26 +660,11 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
if err := json.Unmarshal(jsonData, &existingMap); err != nil {
return fmt.Errorf("error unmarshaling existing record: %w", err)
}
if err := h.hooks.Execute(BeforeUpdate, hookCtx); err != nil {
return err
}
if modifiedData, ok := hookCtx.Data.(map[string]interface{}); ok {
updates = modifiedData
for key, v := range updates {
existingMap[key] = v
}
// Merge non-nil, non-empty values
for key, newValue := range updates {
if newValue == nil {
continue
}
if strVal, ok := newValue.(string); ok && strVal == "" {
continue
}
existingMap[key] = newValue
}
q := tx.NewUpdate().Table(tableName).SetMap(existingMap).
q := tx.NewUpdate().Table(tableName).SetMap(setCols).
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id)
res, err := q.Exec(ctx)
if err != nil {
@@ -711,7 +729,7 @@ func (h *Handler) executeDelete(ctx context.Context, schema, entity, id string)
model = result.Model
tableName := h.getTableName(schema, entity, model)
ctx = withRequestData(ctx, schema, entity, tableName, model, result.ModelPtr)
ctx = withRequestData(h.withModelRules(ctx, schema, entity), schema, entity, tableName, model, result.ModelPtr)
pkName := reflection.GetPrimaryKeyName(model)
@@ -741,9 +759,12 @@ func (h *Handler) executeDelete(ctx context.Context, schema, entity, id string)
return err
}
record := reflect.New(modelType).Interface()
selectQuery := tx.NewSelect().Model(record).
hookCtx.Query = tx.NewSelect().Model(record).
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id)
if err := selectQuery.ScanModel(ctx); err != nil {
if err := h.hooks.Execute(BeforeScan, hookCtx); err != nil {
return err
}
if err := hookCtx.Query.ScanModel(ctx); err != nil {
if err == sql.ErrNoRows {
return fmt.Errorf("record not found")
}