mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-01 19:20:31 +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:
@@ -6,6 +6,8 @@ import (
|
||||
"fmt"
|
||||
|
||||
"github.com/mark3labs/mcp-go/mcp"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
)
|
||||
|
||||
const annotationToolName = "resolvespec_annotate"
|
||||
@@ -50,13 +52,42 @@ func registerAnnotationTool(h *Handler) {
|
||||
})
|
||||
}
|
||||
|
||||
// maxAnnotationKey caps tool_name so the key space cannot be abused as storage.
|
||||
const maxAnnotationKey = 200
|
||||
|
||||
// annotationGate runs the BeforeHandle hooks for an annotation call and returns the hook
|
||||
// context. A hook error (e.g. authentication required) is returned to the caller.
|
||||
func annotationGate(ctx context.Context, h *Handler, operation, toolName string) (*HookContext, error) {
|
||||
if len(toolName) > maxAnnotationKey {
|
||||
return nil, fmt.Errorf("tool_name too long")
|
||||
}
|
||||
hookCtx := &HookContext{
|
||||
Context: ctx,
|
||||
Handler: h,
|
||||
Entity: toolName,
|
||||
Operation: operation,
|
||||
Tx: h.db,
|
||||
}
|
||||
if err := h.hooks.Execute(BeforeHandle, hookCtx); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return hookCtx, nil
|
||||
}
|
||||
|
||||
func executeSetAnnotation(ctx context.Context, h *Handler, toolName string, annotations interface{}) (*mcp.CallToolResult, error) {
|
||||
hookCtx, err := annotationGate(ctx, h, "annotate_set", toolName)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
jsonBytes, err := json.Marshal(annotations)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to marshal annotations: %v", err)), nil
|
||||
}
|
||||
|
||||
_, err = h.db.Exec(ctx, "SELECT resolvespec_set_annotation($1, $2)", toolName, string(jsonBytes))
|
||||
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
|
||||
_, err := tx.Exec(ctx, "SELECT resolvespec_set_annotation($1, $2)", toolName, string(jsonBytes))
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to set annotation: %v", err)), nil
|
||||
}
|
||||
@@ -69,8 +100,14 @@ func executeSetAnnotation(ctx context.Context, h *Handler, toolName string, anno
|
||||
}
|
||||
|
||||
func executeGetAnnotation(ctx context.Context, h *Handler, toolName string) (*mcp.CallToolResult, error) {
|
||||
hookCtx, err := annotationGate(ctx, h, "annotate_get", toolName)
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(err.Error()), nil
|
||||
}
|
||||
var rows []map[string]interface{}
|
||||
err := h.db.Query(ctx, &rows, "SELECT resolvespec_get_annotation($1)", toolName)
|
||||
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
|
||||
return tx.Query(ctx, &rows, "SELECT resolvespec_get_annotation($1)", toolName)
|
||||
})
|
||||
if err != nil {
|
||||
return mcp.NewToolResultError(fmt.Sprintf("failed to get annotation: %v", err)), nil
|
||||
}
|
||||
|
||||
@@ -1,6 +1,11 @@
|
||||
package resolvemcp
|
||||
|
||||
import "context"
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||
)
|
||||
|
||||
type contextKey string
|
||||
|
||||
@@ -69,3 +74,18 @@ func withRequestData(ctx context.Context, schema, entity, tableName string, mode
|
||||
ctx = WithModelPtr(ctx, modelPtr)
|
||||
return ctx
|
||||
}
|
||||
|
||||
// withModelRules puts the handler registry's rules for the model into the context, where the
|
||||
// security hooks look them up first. The handler registry is private, so without this the
|
||||
// hooks would not see rules set by RegisterModelWithRules / SetModelRules.
|
||||
func (h *Handler) withModelRules(ctx context.Context, schema, entity string) context.Context {
|
||||
reg, ok := h.registry.(*modelregistry.DefaultModelRegistry)
|
||||
if !ok {
|
||||
return ctx
|
||||
}
|
||||
rules, err := reg.GetModelRules(buildModelName(schema, entity))
|
||||
if err != nil {
|
||||
return ctx
|
||||
}
|
||||
return context.WithValue(ctx, security.ModelRulesKey, rules)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
package resolvemcp
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||
)
|
||||
|
||||
// Guard returns middleware that requires an authenticated caller on every request.
|
||||
//
|
||||
// The security list's provider decides which credentials are accepted: build it from a
|
||||
// security.ChainAuthenticator over an OAuth bearer token, a session token (header or cookie)
|
||||
// and an API key authenticator. The authenticated security.UserContext is placed in the request
|
||||
// context, which the MCP transports pass on to every tool call, so rules, row security and
|
||||
// OnTxBegin apply to that caller.
|
||||
//
|
||||
// Unlike security.NewAuthMiddleware this guard has no guest or optional mode: it ignores
|
||||
// security.SkipAuth / security.OptionalAuth markers on the request context, and fails closed
|
||||
// (500) when no provider is configured.
|
||||
func Guard(securityList *security.SecurityList) func(http.Handler) http.Handler {
|
||||
return func(next http.Handler) http.Handler {
|
||||
authed := security.NewAuthHandler(securityList, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if uc, ok := security.GetUserContext(r.Context()); !ok || uc == nil {
|
||||
http.Error(w, "Authentication failed", http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
next.ServeHTTP(w, r)
|
||||
}))
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if securityList == nil {
|
||||
http.Error(w, "Security provider not configured", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
authed.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// requireGuard reports whether securityList can guard a route. Setup helpers use it to refuse
|
||||
// to mount an endpoint rather than serve it unauthenticated by mistake.
|
||||
func requireGuard(fn string, securityList *security.SecurityList) bool {
|
||||
if securityList == nil || securityList.Provider() == nil {
|
||||
logger.Error("resolvemcp.%s: no security provider configured; MCP endpoint NOT mounted", fn)
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func warnUnauthenticated(fn string) {
|
||||
logger.Warn("resolvemcp.%s: serving the MCP endpoint WITHOUT authentication; every caller can read and write all registered models", fn)
|
||||
}
|
||||
@@ -0,0 +1,97 @@
|
||||
package resolvemcp
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/providers"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security/sectypes"
|
||||
)
|
||||
|
||||
// tokenAuth accepts the bearer token "good" and nothing else.
|
||||
type tokenAuth struct{ security.Authenticator }
|
||||
|
||||
func (tokenAuth) Authenticate(r *http.Request) (*security.UserContext, error) {
|
||||
if r.Header.Get("Authorization") != "Bearer good" {
|
||||
return nil, errors.New("bad credentials")
|
||||
}
|
||||
return &security.UserContext{UserID: 7, UserName: "kim"}, nil
|
||||
}
|
||||
|
||||
func newTestSecurityList(t *testing.T) *security.SecurityList {
|
||||
t.Helper()
|
||||
p, err := security.NewCompositeSecurityProvider(tokenAuth{},
|
||||
providers.NewConfigColumnSecurityProvider(map[string][]sectypes.ColumnSecurity{}),
|
||||
providers.NewConfigRowSecurityProvider(nil, nil))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
sl, err := security.NewSecurityList(p)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return sl
|
||||
}
|
||||
|
||||
func serve(h http.Handler, auth string, mark func(*http.Request) *http.Request) *httptest.ResponseRecorder {
|
||||
r := httptest.NewRequest(http.MethodGet, "/mcp", nil)
|
||||
if auth != "" {
|
||||
r.Header.Set("Authorization", auth)
|
||||
}
|
||||
if mark != nil {
|
||||
r = mark(r)
|
||||
}
|
||||
w := httptest.NewRecorder()
|
||||
h.ServeHTTP(w, r)
|
||||
return w
|
||||
}
|
||||
|
||||
func TestGuardRejectsUnauthenticated(t *testing.T) {
|
||||
var gotUser int
|
||||
next := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
uc, ok := security.GetUserContext(r.Context())
|
||||
if !ok {
|
||||
t.Error("user context missing downstream")
|
||||
return
|
||||
}
|
||||
gotUser = uc.UserID
|
||||
})
|
||||
g := Guard(newTestSecurityList(t))(next)
|
||||
|
||||
for name, auth := range map[string]string{"none": "", "wrong": "Bearer bad"} {
|
||||
if w := serve(g, auth, nil); w.Code != http.StatusUnauthorized {
|
||||
t.Errorf("%s: status %d, want 401", name, w.Code)
|
||||
}
|
||||
}
|
||||
if w := serve(g, "Bearer good", nil); w.Code != http.StatusOK || gotUser != 7 {
|
||||
t.Errorf("good: status %d user %d, want 200 / 7", w.Code, gotUser)
|
||||
}
|
||||
}
|
||||
|
||||
// Skip/optional markers on the request context must not open the MCP endpoint.
|
||||
func TestGuardIgnoresSkipAndOptionalMarkers(t *testing.T) {
|
||||
called := false
|
||||
g := Guard(newTestSecurityList(t))(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { called = true }))
|
||||
for name, mark := range map[string]func(*http.Request) *http.Request{
|
||||
"skip": func(r *http.Request) *http.Request { return r.WithContext(security.SkipAuth(r.Context())) },
|
||||
"optional": func(r *http.Request) *http.Request { return r.WithContext(security.OptionalAuth(r.Context())) },
|
||||
} {
|
||||
if w := serve(g, "", mark); w.Code != http.StatusUnauthorized || called {
|
||||
t.Errorf("%s: status %d called=%v, want 401 and not called", name, w.Code, called)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestGuardFailsClosedWithoutProvider(t *testing.T) {
|
||||
called := false
|
||||
g := Guard(nil)(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { called = true }))
|
||||
if w := serve(g, "Bearer good", nil); w.Code != http.StatusInternalServerError || called {
|
||||
t.Errorf("status %d called=%v, want 500 and not called", w.Code, called)
|
||||
}
|
||||
if requireGuard("test", nil) {
|
||||
t.Error("requireGuard(nil) must be false")
|
||||
}
|
||||
}
|
||||
+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")
|
||||
}
|
||||
|
||||
@@ -21,6 +21,11 @@ const (
|
||||
BeforeCreate HookType = "before_create"
|
||||
AfterCreate HookType = "after_create"
|
||||
|
||||
// BeforeScan fires on update and delete, with hookCtx.Query set to the select that loads the
|
||||
// target row. Hooks that narrow the query (row security) run here; a row the query does
|
||||
// not return is reported as not found and never written.
|
||||
BeforeScan HookType = "before_scan"
|
||||
|
||||
BeforeUpdate HookType = "before_update"
|
||||
AfterUpdate HookType = "after_update"
|
||||
|
||||
|
||||
@@ -114,25 +114,12 @@ func (h *Handler) mountOAuth2Routes(mux *http.ServeMux) {
|
||||
// context into the request context, making it available to BeforeHandle security hooks.
|
||||
// Unauthenticated requests receive 401 before reaching any MCP tool.
|
||||
func (h *Handler) AuthedSSEServer(securityList *security.SecurityList) http.Handler {
|
||||
return security.NewAuthMiddleware(securityList)(h.SSEServer())
|
||||
}
|
||||
|
||||
// OptionalAuthSSEServer wraps SSEServer with optional authentication middleware.
|
||||
// Unauthenticated requests continue as guest rather than returning 401.
|
||||
// Use together with RegisterSecurityHooks and per-model CanPublicRead/Write rules
|
||||
// to allow mixed public/private access.
|
||||
func (h *Handler) OptionalAuthSSEServer(securityList *security.SecurityList) http.Handler {
|
||||
return security.NewOptionalAuthMiddleware(securityList)(h.SSEServer())
|
||||
return Guard(securityList)(h.SSEServer())
|
||||
}
|
||||
|
||||
// AuthedStreamableHTTPServer wraps StreamableHTTPServer with required authentication middleware.
|
||||
func (h *Handler) AuthedStreamableHTTPServer(securityList *security.SecurityList) http.Handler {
|
||||
return security.NewAuthMiddleware(securityList)(h.StreamableHTTPServer())
|
||||
}
|
||||
|
||||
// OptionalAuthStreamableHTTPServer wraps StreamableHTTPServer with optional authentication middleware.
|
||||
func (h *Handler) OptionalAuthStreamableHTTPServer(securityList *security.SecurityList) http.Handler {
|
||||
return security.NewOptionalAuthMiddleware(securityList)(h.StreamableHTTPServer())
|
||||
return Guard(securityList)(h.StreamableHTTPServer())
|
||||
}
|
||||
|
||||
// --------------------------------------------------------------------------
|
||||
@@ -243,22 +230,3 @@ func SetupMuxOAuth2Routes(muxRouter *mux.Router, auth *security.DatabaseAuthenti
|
||||
OAuth2CallbackHandler(auth, cfg.ProviderName, cfg.AfterLoginRedirect, cookieOpts...),
|
||||
).Methods(http.MethodGet)
|
||||
}
|
||||
|
||||
// SetupMuxRoutesWithAuth mounts the MCP SSE endpoints on a Gorilla Mux router
|
||||
// with required authentication middleware applied.
|
||||
func SetupMuxRoutesWithAuth(muxRouter *mux.Router, handler *Handler, securityList *security.SecurityList) {
|
||||
basePath := handler.config.BasePath
|
||||
h := handler.AuthedSSEServer(securityList)
|
||||
|
||||
muxRouter.Handle(basePath+"/sse", h).Methods(http.MethodGet, http.MethodOptions)
|
||||
muxRouter.Handle(basePath+"/message", h).Methods(http.MethodPost, http.MethodOptions)
|
||||
muxRouter.PathPrefix(basePath).Handler(http.StripPrefix(basePath, h))
|
||||
}
|
||||
|
||||
// SetupMuxStreamableHTTPRoutesWithAuth mounts the MCP streamable HTTP endpoint on a
|
||||
// Gorilla Mux router with required authentication middleware applied.
|
||||
func SetupMuxStreamableHTTPRoutesWithAuth(muxRouter *mux.Router, handler *Handler, securityList *security.SecurityList) {
|
||||
basePath := handler.config.BasePath
|
||||
h := handler.AuthedStreamableHTTPServer(securityList)
|
||||
muxRouter.PathPrefix(basePath).Handler(http.StripPrefix(basePath, h))
|
||||
}
|
||||
|
||||
@@ -12,7 +12,7 @@
|
||||
// handler.RegisterModel("public", "users", &User{})
|
||||
//
|
||||
// r := mux.NewRouter()
|
||||
// resolvemcp.SetupMuxRoutes(r, handler)
|
||||
// resolvemcp.SetupMuxRoutes(r, handler, securityList) // requires an authenticated caller
|
||||
package resolvemcp
|
||||
|
||||
import (
|
||||
@@ -28,6 +28,7 @@ import (
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||
)
|
||||
|
||||
// Config holds configuration for the resolvemcp handler.
|
||||
@@ -39,6 +40,12 @@ type Config struct {
|
||||
// BasePath is the URL path prefix where the MCP endpoints are mounted (e.g. "/mcp").
|
||||
// If empty, the path is detected from each incoming request automatically.
|
||||
BasePath string
|
||||
|
||||
// EnableAnnotations registers the resolvespec_annotate tool. Off by default: annotations
|
||||
// are free text that agents read back, so enabling the tool opens a write channel into
|
||||
// agent-visible text. When on, every call runs the BeforeHandle hooks (operation
|
||||
// "annotate_set" / "annotate_get") and the writes run in a transaction with OnTxBegin.
|
||||
EnableAnnotations bool
|
||||
}
|
||||
|
||||
// NewHandlerWithGORM creates a Handler backed by a GORM database connection.
|
||||
@@ -57,18 +64,29 @@ func NewHandlerWithDB(db common.Database, cfg Config) *Handler {
|
||||
}
|
||||
|
||||
// SetupMuxRoutes mounts the MCP HTTP/SSE endpoints on the given Gorilla Mux router
|
||||
// using the base path from Config.BasePath (falls back to "/mcp" if empty).
|
||||
// using the base path from Config.BasePath, behind Guard(securityList).
|
||||
//
|
||||
// Two routes are registered:
|
||||
// Routes registered:
|
||||
// - GET {basePath}/sse — SSE connection endpoint (client subscribes here)
|
||||
// - POST {basePath}/message — JSON-RPC message endpoint (client sends requests here)
|
||||
//
|
||||
// To protect these routes with authentication, wrap the mux router or apply middleware
|
||||
// before calling SetupMuxRoutes.
|
||||
func SetupMuxRoutes(muxRouter *mux.Router, handler *Handler) {
|
||||
basePath := handler.config.BasePath
|
||||
h := handler.SSEServer()
|
||||
// Nothing is mounted (and an error is logged) when securityList has no provider.
|
||||
func SetupMuxRoutes(muxRouter *mux.Router, handler *Handler, securityList *security.SecurityList) {
|
||||
if !requireGuard("SetupMuxRoutes", securityList) {
|
||||
return
|
||||
}
|
||||
mountMuxSSE(muxRouter, handler, handler.AuthedSSEServer(securityList))
|
||||
}
|
||||
|
||||
// SetupMuxRoutesUnauthenticated is SetupMuxRoutes without the guard. Every caller reaches every
|
||||
// registered model, so use it only behind another trusted layer. A warning is logged.
|
||||
func SetupMuxRoutesUnauthenticated(muxRouter *mux.Router, handler *Handler) {
|
||||
warnUnauthenticated("SetupMuxRoutesUnauthenticated")
|
||||
mountMuxSSE(muxRouter, handler, handler.SSEServer())
|
||||
}
|
||||
|
||||
func mountMuxSSE(muxRouter *mux.Router, handler *Handler, h http.Handler) {
|
||||
basePath := handler.config.BasePath
|
||||
muxRouter.Handle(basePath+"/sse", h).Methods("GET", "OPTIONS")
|
||||
muxRouter.Handle(basePath+"/message", h).Methods("POST", "OPTIONS")
|
||||
|
||||
@@ -78,21 +96,32 @@ func SetupMuxRoutes(muxRouter *mux.Router, handler *Handler) {
|
||||
}
|
||||
|
||||
// SetupBunRouterRoutes mounts the MCP HTTP/SSE endpoints on a bunrouter router
|
||||
// using the base path from Config.BasePath.
|
||||
// using the base path from Config.BasePath, behind Guard(securityList).
|
||||
//
|
||||
// Two routes are registered:
|
||||
// Routes registered:
|
||||
// - GET {basePath}/sse — SSE connection endpoint
|
||||
// - POST {basePath}/message — JSON-RPC message endpoint
|
||||
func SetupBunRouterRoutes(router *bunrouter.Router, handler *Handler) {
|
||||
func SetupBunRouterRoutes(router *bunrouter.Router, handler *Handler, securityList *security.SecurityList) {
|
||||
if !requireGuard("SetupBunRouterRoutes", securityList) {
|
||||
return
|
||||
}
|
||||
mountBunSSE(router, handler, handler.AuthedSSEServer(securityList))
|
||||
}
|
||||
|
||||
// SetupBunRouterRoutesUnauthenticated is SetupBunRouterRoutes without the guard. A warning is logged.
|
||||
func SetupBunRouterRoutesUnauthenticated(router *bunrouter.Router, handler *Handler) {
|
||||
warnUnauthenticated("SetupBunRouterRoutesUnauthenticated")
|
||||
mountBunSSE(router, handler, handler.SSEServer())
|
||||
}
|
||||
|
||||
func mountBunSSE(router *bunrouter.Router, handler *Handler, h http.Handler) {
|
||||
defer func() {
|
||||
if rec := recover(); rec != nil {
|
||||
logger.Error("panic in resolvemcp.SetupBunRouterRoutes: %v\n%s", rec, debug.Stack())
|
||||
logger.Error("panic mounting resolvemcp bunrouter routes: %v\n%s", rec, debug.Stack())
|
||||
}
|
||||
}()
|
||||
|
||||
basePath := handler.config.BasePath
|
||||
h := handler.SSEServer()
|
||||
|
||||
router.GET(basePath+"/sse", bunrouter.HTTPHandler(h))
|
||||
logger.Info("Registered resolvemcp bunrouter route GET %s/sse", basePath)
|
||||
|
||||
@@ -100,45 +129,68 @@ func SetupBunRouterRoutes(router *bunrouter.Router, handler *Handler) {
|
||||
logger.Info("Registered resolvemcp bunrouter route POST %s/message", basePath)
|
||||
}
|
||||
|
||||
// NewSSEServer returns an http.Handler that serves MCP over SSE.
|
||||
// NewSSEServer returns an http.Handler that serves MCP over SSE behind Guard(securityList).
|
||||
// If Config.BasePath is set it is used directly; otherwise the base path is
|
||||
// detected from each incoming request (by stripping the "/sse" or "/message" suffix).
|
||||
//
|
||||
// h := resolvemcp.NewSSEServer(handler)
|
||||
// h := resolvemcp.NewSSEServer(handler, securityList)
|
||||
// http.Handle("/api/mcp/", h)
|
||||
func NewSSEServer(handler *Handler) http.Handler {
|
||||
return handler.SSEServer()
|
||||
func NewSSEServer(handler *Handler, securityList *security.SecurityList) http.Handler {
|
||||
return handler.AuthedSSEServer(securityList)
|
||||
}
|
||||
|
||||
// SetupMuxStreamableHTTPRoutes mounts the MCP streamable HTTP endpoint on the given Gorilla Mux router.
|
||||
// The streamable HTTP transport uses a single endpoint (Config.BasePath) for all communication:
|
||||
// POST for client→server messages, GET for server→client streaming.
|
||||
// SetupMuxStreamableHTTPRoutes mounts the MCP streamable HTTP endpoint on the given Gorilla Mux
|
||||
// router, behind Guard(securityList). The streamable HTTP transport uses a single endpoint
|
||||
// (Config.BasePath) for all communication: POST for client→server messages, GET for
|
||||
// server→client streaming.
|
||||
//
|
||||
// Example:
|
||||
//
|
||||
// resolvemcp.SetupMuxStreamableHTTPRoutes(r, handler) // mounts at Config.BasePath
|
||||
func SetupMuxStreamableHTTPRoutes(muxRouter *mux.Router, handler *Handler) {
|
||||
// Nothing is mounted (and an error is logged) when securityList has no provider.
|
||||
func SetupMuxStreamableHTTPRoutes(muxRouter *mux.Router, handler *Handler, securityList *security.SecurityList) {
|
||||
if !requireGuard("SetupMuxStreamableHTTPRoutes", securityList) {
|
||||
return
|
||||
}
|
||||
basePath := handler.config.BasePath
|
||||
h := handler.StreamableHTTPServer()
|
||||
muxRouter.PathPrefix(basePath).Handler(http.StripPrefix(basePath, h))
|
||||
muxRouter.PathPrefix(basePath).Handler(http.StripPrefix(basePath, handler.AuthedStreamableHTTPServer(securityList)))
|
||||
}
|
||||
|
||||
// SetupBunRouterStreamableHTTPRoutes mounts the MCP streamable HTTP endpoint on a bunrouter router.
|
||||
// The streamable HTTP transport uses a single endpoint (Config.BasePath).
|
||||
func SetupBunRouterStreamableHTTPRoutes(router *bunrouter.Router, handler *Handler) {
|
||||
// SetupMuxStreamableHTTPRoutesUnauthenticated is SetupMuxStreamableHTTPRoutes without the guard.
|
||||
// A warning is logged.
|
||||
func SetupMuxStreamableHTTPRoutesUnauthenticated(muxRouter *mux.Router, handler *Handler) {
|
||||
warnUnauthenticated("SetupMuxStreamableHTTPRoutesUnauthenticated")
|
||||
basePath := handler.config.BasePath
|
||||
muxRouter.PathPrefix(basePath).Handler(http.StripPrefix(basePath, handler.StreamableHTTPServer()))
|
||||
}
|
||||
|
||||
// SetupBunRouterStreamableHTTPRoutes mounts the MCP streamable HTTP endpoint on a bunrouter
|
||||
// router, behind Guard(securityList). The transport uses a single endpoint (Config.BasePath).
|
||||
func SetupBunRouterStreamableHTTPRoutes(router *bunrouter.Router, handler *Handler, securityList *security.SecurityList) {
|
||||
if !requireGuard("SetupBunRouterStreamableHTTPRoutes", securityList) {
|
||||
return
|
||||
}
|
||||
mountBunStreamable(router, handler, handler.AuthedStreamableHTTPServer(securityList))
|
||||
}
|
||||
|
||||
// SetupBunRouterStreamableHTTPRoutesUnauthenticated is SetupBunRouterStreamableHTTPRoutes
|
||||
// without the guard. A warning is logged.
|
||||
func SetupBunRouterStreamableHTTPRoutesUnauthenticated(router *bunrouter.Router, handler *Handler) {
|
||||
warnUnauthenticated("SetupBunRouterStreamableHTTPRoutesUnauthenticated")
|
||||
mountBunStreamable(router, handler, handler.StreamableHTTPServer())
|
||||
}
|
||||
|
||||
func mountBunStreamable(router *bunrouter.Router, handler *Handler, h http.Handler) {
|
||||
basePath := handler.config.BasePath
|
||||
h := handler.StreamableHTTPServer()
|
||||
router.GET(basePath, bunrouter.HTTPHandler(h))
|
||||
router.POST(basePath, bunrouter.HTTPHandler(h))
|
||||
router.DELETE(basePath, bunrouter.HTTPHandler(h))
|
||||
}
|
||||
|
||||
// NewStreamableHTTPHandler returns an http.Handler that serves MCP over the streamable HTTP transport.
|
||||
// Mount it at the desired path; that path becomes the MCP endpoint.
|
||||
// NewStreamableHTTPHandler returns an http.Handler that serves MCP over the streamable HTTP
|
||||
// transport behind Guard(securityList). Mount it at the desired path; that path becomes the
|
||||
// MCP endpoint.
|
||||
//
|
||||
// h := resolvemcp.NewStreamableHTTPHandler(handler)
|
||||
// h := resolvemcp.NewStreamableHTTPHandler(handler, securityList)
|
||||
// http.Handle("/mcp", h)
|
||||
// engine.Any("/mcp", gin.WrapH(h))
|
||||
func NewStreamableHTTPHandler(handler *Handler) http.Handler {
|
||||
return handler.StreamableHTTPServer()
|
||||
func NewStreamableHTTPHandler(handler *Handler, securityList *security.SecurityList) http.Handler {
|
||||
return handler.AuthedStreamableHTTPServer(securityList)
|
||||
}
|
||||
|
||||
@@ -62,6 +62,15 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList
|
||||
return security.ApplyRowSecurity(newSecurityContext(hookCtx), securityList)
|
||||
})
|
||||
|
||||
// BeforeScan: row-level security on the row an update or delete targets. A row the user
|
||||
// cannot see is "not found" and is never written.
|
||||
handler.Hooks().Register(BeforeScan, func(hookCtx *HookContext) error {
|
||||
if err := security.LoadSecurityRules(newSecurityContext(hookCtx), securityList); err != nil {
|
||||
return err
|
||||
}
|
||||
return security.ApplyRowSecurity(newSecurityContext(hookCtx), securityList)
|
||||
})
|
||||
|
||||
// AfterRead (1st): apply column-level security — mask/hide columns in the result.
|
||||
handler.Hooks().Register(AfterRead, func(hookCtx *HookContext) error {
|
||||
return security.ApplyColumnSecurity(newSecurityContext(hookCtx), securityList)
|
||||
@@ -72,6 +81,11 @@ func RegisterSecurityHooks(handler *Handler, securityList *security.SecurityList
|
||||
return security.LogDataAccess(newSecurityContext(hookCtx))
|
||||
})
|
||||
|
||||
// BeforeCreate: enforce CanCreate rule.
|
||||
handler.Hooks().Register(BeforeCreate, func(hookCtx *HookContext) error {
|
||||
return security.CheckModelCreateAllowed(newSecurityContext(hookCtx))
|
||||
})
|
||||
|
||||
// BeforeUpdate: enforce CanUpdate rule.
|
||||
handler.Hooks().Register(BeforeUpdate, func(hookCtx *HookContext) error {
|
||||
return security.CheckModelUpdateAllowed(newSecurityContext(hookCtx))
|
||||
|
||||
@@ -0,0 +1,200 @@
|
||||
package resolvemcp
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/DATA-DOG/go-sqlmock"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||
)
|
||||
|
||||
type wItem struct {
|
||||
ID int `json:"id" bun:"id,pk"`
|
||||
Name string `json:"name" bun:"name"`
|
||||
FullName string `json:"fullName" bun:"full_name"`
|
||||
Note *string `json:"note" bun:"note"`
|
||||
Owner int `json:"-" bun:"-"`
|
||||
}
|
||||
|
||||
func TestWriteColumns(t *testing.T) {
|
||||
m := &wItem{}
|
||||
got, err := writeColumns(m, map[string]interface{}{"fullName": "a", "NAME": "b", "note": nil})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got["full_name"] != "a" || got["name"] != "b" {
|
||||
t.Errorf("json/column names must resolve to columns: %v", got)
|
||||
}
|
||||
if v, ok := got["note"]; !ok || v != nil {
|
||||
t.Errorf("explicit null must be kept: %v", got)
|
||||
}
|
||||
for name, data := range map[string]map[string]interface{}{
|
||||
"unknown": {"nope": 1},
|
||||
"injection": {"name = 'x', id": 1},
|
||||
"unmapped": {"owner": 1},
|
||||
"both forms": {"fullName": 1, "full_name": 2},
|
||||
"empty key": {"": 1},
|
||||
"quoted char": {`"name"`: 1},
|
||||
} {
|
||||
if _, err := writeColumns(m, data); err == nil {
|
||||
t.Errorf("%s: expected error", name)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateRejectsUnknownKeys(t *testing.T) {
|
||||
h, mock, ctx := newTxHarness(t)
|
||||
mock.ExpectBegin()
|
||||
mock.ExpectRollback()
|
||||
_, err := h.executeCreate(ctx, "public", "items", map[string]interface{}{"name": "a", "is_admin": true})
|
||||
if err == nil || !strings.Contains(err.Error(), "is_admin") {
|
||||
t.Fatalf("want unknown field error, got %v", err)
|
||||
}
|
||||
if err := mock.ExpectationsWereMet(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateSetsOnlyGivenKeysAndAllowsNull(t *testing.T) {
|
||||
db := wHarness(t)
|
||||
h, mock, ctx := db.h, db.mock, db.ctx
|
||||
cols := []string{"id", "name", "full_name", "note"}
|
||||
mock.ExpectBegin()
|
||||
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a", "A", "n"))
|
||||
// Only "note" is set (to NULL); the id in the payload addresses the row and is not rewritten.
|
||||
mock.ExpectExec(`UPDATE .* SET "?note"? = \$1 WHERE`).WithArgs(nil, "7").WillReturnResult(sqlmock.NewResult(0, 1))
|
||||
mock.ExpectCommit()
|
||||
mock.ExpectBegin()
|
||||
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a", "A", nil))
|
||||
mock.ExpectCommit()
|
||||
|
||||
if _, err := h.executeUpdate(ctx, "public", "witems", "7", map[string]interface{}{"id": 7, "note": nil}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := mock.ExpectationsWereMet(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateRejectsUnknownKeys(t *testing.T) {
|
||||
db := wHarness(t)
|
||||
db.mock.ExpectBegin()
|
||||
db.mock.ExpectRollback()
|
||||
if _, err := db.h.executeUpdate(db.ctx, "public", "witems", "7", map[string]interface{}{"role": "admin"}); err == nil {
|
||||
t.Fatal("expected error")
|
||||
}
|
||||
if err := db.mock.ExpectationsWereMet(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
type wh struct {
|
||||
h *Handler
|
||||
mock sqlmock.Sqlmock
|
||||
ctx context.Context
|
||||
}
|
||||
|
||||
// wHarness is newTxHarness with a model that has more than id/name.
|
||||
func wHarness(t *testing.T) wh {
|
||||
t.Helper()
|
||||
h, mock, ctx := newTxHarness(t)
|
||||
if err := h.RegisterModel("public", "witems", &wItem{}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return wh{h, mock, ctx}
|
||||
}
|
||||
|
||||
// rlsProvider returns a fixed row security template.
|
||||
type rlsProvider struct{ stubProvider }
|
||||
|
||||
func (rlsProvider) GetRowSecurity(_ context.Context, userRef any, schema, table string) (security.RowSecurity, error) {
|
||||
return security.RowSecurity{Schema: schema, Tablename: table, Template: "owner_id = {UserID}", UserID: userRef}, nil
|
||||
}
|
||||
|
||||
func securedHandler(t *testing.T, prov security.SecurityProvider) (*Handler, sqlmock.Sqlmock, context.Context) {
|
||||
t.Helper()
|
||||
h, mock, ctx := newTxHarness(t)
|
||||
list, err := security.NewSecurityList(prov)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
RegisterSecurityHooks(h, list)
|
||||
ctx = context.WithValue(ctx, security.UserContextKey, &security.UserContext{UserID: 7, UserName: "u"})
|
||||
ctx = context.WithValue(ctx, security.UserIDKey, 7)
|
||||
return h, mock, ctx
|
||||
}
|
||||
|
||||
// A row hidden by row security is "not found" for update and delete, and nothing is written.
|
||||
func TestWritesHonourRowSecurity(t *testing.T) {
|
||||
cols := []string{"id", "name"}
|
||||
t.Run("update", func(t *testing.T) {
|
||||
h, mock, ctx := securedHandler(t, rlsProvider{})
|
||||
mock.ExpectBegin()
|
||||
mock.ExpectQuery(`SELECT .*owner_id`).WillReturnRows(sqlmock.NewRows(cols))
|
||||
mock.ExpectRollback()
|
||||
if _, err := h.executeUpdate(ctx, "public", "items", "7", map[string]interface{}{"name": "x"}); err == nil {
|
||||
t.Fatal("expected not found")
|
||||
}
|
||||
if err := mock.ExpectationsWereMet(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
})
|
||||
t.Run("delete", func(t *testing.T) {
|
||||
h, mock, ctx := securedHandler(t, rlsProvider{})
|
||||
mock.ExpectBegin()
|
||||
mock.ExpectQuery(`SELECT .*owner_id`).WillReturnRows(sqlmock.NewRows(cols))
|
||||
mock.ExpectRollback()
|
||||
if _, err := h.executeDelete(ctx, "public", "items", "7"); err == nil {
|
||||
t.Fatal("expected not found")
|
||||
}
|
||||
if err := mock.ExpectationsWereMet(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// Rules set with RegisterModelWithRules must reach the security hooks.
|
||||
func TestModelRulesReachHooks(t *testing.T) {
|
||||
h, mock, ctx := securedHandler(t, stubProvider{})
|
||||
if err := h.RegisterModelWithRules("public", "locked", &txItem{}, modelregistry.ModelRules{CanRead: true}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
for name, op := range map[string]func() error{
|
||||
"create": func() error {
|
||||
_, err := h.executeCreate(ctx, "public", "locked", map[string]interface{}{"name": "a"})
|
||||
return err
|
||||
},
|
||||
"update": func() error {
|
||||
_, err := h.executeUpdate(ctx, "public", "locked", "7", map[string]interface{}{"name": "a"})
|
||||
return err
|
||||
},
|
||||
"delete": func() error {
|
||||
_, err := h.executeDelete(ctx, "public", "locked", "7")
|
||||
return err
|
||||
},
|
||||
} {
|
||||
// Each denies inside its transaction, before any statement.
|
||||
mock.ExpectBegin()
|
||||
mock.ExpectRollback()
|
||||
if err := op(); err == nil || !strings.Contains(err.Error(), "not allowed") {
|
||||
t.Errorf("%s: want 'not allowed', got %v", name, err)
|
||||
}
|
||||
}
|
||||
if err := mock.ExpectationsWereMet(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAnnotationToolIsOptIn(t *testing.T) {
|
||||
h, _, _ := newTxHarness(t)
|
||||
if h.mcpServer.GetTool(annotationToolName) != nil {
|
||||
t.Fatal("annotation tool must be off by default")
|
||||
}
|
||||
on := NewHandler(h.db, modelregistry.NewModelRegistry(), Config{EnableAnnotations: true})
|
||||
if on.mcpServer.GetTool(annotationToolName) == nil {
|
||||
t.Fatal("annotation tool missing when enabled")
|
||||
}
|
||||
}
|
||||
@@ -297,6 +297,10 @@ func (stubProvider) GetColumnSecurity(context.Context, int, string, string) ([]s
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (stubProvider) GetRowSecurity(context.Context, any, string, string) (security.RowSecurity, error) {
|
||||
return security.RowSecurity{}, nil
|
||||
}
|
||||
|
||||
func TestSecurityHooksStampTxSettingsOnEveryTransaction(t *testing.T) {
|
||||
h, mock, ctx := newTxHarness(t)
|
||||
list, err := security.NewSecurityList(stubProvider{})
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
package resolvemcp
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"reflect"
|
||||
"sort"
|
||||
"strings"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||
)
|
||||
|
||||
// maxKeyEcho caps how much of a rejected client key is echoed back in an error.
|
||||
const maxKeyEcho = 64
|
||||
|
||||
// writeColumns validates the keys of a create/update payload against the model and returns the
|
||||
// values keyed by database column name. A key may be the json name or the column name of a
|
||||
// writable field (case-insensitive); relations, scan-only and unexported fields are not
|
||||
// writable. Unknown keys are rejected rather than dropped, so a client learns that a write
|
||||
// did not take effect, and no client-chosen identifier reaches SQL.
|
||||
func writeColumns(model interface{}, data map[string]interface{}) (map[string]interface{}, error) {
|
||||
modelType := reflect.TypeOf(model)
|
||||
for modelType != nil && (modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice) {
|
||||
modelType = modelType.Elem()
|
||||
}
|
||||
if modelType == nil || modelType.Kind() != reflect.Struct {
|
||||
return nil, fmt.Errorf("invalid model")
|
||||
}
|
||||
|
||||
accepted := make(map[string]string)
|
||||
for jsonKey, col := range reflection.BuildJSONToDBColumnMap(modelType) {
|
||||
accepted[strings.ToLower(jsonKey)] = col
|
||||
accepted[strings.ToLower(col)] = col
|
||||
}
|
||||
|
||||
out := make(map[string]interface{}, len(data))
|
||||
var unknown []string
|
||||
for key, value := range data {
|
||||
col, ok := accepted[strings.ToLower(key)]
|
||||
if !ok {
|
||||
if len(key) > maxKeyEcho {
|
||||
key = key[:maxKeyEcho] + "..."
|
||||
}
|
||||
unknown = append(unknown, key)
|
||||
continue
|
||||
}
|
||||
if _, dup := out[col]; dup {
|
||||
return nil, fmt.Errorf("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 out, nil
|
||||
}
|
||||
@@ -441,6 +441,19 @@ func resolveModelRules(secCtx SecurityContext) (modelregistry.ModelRules, bool)
|
||||
return rules, true
|
||||
}
|
||||
|
||||
// CheckModelCreateAllowed returns an error if CanCreate is false for the model. Rules are read
|
||||
// from context with a fallback to the model registry; an unregistered model is allowed.
|
||||
func CheckModelCreateAllowed(secCtx SecurityContext) error {
|
||||
rules, ok := resolveModelRules(secCtx)
|
||||
if !ok {
|
||||
return nil // model not registered, allow by default
|
||||
}
|
||||
if !rules.CanCreate {
|
||||
return fmt.Errorf("create not allowed for %s", secCtx.GetEntity())
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// CheckModelUpdateAllowed is the public wrapper for checkModelUpdateAllowed.
|
||||
func CheckModelUpdateAllowed(secCtx SecurityContext) error {
|
||||
return checkModelUpdateAllowed(secCtx)
|
||||
|
||||
Reference in New Issue
Block a user