fix(resolvemcp): single transaction for create/update, hook registry mutex, uniform not-found, bounded SSE host cache

This commit is contained in:
Hein
2026-10-01 13:33:11 +02:00
parent ad2f54693f
commit 82f901a49c
15 changed files with 238 additions and 104 deletions
+10 -4
View File
@@ -187,7 +187,7 @@ handler.EnableOAuthServer(security.OAuthServerConfig{
provider, _ := security.NewCompositeSecurityProvider(auth, colSec, rowSec)
securityList, _ := security.NewSecurityList(provider)
security.RegisterSecurityHooks(handler, securityList)
resolvemcp.RegisterSecurityHooks(handler, securityList)
http.ListenAndServe(":8080", handler.HTTPHandler(securityList))
```
@@ -286,7 +286,10 @@ resolvemcp.SetupMuxRoutesWithAuth(r, handler, securityList)
```go
import "github.com/bitechdev/ResolveSpec/pkg/security"
securityList := security.NewSecurityList(mySecurityProvider)
securityList, err := security.NewSecurityList(mySecurityProvider)
if err != nil {
log.Fatal(err)
}
resolvemcp.RegisterSecurityHooks(handler, securityList)
```
@@ -294,10 +297,13 @@ Call `RegisterSecurityHooks` **once**, after creating the handler and before reg
| Hook | Effect |
|---|---|
| `BeforeHandle` | Enforces per-entity operation rules (see below) |
| `OnTxBegin` | Stamps transaction-local settings (RLS GUCs) set with `SecurityList.SetTxSettings` |
| `BeforeHandle` | Enforces per-entity operation rules (see below); preloads column rules for writes |
| `BeforeRead` | Loads RLS/CLS rules, then injects a user-scoped WHERE clause |
| `BeforeScan` | Applies row security to the row an update or delete targets; a row the user cannot see is "not found" |
| `AfterRead` | Masks/hides columns per column-security rules; writes audit log |
| `BeforeUpdate` | Blocks update if `CanUpdate` is false |
| `BeforeCreate` | Blocks create if `CanCreate` is false; drops hidden/masked columns from the payload |
| `BeforeUpdate` | Blocks update if `CanUpdate` is false; drops hidden/masked columns from the payload |
| `BeforeDelete` | Blocks delete if `CanDelete` is false |
### Per-entity operation rules
+63 -41
View File
@@ -4,6 +4,7 @@ import (
"context"
"database/sql"
"encoding/json"
"errors"
"fmt"
"net/http"
"reflect"
@@ -99,7 +100,20 @@ type dynamicSSEHandler struct {
pool map[string]*server.SSEServer
}
// maxSSEPool bounds the per-base-URL server cache; Host and X-Forwarded-Proto are client
// controlled, so without a bound a client could grow it forever.
const maxSSEPool = 32
func (d *dynamicSSEHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
if !d.h.hostAllowed(r.Host) {
http.Error(w, "host not allowed", http.StatusBadRequest)
return
}
proto := r.Header.Get("X-Forwarded-Proto")
if proto != "" && proto != "http" && proto != "https" {
http.Error(w, "invalid forwarded protocol", http.StatusBadRequest)
return
}
baseURL := requestBaseURL(r)
d.mu.Lock()
@@ -108,6 +122,12 @@ func (d *dynamicSSEHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
}
s, ok := d.pool[baseURL]
if !ok {
if len(d.pool) >= maxSSEPool {
d.mu.Unlock()
logger.Warn("resolvemcp: SSE base URL cache full; set Config.BaseURL or Config.AllowedHosts")
http.Error(w, "too many hosts", http.StatusServiceUnavailable)
return
}
s = d.h.newSSEServer(baseURL, d.h.config.BasePath)
d.pool[baseURL] = s
}
@@ -116,6 +136,20 @@ func (d *dynamicSSEHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
s.ServeHTTP(w, r)
}
// hostAllowed reports whether host may be used to build the SSE message URL. With no
// Config.AllowedHosts every host is accepted (the pool cap still applies).
func (h *Handler) hostAllowed(host string) bool {
if len(h.config.AllowedHosts) == 0 {
return true
}
for _, a := range h.config.AllowedHosts {
if strings.EqualFold(a, host) {
return true
}
}
return false
}
// requestBaseURL builds the base URL from an incoming request.
// It honours the X-Forwarded-Proto header for deployments behind a proxy.
func requestBaseURL(r *http.Request) string {
@@ -202,6 +236,10 @@ func (h *Handler) getSchemaAndTable(defaultSchema, entity string, model interfac
return defaultSchema, entity
}
// errRecordNotFound is the one error update and delete return for a row that does not exist,
// is hidden by row security, or vanished mid-write, so ids cannot be enumerated by error text.
var errRecordNotFound = errors.New("record not found")
// recoverPanic catches a panic from the current goroutine and returns it as an error.
// Usage: defer recoverPanic(&returnedErr)
func recoverPanic(err *error) {
@@ -365,7 +403,7 @@ func (h *Handler) readInTx(ctx context.Context, hookCtx *HookContext, model inte
// a destination when the query preloads a has-many relation.
if err := query.ScanModel(ctx); err != nil {
if err == sql.ErrNoRows {
return nil, nil, fmt.Errorf("record not found")
return nil, nil, errRecordNotFound
}
return nil, nil, fmt.Errorf("query error: %w", err)
}
@@ -374,7 +412,7 @@ func (h *Handler) readInTx(ctx context.Context, hookCtx *HookContext, model inte
// for both collection and single-record reads. Extract its one result.
scannedResults := reflect.ValueOf(modelPtr).Elem()
if scannedResults.Len() == 0 {
return nil, nil, fmt.Errorf("record not found")
return nil, nil, errRecordNotFound
}
data = scannedResults.Index(0).Interface()
} else {
@@ -457,11 +495,11 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
modelType = modelType.Elem()
}
// Transaction 1: BeforeCreate + inserts.
var (
single bool
originals []map[string]interface{}
insertedIDs []interface{}
results []interface{}
)
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
if err := h.hooks.Execute(BeforeCreate, hookCtx); err != nil {
@@ -511,22 +549,11 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
}
insertedIDs = append(insertedIDs, returnedID)
}
return nil
})
if err != nil {
if single {
return nil, fmt.Errorf("create error: %w", err)
}
if _, ok := hookCtx.Data.([]interface{}); ok {
return nil, fmt.Errorf("batch create error: %w", err)
}
return nil, err
}
// Transaction 2: re-fetch to capture DB-generated defaults/triggers, then AfterCreate.
results := make([]interface{}, 0, len(insertedIDs))
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
results = results[:0]
// Re-fetch inside the same transaction to capture DB-generated defaults/triggers, then
// AfterCreate: the write is only committed when the whole sequence succeeds, so a
// failure here cannot leave a committed insert behind an error the client may retry.
results = make([]interface{}, 0, len(insertedIDs))
for i, pkVal := range insertedIDs {
if pkVal == nil {
results = append(results, originals[i])
@@ -553,6 +580,12 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
return nil
})
if err != nil {
if single {
return nil, fmt.Errorf("create error: %w", err)
}
if _, ok := hookCtx.Data.([]interface{}); ok {
return nil, fmt.Errorf("batch create error: %w", err)
}
return nil, err
}
if single {
@@ -647,7 +680,7 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
}
if err := hookCtx.Query.ScanModel(ctx); err != nil {
if err == sql.ErrNoRows {
return fmt.Errorf("no records found to update")
return errRecordNotFound
}
return fmt.Errorf("error fetching existing record: %w", err)
}
@@ -671,42 +704,31 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
return fmt.Errorf("error updating record: %w", err)
}
if res.RowsAffected() == 0 {
return fmt.Errorf("no records found to update")
return errRecordNotFound
}
updateResult = existingMap
hookCtx.Result = updateResult
return h.hooks.Execute(AfterUpdate, hookCtx)
})
hookCtx.Result = existingMap
if err != nil {
return nil, err
}
// Transaction 2: re-fetch to capture DB-generated changes.
modelType := reflect.TypeOf(model)
if modelType.Kind() == reflect.Pointer {
modelType = modelType.Elem()
}
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
// Re-fetch inside the same transaction to capture DB-generated changes, then
// AfterUpdate; see executeCreate.
fetchedRecord := reflect.New(modelType).Interface()
if err := tx.NewSelect().Model(fetchedRecord).
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id).
ScanModel(ctx); err == nil {
jsonData, marshalErr := json.Marshal(fetchedRecord)
if marshalErr == nil {
if jsonData, marshalErr := json.Marshal(fetchedRecord); marshalErr == nil {
var fetchedMap map[string]interface{}
if json.Unmarshal(jsonData, &fetchedMap) == nil {
updateResult = fetchedMap
existingMap = fetchedMap
hookCtx.Result = fetchedMap
}
}
}
return nil
updateResult = existingMap
return h.hooks.Execute(AfterUpdate, hookCtx)
})
if err != nil {
return nil, err
}
return updateResult, nil
}
@@ -766,7 +788,7 @@ func (h *Handler) executeDelete(ctx context.Context, schema, entity, id string)
}
if err := hookCtx.Query.ScanModel(ctx); err != nil {
if err == sql.ErrNoRows {
return fmt.Errorf("record not found")
return errRecordNotFound
}
return fmt.Errorf("error fetching record: %w", err)
}
@@ -778,7 +800,7 @@ func (h *Handler) executeDelete(ctx context.Context, schema, entity, id string)
return fmt.Errorf("delete error: %w", err)
}
if res.RowsAffected() == 0 {
return fmt.Errorf("record not found or already deleted")
return errRecordNotFound
}
recordToDelete = record
+80
View File
@@ -0,0 +1,80 @@
package resolvemcp
import (
"fmt"
"net/http"
"net/http/httptest"
"sync"
"testing"
"github.com/DATA-DOG/go-sqlmock"
)
func TestHookRegistryConcurrentUse(t *testing.T) {
r := NewHookRegistry()
var wg sync.WaitGroup
for i := 0; i < 8; i++ {
wg.Add(3)
go func() { defer wg.Done(); r.Register(BeforeRead, func(*HookContext) error { return nil }) }()
go func() { defer wg.Done(); _ = r.Execute(BeforeRead, &HookContext{}) }()
go func() { defer wg.Done(); _ = r.HasHooks(BeforeRead); r.Clear(AfterRead) }()
}
wg.Wait()
}
// Update and delete report a missing row with the same error, so ids cannot be probed.
func TestNotFoundErrorsAreUniform(t *testing.T) {
h, mock, ctx := newTxHarness(t)
empty := sqlmock.NewRows([]string{"id", "name"})
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(empty)
mock.ExpectRollback()
_, errU := h.executeUpdate(ctx, "public", "items", "9", map[string]interface{}{"name": "x"})
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}))
mock.ExpectRollback()
_, errD := h.executeDelete(ctx, "public", "items", "9")
if errU == nil || errD == nil || errU.Error() != errD.Error() {
t.Fatalf("update %v / delete %v must be the same error", errU, errD)
}
}
func TestSSEHostAllowlistAndPoolCap(t *testing.T) {
h, _, _ := newTxHarness(t)
h.config.AllowedHosts = []string{"mcp.example.com"}
d := &dynamicSSEHandler{h: h}
r := httptest.NewRequest(http.MethodPost, "/mcp/message", nil)
r.Host = "evil.example.net"
w := httptest.NewRecorder()
d.ServeHTTP(w, r)
if w.Code != http.StatusBadRequest {
t.Errorf("foreign host: status %d, want 400", w.Code)
}
r = httptest.NewRequest(http.MethodPost, "/mcp/message", nil)
r.Host = "mcp.example.com"
r.Header.Set("X-Forwarded-Proto", "javascript")
w = httptest.NewRecorder()
d.ServeHTTP(w, r)
if w.Code != http.StatusBadRequest {
t.Errorf("bad proto: status %d, want 400", w.Code)
}
h.config.AllowedHosts = nil
for i := 0; i < maxSSEPool+5; i++ {
r = httptest.NewRequest(http.MethodPost, "/mcp/message?sessionId=x", nil)
r.Host = fmt.Sprintf("h%d.example.com", i)
d.ServeHTTP(httptest.NewRecorder(), r)
}
if len(d.pool) > maxSSEPool {
t.Errorf("pool grew to %d, cap is %d", len(d.pool), maxSSEPool)
}
r = httptest.NewRequest(http.MethodPost, "/mcp/message", nil)
r.Host = "one-more.example.com"
w = httptest.NewRecorder()
d.ServeHTTP(w, r)
if w.Code != http.StatusServiceUnavailable {
t.Errorf("full pool: status %d, want 503", w.Code)
}
}
+18 -5
View File
@@ -3,6 +3,7 @@ package resolvemcp
import (
"context"
"fmt"
"sync"
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/logger"
@@ -67,6 +68,7 @@ type HookFunc func(*HookContext) error
// HookRegistry manages all registered hooks
type HookRegistry struct {
mu sync.RWMutex
hooks map[HookType][]HookFunc
}
@@ -77,11 +79,14 @@ func NewHookRegistry() *HookRegistry {
}
func (r *HookRegistry) Register(hookType HookType, hook HookFunc) {
r.mu.Lock()
if r.hooks == nil {
r.hooks = make(map[HookType][]HookFunc)
}
r.hooks[hookType] = append(r.hooks[hookType], hook)
logger.Info("Registered resolvemcp hook for %s (total: %d)", hookType, len(r.hooks[hookType]))
total := len(r.hooks[hookType])
r.mu.Unlock()
logger.Info("Registered resolvemcp hook for %s (total: %d)", hookType, total)
}
func (r *HookRegistry) RegisterMultiple(hookTypes []HookType, hook HookFunc) {
@@ -91,8 +96,11 @@ func (r *HookRegistry) RegisterMultiple(hookTypes []HookType, hook HookFunc) {
}
func (r *HookRegistry) Execute(hookType HookType, ctx *HookContext) error {
hooks, exists := r.hooks[hookType]
if !exists || len(hooks) == 0 {
// Append-only slices: a snapshot of the slice header is safe to iterate without the lock.
r.mu.RLock()
hooks := r.hooks[hookType]
r.mu.RUnlock()
if len(hooks) == 0 {
return nil
}
@@ -114,14 +122,19 @@ func (r *HookRegistry) Execute(hookType HookType, ctx *HookContext) error {
}
func (r *HookRegistry) Clear(hookType HookType) {
r.mu.Lock()
defer r.mu.Unlock()
delete(r.hooks, hookType)
}
func (r *HookRegistry) ClearAll() {
r.mu.Lock()
defer r.mu.Unlock()
r.hooks = make(map[HookType][]HookFunc)
}
func (r *HookRegistry) HasHooks(hookType HookType) bool {
hooks, exists := r.hooks[hookType]
return exists && len(hooks) > 0
r.mu.RLock()
defer r.mu.RUnlock()
return len(r.hooks[hookType]) > 0
}
+5
View File
@@ -41,6 +41,11 @@ type Config struct {
// If empty, the path is detected from each incoming request automatically.
BasePath string
// AllowedHosts restricts the Host header accepted by the SSE transport when BaseURL is
// empty (the message endpoint URL sent to clients is built from it). Empty accepts any
// host, with at most 32 distinct base URLs cached; prefer setting BaseURL.
AllowedHosts []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
-2
View File
@@ -66,8 +66,6 @@ func TestUpdateSetsOnlyGivenKeysAndAllowsNull(t *testing.T) {
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()
+31 -28
View File
@@ -128,14 +128,12 @@ func TestReadRunsInOneTransaction(t *testing.T) {
}
}
func TestCreateSingleUsesTwoTransactions(t *testing.T) {
func TestCreateSingleRunsInOneTransaction(t *testing.T) {
h, mock, ctx := newTxHarness(t)
tr := traceHooks(h, OnTxBegin, BeforeCreate, AfterCreate)
mock.ExpectBegin()
mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(7))
mock.ExpectCommit()
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a"))
mock.ExpectCommit()
@@ -145,24 +143,19 @@ func TestCreateSingleUsesTwoTransactions(t *testing.T) {
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
tr.assertOrder(t, "on_tx_begin", "before_create", "on_tx_begin", "after_create")
if tr.txs["before_create"][0] != tr.txs["on_tx_begin"][0] {
t.Fatal("BeforeCreate must run on the first transaction")
}
if tr.txs["after_create"][0] != tr.txs["on_tx_begin"][1] || tr.txs["on_tx_begin"][0] == tr.txs["on_tx_begin"][1] {
t.Fatal("AfterCreate must run on a second, distinct transaction")
tr.assertOrder(t, "on_tx_begin", "before_create", "after_create")
if tr.txs["before_create"][0] != tr.txs["on_tx_begin"][0] || tr.txs["after_create"][0] != tr.txs["on_tx_begin"][0] {
t.Fatal("BeforeCreate, the re-fetch and AfterCreate must share one transaction")
}
}
func TestCreateBatchRefetchOnSecondTransaction(t *testing.T) {
func TestCreateBatchRefetchInSameTransaction(t *testing.T) {
h, mock, ctx := newTxHarness(t)
tr := traceHooks(h, OnTxBegin, AfterCreate)
mock.ExpectBegin()
mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(1))
mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(2))
mock.ExpectCommit()
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(1, "a"))
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(2, "b"))
mock.ExpectCommit()
@@ -174,10 +167,10 @@ func TestCreateBatchRefetchOnSecondTransaction(t *testing.T) {
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
tr.assertOrder(t, "on_tx_begin", "on_tx_begin", "after_create")
tr.assertOrder(t, "on_tx_begin", "after_create")
}
func TestUpdateRefetchRunsInSecondTransaction(t *testing.T) {
func TestUpdateRefetchRunsInSameTransaction(t *testing.T) {
h, mock, ctx := newTxHarness(t)
tr := traceHooks(h, OnTxBegin, BeforeUpdate, AfterUpdate)
@@ -185,25 +178,38 @@ func TestUpdateRefetchRunsInSecondTransaction(t *testing.T) {
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a"))
mock.ExpectExec(`UPDATE`).WillReturnResult(sqlmock.NewResult(0, 1))
mock.ExpectCommit()
mock.ExpectBegin()
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "b"))
mock.ExpectCommit()
if _, err := h.executeUpdate(ctx, "public", "items", "7", map[string]interface{}{"name": "b"}); err != nil {
res, err := h.executeUpdate(ctx, "public", "items", "7", map[string]interface{}{"name": "b"})
if err != nil {
t.Fatal(err)
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
tr.assertOrder(t, "on_tx_begin", "before_update", "after_update", "on_tx_begin")
if tr.txs["on_tx_begin"][0] == tr.txs["on_tx_begin"][1] {
t.Fatal("re-fetch must run on a second transaction")
if m, _ := res.(map[string]interface{}); m["name"] != "b" {
t.Fatalf("result must be the re-fetched row, got %v", res)
}
for _, ht := range []string{"before_update", "after_update"} {
if tr.txs[ht][0] != tr.txs["on_tx_begin"][0] {
t.Fatalf("%s must run on the first transaction", ht)
}
tr.assertOrder(t, "on_tx_begin", "before_update", "after_update")
}
// A failing AfterCreate must roll the insert back: the client sees an error, so nothing may
// have been committed that a retry would duplicate.
func TestAfterCreateErrorRollsBackInsert(t *testing.T) {
h, mock, ctx := newTxHarness(t)
h.Hooks().Register(AfterCreate, func(*HookContext) error { return sql.ErrConnDone })
mock.ExpectBegin()
mock.ExpectQuery(`INSERT`).WillReturnRows(sqlmock.NewRows([]string{"id"}).AddRow(7))
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(7, "a"))
mock.ExpectRollback()
if _, err := h.executeCreate(ctx, "public", "items", map[string]interface{}{"name": "a"}); err == nil {
t.Fatal("expected error")
}
if err := mock.ExpectationsWereMet(); err != nil {
t.Fatal(err)
}
}
@@ -314,15 +320,12 @@ func TestSecurityHooksStampTxSettingsOnEveryTransaction(t *testing.T) {
ctx = context.WithValue(ctx, security.UserContextKey, &security.UserContext{UserID: 7, UserName: "u"})
ctx = context.WithValue(ctx, security.UserIDKey, 7)
// Update opens two transactions; each must be stamped before any other SQL.
// The transaction is stamped before any other SQL.
cols := []string{"id", "name"}
mock.ExpectBegin()
mock.ExpectExec(`set_config\('app\.user_id'`).WillReturnResult(sqlmock.NewResult(0, 0))
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "a"))
mock.ExpectExec(`UPDATE`).WillReturnResult(sqlmock.NewResult(0, 1))
mock.ExpectCommit()
mock.ExpectBegin()
mock.ExpectExec(`set_config\('app\.user_id'`).WillReturnResult(sqlmock.NewResult(0, 0))
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows(cols).AddRow(7, "b"))
mock.ExpectCommit()