Compare commits

..
4 Commits
Author SHA1 Message Date
Hein 2898b335f8 feat(db): add ApplicationName to connection configuration
* Introduced ApplicationName field to identify clients in DSN
* Set default ApplicationName to "ResolveSpec"
* Updated tests for ApplicationName handling in DSN
2026-09-30 17:12:04 +02:00
Hein 20ba8ed112 feat(update): allow clearing values with "" and null on update
Update handlers skipped empty strings and nulls, so a client could not
blank or null out a column. Every key present in the payload now
overwrites the stored value, including "" and null.

- add common.MergeUpdateValues and use it in resolvespec and restheadspec
- add Handler.SetDisallowNulls to skip null values (""still overwrites)
- websocketspec and mqttspec now write only the keys present in the
  payload via SetMap instead of updating the whole zeroed model, which
  clobbered absent fields
2026-09-30 17:03:42 +02:00
Hein 1214f69e0c feat(middleware): export clientqueue_enqueued_total counter
Counts every request placed in a wait queue regardless of outcome, so the
total ever queued no longer has to be summed from queued, timeout and
canceled.
2026-09-30 16:16:25 +02:00
Hein 89a58ab3a0 feat(middleware): export clientqueue_waiting_clients gauge
Counts clients with at least one queued request, updated whenever a
client's waiter count crosses zero (enqueue, grant, timeout, cancel).
2026-09-30 16:14:27 +02:00
15 changed files with 333 additions and 80 deletions
+23
View File
@@ -0,0 +1,23 @@
package common
// MergeUpdateValues merges the incoming request values into the existing
// record map (in place) and returns it.
//
// Every key present in incoming overwrites the existing value, including empty
// strings and explicit nulls, so clients can clear a column by sending "" or
// null. Keys absent from incoming are left untouched.
//
// When disallowNulls is true, nil values are skipped and the existing value is
// kept. Empty strings are still applied.
func MergeUpdateValues(existing, incoming map[string]interface{}, disallowNulls bool) map[string]interface{} {
if existing == nil {
existing = make(map[string]interface{}, len(incoming))
}
for key, newValue := range incoming {
if newValue == nil && disallowNulls {
continue
}
existing[key] = newValue
}
return existing
}
+29
View File
@@ -0,0 +1,29 @@
package common
import "testing"
func TestMergeUpdateValues(t *testing.T) {
newExisting := func() map[string]interface{} {
return map[string]interface{}{"name": "old", "note": "keep", "other": "x"}
}
incoming := map[string]interface{}{"name": "", "note": nil}
got := MergeUpdateValues(newExisting(), incoming, false)
if got["name"] != "" {
t.Errorf("name = %v, want empty string", got["name"])
}
if v, ok := got["note"]; !ok || v != nil {
t.Errorf("note = %v (present=%v), want nil", v, ok)
}
if got["other"] != "x" {
t.Errorf("absent key changed: %v", got["other"])
}
got = MergeUpdateValues(newExisting(), incoming, true)
if got["name"] != "" {
t.Errorf("disallowNulls: name = %v, want empty string", got["name"])
}
if got["note"] != "keep" {
t.Errorf("disallowNulls: note = %v, want keep", got["note"])
}
}
+4
View File
@@ -55,6 +55,10 @@ type DBConnectionConfig struct {
SSLMode string `mapstructure:"sslmode"` // disable, require, verify-ca, verify-full SSLMode string `mapstructure:"sslmode"` // disable, require, verify-ca, verify-full
Schema string `mapstructure:"schema"` // Default schema Schema string `mapstructure:"schema"` // Default schema
// ApplicationName identifies this client to the server (postgres
// application_name, mssql app name, mongodb appName)
ApplicationName string `mapstructure:"application_name"`
// SQLite specific // SQLite specific
FilePath string `mapstructure:"filepath"` FilePath string `mapstructure:"filepath"`
+1
View File
@@ -271,6 +271,7 @@ db, _ := mgr.GetDefaultDatabase()
| `database` | string | Database name | | `database` | string | Database name |
| `sslmode` | string | SSL mode (postgres/mssql): `disable`, `require`, etc. | | `sslmode` | string | SSL mode (postgres/mssql): `disable`, `require`, etc. |
| `schema` | string | Default schema (postgres/mssql) | | `schema` | string | Default schema (postgres/mssql) |
| `application_name` | string | Client name shown by the server (postgres `application_name`, mssql `app name`, mongodb `appName`); defaults to `ResolveSpec` |
| `filepath` | string | File path (sqlite only) | | `filepath` | string | File path (sqlite only) |
| `auth_source` | string | Auth source (mongodb) | | `auth_source` | string | Auth source (mongodb) |
| `replica_set` | string | Replica set name (mongodb) | | `replica_set` | string | Replica set name (mongodb) |
+22
View File
@@ -72,6 +72,9 @@ type ManagerConfig struct {
EnableAutoReconnect bool `mapstructure:"enable_auto_reconnect"` EnableAutoReconnect bool `mapstructure:"enable_auto_reconnect"`
} }
// DefaultApplicationName is used when a connection does not set ApplicationName.
const DefaultApplicationName = "ResolveSpec"
// ConnectionConfig defines configuration for a single database connection // ConnectionConfig defines configuration for a single database connection
type ConnectionConfig struct { type ConnectionConfig struct {
// Name is the unique name of this connection // Name is the unique name of this connection
@@ -95,6 +98,10 @@ type ConnectionConfig struct {
SSLMode string `mapstructure:"sslmode"` // disable, require, verify-ca, verify-full SSLMode string `mapstructure:"sslmode"` // disable, require, verify-ca, verify-full
Schema string `mapstructure:"schema"` // Default schema Schema string `mapstructure:"schema"` // Default schema
// ApplicationName identifies this client to the server (postgres
// application_name, mssql app name, mongodb appName)
ApplicationName string `mapstructure:"application_name"`
// SQLite specific // SQLite specific
FilePath string `mapstructure:"filepath"` FilePath string `mapstructure:"filepath"`
@@ -225,6 +232,10 @@ func (cc *ConnectionConfig) ApplyDefaults(global *ManagerConfig) {
cc.ConnMaxIdleTime = &idleTime cc.ConnMaxIdleTime = &idleTime
} }
if cc.ApplicationName == "" {
cc.ApplicationName = DefaultApplicationName
}
// Default timeouts // Default timeouts
if cc.ConnectTimeout == 0 { if cc.ConnectTimeout == 0 {
cc.ConnectTimeout = 10 * time.Second cc.ConnectTimeout = 10 * time.Second
@@ -347,6 +358,9 @@ func (cc *ConnectionConfig) buildPostgresDSN() string {
if cc.Schema != "" { if cc.Schema != "" {
q.Set("search_path", cc.Schema) q.Set("search_path", cc.Schema)
} }
if cc.ApplicationName != "" {
q.Set("application_name", cc.ApplicationName)
}
u := url.URL{ u := url.URL{
Scheme: "postgres", Scheme: "postgres",
@@ -405,6 +419,9 @@ func (cc *ConnectionConfig) buildMSSQLDSN() string {
if cc.Schema != "" { if cc.Schema != "" {
q.Set("schema", cc.Schema) q.Set("schema", cc.Schema)
} }
if cc.ApplicationName != "" {
q.Set("app name", cc.ApplicationName)
}
if cc.ConnectTimeout > 0 { if cc.ConnectTimeout > 0 {
sec := strconv.Itoa(int(cc.ConnectTimeout.Seconds())) sec := strconv.Itoa(int(cc.ConnectTimeout.Seconds()))
q.Set("connection timeout", sec) q.Set("connection timeout", sec)
@@ -437,6 +454,9 @@ func (cc *ConnectionConfig) buildMongoDSN() string {
if cc.ReadPreference != "" { if cc.ReadPreference != "" {
q.Set("readPreference", cc.ReadPreference) q.Set("readPreference", cc.ReadPreference)
} }
if cc.ApplicationName != "" {
q.Set("appName", cc.ApplicationName)
}
u := url.URL{ u := url.URL{
Scheme: "mongodb", Scheme: "mongodb",
@@ -480,6 +500,7 @@ func FromConfig(cfg config.DBManagerConfig) ManagerConfig {
Database: connCfg.Database, Database: connCfg.Database,
SSLMode: connCfg.SSLMode, SSLMode: connCfg.SSLMode,
Schema: connCfg.Schema, Schema: connCfg.Schema,
ApplicationName: connCfg.ApplicationName,
FilePath: connCfg.FilePath, FilePath: connCfg.FilePath,
AuthSource: connCfg.AuthSource, AuthSource: connCfg.AuthSource,
ReplicaSet: connCfg.ReplicaSet, ReplicaSet: connCfg.ReplicaSet,
@@ -519,6 +540,7 @@ func (cc *ConnectionConfig) GetConnMaxIdleTime() *time.Duration { return cc.Conn
func (cc *ConnectionConfig) GetQueryTimeout() time.Duration { return cc.QueryTimeout } func (cc *ConnectionConfig) GetQueryTimeout() time.Duration { return cc.QueryTimeout }
func (cc *ConnectionConfig) GetEnableMetrics() bool { return cc.EnableMetrics } func (cc *ConnectionConfig) GetEnableMetrics() bool { return cc.EnableMetrics }
func (cc *ConnectionConfig) GetReadPreference() string { return cc.ReadPreference } func (cc *ConnectionConfig) GetReadPreference() string { return cc.ReadPreference }
func (cc *ConnectionConfig) GetApplicationName() string { return cc.ApplicationName }
func (cc *ConnectionConfig) GetRetryAttempts() int { return cc.RetryAttempts } func (cc *ConnectionConfig) GetRetryAttempts() int { return cc.RetryAttempts }
func (cc *ConnectionConfig) GetRetryDelay() time.Duration { return cc.RetryDelay } func (cc *ConnectionConfig) GetRetryDelay() time.Duration { return cc.RetryDelay }
func (cc *ConnectionConfig) GetRetryMaxDelay() time.Duration { return cc.RetryMaxDelay } func (cc *ConnectionConfig) GetRetryMaxDelay() time.Duration { return cc.RetryMaxDelay }
+30
View File
@@ -107,3 +107,33 @@ func TestSQLiteMemoryPoolPinned(t *testing.T) {
} }
} }
} }
func TestApplicationNameInDSN(t *testing.T) {
cc := ConnectionConfig{Host: "h", Port: 1, Database: "d", ApplicationName: "my app&x=y"}
for name, c := range map[string]struct{ dsn, key string }{
"postgres": {cc.buildPostgresDSN(), "application_name"},
"mssql": {cc.buildMSSQLDSN(), "app name"},
"mongo": {cc.buildMongoDSN(), "appName"},
} {
u, err := url.Parse(c.dsn)
if err != nil {
t.Fatalf("%s: %v", name, err)
}
if got := u.Query().Get(c.key); got != cc.ApplicationName {
t.Errorf("%s: %s = %q, want %q", name, c.key, got, cc.ApplicationName)
}
}
}
func TestApplicationNameDefault(t *testing.T) {
cc := ConnectionConfig{Type: DatabaseTypePostgreSQL, Host: "h", Database: "d"}
cc.ApplyDefaults(nil)
if cc.ApplicationName != "ResolveSpec" {
t.Errorf("ApplicationName = %q, want ResolveSpec", cc.ApplicationName)
}
cc = ConnectionConfig{Type: DatabaseTypePostgreSQL, Host: "h", Database: "d", ApplicationName: "x"}
cc.ApplyDefaults(nil)
if cc.ApplicationName != "x" {
t.Errorf("explicit ApplicationName overwritten: %q", cc.ApplicationName)
}
}
+6
View File
@@ -123,6 +123,12 @@ func buildPGXConfig(cfg ConnectionConfig) (*pgx.ConnConfig, error) {
} }
cc.DialFunc = newDialFunc(cfg.GetConnectTimeout()) cc.DialFunc = newDialFunc(cfg.GetConnectTimeout())
// Also applies to caller-supplied DSNs that do not set application_name.
if name := cfg.GetApplicationName(); name != "" {
if _, set := cc.RuntimeParams["application_name"]; !set {
cc.RuntimeParams["application_name"] = name
}
}
if cfg.GetQueryTimeout() > 0 { if cfg.GetQueryTimeout() > 0 {
if _, set := cc.RuntimeParams["statement_timeout"]; !set { if _, set := cc.RuntimeParams["statement_timeout"]; !set {
cc.RuntimeParams["statement_timeout"] = fmt.Sprintf("%d", cfg.GetQueryTimeout().Milliseconds()) cc.RuntimeParams["statement_timeout"] = fmt.Sprintf("%d", cfg.GetQueryTimeout().Milliseconds())
+1
View File
@@ -59,6 +59,7 @@ type ConnectionConfig interface {
GetConnMaxLifetime() *time.Duration GetConnMaxLifetime() *time.Duration
GetConnMaxIdleTime() *time.Duration GetConnMaxIdleTime() *time.Duration
GetReadPreference() string GetReadPreference() string
GetApplicationName() string
GetRetryAttempts() int GetRetryAttempts() int
GetRetryDelay() time.Duration GetRetryDelay() time.Duration
GetRetryMaxDelay() time.Duration GetRetryMaxDelay() time.Duration
+2
View File
@@ -449,11 +449,13 @@ no per-client labels so cardinality stays bounded.
| Metric | Type | Meaning | | Metric | Type | Meaning |
|---|---|---| |---|---|---|
| `clientqueue_requests_total{result}` | counter | `immediate`, `queued`, `rejected_full`, `timeout`, `canceled` | | `clientqueue_requests_total{result}` | counter | `immediate`, `queued`, `rejected_full`, `timeout`, `canceled` |
| `clientqueue_enqueued_total` | counter | requests ever placed in a wait queue, whatever happened next (ran, timed out, cancelled) |
| `clientqueue_wait_seconds` | histogram | wait for a slot; 0 for requests that ran immediately | | `clientqueue_wait_seconds` | histogram | wait for a slot; 0 for requests that ran immediately |
| `clientqueue_wait_max_seconds` | gauge | longest wait since process start | | `clientqueue_wait_max_seconds` | gauge | longest wait since process start |
| `clientqueue_burst_size` | histogram | peak outstanding requests per client busy period | | `clientqueue_burst_size` | histogram | peak outstanding requests per client busy period |
| `clientqueue_burst_max` | gauge | largest burst since process start | | `clientqueue_burst_max` | gauge | largest burst since process start |
| `clientqueue_active` / `clientqueue_queue_depth` | gauge | running / waiting now | | `clientqueue_active` / `clientqueue_queue_depth` | gauge | running / waiting now |
| `clientqueue_waiting_clients` | gauge | clients with at least one request waiting (`queue_depth` counts requests, this counts clients) |
| `clientqueue_clients` | gauge | clients currently tracked | | `clientqueue_clients` | gauge | clients currently tracked |
A **burst** is one client's busy period: the peak number of its requests running plus waiting between A **burst** is one client's busy period: the peak number of its requests running plus waiting between
+20
View File
@@ -70,6 +70,16 @@ var (
Help: "Requests currently waiting for a slot", Help: "Requests currently waiting for a slot",
}) })
queueEnqueued = promauto.NewCounter(prometheus.CounterOpts{
Name: "clientqueue_enqueued_total",
Help: "Requests ever placed in a wait queue, whatever their outcome (ran, timed out or cancelled)",
})
queueWaitingClients = promauto.NewGauge(prometheus.GaugeOpts{
Name: "clientqueue_waiting_clients",
Help: "Clients currently with at least one request waiting for a slot",
})
queueClients = promauto.NewGauge(prometheus.GaugeOpts{ queueClients = promauto.NewGauge(prometheus.GaugeOpts{
Name: "clientqueue_clients", Name: "clientqueue_clients",
Help: "Clients currently tracked by the queue", Help: "Clients currently tracked by the queue",
@@ -275,9 +285,13 @@ func (q *ClientQueue) acquire(ctx context.Context, key string) error {
} }
w := &queueWaiter{ready: make(chan struct{})} w := &queueWaiter{ready: make(chan struct{})}
w.elem = c.waiters.PushBack(w) w.elem = c.waiters.PushBack(w)
if c.waiters.Len() == 1 {
queueWaitingClients.Inc()
}
c.enter() c.enter()
q.mu.Unlock() q.mu.Unlock()
queueDepth.Inc() queueDepth.Inc()
queueEnqueued.Inc()
timer := time.NewTimer(q.cfg.MaxWait) timer := time.NewTimer(q.cfg.MaxWait)
defer timer.Stop() defer timer.Stop()
@@ -302,6 +316,9 @@ func (q *ClientQueue) acquire(ctx context.Context, key string) error {
q.releaseLocked(key) q.releaseLocked(key)
} else { } else {
c.waiters.Remove(w.elem) c.waiters.Remove(w.elem)
if c.waiters.Len() == 0 {
queueWaitingClients.Dec()
}
c.leave() c.leave()
queueDepth.Dec() queueDepth.Dec()
} }
@@ -331,6 +348,9 @@ func (q *ClientQueue) releaseLocked(key string) {
queueActive.Dec() queueActive.Dec()
if front := c.waiters.Front(); front != nil { if front := c.waiters.Front(); front != nil {
w := c.waiters.Remove(front).(*queueWaiter) w := c.waiters.Remove(front).(*queueWaiter)
if c.waiters.Len() == 0 {
queueWaitingClients.Dec()
}
w.granted = true w.granted = true
queueDepth.Dec() queueDepth.Dec()
queueActive.Inc() queueActive.Inc()
+101
View File
@@ -2,6 +2,7 @@ package middleware
import ( import (
"context" "context"
"errors"
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
"strings" "strings"
@@ -278,6 +279,7 @@ func TestClientQueueMetrics(t *testing.T) {
q := newTestQueue(t, ClientQueueConfig{MaxConcurrent: 2}) q := newTestQueue(t, ClientQueueConfig{MaxConcurrent: 2})
imm0 := gaugeVal(t, queueRequests.WithLabelValues("immediate")) imm0 := gaugeVal(t, queueRequests.WithLabelValues("immediate"))
que0 := gaugeVal(t, queueRequests.WithLabelValues("queued")) que0 := gaugeVal(t, queueRequests.WithLabelValues("queued"))
enq0 := gaugeVal(t, queueEnqueued)
bc0, bs0 := histVal(t, queueBurst) bc0, bs0 := histVal(t, queueBurst)
wc0, _ := histVal(t, queueWait) wc0, _ := histVal(t, queueWait)
@@ -317,6 +319,9 @@ func TestClientQueueMetrics(t *testing.T) {
if got := gaugeVal(t, queueRequests.WithLabelValues("immediate")) - imm0; got != 2 { if got := gaugeVal(t, queueRequests.WithLabelValues("immediate")) - imm0; got != 2 {
t.Errorf("immediate = %v, want 2", got) t.Errorf("immediate = %v, want 2", got)
} }
if got := gaugeVal(t, queueEnqueued) - enq0; got != 3 {
t.Errorf("enqueued = %v, want 3", got)
}
if got := gaugeVal(t, queueRequests.WithLabelValues("queued")) - que0; got != 3 { if got := gaugeVal(t, queueRequests.WithLabelValues("queued")) - que0; got != 3 {
t.Errorf("queued = %v, want 3", got) t.Errorf("queued = %v, want 3", got)
} }
@@ -331,3 +336,99 @@ func TestClientQueueMetrics(t *testing.T) {
t.Errorf("burst max = %v, want >= 5", got) t.Errorf("burst max = %v, want >= 5", got)
} }
} }
func TestClientQueueWaitingClientsGauge(t *testing.T) {
q := newTestQueue(t, ClientQueueConfig{MaxConcurrent: 1, MaxWait: 100 * time.Millisecond})
base := gaugeVal(t, queueWaitingClients)
waiting := func() float64 { return gaugeVal(t, queueWaitingClients) - base }
waitFor := func(want float64) {
t.Helper()
deadline := time.Now().Add(2 * time.Second)
for waiting() != want {
if time.Now().After(deadline) {
t.Fatalf("waiting clients = %v, want %v", waiting(), want)
}
time.Sleep(time.Millisecond)
}
}
// Two clients each hold their only slot.
for _, k := range []string{"a", "b"} {
if err := q.acquire(t.Context(), k); err != nil {
t.Fatal(err)
}
}
if waiting() != 0 {
t.Fatalf("waiting = %v with nobody queued", waiting())
}
// Two waiters for "a" count as one waiting client, not two.
bg := func(ctx context.Context, k string) chan error {
ch := make(chan error, 1)
go func() { ch <- q.acquire(ctx, k) }()
return ch
}
a1 := bg(t.Context(), "a")
waitFor(1)
a2 := bg(t.Context(), "a")
time.Sleep(10 * time.Millisecond)
waitFor(1)
// A waiter for "b" adds a second waiting client, and cancelling it drops it.
cctx, cancel := context.WithCancel(t.Context())
b1 := bg(cctx, "b")
waitFor(2)
cancel()
if err := <-b1; err == nil {
t.Fatal("expected cancellation")
}
waitFor(1)
// Draining "a": the first release hands over to a1 (one still waits),
// the second hands over to a2 and "a" stops waiting.
q.release("a")
if err := <-a1; err != nil {
t.Fatal(err)
}
waitFor(1)
q.release("a")
if err := <-a2; err != nil {
t.Fatal(err)
}
waitFor(0)
q.release("a")
q.release("a")
// Timeout: "b" still holds its slot, so a new waiter times out.
b2 := bg(t.Context(), "b")
waitFor(1)
if err := <-b2; !errors.Is(err, errQueueWait) {
t.Fatalf("err = %v, want timeout", err)
}
waitFor(0)
q.release("b")
}
func TestClientQueueEnqueuedCountsAllOutcomes(t *testing.T) {
q := newTestQueue(t, ClientQueueConfig{MaxConcurrent: 1, MaxQueue: 1, MaxWait: 50 * time.Millisecond})
base := gaugeVal(t, queueEnqueued)
if err := q.acquire(t.Context(), "k"); err != nil { // immediate: not enqueued
t.Fatal(err)
}
if err := q.acquire(t.Context(), "k"); !errors.Is(err, errQueueWait) { // enqueued, times out
t.Fatalf("err = %v, want timeout", err)
}
cctx, cancel := context.WithCancel(t.Context())
done := make(chan error, 1)
go func() { done <- q.acquire(cctx, "k") }() // enqueued, cancelled
time.Sleep(10 * time.Millisecond)
if err := q.acquire(t.Context(), "k"); !errors.Is(err, errQueueFull) { // rejected: not enqueued
t.Fatalf("err = %v, want queue full", err)
}
cancel()
<-done
q.release("k")
if got := gaugeVal(t, queueEnqueued) - base; got != 2 {
t.Fatalf("enqueued = %v, want 2 (timeout + cancel; not immediate or rejected)", got)
}
}
+28 -9
View File
@@ -45,6 +45,9 @@ type Handler struct {
// Started flag // Started flag
started bool started bool
mu sync.RWMutex mu sync.RWMutex
// disallowNulls skips null values in update payloads instead of applying them
disallowNulls bool
} }
// NewHandler creates a new MQTT handler // NewHandler creates a new MQTT handler
@@ -812,6 +815,14 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata
} }
// create creates a new record // create creates a new record
// SetDisallowNulls controls whether explicit null values in update payloads are
// ignored. By default a key present in the payload overwrites the stored value,
// including "" and null. When true, null values are skipped and the existing
// value is kept ("" still overwrites).
func (h *Handler) SetDisallowNulls(disallow bool) {
h.disallowNulls = disallow
}
func (h *Handler) create(hookCtx *HookContext) (interface{}, error) { func (h *Handler) create(hookCtx *HookContext) (interface{}, error) {
// Marshal and unmarshal data into model // Marshal and unmarshal data into model
dataBytes, err := json.Marshal(hookCtx.Data) dataBytes, err := json.Marshal(hookCtx.Data)
@@ -840,26 +851,34 @@ func (h *Handler) create(hookCtx *HookContext) (interface{}, error) {
// update updates an existing record // update updates an existing record
func (h *Handler) update(hookCtx *HookContext) (interface{}, error) { func (h *Handler) update(hookCtx *HookContext) (interface{}, error) {
// Marshal and unmarshal data into model // Convert request data to a map
var updates map[string]interface{}
if m, ok := hookCtx.Data.(map[string]interface{}); ok {
updates = m
} else {
dataBytes, err := json.Marshal(hookCtx.Data) dataBytes, err := json.Marshal(hookCtx.Data)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to marshal data: %w", err) return nil, fmt.Errorf("failed to marshal data: %w", err)
} }
if err := json.Unmarshal(dataBytes, &updates); err != nil {
if err := json.Unmarshal(dataBytes, hookCtx.ModelPtr); err != nil { return nil, fmt.Errorf("failed to unmarshal data into map: %w", err)
return nil, fmt.Errorf("failed to unmarshal data into model: %w", err) }
} }
// Update record
query := h.db.NewUpdate().Model(hookCtx.ModelPtr).Table(hookCtx.TableName)
// Add ID filter
pkName := reflection.GetPrimaryKeyName(hookCtx.Model) pkName := reflection.GetPrimaryKeyName(hookCtx.Model)
query = query.Where(fmt.Sprintf("%s = ?", pkName), hookCtx.ID)
// Only the keys present in the request are written. "" and null overwrite
// the stored value unless disallowNulls is set, in which case null is skipped.
values := common.MergeUpdateValues(make(map[string]interface{}, len(updates)), updates, h.disallowNulls)
if len(values) > 0 {
query := h.db.NewUpdate().Table(hookCtx.TableName).SetMap(values).
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), hookCtx.ID)
if _, err := query.Exec(hookCtx.Context); err != nil { if _, err := query.Exec(hookCtx.Context); err != nil {
return nil, fmt.Errorf("failed to update record: %w", err) return nil, fmt.Errorf("failed to update record: %w", err)
} }
}
// Fetch updated record // Fetch updated record
return h.readByID(hookCtx) return h.readByID(hookCtx)
+15 -35
View File
@@ -32,6 +32,7 @@ type Handler struct {
fallbackHandler FallbackHandler fallbackHandler FallbackHandler
openAPIGenerator func() (string, error) openAPIGenerator func() (string, error)
defaultSort map[string][]common.SortOption defaultSort map[string][]common.SortOption
disallowNulls bool
} }
// NewHandler creates a new API handler with database and registry abstractions // NewHandler creates a new API handler with database and registry abstractions
@@ -52,6 +53,14 @@ func (h *Handler) Hooks() *HookRegistry {
return h.hooks return h.hooks
} }
// SetDisallowNulls controls whether explicit null values in update payloads are
// ignored. By default a key present in the payload overwrites the stored value,
// including "" and null. When true, null values are skipped and the existing
// value is kept ("" still overwrites).
func (h *Handler) SetDisallowNulls(disallow bool) {
h.disallowNulls = disallow
}
// SetFallbackHandler sets a fallback handler to be called when no model is found // SetFallbackHandler sets a fallback handler to be called when no model is found
// If not set, the handler will simply return (pass through to next route) // If not set, the handler will simply return (pass through to next route)
func (h *Handler) SetFallbackHandler(fallback FallbackHandler) { func (h *Handler) SetFallbackHandler(fallback FallbackHandler) {
@@ -1236,21 +1245,8 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
return fmt.Errorf("error unmarshaling existing record: %w", err) return fmt.Errorf("error unmarshaling existing record: %w", err)
} }
// Merge only non-null and non-empty values from the incoming request into the existing record // Overwrite with every key present in the request (including "" and null unless disallowed)
for key, newValue := range updates { common.MergeUpdateValues(existingMap, updates, h.disallowNulls)
// Skip if the value is nil
if newValue == nil {
continue
}
// Skip if the value is an empty string
if strVal, ok := newValue.(string); ok && strVal == "" {
continue
}
// Update the existing map with the new value
existingMap[key] = newValue
}
// Build update query with merged data // Build update query with merged data
query := tx.NewUpdate().Table(tableName).SetMap(existingMap) query := tx.NewUpdate().Table(tableName).SetMap(existingMap)
@@ -1421,16 +1417,8 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
item = modifiedData item = modifiedData
} }
// Merge only non-null and non-empty values // Overwrite with every key present in the request (including "" and null unless disallowed)
for key, newValue := range item { common.MergeUpdateValues(existingMap, item, h.disallowNulls)
if newValue == nil {
continue
}
if strVal, ok := newValue.(string); ok && strVal == "" {
continue
}
existingMap[key] = newValue
}
txQuery := tx.NewUpdate().Table(tableName).SetMap(existingMap).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID) txQuery := tx.NewUpdate().Table(tableName).SetMap(existingMap).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID)
if _, err := txQuery.Exec(ctx); err != nil { if _, err := txQuery.Exec(ctx); err != nil {
@@ -1578,16 +1566,8 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
itemMap = modifiedData itemMap = modifiedData
} }
// Merge only non-null and non-empty values // Overwrite with every key present in the request (including "" and null unless disallowed)
for key, newValue := range itemMap { common.MergeUpdateValues(existingMap, itemMap, h.disallowNulls)
if newValue == nil {
continue
}
if strVal, ok := newValue.(string); ok && strVal == "" {
continue
}
existingMap[key] = newValue
}
txQuery := tx.NewUpdate().Table(tableName).SetMap(existingMap).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID) txQuery := tx.NewUpdate().Table(tableName).SetMap(existingMap).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID)
if _, err := txQuery.Exec(ctx); err != nil { if _, err := txQuery.Exec(ctx); err != nil {
+11 -15
View File
@@ -33,6 +33,7 @@ type Handler struct {
fallbackHandler FallbackHandler fallbackHandler FallbackHandler
openAPIGenerator func() (string, error) openAPIGenerator func() (string, error)
defaultSort map[string][]common.SortOption defaultSort map[string][]common.SortOption
disallowNulls bool
} }
// NewHandler creates a new API handler with database and registry abstractions // NewHandler creates a new API handler with database and registry abstractions
@@ -59,6 +60,14 @@ func (h *Handler) Hooks() *HookRegistry {
return h.hooks return h.hooks
} }
// SetDisallowNulls controls whether explicit null values in update payloads are
// ignored. By default a key present in the payload overwrites the stored value,
// including "" and null. When true, null values are skipped and the existing
// value is kept ("" still overwrites).
func (h *Handler) SetDisallowNulls(disallow bool) {
h.disallowNulls = disallow
}
// SetFallbackHandler sets a fallback handler to be called when no model is found // SetFallbackHandler sets a fallback handler to be called when no model is found
// If not set, the handler will simply return (pass through to next route) // If not set, the handler will simply return (pass through to next route)
func (h *Handler) SetFallbackHandler(fallback FallbackHandler) { func (h *Handler) SetFallbackHandler(fallback FallbackHandler) {
@@ -1597,21 +1606,8 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
nestedRelations = relations nestedRelations = relations
} }
// Merge only non-null and non-empty values from the incoming request into the existing record // Overwrite with every key present in the request (including "" and null unless disallowed)
for key, newValue := range dataMap { common.MergeUpdateValues(existingMap, dataMap, h.disallowNulls)
// Skip if the value is nil
if newValue == nil {
continue
}
// Skip if the value is an empty string
if strVal, ok := newValue.(string); ok && strVal == "" {
continue
}
// Update the existing map with the new value
existingMap[key] = newValue
}
// Ensure ID is in the data map for the update // Ensure ID is in the data map for the update
existingMap[pkName] = targetID existingMap[pkName] = targetID
+28 -9
View File
@@ -38,6 +38,9 @@ type Handler struct {
subscriptionManager *SubscriptionManager subscriptionManager *SubscriptionManager
upgrader websocket.Upgrader upgrader websocket.Upgrader
ctx context.Context ctx context.Context
// disallowNulls skips null values in update payloads instead of applying them
disallowNulls bool
} }
// NewHandler creates a new WebSocket handler // NewHandler creates a new WebSocket handler
@@ -682,6 +685,14 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata
return hookCtx.ModelPtr, metadata, nil return hookCtx.ModelPtr, metadata, nil
} }
// SetDisallowNulls controls whether explicit null values in update payloads are
// ignored. By default a key present in the payload overwrites the stored value,
// including "" and null. When true, null values are skipped and the existing
// value is kept ("" still overwrites).
func (h *Handler) SetDisallowNulls(disallow bool) {
h.disallowNulls = disallow
}
func (h *Handler) create(hookCtx *HookContext) (interface{}, error) { func (h *Handler) create(hookCtx *HookContext) (interface{}, error) {
// Marshal and unmarshal data into model // Marshal and unmarshal data into model
dataBytes, err := json.Marshal(hookCtx.Data) dataBytes, err := json.Marshal(hookCtx.Data)
@@ -709,26 +720,34 @@ func (h *Handler) create(hookCtx *HookContext) (interface{}, error) {
} }
func (h *Handler) update(hookCtx *HookContext) (interface{}, error) { func (h *Handler) update(hookCtx *HookContext) (interface{}, error) {
// Marshal and unmarshal data into model // Convert request data to a map
var updates map[string]interface{}
if m, ok := hookCtx.Data.(map[string]interface{}); ok {
updates = m
} else {
dataBytes, err := json.Marshal(hookCtx.Data) dataBytes, err := json.Marshal(hookCtx.Data)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to marshal data: %w", err) return nil, fmt.Errorf("failed to marshal data: %w", err)
} }
if err := json.Unmarshal(dataBytes, &updates); err != nil {
if err := json.Unmarshal(dataBytes, hookCtx.ModelPtr); err != nil { return nil, fmt.Errorf("failed to unmarshal data into map: %w", err)
return nil, fmt.Errorf("failed to unmarshal data into model: %w", err) }
} }
// Update record
query := h.db.NewUpdate().Model(hookCtx.ModelPtr).Table(hookCtx.TableName)
// Add ID filter
pkName := reflection.GetPrimaryKeyName(hookCtx.Model) pkName := reflection.GetPrimaryKeyName(hookCtx.Model)
query = query.Where(fmt.Sprintf("%s = ?", pkName), hookCtx.ID)
// Only the keys present in the request are written. "" and null overwrite
// the stored value unless disallowNulls is set, in which case null is skipped.
values := common.MergeUpdateValues(make(map[string]interface{}, len(updates)), updates, h.disallowNulls)
if len(values) > 0 {
query := h.db.NewUpdate().Table(hookCtx.TableName).SetMap(values).
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), hookCtx.ID)
if _, err := query.Exec(hookCtx.Context); err != nil { if _, err := query.Exec(hookCtx.Context); err != nil {
return nil, fmt.Errorf("failed to update record: %w", err) return nil, fmt.Errorf("failed to update record: %w", err)
} }
}
// Fetch updated record // Fetch updated record
return h.readByID(hookCtx) return h.readByID(hookCtx)