From 82f901a49cfee57d6991b3dac50b79d163c39e46 Mon Sep 17 00:00:00 2001 From: Hein Date: Thu, 1 Oct 2026 13:33:11 +0200 Subject: [PATCH] fix(resolvemcp): single transaction for create/update, hook registry mutex, uniform not-found, bounded SSE host cache --- pkg/dbtrace/README.md | 2 +- pkg/resolvemcp/README.md | 14 ++- pkg/resolvemcp/handler.go | 104 +++++++++++------- pkg/resolvemcp/hardening_test.go | 80 ++++++++++++++ pkg/resolvemcp/hooks.go | 23 +++- pkg/resolvemcp/resolvemcp.go | 5 + pkg/resolvemcp/security_test.go | 2 - pkg/resolvemcp/tx_test.go | 59 +++++----- pkg/security/OAUTH2.md | 2 +- .../OAUTH2_REFRESH_QUICK_REFERENCE.md | 2 +- .../OAUTH2_REFRESH_TOKEN_IMPLEMENTATION.md | 10 +- pkg/security/PASSKEY_QUICK_REFERENCE.md | 2 +- pkg/security/QUICK_REFERENCE.md | 11 +- pkg/security/README.md | 17 ++- pkg/security/breaking_changes.md | 9 +- 15 files changed, 238 insertions(+), 104 deletions(-) create mode 100644 pkg/resolvemcp/hardening_test.go diff --git a/pkg/dbtrace/README.md b/pkg/dbtrace/README.md index 5a204b0..d458d2c 100644 --- a/pkg/dbtrace/README.md +++ b/pkg/dbtrace/README.md @@ -15,7 +15,7 @@ Wire: `dbtrace.Configure(dbtrace.FromConfig(cfg.DBTrace))` and wrap handlers wit ## Log fields - `tx` transactions begun · `tx_queries` adapter queries inside `RunInTransaction` (share the tx connection) - `pooled` adapter queries outside a tx (each takes a pool connection) -- `raw` direct `*sql.DB` calls, with kinds: `auth.session`, `auth.activity`, `security.column`, `security.row`, `probe.pg_proc`, `keystore.validate` +- `raw` direct `*sql.DB` calls, with kinds: `auth.session`, `auth.activity`, `security.column`, `security.row`, `probe.pg_proc` (lookup `ModeAuto` only), `keystore.validate` - Connections used ≈ `tx + pooled + raw` ## Pool log diff --git a/pkg/resolvemcp/README.md b/pkg/resolvemcp/README.md index 2e3d99f..206fdf8 100644 --- a/pkg/resolvemcp/README.md +++ b/pkg/resolvemcp/README.md @@ -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 diff --git a/pkg/resolvemcp/handler.go b/pkg/resolvemcp/handler.go index 1a4f6f9..a38cc66 100644 --- a/pkg/resolvemcp/handler.go +++ b/pkg/resolvemcp/handler.go @@ -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 diff --git a/pkg/resolvemcp/hardening_test.go b/pkg/resolvemcp/hardening_test.go new file mode 100644 index 0000000..c132449 --- /dev/null +++ b/pkg/resolvemcp/hardening_test.go @@ -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) + } +} diff --git a/pkg/resolvemcp/hooks.go b/pkg/resolvemcp/hooks.go index 27dfcbc..5187d81 100644 --- a/pkg/resolvemcp/hooks.go +++ b/pkg/resolvemcp/hooks.go @@ -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 } diff --git a/pkg/resolvemcp/resolvemcp.go b/pkg/resolvemcp/resolvemcp.go index 00ac54b..fc7a996 100644 --- a/pkg/resolvemcp/resolvemcp.go +++ b/pkg/resolvemcp/resolvemcp.go @@ -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 diff --git a/pkg/resolvemcp/security_test.go b/pkg/resolvemcp/security_test.go index b354514..96162ae 100644 --- a/pkg/resolvemcp/security_test.go +++ b/pkg/resolvemcp/security_test.go @@ -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() diff --git a/pkg/resolvemcp/tx_test.go b/pkg/resolvemcp/tx_test.go index d00421b..899ce62 100644 --- a/pkg/resolvemcp/tx_test.go +++ b/pkg/resolvemcp/tx_test.go @@ -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() diff --git a/pkg/security/OAUTH2.md b/pkg/security/OAUTH2.md index 7dd5075..7072c14 100644 --- a/pkg/security/OAUTH2.md +++ b/pkg/security/OAUTH2.md @@ -397,7 +397,7 @@ UserInfoParser: func(userInfo map[string]any) (*security.UserContext, error) { ## Implementation Details -All database operations use stored procedures for consistency and security: +On PostgreSQL, database operations use stored procedures by default (other dialects use direct SQL through `pkg/security/lookup`): - `resolvespec_oauth_getorcreateuser` - Find or create OAuth2 user - `resolvespec_oauth_createsession` - Create OAuth2 session - `resolvespec_oauth_getsession` - Validate and retrieve session diff --git a/pkg/security/OAUTH2_REFRESH_QUICK_REFERENCE.md b/pkg/security/OAUTH2_REFRESH_QUICK_REFERENCE.md index bf08888..7d34ce6 100644 --- a/pkg/security/OAUTH2_REFRESH_QUICK_REFERENCE.md +++ b/pkg/security/OAUTH2_REFRESH_QUICK_REFERENCE.md @@ -276,6 +276,6 @@ authURL += "&access_type=offline&prompt=consent" ## Complete Example -See `/pkg/security/oauth2_examples.go` line 250 for full working example. +See `/pkg/security/oauth2_examples.go` (`ExampleOAuth2TokenRefresh`) for full working example. For detailed documentation see `/pkg/security/OAUTH2_REFRESH_TOKEN_IMPLEMENTATION.md`. diff --git a/pkg/security/OAUTH2_REFRESH_TOKEN_IMPLEMENTATION.md b/pkg/security/OAUTH2_REFRESH_TOKEN_IMPLEMENTATION.md index 13a3f3e..a8afe8e 100644 --- a/pkg/security/OAUTH2_REFRESH_TOKEN_IMPLEMENTATION.md +++ b/pkg/security/OAUTH2_REFRESH_TOKEN_IMPLEMENTATION.md @@ -43,16 +43,16 @@ CREATE TABLE IF NOT EXISTS user_sessions ( **`resolvespec_oauth_getrefreshtoken(p_refresh_token)`** - Gets OAuth2 session data by refresh token - Returns: `{user_id, access_token, token_type, expiry}` -- Location: `lookup/database_schema.sql:714` +- Location: `lookup/database_schema.sql` (section 15); direct mode: `lookup/direct` `OAuthUserStore.GetByRefreshToken` **`resolvespec_oauth_updaterefreshtoken(p_update_data)`** - Updates session with new tokens after refresh - Input: `{user_id, old_refresh_token, new_session_token, new_access_token, new_refresh_token, expires_at}` -- Location: `lookup/database_schema.sql:752` +- Location: `lookup/database_schema.sql` (section 16); direct mode: `lookup/direct` `OAuthUserStore.UpdateRefreshToken` **`resolvespec_oauth_getuser(p_user_id)`** - Gets user data by ID for building UserContext -- Location: `lookup/database_schema.sql:791` +- Location: `lookup/database_schema.sql` (section 17); direct mode: `lookup/direct` `OAuthUserStore.GetUser` --- @@ -68,7 +68,7 @@ func (a *DatabaseAuthenticator) OAuth2RefreshToken( ) (*LoginResponse, error) ``` -**Location:** `pkg/security/oauth2_methods.go:375` +**Location:** `pkg/security/oauth2_methods.go` (`OAuth2RefreshToken`) ### Implementation Flow @@ -476,7 +476,7 @@ auth.OAuth2RefreshToken(ctx, token, "google") // Must match ProviderName ## 8. Complete Working Example -See `pkg/security/oauth2_examples.go:250` for full working example with token refresh. +See `pkg/security/oauth2_examples.go` (`ExampleOAuth2TokenRefresh`) for full working example with token refresh. --- diff --git a/pkg/security/PASSKEY_QUICK_REFERENCE.md b/pkg/security/PASSKEY_QUICK_REFERENCE.md index 4bdbe84..b29b25f 100644 --- a/pkg/security/PASSKEY_QUICK_REFERENCE.md +++ b/pkg/security/PASSKEY_QUICK_REFERENCE.md @@ -8,7 +8,7 @@ Passkey authentication (WebAuthn/FIDO2) is now integrated into the DatabaseAuthe ### Database Schema Run the passkey SQL schema (in lookup/database_schema.sql): - Creates `user_passkey_credentials` table -- Adds stored procedures for passkey operations +- Adds stored procedures for passkey operations (Postgres procedure backend; other dialects use `lookup/ddl` tables with direct SQL) ### Go Code ```go diff --git a/pkg/security/QUICK_REFERENCE.md b/pkg/security/QUICK_REFERENCE.md index 36e54db..b2fc72f 100644 --- a/pkg/security/QUICK_REFERENCE.md +++ b/pkg/security/QUICK_REFERENCE.md @@ -16,7 +16,8 @@ rowSec := security.NewDatabaseRowSecurityProvider(db) provider, _ := security.NewCompositeSecurityProvider(auth, colSec, rowSec) // Step 3: Setup and apply middleware -securityList, _ := security.SetupSecurityProvider(handler, provider) +securityList, _ := security.NewSecurityList(provider) +restheadspec.RegisterSecurityHooks(handler, securityList) router.Use(security.NewAuthMiddleware(securityList)) router.Use(security.SetSecurityMiddleware(securityList)) ``` @@ -25,7 +26,7 @@ router.Use(security.SetSecurityMiddleware(securityList)) ## Stored Procedures -**All database operations use PostgreSQL stored procedures** with `resolvespec_*` naming: +**On PostgreSQL, database operations use stored procedures by default** with `resolvespec_*` naming (other dialects use direct SQL; see `lookup.Config` in README.md): ### Database Authenticators ```go @@ -632,7 +633,8 @@ func main() { // Setup security provider := &SimpleProvider{} - securityList := security.SetupSecurityProvider(handler, provider) + securityList, _ := security.NewSecurityList(provider) + restheadspec.RegisterSecurityHooks(handler, securityList) // Apply middleware router := mux.NewRouter() @@ -761,7 +763,8 @@ auth := security.NewJWTAuthenticator("secret", db) colSec := security.NewDatabaseColumnSecurityProvider(db) rowSec := security.NewDatabaseRowSecurityProvider(db) provider := security.NewCompositeSecurityProvider(auth, colSec, rowSec) -securityList := security.SetupSecurityProvider(handler, provider) +securityList, _ := security.NewSecurityList(provider) +restheadspec.RegisterSecurityHooks(handler, securityList) // ===== INTERFACE METHODS ===== Authenticate(r *http.Request) (*UserContext, error) diff --git a/pkg/security/README.md b/pkg/security/README.md index e6e8cbc..ff0f3ac 100644 --- a/pkg/security/README.md +++ b/pkg/security/README.md @@ -18,7 +18,7 @@ Type-safe, composable security system for ResolveSpec with support for authentic ## Stored Procedure Architecture -**All database-backed security providers use PostgreSQL stored procedures exclusively.** No raw SQL queries are executed from Go code. +**On PostgreSQL, database-backed security providers use stored procedures by default.** `pkg/security` itself contains no SQL; all database access lives in [`pkg/security/lookup`](lookup), which can also run the same operations as direct SQL on tables (see [Database access (lookup)](#database-access-lookup)). ### Benefits @@ -139,7 +139,7 @@ Read them from Go with `ddl.SQL("sqlite")` or, for drivers that reject multi-sta - Direct mode stores `bytea` / array / `jsonb` values (passkey credentials, OAuth client lists, key meta) as base64 / JSON text; the Go API is unchanged. - OAuth authorization codes are consumed atomically. - Adding a database: implement `dialect.Dialect`, register it with `dialect.Register`, then set `Config.Dialect`. -- Backend conformance: `lookup/conformance` is one behavioural suite run against every backend (`go test ./pkg/security/lookup/backends -run TestConformance`). SQLite runs always; Postgres (procedure and direct), MySQL and SQL Server run when `RESOLVESPEC_TEST_PG_DSN`, `RESOLVESPEC_TEST_PG_DIRECT_DSN`, `RESOLVESPEC_TEST_MYSQL_DSN` or `RESOLVESPEC_TEST_MSSQL_DSN` is set (see the comment in `backends/conformance_test.go`). Rows are prefixed and removed afterwards. +- Backend conformance: `lookup/conformance` is one behavioural suite run against every backend (`go test ./pkg/security/lookup/backends -run TestConformance`). SQLite runs always; Postgres (procedure and direct), MySQL and SQL Server run when `RESOLVESPEC_TEST_PG_DSN`, `RESOLVESPEC_TEST_PG_DIRECT_DSN`, `RESOLVESPEC_TEST_MYSQL_DSN` or `RESOLVESPEC_TEST_MSSQL_DSN` is set (see the comment in `backends/conformance_test.go`). Rows are prefixed and removed afterwards. With `RESOLVESPEC_TEST_CONTAINERS=1` (and not `-short`) the container tests start a throwaway database with podman or docker (podman first), run the suite and remove the container, so no DSN is needed. - Migration from the old `SQLNames` / `TableNames` / `QueryMode` API: see `breaking_changes.md`. ## Quick Start @@ -936,7 +936,8 @@ func TestMyHandler(t *testing.T) { &MockRowSecurity{}, ) - securityList := security.SetupSecurityProvider(handler, provider) + securityList, _ := security.NewSecurityList(provider) + restheadspec.RegisterSecurityHooks(handler, securityList) // ... test your handler } ``` @@ -1227,9 +1228,13 @@ The main changes: | File | Description | |------|-------------| | **QUICK_REFERENCE.md** | Quick reference guide with examples | -| **INTERFACE_GUIDE.md** | Complete implementation guide | -| **examples.go** | Working provider implementations | -| **setup_example.go** | 6 complete integration examples | +| **KEYSTORE.md** | Per-user auth keys and key stores | +| **OAUTH2.md** | OAuth2 client login and the authorization server | +| **OAUTH2_REFRESH_QUICK_REFERENCE.md** / **OAUTH2_REFRESH_TOKEN_IMPLEMENTATION.md** | OAuth2 refresh tokens | +| **PASSKEY_QUICK_REFERENCE.md** | WebAuthn passkeys | +| **SECURITY_FEATURES.md** | Security feature overview | +| **breaking_changes.md** | Migration notes for the `lookup` refactor | +| **examples.go**, **examples_funcspec.go**, **oauth2_examples.go**, **passkey_examples.go** | Working provider implementations | ## API Reference diff --git a/pkg/security/breaking_changes.md b/pkg/security/breaking_changes.md index a564837..6348f17 100644 --- a/pkg/security/breaking_changes.md +++ b/pkg/security/breaking_changes.md @@ -27,9 +27,8 @@ Import `github.com/bitechdev/ResolveSpec/pkg/security/totp`. No aliases (import `totp.NewAuthenticator` takes a `totp.BaseAuthenticator` (Login, Logout, Authenticate) instead of `security.Authenticator`; any `security.Authenticator` satisfies it. -`DatabaseTwoFactorProvider` stays in `security` for now (it uses the core SQL internals) and moves -into `totp` once the lookup `TOTPStore` replaces them. Until then core imports `totp`, so `totp` -must not import `security`. +`DatabaseTwoFactorProvider` stays in `security` (it now calls the lookup `TOTPStore`). Core imports +`totp`, so `totp` must not import `security`. ## Step 0b (providers, first part): moved to `pkg/security/providers` @@ -45,7 +44,7 @@ Import `github.com/bitechdev/ResolveSpec/pkg/security/providers`. Names unchange The SHA-256 key hash helper is now `sectypes.HashKey`. The database-backed providers (`DatabaseAuthenticator`, `JWTAuthenticator`, `DatabaseKeyStore`, `DatabaseColumn/RowSecurityProvider`) -stay in `security` until the lookup stores replace their SQL. +stay in `security`; they call the lookup stores (see step 5). ## Additions (no action needed) @@ -67,7 +66,7 @@ stay in `security` until the lookup stores replace their SQL. - New `lookup/direct` package: table-backed stores for auth, keys, OAuth (client + user), passkey, TOTP and policy, built from `lookup.Schema` and the dialect. Nothing in `pkg/security` calls it - yet (wiring is step 5), so no existing API changes here. + yet (wiring happens in step 5), so no existing API changes here. - Direct `LoginAPIKey` is new: `header_api` / `api` keys only; unknown, expired, inactive and wrong-type keys (and inactive users) all return `lookup.ErrInvalidAPIKey`. - Policy tables (`sec_group_members`, `sec_column_rules`, `sec_row_rules`) are required for the