Compare commits

...
6 Commits
Author SHA1 Message Date
Hein 20ba8ed112 feat(update): allow clearing values with "" and null on update
Update handlers skipped empty strings and nulls, so a client could not
blank or null out a column. Every key present in the payload now
overwrites the stored value, including "" and null.

- add common.MergeUpdateValues and use it in resolvespec and restheadspec
- add Handler.SetDisallowNulls to skip null values (""still overwrites)
- websocketspec and mqttspec now write only the keys present in the
  payload via SetMap instead of updating the whole zeroed model, which
  clobbered absent fields
2026-09-30 17:03:42 +02:00
Hein 1214f69e0c feat(middleware): export clientqueue_enqueued_total counter
Counts every request placed in a wait queue regardless of outcome, so the
total ever queued no longer has to be summed from queued, timeout and
canceled.
2026-09-30 16:16:25 +02:00
Hein 89a58ab3a0 feat(middleware): export clientqueue_waiting_clients gauge
Counts clients with at least one queued request, updated whenever a
client's waiter count crosses zero (enqueue, grant, timeout, cancel).
2026-09-30 16:14:27 +02:00
Hein 48081b4aa4 docs: document client queue and dbmanager pool stats 2026-09-30 15:46:43 +02:00
Hein 467dbd66c8 feat(middleware): add per-client FIFO request queue with burst metrics
ClientQueue limits concurrent requests per client (X-Client-Id, then
Authorization, session, IP) and queues the rest first-in-first-out, with
bounded depth, max wait and idle eviction. Chain composes it with auth for
the Setup*Routes middleware slot. Exports burst, wait and depth metrics.
The test server uses it with a limit of 10.
2026-09-30 15:46:42 +02:00
Hein 178d40587d feat(dbmanager): report pool max size and total connections opened
Add MaxOpenConnections to provider and connection stats, exposed as the
max state of dbmanager_connection_pool_size. Count physical dials through a
counting driver connector and export dbmanager_connections_opened_total;
existing_db pools derive an approximate count from sql.DBStats.
2026-09-30 15:46:41 +02:00
22 changed files with 1302 additions and 137 deletions
+11 -3
View File
@@ -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
+7 -1
View File
@@ -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()
+23
View File
@@ -0,0 +1,23 @@
package common
// MergeUpdateValues merges the incoming request values into the existing
// record map (in place) and returns it.
//
// Every key present in incoming overwrites the existing value, including empty
// strings and explicit nulls, so clients can clear a column by sending "" or
// null. Keys absent from incoming are left untouched.
//
// When disallowNulls is true, nil values are skipped and the existing value is
// kept. Empty strings are still applied.
func MergeUpdateValues(existing, incoming map[string]interface{}, disallowNulls bool) map[string]interface{} {
if existing == nil {
existing = make(map[string]interface{}, len(incoming))
}
for key, newValue := range incoming {
if newValue == nil && disallowNulls {
continue
}
existing[key] = newValue
}
return existing
}
+29
View File
@@ -0,0 +1,29 @@
package common
import "testing"
func TestMergeUpdateValues(t *testing.T) {
newExisting := func() map[string]interface{} {
return map[string]interface{}{"name": "old", "note": "keep", "other": "x"}
}
incoming := map[string]interface{}{"name": "", "note": nil}
got := MergeUpdateValues(newExisting(), incoming, false)
if got["name"] != "" {
t.Errorf("name = %v, want empty string", got["name"])
}
if v, ok := got["note"]; !ok || v != nil {
t.Errorf("note = %v (present=%v), want nil", v, ok)
}
if got["other"] != "x" {
t.Errorf("absent key changed: %v", got["other"])
}
got = MergeUpdateValues(newExisting(), incoming, true)
if got["name"] != "" {
t.Errorf("disallowNulls: name = %v, want empty string", got["name"])
}
if got["note"] != "keep" {
t.Errorf("disallowNulls: note = %v, want keep", got["note"])
}
}
+4 -2
View File
@@ -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
+14
View File
@@ -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)
}
+11 -7
View File
@@ -55,13 +55,15 @@ type ConnectionStats struct {
HealthCheckStatus string
// SQL connection pool stats
OpenConnections int
InUse int
Idle int
WaitCount int64
WaitDuration time.Duration
MaxIdleClosed int64
MaxLifetimeClosed int64
OpenConnections int
MaxOpenConnections int
TotalOpened int64 // physical connections ever opened
InUse int
Idle int
WaitCount int64
WaitDuration time.Duration
MaxIdleClosed int64
MaxLifetimeClosed int64
}
// sqlConnection implements Connection for SQL databases (PostgreSQL, SQLite, MSSQL)
@@ -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
+13 -2
View File
@@ -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
+51
View File
@@ -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
}
+6
View File
@@ -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)
}
}
+15 -11
View File
@@ -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() {
@@ -169,15 +171,17 @@ func (p *MSSQLProvider) Stats() *ConnectionStats {
stats := p.db.Stats()
return &ConnectionStats{
Name: p.config.GetName(),
Type: "mssql",
Connected: true,
OpenConnections: stats.OpenConnections,
InUse: stats.InUse,
Idle: stats.Idle,
WaitCount: stats.WaitCount,
WaitDuration: stats.WaitDuration,
MaxIdleClosed: stats.MaxIdleClosed,
MaxLifetimeClosed: stats.MaxLifetimeClosed,
Name: p.config.GetName(),
Type: "mssql",
Connected: true,
OpenConnections: stats.OpenConnections,
MaxOpenConnections: stats.MaxOpenConnections,
TotalOpened: p.opened.Load(),
InUse: stats.InUse,
Idle: stats.Idle,
WaitCount: stats.WaitCount,
WaitDuration: stats.WaitDuration,
MaxIdleClosed: stats.MaxIdleClosed,
MaxLifetimeClosed: stats.MaxLifetimeClosed,
}
}
+15 -11
View File
@@ -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
@@ -201,16 +203,18 @@ func (p *PostgresProvider) Stats() *ConnectionStats {
stats := p.db.Stats()
return &ConnectionStats{
Name: p.config.GetName(),
Type: "postgres",
Connected: true,
OpenConnections: stats.OpenConnections,
InUse: stats.InUse,
Idle: stats.Idle,
WaitCount: stats.WaitCount,
WaitDuration: stats.WaitDuration,
MaxIdleClosed: stats.MaxIdleClosed,
MaxLifetimeClosed: stats.MaxLifetimeClosed,
Name: p.config.GetName(),
Type: "postgres",
Connected: true,
OpenConnections: stats.OpenConnections,
MaxOpenConnections: stats.MaxOpenConnections,
TotalOpened: p.opened.Load(),
InUse: stats.InUse,
Idle: stats.Idle,
WaitCount: stats.WaitCount,
WaitDuration: stats.WaitDuration,
MaxIdleClosed: stats.MaxIdleClosed,
MaxLifetimeClosed: stats.MaxLifetimeClosed,
}
}
+9 -7
View File
@@ -27,13 +27,15 @@ type ConnectionStats struct {
HealthCheckStatus string
// SQL connection pool stats
OpenConnections int
InUse int
Idle int
WaitCount int64
WaitDuration time.Duration
MaxIdleClosed int64
MaxLifetimeClosed int64
OpenConnections int
MaxOpenConnections int
TotalOpened int64 // physical connections ever dialled (0 when unknown)
InUse int
Idle int
WaitCount int64
WaitDuration time.Duration
MaxIdleClosed int64
MaxLifetimeClosed int64
}
// ConnectionConfig is a minimal interface for configuration
+15 -11
View File
@@ -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)
}
@@ -178,15 +180,17 @@ func (p *SQLiteProvider) Stats() *ConnectionStats {
stats := p.db.Stats()
return &ConnectionStats{
Name: p.config.GetName(),
Type: "sqlite",
Connected: true,
OpenConnections: stats.OpenConnections,
InUse: stats.InUse,
Idle: stats.Idle,
WaitCount: stats.WaitCount,
WaitDuration: stats.WaitDuration,
MaxIdleClosed: stats.MaxIdleClosed,
MaxLifetimeClosed: stats.MaxLifetimeClosed,
Name: p.config.GetName(),
Type: "sqlite",
Connected: true,
OpenConnections: stats.OpenConnections,
MaxOpenConnections: stats.MaxOpenConnections,
TotalOpened: p.opened.Load(),
InUse: stats.InUse,
Idle: stats.Idle,
WaitCount: stats.WaitCount,
WaitDuration: stats.WaitDuration,
MaxIdleClosed: stats.MaxIdleClosed,
MaxLifetimeClosed: stats.MaxLifetimeClosed,
}
}
+108 -2
View File
@@ -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.
+405
View File
@@ -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
}
}
+434
View File
@@ -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)
}
}
+34 -15
View File
@@ -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,25 +851,33 @@ 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
dataBytes, err := json.Marshal(hookCtx.Data)
if err != nil {
return nil, fmt.Errorf("failed to marshal data: %w", err)
// 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, &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)
query = query.Where(fmt.Sprintf("%s = ?", pkName), hookCtx.ID)
if _, err := query.Exec(hookCtx.Context); err != nil {
return nil, fmt.Errorf("failed to update record: %w", err)
// 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
+15 -35
View File
@@ -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
View File
@@ -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
+34 -15
View File
@@ -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,25 +720,33 @@ func (h *Handler) create(hookCtx *HookContext) (interface{}, error) {
}
func (h *Handler) update(hookCtx *HookContext) (interface{}, error) {
// Marshal and unmarshal data into model
dataBytes, err := json.Marshal(hookCtx.Data)
if err != nil {
return nil, fmt.Errorf("failed to marshal data: %w", err)
// 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, &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)
query = query.Where(fmt.Sprintf("%s = ?", pkName), hookCtx.ID)
if _, err := query.Exec(hookCtx.Context); err != nil {
return nil, fmt.Errorf("failed to update record: %w", err)
// 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