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