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