mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-01 12:31:59 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
20ba8ed112 | ||
|
|
1214f69e0c | ||
|
|
89a58ab3a0 | ||
|
|
48081b4aa4 | ||
|
|
467dbd66c8 | ||
|
|
178d40587d |
@@ -579,7 +579,7 @@ Centralized management of multiple database connections with support for Postgre
|
||||
- Background health checks (report status; they never close the pool)
|
||||
- Prometheus metrics for monitoring
|
||||
- Configuration-driven via YAML
|
||||
- Per-connection statistics and management
|
||||
- Per-connection statistics and management, including pool limit (`max`) and total connections ever opened (`dbmanager_connection_pool_size{state="max"}`, `dbmanager_connections_opened_total`)
|
||||
|
||||
**How to use it correctly**:
|
||||
|
||||
@@ -618,7 +618,15 @@ For documentation, see [pkg/security/README.md](pkg/security/README.md) (see "Di
|
||||
|
||||
#### Middleware
|
||||
|
||||
HTTP middleware collection for common tasks (CORS, logging, metrics, etc.).
|
||||
HTTP middleware collection for common tasks (CORS, logging, metrics, rate limiting, etc.).
|
||||
|
||||
**Client request queue** (`middleware.ClientQueue`): limits how many requests each client runs concurrently and queues the rest first-in-first-out, smoothing bursts such as a page load that fires ~15 requests at once. Clients are identified by `X-Client-Id`, then `Authorization`, then the built-in session, then IP, so no client changes are required. Exposes burst, wait and queue-depth Prometheus metrics. Add it through the middleware slot of `SetupMuxRoutes` / `SetupBunRouterRoutes`:
|
||||
|
||||
```go
|
||||
q := middleware.NewClientQueue(middleware.ClientQueueConfig{MaxConcurrent: 10})
|
||||
defer q.Close()
|
||||
restheadspec.SetupMuxRoutes(router, handler, middleware.Chain(authMiddleware, q.Middleware))
|
||||
```
|
||||
|
||||
For documentation, see [pkg/middleware/README.md](pkg/middleware/README.md).
|
||||
|
||||
@@ -657,7 +665,7 @@ For documentation, see [pkg/config/README.md](pkg/config/README.md).
|
||||
* Implement proper authentication and authorization
|
||||
* Validate all input parameters
|
||||
* Use prepared statements (handled by GORM/Bun/your ORM)
|
||||
* Implement rate limiting
|
||||
* Implement rate limiting (`middleware.RateLimiter`) and per-client request queueing (`middleware.ClientQueue`)
|
||||
* Control access at schema/entity level
|
||||
* **New**: Database abstraction layer provides additional security through interface boundaries
|
||||
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"github.com/bitechdev/ResolveSpec/pkg/config"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/dbmanager"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/middleware"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/server"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/testmodels"
|
||||
@@ -67,8 +68,13 @@ func main() {
|
||||
handler.RegisterModel("public", modelNames[i], model)
|
||||
}
|
||||
|
||||
// Queue requests per client (X-Client-Id, Authorization, session, then IP)
|
||||
// so a burst such as a page load cannot flood the database pool.
|
||||
queue := middleware.NewClientQueue(middleware.ClientQueueConfig{MaxConcurrent: 10})
|
||||
defer queue.Close()
|
||||
|
||||
// Setup routes using new SetupMuxRoutes function (without authentication)
|
||||
resolvespec.SetupMuxRoutes(r, handler, nil)
|
||||
resolvespec.SetupMuxRoutes(r, handler, middleware.Chain(queue.Middleware))
|
||||
|
||||
// Create server manager
|
||||
mgr := server.NewManager()
|
||||
|
||||
@@ -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"])
|
||||
}
|
||||
}
|
||||
@@ -319,9 +319,10 @@ fmt.Printf("Healthy: %d, Unhealthy: %d\n", stats.HealthyCount, stats.UnhealthyCo
|
||||
|
||||
// Per-connection stats
|
||||
for name, connStats := range stats.ConnectionStats {
|
||||
fmt.Printf("%s: %d open, %d in use, %d idle\n",
|
||||
fmt.Printf("%s: %d open (max %d, 0 = unlimited), %d in use, %d idle\n",
|
||||
name,
|
||||
connStats.OpenConnections,
|
||||
connStats.MaxOpenConnections,
|
||||
connStats.InUse,
|
||||
connStats.Idle)
|
||||
}
|
||||
@@ -340,7 +341,8 @@ The package automatically exports Prometheus metrics:
|
||||
|
||||
- `dbmanager_connections_total` - Total configured connections by type
|
||||
- `dbmanager_connection_status` - Connection health status (1=healthy, 0=unhealthy)
|
||||
- `dbmanager_connection_pool_size` - Connection pool statistics by state
|
||||
- `dbmanager_connections_opened_total` - Physical connections ever opened (counter; for a pool wrapped via `existing_db` this is derived from open + closed counts and can undercount)
|
||||
- `dbmanager_connection_pool_size` - Connection pool statistics by state (`open`, `idle`, `in_use`, `max`; `max` 0 = unlimited)
|
||||
- `dbmanager_connection_wait_count` - Times connections waited for availability
|
||||
- `dbmanager_connection_wait_duration_seconds` - Total wait duration
|
||||
- `dbmanager_health_check_duration_seconds` - Health check execution time
|
||||
|
||||
@@ -5,6 +5,8 @@ import (
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
dto "github.com/prometheus/client_model/go"
|
||||
)
|
||||
|
||||
func TestPostgresDSNEscapesCredentials(t *testing.T) {
|
||||
@@ -83,6 +85,18 @@ func TestSQLiteMemoryPoolPinned(t *testing.T) {
|
||||
if got := db.Stats().MaxOpenConnections; got != 1 {
|
||||
t.Fatalf("MaxOpenConnections = %d, want 1 for :memory:", got)
|
||||
}
|
||||
if got := conn.Stats().MaxOpenConnections; got != 1 {
|
||||
t.Fatalf("ConnectionStats.MaxOpenConnections = %d, want 1", got)
|
||||
}
|
||||
mgr.(*connectionManager).PublishMetrics()
|
||||
name := conn.Name()
|
||||
var m dto.Metric
|
||||
if err := connectionPoolSize.WithLabelValues(name, string(conn.Stats().Type), "max").Write(&m); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := m.GetGauge().GetValue(); got != 1 {
|
||||
t.Fatalf("pool_size{state=max} = %v, want 1", got)
|
||||
}
|
||||
if _, err := db.Exec("CREATE TABLE t(a int)"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
@@ -56,6 +56,8 @@ type ConnectionStats struct {
|
||||
|
||||
// SQL connection pool stats
|
||||
OpenConnections int
|
||||
MaxOpenConnections int
|
||||
TotalOpened int64 // physical connections ever opened
|
||||
InUse int
|
||||
Idle int
|
||||
WaitCount int64
|
||||
@@ -415,6 +417,8 @@ func (c *sqlConnection) Stats() *ConnectionStats {
|
||||
if c.connected && c.provider != nil {
|
||||
if providerStats := c.provider.Stats(); providerStats != nil {
|
||||
stats.OpenConnections = providerStats.OpenConnections
|
||||
stats.MaxOpenConnections = providerStats.MaxOpenConnections
|
||||
stats.TotalOpened = providerStats.TotalOpened
|
||||
stats.InUse = providerStats.InUse
|
||||
stats.Idle = providerStats.Idle
|
||||
stats.WaitCount = providerStats.WaitCount
|
||||
|
||||
@@ -32,7 +32,7 @@ var (
|
||||
Name: "dbmanager_connection_pool_size",
|
||||
Help: "Current connection pool size",
|
||||
},
|
||||
[]string{"name", "type", "state"}, // state: open, idle, in_use
|
||||
[]string{"name", "type", "state"}, // state: open, idle, in_use, max
|
||||
)
|
||||
|
||||
// connectionWaitCount tracks how many times connections had to wait for availability
|
||||
@@ -71,6 +71,15 @@ var (
|
||||
[]string{"name", "type"},
|
||||
)
|
||||
|
||||
// connectionOpened tracks physical connections ever opened
|
||||
connectionOpened = promauto.NewCounterVec(
|
||||
prometheus.CounterOpts{
|
||||
Name: "dbmanager_connections_opened_total",
|
||||
Help: "Total physical connections opened by the pool",
|
||||
},
|
||||
[]string{"name", "type"},
|
||||
)
|
||||
|
||||
// connectionIdleClosed tracks connections closed due to max idle time
|
||||
connectionIdleClosed = promauto.NewCounterVec(
|
||||
prometheus.CounterOpts{
|
||||
@@ -115,6 +124,7 @@ func (m *connectionManager) PublishMetrics() {
|
||||
connectionPoolSize.WithLabelValues(name, string(connStats.Type), "open").Set(float64(connStats.OpenConnections))
|
||||
connectionPoolSize.WithLabelValues(name, string(connStats.Type), "idle").Set(float64(connStats.Idle))
|
||||
connectionPoolSize.WithLabelValues(name, string(connStats.Type), "in_use").Set(float64(connStats.InUse))
|
||||
connectionPoolSize.WithLabelValues(name, string(connStats.Type), "max").Set(float64(connStats.MaxOpenConnections))
|
||||
|
||||
// sql.DBStats values are cumulative, so add only the growth since
|
||||
// the last publish to keep these true counters.
|
||||
@@ -122,6 +132,7 @@ func (m *connectionManager) PublishMetrics() {
|
||||
connectionWaitCount.With(labels).Add(float64(connStats.WaitCount - prev.WaitCount))
|
||||
connectionWaitDuration.With(labels).Add((connStats.WaitDuration - prev.WaitDuration).Seconds())
|
||||
connectionLifetimeClosed.With(labels).Add(float64(connStats.MaxLifetimeClosed - prev.MaxLifetimeClosed))
|
||||
connectionOpened.With(labels).Add(float64(connStats.TotalOpened - prev.TotalOpened))
|
||||
connectionIdleClosed.With(labels).Add(float64(connStats.MaxIdleClosed - prev.MaxIdleClosed))
|
||||
}
|
||||
}
|
||||
@@ -152,7 +163,7 @@ func (p *publishedStats) swap(name string, cur *ConnectionStats) ConnectionStats
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
prev := p.last[name]
|
||||
if cur.WaitCount < prev.WaitCount || cur.MaxIdleClosed < prev.MaxIdleClosed || cur.MaxLifetimeClosed < prev.MaxLifetimeClosed {
|
||||
if cur.WaitCount < prev.WaitCount || cur.MaxIdleClosed < prev.MaxIdleClosed || cur.MaxLifetimeClosed < prev.MaxLifetimeClosed || cur.TotalOpened < prev.TotalOpened {
|
||||
prev = ConnectionStats{}
|
||||
}
|
||||
p.last[name] = *cur
|
||||
|
||||
@@ -0,0 +1,51 @@
|
||||
package providers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"database/sql/driver"
|
||||
"sync/atomic"
|
||||
)
|
||||
|
||||
// countingConnector wraps a driver.Connector and counts every physical
|
||||
// connection it successfully dials. sql.DBStats has no such field, so this is
|
||||
// the only way to report the total number of connections ever opened.
|
||||
type countingConnector struct {
|
||||
driver.Connector
|
||||
opened *atomic.Int64
|
||||
}
|
||||
|
||||
func (c *countingConnector) Connect(ctx context.Context) (driver.Conn, error) {
|
||||
conn, err := c.Connector.Connect(ctx)
|
||||
if err == nil {
|
||||
c.opened.Add(1)
|
||||
}
|
||||
return conn, err
|
||||
}
|
||||
|
||||
// dsnConnector adapts a plain driver.Driver to driver.Connector.
|
||||
type dsnConnector struct {
|
||||
dsn string
|
||||
drv driver.Driver
|
||||
}
|
||||
|
||||
func (c dsnConnector) Connect(context.Context) (driver.Conn, error) { return c.drv.Open(c.dsn) }
|
||||
func (c dsnConnector) Driver() driver.Driver { return c.drv }
|
||||
|
||||
// openCounted is sql.Open with every dialled connection counted in opened.
|
||||
func openCounted(driverName, dsn string, opened *atomic.Int64) (*sql.DB, error) {
|
||||
probe, err := sql.Open(driverName, dsn)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
drv := probe.Driver()
|
||||
probe.Close() //nolint:gosec // G104: probe handle never dialled
|
||||
|
||||
var connector driver.Connector = dsnConnector{dsn: dsn, drv: drv}
|
||||
if dc, ok := drv.(driver.DriverContext); ok {
|
||||
if connector, err = dc.OpenConnector(dsn); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
return sql.OpenDB(&countingConnector{Connector: connector, opened: opened}), nil
|
||||
}
|
||||
@@ -112,6 +112,12 @@ func (p *ExistingDBProvider) Stats() *ConnectionStats {
|
||||
if p.db != nil {
|
||||
dbStats := p.db.Stats()
|
||||
stats.OpenConnections = dbStats.OpenConnections
|
||||
stats.MaxOpenConnections = dbStats.MaxOpenConnections
|
||||
// The pool was opened outside dbmanager so dials cannot be counted.
|
||||
// Open plus every connection database/sql retired for idle or
|
||||
// lifetime limits is a close lower bound (it misses connections
|
||||
// dropped as broken).
|
||||
stats.TotalOpened = int64(dbStats.OpenConnections) + dbStats.MaxIdleClosed + dbStats.MaxIdleTimeClosed + dbStats.MaxLifetimeClosed
|
||||
stats.InUse = dbStats.InUse
|
||||
stats.Idle = dbStats.Idle
|
||||
stats.WaitCount = dbStats.WaitCount
|
||||
|
||||
@@ -3,9 +3,11 @@ package providers
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
_ "github.com/glebarez/sqlite"
|
||||
_ "github.com/mattn/go-sqlite3"
|
||||
)
|
||||
|
||||
@@ -159,6 +161,10 @@ func TestExistingDBProvider_Stats(t *testing.T) {
|
||||
t.Errorf("Expected stats.Type to be 'sql', got '%s'", stats.Type)
|
||||
}
|
||||
|
||||
if stats.MaxOpenConnections != 10 {
|
||||
t.Errorf("Expected stats.MaxOpenConnections to be 10, got %d", stats.MaxOpenConnections)
|
||||
}
|
||||
|
||||
if !stats.Connected {
|
||||
t.Error("Expected stats.Connected to be true")
|
||||
}
|
||||
@@ -192,3 +198,35 @@ func TestExistingDBProvider_Close_NilDB(t *testing.T) {
|
||||
t.Errorf("Expected Close to succeed with nil database, got error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOpenCountedCountsDials(t *testing.T) {
|
||||
var opened atomic.Int64
|
||||
db, err := openCounted("sqlite", ":memory:", &opened)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer db.Close()
|
||||
db.SetMaxOpenConns(1)
|
||||
|
||||
if err := db.PingContext(context.Background()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := opened.Load(); got != 1 {
|
||||
t.Fatalf("opened = %d after first ping, want 1", got)
|
||||
}
|
||||
// Reusing the pooled connection must not count as a new dial.
|
||||
if err := db.PingContext(context.Background()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := opened.Load(); got != 1 {
|
||||
t.Fatalf("opened = %d after reuse, want 1", got)
|
||||
}
|
||||
// Dropping idle connections forces a fresh dial.
|
||||
db.SetMaxIdleConns(0)
|
||||
if err := db.PingContext(context.Background()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := opened.Load(); got != 2 {
|
||||
t.Fatalf("opened = %d after redial, want 2", got)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
_ "github.com/microsoft/go-mssqldb" // MSSQL driver
|
||||
@@ -16,6 +17,7 @@ import (
|
||||
type MSSQLProvider struct {
|
||||
db *sql.DB
|
||||
config ConnectionConfig
|
||||
opened atomic.Int64
|
||||
}
|
||||
|
||||
// NewMSSQLProvider creates a new MSSQL provider
|
||||
@@ -52,7 +54,7 @@ func (p *MSSQLProvider) Connect(ctx context.Context, cfg ConnectionConfig) error
|
||||
}
|
||||
|
||||
// Open database connection
|
||||
db, err = sql.Open("sqlserver", dsn)
|
||||
db, err = openCounted("sqlserver", dsn, &p.opened)
|
||||
if err != nil {
|
||||
lastErr = err
|
||||
if cfg.GetEnableLogging() {
|
||||
@@ -173,6 +175,8 @@ func (p *MSSQLProvider) Stats() *ConnectionStats {
|
||||
Type: "mssql",
|
||||
Connected: true,
|
||||
OpenConnections: stats.OpenConnections,
|
||||
MaxOpenConnections: stats.MaxOpenConnections,
|
||||
TotalOpened: p.opened.Load(),
|
||||
InUse: stats.InUse,
|
||||
Idle: stats.Idle,
|
||||
WaitCount: stats.WaitCount,
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"fmt"
|
||||
"math"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"go.mongodb.org/mongo-driver/mongo"
|
||||
@@ -21,6 +22,7 @@ type PostgresProvider struct {
|
||||
config ConnectionConfig
|
||||
listener *PostgresListener
|
||||
mu sync.Mutex
|
||||
opened atomic.Int64
|
||||
}
|
||||
|
||||
// NewPostgresProvider creates a new PostgreSQL provider
|
||||
@@ -38,7 +40,7 @@ func (p *PostgresProvider) Connect(ctx context.Context, cfg ConnectionConfig) er
|
||||
// The connector and *sql.DB are created once; the pool is never closed to
|
||||
// recover from errors (see Refresh).
|
||||
connector := newPGConnector(connCfg)
|
||||
db := sql.OpenDB(connector)
|
||||
db := sql.OpenDB(&countingConnector{Connector: connector, opened: &p.opened})
|
||||
|
||||
// Connect with retry logic
|
||||
var lastErr error
|
||||
@@ -205,6 +207,8 @@ func (p *PostgresProvider) Stats() *ConnectionStats {
|
||||
Type: "postgres",
|
||||
Connected: true,
|
||||
OpenConnections: stats.OpenConnections,
|
||||
MaxOpenConnections: stats.MaxOpenConnections,
|
||||
TotalOpened: p.opened.Load(),
|
||||
InUse: stats.InUse,
|
||||
Idle: stats.Idle,
|
||||
WaitCount: stats.WaitCount,
|
||||
|
||||
@@ -28,6 +28,8 @@ type ConnectionStats struct {
|
||||
|
||||
// SQL connection pool stats
|
||||
OpenConnections int
|
||||
MaxOpenConnections int
|
||||
TotalOpened int64 // physical connections ever dialled (0 when unknown)
|
||||
InUse int
|
||||
Idle int
|
||||
WaitCount int64
|
||||
|
||||
@@ -6,6 +6,7 @@ import (
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
_ "github.com/glebarez/sqlite" // Pure Go SQLite driver
|
||||
@@ -19,6 +20,7 @@ type SQLiteProvider struct {
|
||||
db *sql.DB
|
||||
dbMu sync.RWMutex
|
||||
config ConnectionConfig
|
||||
opened atomic.Int64
|
||||
}
|
||||
|
||||
// NewSQLiteProvider creates a new SQLite provider
|
||||
@@ -51,7 +53,7 @@ func (p *SQLiteProvider) Connect(ctx context.Context, cfg ConnectionConfig) erro
|
||||
}
|
||||
|
||||
// Open database connection
|
||||
db, err := sql.Open("sqlite", dsn)
|
||||
db, err := openCounted("sqlite", dsn, &p.opened)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to open SQLite connection: %w", err)
|
||||
}
|
||||
@@ -182,6 +184,8 @@ func (p *SQLiteProvider) Stats() *ConnectionStats {
|
||||
Type: "sqlite",
|
||||
Connected: true,
|
||||
OpenConnections: stats.OpenConnections,
|
||||
MaxOpenConnections: stats.MaxOpenConnections,
|
||||
TotalOpened: p.opened.Load(),
|
||||
InUse: stats.InUse,
|
||||
Idle: stats.Idle,
|
||||
WaitCount: stats.WaitCount,
|
||||
|
||||
+108
-2
@@ -5,8 +5,9 @@ HTTP middleware utilities for security and performance.
|
||||
## Table of Contents
|
||||
|
||||
1. [Rate Limiting](#rate-limiting)
|
||||
2. [Request Size Limits](#request-size-limits)
|
||||
3. [Input Sanitization](#input-sanitization)
|
||||
2. [Client Request Queue](#client-request-queue)
|
||||
3. [Request Size Limits](#request-size-limits)
|
||||
4. [Input Sanitization](#input-sanitization)
|
||||
|
||||
---
|
||||
|
||||
@@ -381,6 +382,111 @@ func healthHandler(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
---
|
||||
|
||||
## Client Request Queue
|
||||
|
||||
`ClientQueue` smooths bursts (for example a page load that fires ~15 requests at once) by
|
||||
limiting how many requests each client runs concurrently and queueing the rest first-in-first-out.
|
||||
Unlike the rate limiter it never rejects a request that can be served shortly.
|
||||
|
||||
```go
|
||||
q := middleware.NewClientQueue(middleware.ClientQueueConfig{
|
||||
MaxConcurrent: 4, // running at once, per client
|
||||
MaxQueue: 50, // waiting, per client; beyond this -> 429
|
||||
MaxWait: 30 * time.Second, // waiting too long -> 503 + Retry-After
|
||||
})
|
||||
defer q.Close()
|
||||
|
||||
router.Use(q.Middleware) // gorilla/mux; or wrap any http.Handler
|
||||
```
|
||||
|
||||
### Adding it to restheadspec / resolvespec
|
||||
|
||||
Both packages' `SetupMuxRoutes` and `SetupBunRouterRoutes` take one middleware and apply it to every
|
||||
route, so pass the queue there. Use `Chain` to combine it with auth (first is outermost, so requests
|
||||
are authenticated before they can take a queue slot):
|
||||
|
||||
```go
|
||||
q := middleware.NewClientQueue(middleware.ClientQueueConfig{MaxConcurrent: 4})
|
||||
defer q.Close()
|
||||
|
||||
mw := middleware.Chain(authMiddleware, q.Middleware) // authMiddleware may be nil
|
||||
|
||||
restheadspec.SetupMuxRoutes(muxRouter, headHandler, mw)
|
||||
resolvespec.SetupMuxRoutes(muxRouter, resolveHandler, mw)
|
||||
// or SetupBunRouterRoutes(router, handler, mw)
|
||||
```
|
||||
|
||||
Share one `ClientQueue` across packages to give a client a single limit over all of them.
|
||||
|
||||
### Client identification
|
||||
|
||||
The first of these that is present is used, in order:
|
||||
|
||||
1. `X-Client-Id` header
|
||||
2. `Authorization` header
|
||||
3. the built-in server session (`security.GetSessionID` from the request context, else the session cookie)
|
||||
4. the client IP
|
||||
|
||||
Nothing is required from the client: without an id it falls back to the session, then to the IP.
|
||||
Secrets are hashed and never stored. Ids longer than 128 characters are ignored.
|
||||
|
||||
The session is only in the request context if the `security` auth middleware has run first, so put
|
||||
auth before the queue with `Chain` (below). Without it the cookie is still used.
|
||||
|
||||
To get one queue per browser tab (rather than per session), have the client send an id generated once
|
||||
per tab; the server cannot read `sessionStorage` itself:
|
||||
|
||||
```js
|
||||
const id = sessionStorage.clientId ??= crypto.randomUUID();
|
||||
fetch(url, { headers: { "X-Client-Id": id } });
|
||||
```
|
||||
|
||||
### Metrics
|
||||
|
||||
Registered on the default Prometheus registry (the same one `pkg/metrics` and `dbmanager` use), with
|
||||
no per-client labels so cardinality stays bounded.
|
||||
|
||||
| Metric | Type | Meaning |
|
||||
|---|---|---|
|
||||
| `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_max_seconds` | gauge | longest wait since process start |
|
||||
| `clientqueue_burst_size` | histogram | peak outstanding requests per client busy period |
|
||||
| `clientqueue_burst_max` | gauge | largest burst since process start |
|
||||
| `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 |
|
||||
|
||||
A **burst** is one client's busy period: the peak number of its requests running plus waiting between
|
||||
going from idle to busy and back to idle. A page load firing 15 requests at once is a burst of 15.
|
||||
|
||||
```promql
|
||||
# average burst size
|
||||
rate(clientqueue_burst_size_sum[5m]) / rate(clientqueue_burst_size_count[5m])
|
||||
# bursts larger than the concurrency limit (need queueing)
|
||||
1 - (sum(rate(clientqueue_burst_size_bucket{le="10"}[5m])) / sum(rate(clientqueue_burst_size_count[5m])))
|
||||
# average and p95 wait
|
||||
rate(clientqueue_wait_seconds_sum[5m]) / rate(clientqueue_wait_seconds_count[5m])
|
||||
histogram_quantile(0.95, sum(rate(clientqueue_wait_seconds_bucket[5m])) by (le))
|
||||
# share of requests that had to queue
|
||||
sum(rate(clientqueue_requests_total{result="queued"}[5m])) / sum(rate(clientqueue_requests_total[5m]))
|
||||
```
|
||||
|
||||
Use these to pick `MaxConcurrent`: if the p95 burst is well above it and waits are short, it is doing its
|
||||
job; if waits grow, the limit is too low or the database is the bottleneck. The `_max` gauges reset on
|
||||
restart; use the histograms for anything over time. A client that never goes idle produces one long
|
||||
busy period, so its burst is only recorded when it finally drains.
|
||||
|
||||
### Behaviour
|
||||
|
||||
- Queued requests are dropped if the client disconnects, so abandoned requests never run.
|
||||
- CORS preflights (`OPTIONS`) and connection upgrades (websockets) bypass the queue.
|
||||
- Idle clients are forgotten after `IdleTimeout` (default 5m).
|
||||
- The client id is client-supplied, so a hostile client can rotate ids to get more slots. It smooths
|
||||
well-behaved clients; it is not an abuse control. Combine it with `RateLimiter` for that.
|
||||
- Queue time counts against your own request timeouts; keep `MaxWait` below them.
|
||||
|
||||
## Request Size Limits
|
||||
|
||||
Protect against oversized request bodies with configurable size limits.
|
||||
|
||||
@@ -0,0 +1,405 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"container/list"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/prometheus/client_golang/prometheus"
|
||||
"github.com/prometheus/client_golang/prometheus/promauto"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||
)
|
||||
|
||||
// DefaultClientIDHeader is the header a client (one per browser tab) sends to
|
||||
// identify itself to the ClientQueue.
|
||||
const DefaultClientIDHeader = "X-Client-Id"
|
||||
|
||||
const maxClientIDLen = 128
|
||||
|
||||
// Client queue metrics. They are aggregated over all clients (no per-client
|
||||
// labels, to keep cardinality bounded) and over all ClientQueue instances.
|
||||
//
|
||||
// A burst is one client's busy period: the peak number of its requests
|
||||
// outstanding (running plus waiting) between the moment it goes from idle to
|
||||
// busy and the moment it is idle again. A page load firing 15 requests at
|
||||
// once is a burst of 15. Average burst is burst_size_sum / burst_size_count;
|
||||
// the highest burst seen is clientqueue_burst_max.
|
||||
var (
|
||||
queueRequests = promauto.NewCounterVec(prometheus.CounterOpts{
|
||||
Name: "clientqueue_requests_total",
|
||||
Help: "Requests by outcome: immediate (ran without waiting), queued (waited, then ran), rejected_full, timeout, canceled",
|
||||
}, []string{"result"})
|
||||
|
||||
queueWait = promauto.NewHistogram(prometheus.HistogramOpts{
|
||||
Name: "clientqueue_wait_seconds",
|
||||
Help: "Time admitted requests spent waiting for a slot (0 for immediate ones)",
|
||||
Buckets: []float64{0.001, 0.005, 0.01, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5, 10, 30},
|
||||
})
|
||||
|
||||
queueBurst = promauto.NewHistogram(prometheus.HistogramOpts{
|
||||
Name: "clientqueue_burst_size",
|
||||
Help: "Peak outstanding requests (running + waiting) per client busy period",
|
||||
Buckets: []float64{1, 2, 3, 4, 5, 8, 10, 15, 20, 30, 50, 100},
|
||||
})
|
||||
|
||||
queueBurstMax = promauto.NewGauge(prometheus.GaugeOpts{
|
||||
Name: "clientqueue_burst_max",
|
||||
Help: "Largest burst seen since the process started",
|
||||
})
|
||||
|
||||
queueWaitMax = promauto.NewGauge(prometheus.GaugeOpts{
|
||||
Name: "clientqueue_wait_max_seconds",
|
||||
Help: "Longest queue wait seen since the process started",
|
||||
})
|
||||
|
||||
queueActive = promauto.NewGauge(prometheus.GaugeOpts{
|
||||
Name: "clientqueue_active",
|
||||
Help: "Requests currently running through the queue",
|
||||
})
|
||||
|
||||
queueDepth = promauto.NewGauge(prometheus.GaugeOpts{
|
||||
Name: "clientqueue_queue_depth",
|
||||
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{
|
||||
Name: "clientqueue_clients",
|
||||
Help: "Clients currently tracked by the queue",
|
||||
})
|
||||
|
||||
maxMu sync.Mutex
|
||||
maxBurst float64
|
||||
maxWaitSec float64
|
||||
)
|
||||
|
||||
// observeMax raises a high-water gauge when v exceeds the recorded maximum.
|
||||
func observeMax(g prometheus.Gauge, cur *float64, v float64) {
|
||||
maxMu.Lock()
|
||||
if v > *cur {
|
||||
*cur = v
|
||||
g.Set(v)
|
||||
}
|
||||
maxMu.Unlock()
|
||||
}
|
||||
|
||||
var (
|
||||
errQueueFull = errors.New("client queue full")
|
||||
errQueueWait = errors.New("client queue wait exceeded")
|
||||
)
|
||||
|
||||
// ClientQueueConfig configures a ClientQueue. Zero values use the defaults.
|
||||
type ClientQueueConfig struct {
|
||||
// MaxConcurrent is how many requests one client may have running at once.
|
||||
// Default 4.
|
||||
MaxConcurrent int
|
||||
// MaxQueue is how many further requests one client may have waiting.
|
||||
// Requests beyond it are rejected with 429. Default 50.
|
||||
MaxQueue int
|
||||
// MaxWait is the longest a request waits for a slot before it is rejected
|
||||
// with 503. Default 30s.
|
||||
MaxWait time.Duration
|
||||
// IdleTimeout is how long a client with no requests is remembered.
|
||||
// Default 5m.
|
||||
IdleTimeout time.Duration
|
||||
// HeaderName is the client id header. Default X-Client-Id.
|
||||
HeaderName string
|
||||
}
|
||||
|
||||
func (c *ClientQueueConfig) applyDefaults() {
|
||||
if c.MaxConcurrent <= 0 {
|
||||
c.MaxConcurrent = 4
|
||||
}
|
||||
if c.MaxQueue <= 0 {
|
||||
c.MaxQueue = 50
|
||||
}
|
||||
if c.MaxWait <= 0 {
|
||||
c.MaxWait = 30 * time.Second
|
||||
}
|
||||
if c.IdleTimeout <= 0 {
|
||||
c.IdleTimeout = 5 * time.Minute
|
||||
}
|
||||
if c.HeaderName == "" {
|
||||
c.HeaderName = DefaultClientIDHeader
|
||||
}
|
||||
}
|
||||
|
||||
// ClientQueue limits how many requests each client runs concurrently and
|
||||
// queues the rest first-in-first-out, smoothing bursts such as a page load
|
||||
// that fires many requests at once. A client is identified by, in order:
|
||||
// the client id header, the Authorization header (session), then the IP.
|
||||
type ClientQueue struct {
|
||||
cfg ClientQueueConfig
|
||||
mu sync.Mutex
|
||||
clients map[string]*queueClient
|
||||
stop chan struct{}
|
||||
once sync.Once
|
||||
}
|
||||
|
||||
type queueClient struct {
|
||||
active int
|
||||
waiters list.List // of *queueWaiter
|
||||
lastUsed time.Time
|
||||
outstanding int // running + waiting
|
||||
peak int // highest outstanding in the current busy period
|
||||
}
|
||||
|
||||
// enter records a request entering the client's outstanding set.
|
||||
func (c *queueClient) enter() {
|
||||
c.outstanding++
|
||||
if c.outstanding > c.peak {
|
||||
c.peak = c.outstanding
|
||||
}
|
||||
}
|
||||
|
||||
// leave records a request leaving it, closing the burst when the client goes idle.
|
||||
func (c *queueClient) leave() {
|
||||
c.outstanding--
|
||||
if c.outstanding == 0 {
|
||||
queueBurst.Observe(float64(c.peak))
|
||||
observeMax(queueBurstMax, &maxBurst, float64(c.peak))
|
||||
c.peak = 0
|
||||
}
|
||||
}
|
||||
|
||||
type queueWaiter struct {
|
||||
ready chan struct{}
|
||||
granted bool
|
||||
elem *list.Element
|
||||
}
|
||||
|
||||
// NewClientQueue creates a ClientQueue and starts its idle-client cleanup.
|
||||
// Call Close to stop it.
|
||||
func NewClientQueue(cfg ClientQueueConfig) *ClientQueue {
|
||||
cfg.applyDefaults()
|
||||
q := &ClientQueue{
|
||||
cfg: cfg,
|
||||
clients: make(map[string]*queueClient),
|
||||
stop: make(chan struct{}),
|
||||
}
|
||||
go q.cleanupRoutine()
|
||||
return q
|
||||
}
|
||||
|
||||
// Close stops the cleanup goroutine. Requests in flight are unaffected.
|
||||
func (q *ClientQueue) Close() { q.once.Do(func() { close(q.stop) }) }
|
||||
|
||||
func (q *ClientQueue) cleanupRoutine() {
|
||||
interval := q.cfg.IdleTimeout / 2
|
||||
if interval < time.Second {
|
||||
interval = time.Second
|
||||
}
|
||||
ticker := time.NewTicker(interval)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-q.stop:
|
||||
return
|
||||
case now := <-ticker.C:
|
||||
q.evictIdle(now)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (q *ClientQueue) evictIdle(now time.Time) {
|
||||
q.mu.Lock()
|
||||
defer q.mu.Unlock()
|
||||
for key, c := range q.clients {
|
||||
if c.active == 0 && c.waiters.Len() == 0 && now.Sub(c.lastUsed) >= q.cfg.IdleTimeout {
|
||||
delete(q.clients, key)
|
||||
queueClients.Dec()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// clientKey identifies the caller by the first of these that is present:
|
||||
// the client id header, the Authorization header, the built-in server session
|
||||
// (from the request context once security auth has run, else the session
|
||||
// cookie), then the IP. Secrets are hashed so they are never held in memory
|
||||
// or logs. The session is only in the context if the auth middleware runs
|
||||
// before this one (see Chain).
|
||||
func (q *ClientQueue) clientKey(r *http.Request) string {
|
||||
id := strings.TrimSpace(r.Header.Get(q.cfg.HeaderName))
|
||||
if id != "" && len(id) <= maxClientIDLen && !strings.ContainsAny(id, "\r\n\x00") {
|
||||
return "cid:" + id
|
||||
}
|
||||
if a := r.Header.Get("Authorization"); a != "" {
|
||||
return "auth:" + hashKey(a)
|
||||
}
|
||||
if sid, ok := security.GetSessionID(r.Context()); ok && sid != "" {
|
||||
return "session:" + hashKey(sid)
|
||||
}
|
||||
if c := security.GetSessionCookie(r); c != "" {
|
||||
return "session:" + hashKey(c)
|
||||
}
|
||||
return "ip:" + getClientIP(r)
|
||||
}
|
||||
|
||||
func hashKey(v string) string {
|
||||
sum := sha256.Sum256([]byte(v))
|
||||
return hex.EncodeToString(sum[:16])
|
||||
}
|
||||
|
||||
// acquire blocks until key has a free slot, the queue is full, ctx ends, or
|
||||
// MaxWait passes. On success the caller must call release(key).
|
||||
func (q *ClientQueue) acquire(ctx context.Context, key string) error {
|
||||
start := time.Now()
|
||||
q.mu.Lock()
|
||||
c := q.clients[key]
|
||||
if c == nil {
|
||||
c = &queueClient{}
|
||||
q.clients[key] = c
|
||||
queueClients.Inc()
|
||||
}
|
||||
c.lastUsed = start
|
||||
if c.active < q.cfg.MaxConcurrent && c.waiters.Len() == 0 {
|
||||
c.active++
|
||||
c.enter()
|
||||
q.mu.Unlock()
|
||||
queueActive.Inc()
|
||||
queueWait.Observe(0)
|
||||
queueRequests.WithLabelValues("immediate").Inc()
|
||||
return nil
|
||||
}
|
||||
if c.waiters.Len() >= q.cfg.MaxQueue {
|
||||
q.mu.Unlock()
|
||||
queueRequests.WithLabelValues("rejected_full").Inc()
|
||||
return errQueueFull
|
||||
}
|
||||
w := &queueWaiter{ready: make(chan struct{})}
|
||||
w.elem = c.waiters.PushBack(w)
|
||||
if c.waiters.Len() == 1 {
|
||||
queueWaitingClients.Inc()
|
||||
}
|
||||
c.enter()
|
||||
q.mu.Unlock()
|
||||
queueDepth.Inc()
|
||||
queueEnqueued.Inc()
|
||||
|
||||
timer := time.NewTimer(q.cfg.MaxWait)
|
||||
defer timer.Stop()
|
||||
|
||||
var err error
|
||||
select {
|
||||
case <-w.ready:
|
||||
waited := time.Since(start).Seconds()
|
||||
queueWait.Observe(waited)
|
||||
observeMax(queueWaitMax, &maxWaitSec, waited)
|
||||
queueRequests.WithLabelValues("queued").Inc()
|
||||
return nil
|
||||
case <-ctx.Done():
|
||||
err = ctx.Err()
|
||||
case <-timer.C:
|
||||
err = errQueueWait
|
||||
}
|
||||
|
||||
q.mu.Lock()
|
||||
if w.granted {
|
||||
// A slot was handed over as we gave up; pass it on.
|
||||
q.releaseLocked(key)
|
||||
} else {
|
||||
c.waiters.Remove(w.elem)
|
||||
if c.waiters.Len() == 0 {
|
||||
queueWaitingClients.Dec()
|
||||
}
|
||||
c.leave()
|
||||
queueDepth.Dec()
|
||||
}
|
||||
q.mu.Unlock()
|
||||
if errors.Is(err, errQueueWait) {
|
||||
queueRequests.WithLabelValues("timeout").Inc()
|
||||
} else {
|
||||
queueRequests.WithLabelValues("canceled").Inc()
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func (q *ClientQueue) release(key string) {
|
||||
q.mu.Lock()
|
||||
q.releaseLocked(key)
|
||||
q.mu.Unlock()
|
||||
}
|
||||
|
||||
// releaseLocked hands the slot to the longest-waiting request, or frees it.
|
||||
func (q *ClientQueue) releaseLocked(key string) {
|
||||
c := q.clients[key]
|
||||
if c == nil {
|
||||
return
|
||||
}
|
||||
c.lastUsed = time.Now()
|
||||
c.leave()
|
||||
queueActive.Dec()
|
||||
if front := c.waiters.Front(); front != nil {
|
||||
w := c.waiters.Remove(front).(*queueWaiter)
|
||||
if c.waiters.Len() == 0 {
|
||||
queueWaitingClients.Dec()
|
||||
}
|
||||
w.granted = true
|
||||
queueDepth.Dec()
|
||||
queueActive.Inc()
|
||||
close(w.ready)
|
||||
return
|
||||
}
|
||||
c.active--
|
||||
}
|
||||
|
||||
// Middleware queues requests per client. CORS preflights and connection
|
||||
// upgrades (websockets) bypass the queue, since they are short or long-lived
|
||||
// respectively and should not hold or wait for a slot.
|
||||
func (q *ClientQueue) Middleware(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method == http.MethodOptions || r.Header.Get("Upgrade") != "" {
|
||||
next.ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
|
||||
key := q.clientKey(r)
|
||||
if err := q.acquire(r.Context(), key); err != nil {
|
||||
switch {
|
||||
case errors.Is(err, errQueueFull):
|
||||
w.Header().Set("Retry-After", "1")
|
||||
http.Error(w, `{"error":"queue_full","message":"Too many queued requests"}`, http.StatusTooManyRequests)
|
||||
case errors.Is(err, errQueueWait):
|
||||
w.Header().Set("Retry-After", strconv.Itoa(int(q.cfg.MaxWait.Seconds())))
|
||||
http.Error(w, `{"error":"queue_timeout","message":"Timed out waiting in the request queue"}`, http.StatusServiceUnavailable)
|
||||
}
|
||||
// Otherwise the client went away; there is no one to answer.
|
||||
return
|
||||
}
|
||||
defer q.release(key)
|
||||
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
// Chain composes middlewares into one, for the single middleware slot of
|
||||
// restheadspec.SetupMuxRoutes and friends. The first argument is the
|
||||
// outermost: Chain(auth, q.Middleware) authenticates before a request may
|
||||
// occupy a queue slot, so unauthenticated traffic cannot fill queues.
|
||||
func Chain(mws ...func(http.Handler) http.Handler) func(http.Handler) http.Handler {
|
||||
return func(next http.Handler) http.Handler {
|
||||
for i := len(mws) - 1; i >= 0; i-- {
|
||||
if mws[i] != nil {
|
||||
next = mws[i](next)
|
||||
}
|
||||
}
|
||||
return next
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,434 @@
|
||||
package middleware
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||
"github.com/prometheus/client_golang/prometheus"
|
||||
dto "github.com/prometheus/client_model/go"
|
||||
)
|
||||
|
||||
func newTestQueue(t *testing.T, cfg ClientQueueConfig) *ClientQueue {
|
||||
q := NewClientQueue(cfg)
|
||||
t.Cleanup(q.Close)
|
||||
return q
|
||||
}
|
||||
|
||||
func TestClientQueueKey(t *testing.T) {
|
||||
q := newTestQueue(t, ClientQueueConfig{})
|
||||
mk := func(id, auth string) *http.Request {
|
||||
r := httptest.NewRequest("GET", "/", nil)
|
||||
r.RemoteAddr = "10.0.0.1:1234"
|
||||
if id != "" {
|
||||
r.Header.Set("X-Client-Id", id)
|
||||
}
|
||||
if auth != "" {
|
||||
r.Header.Set("Authorization", auth)
|
||||
}
|
||||
return r
|
||||
}
|
||||
withCtxSession := func(r *http.Request, sid string) *http.Request {
|
||||
return r.WithContext(context.WithValue(r.Context(), security.SessionIDKey, sid))
|
||||
}
|
||||
|
||||
// 1. client id wins over everything else.
|
||||
r := withCtxSession(mk("tab1", "Bearer a"), "sess")
|
||||
r.AddCookie(&http.Cookie{Name: "session_token", Value: "cookie"})
|
||||
if got := q.clientKey(r); got != "cid:tab1" {
|
||||
t.Errorf("id: %q", got)
|
||||
}
|
||||
// 2. Authorization next, hashed.
|
||||
got := q.clientKey(withCtxSession(mk("", "Bearer secret"), "sess"))
|
||||
if !strings.HasPrefix(got, "auth:") || strings.Contains(got, "secret") {
|
||||
t.Errorf("auth must be hashed: %q", got)
|
||||
}
|
||||
// 3. built-in session from the context, then the cookie.
|
||||
got = q.clientKey(withCtxSession(mk("", ""), "sess"))
|
||||
if !strings.HasPrefix(got, "session:") || strings.Contains(strings.TrimPrefix(got, "session:"), "sess") {
|
||||
t.Errorf("context session: %q", got)
|
||||
}
|
||||
rc := mk("", "")
|
||||
rc.AddCookie(&http.Cookie{Name: "session_token", Value: "cookie"})
|
||||
if got := q.clientKey(rc); !strings.HasPrefix(got, "session:") {
|
||||
t.Errorf("cookie session: %q", got)
|
||||
}
|
||||
// 4. IP last.
|
||||
if got := q.clientKey(mk("", "")); got != "ip:10.0.0.1" {
|
||||
t.Errorf("ip: %q", got)
|
||||
}
|
||||
// Oversized ids are ignored and fall through.
|
||||
if got := q.clientKey(mk(strings.Repeat("x", 200), "")); got != "ip:10.0.0.1" {
|
||||
t.Errorf("oversized id should be ignored: %q", got)
|
||||
}
|
||||
// Tabs with different ids get separate queues; same id shares one.
|
||||
if q.clientKey(mk("tab1", "Bearer a")) == q.clientKey(mk("tab2", "Bearer a")) {
|
||||
t.Error("different ids must not share a queue")
|
||||
}
|
||||
if q.clientKey(mk("", "Bearer a")) != q.clientKey(mk("", "Bearer a")) {
|
||||
t.Error("same session without id must share a queue")
|
||||
}
|
||||
}
|
||||
|
||||
func TestClientQueueLimitsConcurrencyPerClient(t *testing.T) {
|
||||
q := newTestQueue(t, ClientQueueConfig{MaxConcurrent: 3})
|
||||
var cur, peak atomic.Int32
|
||||
h := q.Middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
n := cur.Add(1)
|
||||
for {
|
||||
p := peak.Load()
|
||||
if n <= p || peak.CompareAndSwap(p, n) {
|
||||
break
|
||||
}
|
||||
}
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
cur.Add(-1)
|
||||
}))
|
||||
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < 15; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
r := httptest.NewRequest("GET", "/", nil)
|
||||
r.Header.Set("X-Client-Id", "tab1")
|
||||
w := httptest.NewRecorder()
|
||||
h.ServeHTTP(w, r)
|
||||
if w.Code != http.StatusOK {
|
||||
t.Errorf("code %d", w.Code)
|
||||
}
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
if p := peak.Load(); p > 3 {
|
||||
t.Fatalf("peak concurrency %d, want <= 3", p)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClientQueueClientsAreIndependent(t *testing.T) {
|
||||
q := newTestQueue(t, ClientQueueConfig{MaxConcurrent: 1})
|
||||
block := make(chan struct{})
|
||||
started := make(chan struct{}, 2)
|
||||
h := q.Middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
started <- struct{}{}
|
||||
<-block
|
||||
}))
|
||||
for _, id := range []string{"a", "b"} {
|
||||
go func(id string) {
|
||||
r := httptest.NewRequest("GET", "/", nil)
|
||||
r.Header.Set("X-Client-Id", id)
|
||||
h.ServeHTTP(httptest.NewRecorder(), r)
|
||||
}(id)
|
||||
}
|
||||
for i := 0; i < 2; i++ {
|
||||
select {
|
||||
case <-started:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("a second client was blocked by the first")
|
||||
}
|
||||
}
|
||||
close(block)
|
||||
}
|
||||
|
||||
func TestClientQueueFullAndTimeout(t *testing.T) {
|
||||
q := newTestQueue(t, ClientQueueConfig{MaxConcurrent: 1, MaxQueue: 1, MaxWait: 50 * time.Millisecond})
|
||||
block := make(chan struct{})
|
||||
started := make(chan struct{}, 1)
|
||||
h := q.Middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
started <- struct{}{}
|
||||
<-block
|
||||
}))
|
||||
do := func() *httptest.ResponseRecorder {
|
||||
r := httptest.NewRequest("GET", "/", nil)
|
||||
r.Header.Set("X-Client-Id", "tab1")
|
||||
w := httptest.NewRecorder()
|
||||
h.ServeHTTP(w, r)
|
||||
return w
|
||||
}
|
||||
|
||||
go do() // takes the only slot
|
||||
<-started
|
||||
|
||||
queued := make(chan *httptest.ResponseRecorder, 1)
|
||||
go func() { queued <- do() }()
|
||||
time.Sleep(10 * time.Millisecond) // let it enter the queue
|
||||
|
||||
if w := do(); w.Code != http.StatusTooManyRequests {
|
||||
t.Errorf("queue full: got %d, want 429", w.Code)
|
||||
}
|
||||
if w := <-queued; w.Code != http.StatusServiceUnavailable || w.Header().Get("Retry-After") == "" {
|
||||
t.Errorf("wait timeout: got %d, want 503 with Retry-After", w.Code)
|
||||
}
|
||||
close(block)
|
||||
}
|
||||
|
||||
func TestClientQueueCancelledWaiterFreesQueue(t *testing.T) {
|
||||
q := newTestQueue(t, ClientQueueConfig{MaxConcurrent: 1})
|
||||
if err := q.acquire(t.Context(), "k"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
cctx, ccancel := context.WithCancel(t.Context())
|
||||
done := make(chan error, 1)
|
||||
go func() { done <- q.acquire(cctx, "k") }()
|
||||
time.Sleep(10 * time.Millisecond)
|
||||
ccancel()
|
||||
if err := <-done; err == nil {
|
||||
t.Fatal("expected cancellation error")
|
||||
}
|
||||
q.release("k")
|
||||
|
||||
q.mu.Lock()
|
||||
c := q.clients["k"]
|
||||
active, waiting := c.active, c.waiters.Len()
|
||||
q.mu.Unlock()
|
||||
if active != 0 || waiting != 0 {
|
||||
t.Fatalf("leaked state: active=%d waiting=%d", active, waiting)
|
||||
}
|
||||
}
|
||||
|
||||
func TestClientQueueBypassesPreflightAndEvictsIdle(t *testing.T) {
|
||||
q := newTestQueue(t, ClientQueueConfig{MaxConcurrent: 1, IdleTimeout: time.Minute})
|
||||
block := make(chan struct{})
|
||||
started := make(chan struct{}, 1)
|
||||
h := q.Middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method == http.MethodGet {
|
||||
started <- struct{}{}
|
||||
<-block
|
||||
}
|
||||
}))
|
||||
go func() {
|
||||
r := httptest.NewRequest("GET", "/", nil)
|
||||
r.Header.Set("X-Client-Id", "tab1")
|
||||
h.ServeHTTP(httptest.NewRecorder(), r)
|
||||
}()
|
||||
<-started
|
||||
|
||||
r := httptest.NewRequest("OPTIONS", "/", nil)
|
||||
r.Header.Set("X-Client-Id", "tab1")
|
||||
done := make(chan struct{})
|
||||
go func() { h.ServeHTTP(httptest.NewRecorder(), r); close(done) }()
|
||||
select {
|
||||
case <-done:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("preflight was queued")
|
||||
}
|
||||
close(block)
|
||||
time.Sleep(20 * time.Millisecond)
|
||||
|
||||
q.evictIdle(time.Now().Add(2 * time.Minute))
|
||||
q.mu.Lock()
|
||||
n := len(q.clients)
|
||||
q.mu.Unlock()
|
||||
if n != 0 {
|
||||
t.Fatalf("idle client not evicted: %d tracked", n)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChainOrder(t *testing.T) {
|
||||
var order []string
|
||||
tag := func(name string) func(http.Handler) http.Handler {
|
||||
return func(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
order = append(order, name)
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
}
|
||||
h := Chain(tag("auth"), nil, tag("queue"))(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
order = append(order, "handler")
|
||||
}))
|
||||
h.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest("GET", "/", nil))
|
||||
if got := len(order); got != 3 || order[0] != "auth" || order[1] != "queue" || order[2] != "handler" {
|
||||
t.Fatalf("order = %v", order)
|
||||
}
|
||||
}
|
||||
|
||||
func gaugeVal(t *testing.T, m prometheus.Metric) float64 {
|
||||
t.Helper()
|
||||
var d dto.Metric
|
||||
if err := m.Write(&d); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
switch {
|
||||
case d.Gauge != nil:
|
||||
return d.Gauge.GetValue()
|
||||
case d.Counter != nil:
|
||||
return d.Counter.GetValue()
|
||||
}
|
||||
t.Fatal("not a gauge or counter")
|
||||
return 0
|
||||
}
|
||||
|
||||
func histVal(t *testing.T, h prometheus.Histogram) (count uint64, sum float64) {
|
||||
t.Helper()
|
||||
var d dto.Metric
|
||||
if err := h.Write(&d); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return d.Histogram.GetSampleCount(), d.Histogram.GetSampleSum()
|
||||
}
|
||||
|
||||
func TestClientQueueMetrics(t *testing.T) {
|
||||
q := newTestQueue(t, ClientQueueConfig{MaxConcurrent: 2})
|
||||
imm0 := gaugeVal(t, queueRequests.WithLabelValues("immediate"))
|
||||
que0 := gaugeVal(t, queueRequests.WithLabelValues("queued"))
|
||||
enq0 := gaugeVal(t, queueEnqueued)
|
||||
bc0, bs0 := histVal(t, queueBurst)
|
||||
wc0, _ := histVal(t, queueWait)
|
||||
|
||||
block := make(chan struct{})
|
||||
h := q.Middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { <-block }))
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < 5; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
r := httptest.NewRequest("GET", "/", nil)
|
||||
r.Header.Set("X-Client-Id", "metrics-tab")
|
||||
h.ServeHTTP(httptest.NewRecorder(), r)
|
||||
}()
|
||||
}
|
||||
// Wait for 2 running and 3 queued.
|
||||
deadline := time.Now().Add(2 * time.Second)
|
||||
for {
|
||||
q.mu.Lock()
|
||||
c := q.clients["cid:metrics-tab"]
|
||||
ok := c != nil && c.active == 2 && c.waiters.Len() == 3
|
||||
q.mu.Unlock()
|
||||
if ok {
|
||||
break
|
||||
}
|
||||
if time.Now().After(deadline) {
|
||||
t.Fatal("requests did not reach 2 running + 3 queued")
|
||||
}
|
||||
time.Sleep(time.Millisecond)
|
||||
}
|
||||
if got := gaugeVal(t, queueDepth); got < 3 {
|
||||
t.Errorf("queue depth = %v, want >= 3", got)
|
||||
}
|
||||
close(block)
|
||||
wg.Wait()
|
||||
|
||||
if got := gaugeVal(t, queueRequests.WithLabelValues("immediate")) - imm0; got != 2 {
|
||||
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 {
|
||||
t.Errorf("queued = %v, want 3", got)
|
||||
}
|
||||
bc, bs := histVal(t, queueBurst)
|
||||
if bc-bc0 != 1 || bs-bs0 != 5 {
|
||||
t.Errorf("burst observations = %d (sum %v), want 1 burst of size 5", bc-bc0, bs-bs0)
|
||||
}
|
||||
if wc, _ := histVal(t, queueWait); wc-wc0 != 5 {
|
||||
t.Errorf("wait observations = %d, want 5", wc-wc0)
|
||||
}
|
||||
if got := gaugeVal(t, queueBurstMax); got < 5 {
|
||||
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
@@ -45,6 +45,9 @@ type Handler struct {
|
||||
// Started flag
|
||||
started bool
|
||||
mu sync.RWMutex
|
||||
|
||||
// disallowNulls skips null values in update payloads instead of applying them
|
||||
disallowNulls bool
|
||||
}
|
||||
|
||||
// NewHandler creates a new MQTT handler
|
||||
@@ -812,6 +815,14 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata
|
||||
}
|
||||
|
||||
// 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) {
|
||||
// Marshal and unmarshal data into model
|
||||
dataBytes, err := json.Marshal(hookCtx.Data)
|
||||
@@ -840,26 +851,34 @@ func (h *Handler) create(hookCtx *HookContext) (interface{}, error) {
|
||||
|
||||
// update updates an existing record
|
||||
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)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal data: %w", err)
|
||||
}
|
||||
|
||||
if err := json.Unmarshal(dataBytes, hookCtx.ModelPtr); err != nil {
|
||||
return nil, fmt.Errorf("failed to unmarshal data into model: %w", err)
|
||||
if err := json.Unmarshal(dataBytes, &updates); err != nil {
|
||||
return nil, fmt.Errorf("failed to unmarshal data into map: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Update record
|
||||
query := h.db.NewUpdate().Model(hookCtx.ModelPtr).Table(hookCtx.TableName)
|
||||
|
||||
// Add ID filter
|
||||
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 {
|
||||
return nil, fmt.Errorf("failed to update record: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Fetch updated record
|
||||
return h.readByID(hookCtx)
|
||||
|
||||
+15
-35
@@ -32,6 +32,7 @@ type Handler struct {
|
||||
fallbackHandler FallbackHandler
|
||||
openAPIGenerator func() (string, error)
|
||||
defaultSort map[string][]common.SortOption
|
||||
disallowNulls bool
|
||||
}
|
||||
|
||||
// NewHandler creates a new API handler with database and registry abstractions
|
||||
@@ -52,6 +53,14 @@ func (h *Handler) Hooks() *HookRegistry {
|
||||
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
|
||||
// If not set, the handler will simply return (pass through to next route)
|
||||
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)
|
||||
}
|
||||
|
||||
// Merge only non-null and non-empty values from the incoming request into the existing record
|
||||
for key, newValue := range updates {
|
||||
// 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
|
||||
}
|
||||
// Overwrite with every key present in the request (including "" and null unless disallowed)
|
||||
common.MergeUpdateValues(existingMap, updates, h.disallowNulls)
|
||||
|
||||
// Build update query with merged data
|
||||
query := tx.NewUpdate().Table(tableName).SetMap(existingMap)
|
||||
@@ -1421,16 +1417,8 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
|
||||
item = modifiedData
|
||||
}
|
||||
|
||||
// Merge only non-null and non-empty values
|
||||
for key, newValue := range item {
|
||||
if newValue == nil {
|
||||
continue
|
||||
}
|
||||
if strVal, ok := newValue.(string); ok && strVal == "" {
|
||||
continue
|
||||
}
|
||||
existingMap[key] = newValue
|
||||
}
|
||||
// Overwrite with every key present in the request (including "" and null unless disallowed)
|
||||
common.MergeUpdateValues(existingMap, item, h.disallowNulls)
|
||||
|
||||
txQuery := tx.NewUpdate().Table(tableName).SetMap(existingMap).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID)
|
||||
if _, err := txQuery.Exec(ctx); err != nil {
|
||||
@@ -1578,16 +1566,8 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
|
||||
itemMap = modifiedData
|
||||
}
|
||||
|
||||
// Merge only non-null and non-empty values
|
||||
for key, newValue := range itemMap {
|
||||
if newValue == nil {
|
||||
continue
|
||||
}
|
||||
if strVal, ok := newValue.(string); ok && strVal == "" {
|
||||
continue
|
||||
}
|
||||
existingMap[key] = newValue
|
||||
}
|
||||
// Overwrite with every key present in the request (including "" and null unless disallowed)
|
||||
common.MergeUpdateValues(existingMap, itemMap, h.disallowNulls)
|
||||
|
||||
txQuery := tx.NewUpdate().Table(tableName).SetMap(existingMap).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID)
|
||||
if _, err := txQuery.Exec(ctx); err != nil {
|
||||
|
||||
+11
-15
@@ -33,6 +33,7 @@ type Handler struct {
|
||||
fallbackHandler FallbackHandler
|
||||
openAPIGenerator func() (string, error)
|
||||
defaultSort map[string][]common.SortOption
|
||||
disallowNulls bool
|
||||
}
|
||||
|
||||
// NewHandler creates a new API handler with database and registry abstractions
|
||||
@@ -59,6 +60,14 @@ func (h *Handler) Hooks() *HookRegistry {
|
||||
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
|
||||
// If not set, the handler will simply return (pass through to next route)
|
||||
func (h *Handler) SetFallbackHandler(fallback FallbackHandler) {
|
||||
@@ -1597,21 +1606,8 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
|
||||
nestedRelations = relations
|
||||
}
|
||||
|
||||
// Merge only non-null and non-empty values from the incoming request into the existing record
|
||||
for key, newValue := range dataMap {
|
||||
// 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
|
||||
}
|
||||
// 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
|
||||
|
||||
@@ -38,6 +38,9 @@ type Handler struct {
|
||||
subscriptionManager *SubscriptionManager
|
||||
upgrader websocket.Upgrader
|
||||
ctx context.Context
|
||||
|
||||
// disallowNulls skips null values in update payloads instead of applying them
|
||||
disallowNulls bool
|
||||
}
|
||||
|
||||
// NewHandler creates a new WebSocket handler
|
||||
@@ -682,6 +685,14 @@ func (h *Handler) readMultiple(hookCtx *HookContext) (data interface{}, metadata
|
||||
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) {
|
||||
// Marshal and unmarshal data into model
|
||||
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) {
|
||||
// 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)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal data: %w", err)
|
||||
}
|
||||
|
||||
if err := json.Unmarshal(dataBytes, hookCtx.ModelPtr); err != nil {
|
||||
return nil, fmt.Errorf("failed to unmarshal data into model: %w", err)
|
||||
if err := json.Unmarshal(dataBytes, &updates); err != nil {
|
||||
return nil, fmt.Errorf("failed to unmarshal data into map: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Update record
|
||||
query := h.db.NewUpdate().Model(hookCtx.ModelPtr).Table(hookCtx.TableName)
|
||||
|
||||
// Add ID filter
|
||||
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 {
|
||||
return nil, fmt.Errorf("failed to update record: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Fetch updated record
|
||||
return h.readByID(hookCtx)
|
||||
|
||||
Reference in New Issue
Block a user