mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-01 19:20:31 +00:00
fix(resolvemcp): single transaction for create/update, hook registry mutex, uniform not-found, bounded SSE host cache
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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()
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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`.
|
||||
|
||||
@@ -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.
|
||||
|
||||
---
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
+11
-6
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user