diff --git a/pkg/resolvespec/handler.go b/pkg/resolvespec/handler.go index da30452..47b0a4f 100644 --- a/pkg/resolvespec/handler.go +++ b/pkg/resolvespec/handler.go @@ -1244,6 +1244,16 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url 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 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. @@ -1337,6 +1347,14 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url 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 hookCtx.Result = updates 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() 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)...) - 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) } else if reqID != nil { switch id := reqID.(type) { diff --git a/pkg/restheadspec/handler.go b/pkg/restheadspec/handler.go index 29a59ac..9beb02c 100644 --- a/pkg/restheadspec/handler.go +++ b/pkg/restheadspec/handler.go @@ -1539,6 +1539,10 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id // Variable to store the updated record 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 hookCtx := &HookContext{ Context: ctx, @@ -1604,11 +1608,25 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id 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) common.MergeUpdateValues(existingMap, dataMap, h.disallowNulls) - // Ensure ID is in the data map for the update - existingMap[pkName] = targetID + // Ensure ID is in the data map for the update (new value if the PK is being changed) + if pkChanged { + existingMap[pkName] = newPK + finalID = newPK + } else { + existingMap[pkName] = targetID + } dataMap = existingMap // 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 + + // 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 }) @@ -1667,7 +1699,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id var errCode, errMsg string err = h.runInTx(ctx, hookCtx, func(tx common.Database) error { 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 // 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 } - logger.Info("Successfully updated record with ID: %v", targetID) + logger.Info("Successfully updated record with ID: %v", finalID) // Invalidate cache for this table cacheTags := buildCacheTags(schema, tableName) 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) } +// 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{}) { // Capture panics and return error response defer func() { diff --git a/pkg/restheadspec/pk_update_integration_test.go b/pkg/restheadspec/pk_update_integration_test.go new file mode 100644 index 0000000..cba6e69 --- /dev/null +++ b/pkg/restheadspec/pk_update_integration_test.go @@ -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) +}