mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-01 12:31:59 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
2898b335f8 | ||
|
|
20ba8ed112 | ||
|
|
1214f69e0c | ||
|
|
89a58ab3a0 |
@@ -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
|
||||||
|
}
|
||||||
@@ -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"])
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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"`
|
||||||
|
|
||||||
|
|||||||
@@ -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) |
|
||||||
|
|||||||
@@ -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 }
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -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())
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+34
-15
@@ -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,25 +851,33 @@ 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
|
||||||
dataBytes, err := json.Marshal(hookCtx.Data)
|
var updates map[string]interface{}
|
||||||
if err != nil {
|
if m, ok := hookCtx.Data.(map[string]interface{}); ok {
|
||||||
return nil, fmt.Errorf("failed to marshal data: %w", err)
|
updates = m
|
||||||
|
} else {
|
||||||
|
dataBytes, err := json.Marshal(hookCtx.Data)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to marshal data: %w", err)
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(dataBytes, &updates); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to unmarshal data into map: %w", err)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := json.Unmarshal(dataBytes, hookCtx.ModelPtr); err != nil {
|
|
||||||
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)
|
|
||||||
|
|
||||||
if _, err := query.Exec(hookCtx.Context); err != nil {
|
// Only the keys present in the request are written. "" and null overwrite
|
||||||
return nil, fmt.Errorf("failed to update record: %w", err)
|
// 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 {
|
||||||
|
return nil, fmt.Errorf("failed to update record: %w", err)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Fetch updated record
|
// Fetch updated record
|
||||||
|
|||||||
+15
-35
@@ -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
@@ -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
|
||||||
|
|||||||
@@ -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,25 +720,33 @@ 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
|
||||||
dataBytes, err := json.Marshal(hookCtx.Data)
|
var updates map[string]interface{}
|
||||||
if err != nil {
|
if m, ok := hookCtx.Data.(map[string]interface{}); ok {
|
||||||
return nil, fmt.Errorf("failed to marshal data: %w", err)
|
updates = m
|
||||||
|
} else {
|
||||||
|
dataBytes, err := json.Marshal(hookCtx.Data)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to marshal data: %w", err)
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(dataBytes, &updates); err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to unmarshal data into map: %w", err)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := json.Unmarshal(dataBytes, hookCtx.ModelPtr); err != nil {
|
|
||||||
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)
|
|
||||||
|
|
||||||
if _, err := query.Exec(hookCtx.Context); err != nil {
|
// Only the keys present in the request are written. "" and null overwrite
|
||||||
return nil, fmt.Errorf("failed to update record: %w", err)
|
// 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 {
|
||||||
|
return nil, fmt.Errorf("failed to update record: %w", err)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Fetch updated record
|
// Fetch updated record
|
||||||
|
|||||||
Reference in New Issue
Block a user