mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-06 05:16:27 +00:00
fix(restheadspec): handle primary key changes in updates
* Allow primary key changes when specified in the request body. * Ensure correct record fetching after primary key updates. * Add integration tests for primary key update scenarios.
This commit is contained in:
@@ -1244,6 +1244,16 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// A primary key change is only honoured when the ID was given in the URL and
|
||||||
|
// the body carries a different, non-null primary key value.
|
||||||
|
var newPK interface{}
|
||||||
|
pkChanged := false
|
||||||
|
if urlID != "" {
|
||||||
|
if v, ok := updates[pkName]; ok && v != nil && !reflection.IsEmptyValue(v) && fmt.Sprintf("%v", v) != urlID {
|
||||||
|
newPK, pkChanged = v, true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// Wrap in transaction to ensure BeforeUpdate hook is inside transaction
|
// Wrap in transaction to ensure BeforeUpdate hook is inside transaction
|
||||||
err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error {
|
err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error {
|
||||||
// Execute BeforeUpdate hooks inside transaction, before any queries run.
|
// Execute BeforeUpdate hooks inside transaction, before any queries run.
|
||||||
@@ -1337,6 +1347,14 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
|
|||||||
return fmt.Errorf("no records found to update")
|
return fmt.Errorf("no records found to update")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// SetMap skips primary key columns, so apply a PK change explicitly.
|
||||||
|
if pkChanged {
|
||||||
|
if _, err := tx.NewUpdate().Table(tableName).Set(pkName, newPK).
|
||||||
|
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID).Exec(ctx); err != nil {
|
||||||
|
return fmt.Errorf("error updating primary key: %w", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// Execute AfterUpdate hooks inside transaction
|
// Execute AfterUpdate hooks inside transaction
|
||||||
hookCtx.Result = updates
|
hookCtx.Result = updates
|
||||||
hookCtx.Error = nil
|
hookCtx.Error = nil
|
||||||
@@ -1362,7 +1380,9 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
|
|||||||
updatedRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface()
|
updatedRecord := reflect.New(reflection.GetPointerElement(reflect.TypeOf(model))).Interface()
|
||||||
if err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error {
|
if err := h.runInTx(ctx, h.newTxHookContext(ctx, schema, entity, model, "update", options, w), func(tx common.Database) error {
|
||||||
fetchQuery := tx.NewSelect().Model(updatedRecord).Column(reflection.GetSQLModelColumns(model)...)
|
fetchQuery := tx.NewSelect().Model(updatedRecord).Column(reflection.GetSQLModelColumns(model)...)
|
||||||
if urlID != "" {
|
if pkChanged {
|
||||||
|
fetchQuery = fetchQuery.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), newPK)
|
||||||
|
} else if urlID != "" {
|
||||||
fetchQuery = fetchQuery.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), urlID)
|
fetchQuery = fetchQuery.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), urlID)
|
||||||
} else if reqID != nil {
|
} else if reqID != nil {
|
||||||
switch id := reqID.(type) {
|
switch id := reqID.(type) {
|
||||||
|
|||||||
@@ -1539,6 +1539,10 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
|
|||||||
// Variable to store the updated record
|
// Variable to store the updated record
|
||||||
var updatedRecord interface{}
|
var updatedRecord interface{}
|
||||||
|
|
||||||
|
// ID used to re-fetch the record after the update; differs from targetID
|
||||||
|
// when the request changes the primary key.
|
||||||
|
finalID := targetID
|
||||||
|
|
||||||
// Hook context used inside and outside transaction
|
// Hook context used inside and outside transaction
|
||||||
hookCtx := &HookContext{
|
hookCtx := &HookContext{
|
||||||
Context: ctx,
|
Context: ctx,
|
||||||
@@ -1604,11 +1608,25 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
|
|||||||
nestedRelations = relations
|
nestedRelations = relations
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// Capture a changed primary key from the request before merging. The row is
|
||||||
|
// located by the original targetID (WHERE), while the new value is written via SET.
|
||||||
|
// Only honoured when an ID was given in the URL (id != "").
|
||||||
|
var newPK interface{}
|
||||||
|
var pkChanged bool
|
||||||
|
if id != "" {
|
||||||
|
newPK, pkChanged = h.requestedPrimaryKey(model, pkName, dataMap, targetID)
|
||||||
|
}
|
||||||
|
|
||||||
// Overwrite with every key present in the request (including "" and null unless disallowed)
|
// Overwrite with every key present in the request (including "" and null unless disallowed)
|
||||||
common.MergeUpdateValues(existingMap, dataMap, h.disallowNulls)
|
common.MergeUpdateValues(existingMap, dataMap, h.disallowNulls)
|
||||||
|
|
||||||
// Ensure ID is in the data map for the update
|
// Ensure ID is in the data map for the update (new value if the PK is being changed)
|
||||||
existingMap[pkName] = targetID
|
if pkChanged {
|
||||||
|
existingMap[pkName] = newPK
|
||||||
|
finalID = newPK
|
||||||
|
} else {
|
||||||
|
existingMap[pkName] = targetID
|
||||||
|
}
|
||||||
dataMap = existingMap
|
dataMap = existingMap
|
||||||
|
|
||||||
// Populate model instance from dataMap to preserve custom types (like SqlJSONB)
|
// Populate model instance from dataMap to preserve custom types (like SqlJSONB)
|
||||||
@@ -1651,6 +1669,20 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
|
|||||||
}
|
}
|
||||||
|
|
||||||
_ = result
|
_ = result
|
||||||
|
|
||||||
|
// Primary key changes are not part of the struct SET, so apply them explicitly.
|
||||||
|
if pkChanged {
|
||||||
|
pkResult, err := tx.NewUpdate().Table(tableName).
|
||||||
|
Set(pkName, newPK).
|
||||||
|
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID).
|
||||||
|
Exec(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("failed to update primary key: %w", err)
|
||||||
|
}
|
||||||
|
if pkResult.RowsAffected() == 0 {
|
||||||
|
return fmt.Errorf("primary key update affected no rows for ID: %v", targetID)
|
||||||
|
}
|
||||||
|
}
|
||||||
return nil
|
return nil
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -1667,7 +1699,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
|
|||||||
var errCode, errMsg string
|
var errCode, errMsg string
|
||||||
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
|
err = h.runInTx(ctx, hookCtx, func(tx common.Database) error {
|
||||||
fetchedRecord := reflect.New(reflect.TypeOf(model)).Interface()
|
fetchedRecord := reflect.New(reflect.TypeOf(model)).Interface()
|
||||||
selectQuery := tx.NewSelect().Model(fetchedRecord).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID)
|
selectQuery := tx.NewSelect().Model(fetchedRecord).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), finalID)
|
||||||
|
|
||||||
// Execute BeforeScan hooks so row security is re-applied to the post-update
|
// Execute BeforeScan hooks so row security is re-applied to the post-update
|
||||||
// re-fetch, same as it is for the initial read and the update query itself.
|
// re-fetch, same as it is for the initial read and the update query itself.
|
||||||
@@ -1711,7 +1743,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
logger.Info("Successfully updated record with ID: %v", targetID)
|
logger.Info("Successfully updated record with ID: %v", finalID)
|
||||||
// Invalidate cache for this table
|
// Invalidate cache for this table
|
||||||
cacheTags := buildCacheTags(schema, tableName)
|
cacheTags := buildCacheTags(schema, tableName)
|
||||||
if err := invalidateCacheForTags(ctx, cacheTags); err != nil {
|
if err := invalidateCacheForTags(ctx, cacheTags); err != nil {
|
||||||
@@ -1720,6 +1752,28 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
|
|||||||
h.sendResponseWithOptions(w, mergedData, nil, &options)
|
h.sendResponseWithOptions(w, mergedData, nil, &options)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// requestedPrimaryKey returns the primary key value carried in the request body
|
||||||
|
// (by column name or JSON key) when it differs from the current target ID.
|
||||||
|
func (h *Handler) requestedPrimaryKey(model interface{}, pkName string, dataMap map[string]interface{}, targetID interface{}) (interface{}, bool) {
|
||||||
|
val, exists := dataMap[pkName]
|
||||||
|
if !exists {
|
||||||
|
modelType := reflection.GetPointerElement(reflect.TypeOf(model))
|
||||||
|
for jsonKey, col := range reflection.BuildJSONToDBColumnMap(modelType) {
|
||||||
|
if col == pkName {
|
||||||
|
val, exists = dataMap[jsonKey]
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !exists || val == nil || reflection.IsEmptyValue(val) {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
if fmt.Sprintf("%v", val) == fmt.Sprintf("%v", targetID) {
|
||||||
|
return nil, false
|
||||||
|
}
|
||||||
|
return val, true
|
||||||
|
}
|
||||||
|
|
||||||
func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id string, data interface{}) {
|
func (h *Handler) handleDelete(ctx context.Context, w common.ResponseWriter, id string, data interface{}) {
|
||||||
// Capture panics and return error response
|
// Capture panics and return error response
|
||||||
defer func() {
|
defer func() {
|
||||||
|
|||||||
@@ -0,0 +1,125 @@
|
|||||||
|
//go:build integration
|
||||||
|
|
||||||
|
package restheadspec
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
_ "github.com/jackc/pgx/v5/stdlib"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
"github.com/testcontainers/testcontainers-go"
|
||||||
|
"github.com/testcontainers/testcontainers-go/wait"
|
||||||
|
"github.com/uptrace/bun"
|
||||||
|
"github.com/uptrace/bun/dialect/pgdialect"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||||
|
)
|
||||||
|
|
||||||
|
type pkAsset struct {
|
||||||
|
bun.BaseModel `bun:"table:public.t_pkasset,alias:t_pkasset"`
|
||||||
|
Category string `json:"category" bun:"category,type:citext"`
|
||||||
|
Description string `json:"description" bun:"description,type:citext,pk"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func setupPKTestDB(t *testing.T) *sql.DB {
|
||||||
|
t.Helper()
|
||||||
|
ctx := context.Background()
|
||||||
|
pg, err := testcontainers.GenericContainer(ctx, testcontainers.GenericContainerRequest{
|
||||||
|
ContainerRequest: testcontainers.ContainerRequest{
|
||||||
|
Image: "postgres:15-alpine",
|
||||||
|
ExposedPorts: []string{"5432/tcp"},
|
||||||
|
Env: map[string]string{
|
||||||
|
"POSTGRES_USER": "testuser", "POSTGRES_PASSWORD": "testpass", "POSTGRES_DB": "testdb",
|
||||||
|
},
|
||||||
|
WaitingFor: wait.ForLog("database system is ready to accept connections").
|
||||||
|
WithOccurrence(2).WithStartupTimeout(60 * time.Second),
|
||||||
|
},
|
||||||
|
Started: true,
|
||||||
|
})
|
||||||
|
require.NoError(t, err)
|
||||||
|
t.Cleanup(func() { _ = pg.Terminate(ctx) })
|
||||||
|
|
||||||
|
host, err := pg.Host(ctx)
|
||||||
|
require.NoError(t, err)
|
||||||
|
port, err := pg.MappedPort(ctx, "5432")
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
db, err := sql.Open("pgx", fmt.Sprintf("postgres://testuser:testpass@%s:%s/testdb?sslmode=disable", host, port.Port()))
|
||||||
|
require.NoError(t, err)
|
||||||
|
t.Cleanup(func() { _ = db.Close() })
|
||||||
|
require.NoError(t, db.Ping())
|
||||||
|
|
||||||
|
_, err = db.Exec(`
|
||||||
|
CREATE EXTENSION IF NOT EXISTS citext;
|
||||||
|
CREATE TABLE public.t_pkasset (
|
||||||
|
category citext,
|
||||||
|
description citext PRIMARY KEY
|
||||||
|
);
|
||||||
|
INSERT INTO public.t_pkasset VALUES ('cat', 'old-pk');`)
|
||||||
|
require.NoError(t, err)
|
||||||
|
return db
|
||||||
|
}
|
||||||
|
|
||||||
|
func pkUpdateCtx(base context.Context) context.Context {
|
||||||
|
ctx := WithSchema(base, "public")
|
||||||
|
ctx = WithEntity(ctx, "t_pkasset")
|
||||||
|
ctx = WithTableName(ctx, "t_pkasset")
|
||||||
|
return WithModel(ctx, pkAsset{})
|
||||||
|
}
|
||||||
|
|
||||||
|
func countPK(t *testing.T, db *sql.DB, pk string) int {
|
||||||
|
t.Helper()
|
||||||
|
var n int
|
||||||
|
require.NoError(t, db.QueryRow(`SELECT count(*) FROM public.t_pkasset WHERE description = $1`, pk).Scan(&n))
|
||||||
|
return n
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUpdateChangesPrimaryKeyWhenURLIDGiven(t *testing.T) {
|
||||||
|
db := setupPKTestDB(t)
|
||||||
|
h := NewHandler(database.NewBunAdapter(bun.NewDB(db, pgdialect.New())), modelregistry.NewModelRegistry())
|
||||||
|
|
||||||
|
update := func(id string, body map[string]interface{}) *httptest.ResponseRecorder {
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodPut, "/", nil))
|
||||||
|
base, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
h.handleUpdate(pkUpdateCtx(base), w, id, nil, body, ExtendedRequestOptions{})
|
||||||
|
return rec
|
||||||
|
}
|
||||||
|
|
||||||
|
// URL id = old PK, body carries a different PK: the PK is changed.
|
||||||
|
rec := update("old-pk", map[string]interface{}{"description": "new-pk", "category": "cat2"})
|
||||||
|
require.Equal(t, http.StatusOK, rec.Code, rec.Body.String())
|
||||||
|
require.Equal(t, 0, countPK(t, db, "old-pk"))
|
||||||
|
require.Equal(t, 1, countPK(t, db, "new-pk"))
|
||||||
|
var cat string
|
||||||
|
require.NoError(t, db.QueryRow(`SELECT category FROM public.t_pkasset WHERE description = 'new-pk'`).Scan(&cat))
|
||||||
|
require.Equal(t, "cat2", cat)
|
||||||
|
|
||||||
|
// Body PK equal to the URL id: ordinary update, PK untouched.
|
||||||
|
rec = update("new-pk", map[string]interface{}{"description": "new-pk", "category": "cat3"})
|
||||||
|
require.Equal(t, http.StatusOK, rec.Code, rec.Body.String())
|
||||||
|
require.Equal(t, 1, countPK(t, db, "new-pk"))
|
||||||
|
|
||||||
|
// Body without a PK: ordinary update.
|
||||||
|
rec = update("new-pk", map[string]interface{}{"category": "cat4"})
|
||||||
|
require.Equal(t, http.StatusOK, rec.Code, rec.Body.String())
|
||||||
|
|
||||||
|
// Null PK in the body is ignored, not applied.
|
||||||
|
rec = update("new-pk", map[string]interface{}{"description": nil, "category": "cat5"})
|
||||||
|
require.Equal(t, http.StatusOK, rec.Code, rec.Body.String())
|
||||||
|
require.Equal(t, 1, countPK(t, db, "new-pk"))
|
||||||
|
var total int
|
||||||
|
require.NoError(t, db.QueryRow(`SELECT count(*) FROM public.t_pkasset`).Scan(&total))
|
||||||
|
require.Equal(t, 1, total)
|
||||||
|
require.NoError(t, db.QueryRow(`SELECT category FROM public.t_pkasset WHERE description = 'new-pk'`).Scan(&cat))
|
||||||
|
require.Equal(t, "cat5", cat)
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user