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

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