From ad2f54693fac6aa34407ffa57173ee9f8bbb01cf Mon Sep 17 00:00:00 2001 From: Hein Date: Thu, 1 Oct 2026 13:31:13 +0200 Subject: [PATCH] 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. --- pkg/resolvemcp/annotation.go | 41 ++++++- pkg/resolvemcp/context.go | 22 +++- pkg/resolvemcp/guard.go | 52 ++++++++ pkg/resolvemcp/guard_test.go | 97 +++++++++++++++ pkg/resolvemcp/handler.go | 83 ++++++++----- pkg/resolvemcp/hooks.go | 5 + pkg/resolvemcp/oauth2.go | 36 +----- pkg/resolvemcp/resolvemcp.go | 124 +++++++++++++------ pkg/resolvemcp/security_hooks.go | 14 +++ pkg/resolvemcp/security_test.go | 200 +++++++++++++++++++++++++++++++ pkg/resolvemcp/tx_test.go | 4 + pkg/resolvemcp/writecols.go | 56 +++++++++ pkg/security/hooks.go | 13 ++ 13 files changed, 643 insertions(+), 104 deletions(-) create mode 100644 pkg/resolvemcp/guard.go create mode 100644 pkg/resolvemcp/guard_test.go create mode 100644 pkg/resolvemcp/security_test.go create mode 100644 pkg/resolvemcp/writecols.go diff --git a/pkg/resolvemcp/annotation.go b/pkg/resolvemcp/annotation.go index af56fee..6ea9c65 100644 --- a/pkg/resolvemcp/annotation.go +++ b/pkg/resolvemcp/annotation.go @@ -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 } diff --git a/pkg/resolvemcp/context.go b/pkg/resolvemcp/context.go index f8e97f7..1bafd55 100644 --- a/pkg/resolvemcp/context.go +++ b/pkg/resolvemcp/context.go @@ -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) +} diff --git a/pkg/resolvemcp/guard.go b/pkg/resolvemcp/guard.go new file mode 100644 index 0000000..d27d842 --- /dev/null +++ b/pkg/resolvemcp/guard.go @@ -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) +} diff --git a/pkg/resolvemcp/guard_test.go b/pkg/resolvemcp/guard_test.go new file mode 100644 index 0000000..c248baa --- /dev/null +++ b/pkg/resolvemcp/guard_test.go @@ -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") + } +} diff --git a/pkg/resolvemcp/handler.go b/pkg/resolvemcp/handler.go index 4feb95a..1a4f6f9 100644 --- a/pkg/resolvemcp/handler.go +++ b/pkg/resolvemcp/handler.go @@ -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") } diff --git a/pkg/resolvemcp/hooks.go b/pkg/resolvemcp/hooks.go index 11aa4fb..27dfcbc 100644 --- a/pkg/resolvemcp/hooks.go +++ b/pkg/resolvemcp/hooks.go @@ -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" diff --git a/pkg/resolvemcp/oauth2.go b/pkg/resolvemcp/oauth2.go index 948b5d8..4d1f51a 100644 --- a/pkg/resolvemcp/oauth2.go +++ b/pkg/resolvemcp/oauth2.go @@ -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)) -} diff --git a/pkg/resolvemcp/resolvemcp.go b/pkg/resolvemcp/resolvemcp.go index 32fa8ee..00ac54b 100644 --- a/pkg/resolvemcp/resolvemcp.go +++ b/pkg/resolvemcp/resolvemcp.go @@ -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) } diff --git a/pkg/resolvemcp/security_hooks.go b/pkg/resolvemcp/security_hooks.go index a26b109..8af2406 100644 --- a/pkg/resolvemcp/security_hooks.go +++ b/pkg/resolvemcp/security_hooks.go @@ -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)) diff --git a/pkg/resolvemcp/security_test.go b/pkg/resolvemcp/security_test.go new file mode 100644 index 0000000..b354514 --- /dev/null +++ b/pkg/resolvemcp/security_test.go @@ -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") + } +} diff --git a/pkg/resolvemcp/tx_test.go b/pkg/resolvemcp/tx_test.go index 2279306..d00421b 100644 --- a/pkg/resolvemcp/tx_test.go +++ b/pkg/resolvemcp/tx_test.go @@ -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{}) diff --git a/pkg/resolvemcp/writecols.go b/pkg/resolvemcp/writecols.go new file mode 100644 index 0000000..f6fcf16 --- /dev/null +++ b/pkg/resolvemcp/writecols.go @@ -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 +} diff --git a/pkg/security/hooks.go b/pkg/security/hooks.go index 5bdf0e8..7e210e4 100644 --- a/pkg/security/hooks.go +++ b/pkg/security/hooks.go @@ -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)