mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-07 13:56:29 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
234aac9770 | ||
|
|
8cff3bde85 | ||
|
|
3e6224698c | ||
|
|
aec87a81e7 |
@@ -109,8 +109,8 @@ require (
|
||||
github.com/pkg/errors v0.9.1 // indirect
|
||||
github.com/pmezard/go-difflib v1.0.0 // indirect
|
||||
github.com/power-devops/perfstat v0.0.0-20210106213030-5aafc221ea8c // indirect
|
||||
github.com/prometheus/client_model v0.6.2 // indirect
|
||||
github.com/prometheus/common v0.67.5 // indirect
|
||||
github.com/prometheus/client_model v0.6.2
|
||||
github.com/prometheus/common v0.67.5
|
||||
github.com/prometheus/procfs v0.20.1 // indirect
|
||||
github.com/puzpuzpuz/xsync/v3 v3.5.1 // indirect
|
||||
github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect
|
||||
|
||||
@@ -10,6 +10,7 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/uptrace/bun"
|
||||
"github.com/uptrace/bun/schema"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/dbtrace"
|
||||
@@ -1507,8 +1508,30 @@ func (b *BunInsertQuery) OnConflict(action string) common.InsertQuery {
|
||||
return b
|
||||
}
|
||||
|
||||
// bunWritableExcludes drops columns bun already leaves out of INSERT/UPDATE
|
||||
// (scanonly fields) or does not know, since bun's ExcludeColumn errors with
|
||||
// "can't find column" for anything that is not in the table's writable fields.
|
||||
func bunWritableExcludes(model bun.Model, columns []string) []string {
|
||||
tm, ok := model.(interface{ Table() *schema.Table })
|
||||
if !ok || tm.Table() == nil {
|
||||
return columns
|
||||
}
|
||||
table := tm.Table()
|
||||
writable := make(map[string]struct{}, len(table.Fields))
|
||||
for _, f := range table.Fields {
|
||||
writable[f.Name] = struct{}{}
|
||||
}
|
||||
out := make([]string, 0, len(columns))
|
||||
for _, c := range columns {
|
||||
if _, ok := writable[c]; ok || c == "*" {
|
||||
out = append(out, c)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func (b *BunInsertQuery) ExcludeColumn(columns ...string) common.InsertQuery {
|
||||
if len(columns) > 0 {
|
||||
if columns = bunWritableExcludes(b.query.GetModel(), columns); len(columns) > 0 {
|
||||
b.query = b.query.ExcludeColumn(columns...)
|
||||
}
|
||||
return b
|
||||
@@ -1627,7 +1650,7 @@ func (b *BunUpdateQuery) SetMap(values map[string]interface{}) common.UpdateQuer
|
||||
}
|
||||
|
||||
func (b *BunUpdateQuery) ExcludeColumn(columns ...string) common.UpdateQuery {
|
||||
if len(columns) > 0 {
|
||||
if columns = bunWritableExcludes(b.query.GetModel(), columns); len(columns) > 0 {
|
||||
b.query = b.query.ExcludeColumn(columns...)
|
||||
}
|
||||
return b
|
||||
|
||||
@@ -0,0 +1,100 @@
|
||||
package database
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/uptrace/bun"
|
||||
"github.com/uptrace/bun/dialect/pgdialect"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||
)
|
||||
|
||||
// adhocBuffer mirrors the real-world DBAdhocBuffer: scanonly fields with both
|
||||
// bun and gorm read-only tags.
|
||||
type adhocBuffer struct {
|
||||
CQL1 string `json:"cql1,omitempty" gorm:"->" bun:",scanonly"`
|
||||
CQL2 string `json:"cql2,omitempty" gorm:"->" bun:",scanonly"`
|
||||
RowNumber int64 `json:"_rownumber,omitempty" gorm:"-" bun:",scanonly"`
|
||||
RecordError string `json:"_error,omitempty" gorm:"-" bun:",scanonly"`
|
||||
}
|
||||
|
||||
type excludeModel struct {
|
||||
bun.BaseModel `bun:"table:public.crmnote,alias:crmnote"`
|
||||
ID int `json:"id" bun:"id,pk"`
|
||||
Note string `json:"note" bun:"note,type:citext,"`
|
||||
Norm string `json:"norm" bun:"norm,generated"`
|
||||
|
||||
adhocBuffer `json:",omitempty" bun:",scanonly"`
|
||||
}
|
||||
|
||||
func newExcludeDB() *bun.DB {
|
||||
return bun.NewDB(&sql.DB{}, pgdialect.New())
|
||||
}
|
||||
|
||||
// TestBunExcludeColumnWithNonWritableColumns feeds the reflection output
|
||||
// straight into the adapter, as the handlers do, for insert and update.
|
||||
func TestBunExcludeColumnWithNonWritableColumns(t *testing.T) {
|
||||
db := newExcludeDB()
|
||||
m := &excludeModel{}
|
||||
cols := reflection.NonWritableColumns(m)
|
||||
if len(cols) == 0 {
|
||||
t.Fatal("expected non-writable columns")
|
||||
}
|
||||
|
||||
ins := &BunInsertQuery{query: db.NewInsert().Model(m)}
|
||||
ins.ExcludeColumn(cols...)
|
||||
insSQL, err := ins.query.AppendQuery(db.QueryGen(), nil)
|
||||
if err != nil {
|
||||
t.Fatalf("insert: %v", err)
|
||||
}
|
||||
|
||||
upd := &BunUpdateQuery{query: db.NewUpdate().Model(m).Where("id = 1")}
|
||||
upd.ExcludeColumn(cols...)
|
||||
updSQL, err := upd.query.AppendQuery(db.QueryGen(), nil)
|
||||
if err != nil {
|
||||
t.Fatalf("update: %v", err)
|
||||
}
|
||||
|
||||
for name, q := range map[string]string{"insert": string(insSQL), "update": string(updSQL)} {
|
||||
for _, bad := range []string{"cql1", "cql2", "_rownumber", "_error", "norm"} {
|
||||
if strings.Contains(q, `"`+bad+`"`) {
|
||||
t.Errorf("%s writes non-writable column %s: %s", name, bad, q)
|
||||
}
|
||||
}
|
||||
if !strings.Contains(q, `"note"`) {
|
||||
t.Errorf("%s dropped writable column note: %s", name, q)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestBunExcludeColumnIgnoresUnknownAndKeepsWritable(t *testing.T) {
|
||||
db := newExcludeDB()
|
||||
m := &excludeModel{}
|
||||
|
||||
ins := &BunInsertQuery{query: db.NewInsert().Model(m)}
|
||||
ins.ExcludeColumn("does_not_exist", "note")
|
||||
q, err := ins.query.AppendQuery(db.QueryGen(), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if strings.Contains(string(q), `"note"`) {
|
||||
t.Errorf("writable column note should have been excluded: %s", q)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBunExcludeColumnOnlyNonWritable(t *testing.T) {
|
||||
db := newExcludeDB()
|
||||
ins := &BunInsertQuery{query: db.NewInsert().Model(&excludeModel{})}
|
||||
ins.ExcludeColumn("cql1") // everything filtered out: must not error or panic
|
||||
if _, err := ins.query.AppendQuery(db.QueryGen(), nil); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBunExcludeColumnWithoutModel(t *testing.T) {
|
||||
db := newExcludeDB()
|
||||
ins := &BunInsertQuery{query: db.NewInsert()}
|
||||
ins.ExcludeColumn("cql1") // no model yet: must not panic
|
||||
}
|
||||
+50
-3
@@ -48,11 +48,59 @@ metrics.SetProvider(provider)
|
||||
| `Namespace` | `string` | `""` | Prefix for all metric names |
|
||||
| `HTTPRequestBuckets` | `[]float64` | See below | Histogram buckets for HTTP duration (seconds) |
|
||||
| `DBQueryBuckets` | `[]float64` | See below | Histogram buckets for DB query duration (seconds) |
|
||||
| `HTTPMaxPaths` | `int` | `1024` | Max distinct `path` label values; extras become `"other"` (negative disables) |
|
||||
| `HTTPPathNormalizer` | `func(*http.Request) string` | `nil` | Custom request → `path` label mapping (return `""` to use the default) |
|
||||
|
||||
**HTTP `path` label:** the middleware uses, in order: `HTTPPathNormalizer`, the matched `http.ServeMux` pattern (`r.Pattern`, e.g. `/users/{id}`), then the raw path with numeric/UUID/hex/opaque-token segments replaced by `:id`. For routers other than `ServeMux`, supply `HTTPPathNormalizer` with your route template. The `HTTPMaxPaths` cap applies on top.
|
||||
|
||||
**Default HTTP Request Buckets:** `[0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5, 10]`
|
||||
|
||||
**Default DB Query Buckets:** `[0.001, 0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5]`
|
||||
|
||||
### Enabled flag and JSON pull
|
||||
|
||||
`Config.Enabled` is honoured: a disabled provider records nothing, `Middleware` passes requests straight through, `Handler()`/`JSONHandler()` answer 404, and push loops are not started (manual pushes return an error). Note a `&metrics.Config{}` literal has `Enabled: false`; use `DefaultConfig()` or set `Enabled: true`. `NewPrometheusProvider(nil)` is enabled.
|
||||
|
||||
`provider.JSONHandler()` serves the same JSON as the push `json` format on `GET`/`HEAD`:
|
||||
|
||||
```go
|
||||
http.Handle("/metrics", provider.Handler()) // Prometheus text
|
||||
http.Handle("/metrics.json", provider.JSONHandler()) // JSON
|
||||
```
|
||||
|
||||
### Resetting Stats
|
||||
|
||||
- `provider.Reset()` clears counters, histograms and the cache-size gauge (live gauges such as in-flight requests are kept). Package-level `metrics.Reset()` does the same for the current provider if it implements `metrics.Resetter`.
|
||||
- `provider.PushAndReset()` pushes to the Pushgateway and resets only if the push succeeded (errors if no Pushgateway is configured).
|
||||
- `Config.PushgatewayResetOnPush: true` makes the automatic push loop do this on every tick.
|
||||
- `provider.ResetHandler()` is a `POST`-only endpoint (`?push=true` to push first). It has no auth: mount it on an internal route.
|
||||
|
||||
```go
|
||||
http.Handle("/metrics/reset", provider.ResetHandler())
|
||||
```
|
||||
|
||||
Note: the normal `/metrics` scrape is read-only and never clears anything. Observations recorded between a push and its reset are lost. Prometheus handles the counter drop as a reset, but if you reset often, prefer `increase()`/`rate()` over raw counter values.
|
||||
|
||||
### Custom Push Endpoint (Optional)
|
||||
|
||||
POST metrics to your own server, optionally clearing local stats after a 2xx reply:
|
||||
|
||||
```go
|
||||
provider := metrics.NewPrometheusProvider(&metrics.Config{
|
||||
PushEndpointURL: "https://collector.example.com/metrics",
|
||||
PushEndpointFormat: "json", // or "text" (Prometheus exposition, default)
|
||||
PushEndpointHeaders: map[string]string{"Authorization": "Bearer token"},
|
||||
PushEndpointInterval: 30, // seconds; 0 = manual only
|
||||
PushEndpointTimeout: 10, // seconds (default 10)
|
||||
PushEndpointResetOnSuccess: true, // clear local stats after a 2xx
|
||||
})
|
||||
|
||||
err := provider.PushToEndpoint(ctx) // manual push; also honours ResetOnSuccess
|
||||
provider.StopAutoPush() // stops the Pushgateway and endpoint loops
|
||||
```
|
||||
|
||||
The `json` body is a list of `{name, help, type, metrics:[{labels, value | count, sum, buckets}]}`. Failures (non-2xx, network, timeout) are logged and never reset stats, so the next tick retries with the accumulated data. The payload covers everything in the default Prometheus registry, including Go runtime metrics.
|
||||
|
||||
### Pushgateway Configuration (Optional)
|
||||
|
||||
For batch jobs, cron tasks, or short-lived processes, you can push metrics to Prometheus Pushgateway:
|
||||
@@ -457,10 +505,9 @@ scrape_configs:
|
||||
- ✅ Good: `method`, `status_code`
|
||||
- ❌ Bad: `user_id`, `timestamp`
|
||||
|
||||
2. **Path Normalization**: Normalize dynamic paths
|
||||
2. **Path Normalization**: Done automatically for the `path` label (see Configuration Options)
|
||||
```go
|
||||
// Instead of /api/users/123
|
||||
// Use /api/users/:id
|
||||
// /api/users/123 is recorded as /api/users/:id
|
||||
```
|
||||
|
||||
3. **Metric Naming**: Follow Prometheus conventions
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
package metrics
|
||||
|
||||
import "net/http"
|
||||
|
||||
// Config holds configuration for the metrics provider
|
||||
type Config struct {
|
||||
// Enabled determines whether metrics collection is enabled
|
||||
@@ -19,6 +21,17 @@ type Config struct {
|
||||
// Default: [0.001, 0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5]
|
||||
DBQueryBuckets []float64 `mapstructure:"db_query_buckets"`
|
||||
|
||||
// HTTPMaxPaths caps the number of distinct values of the "path" label on HTTP
|
||||
// metrics. Paths beyond the cap are reported as "other". Paths are already
|
||||
// normalized (route pattern, or dynamic segments replaced with ":id").
|
||||
// Default: 1024. Set to a negative value to disable the cap.
|
||||
HTTPMaxPaths int `mapstructure:"http_max_paths"`
|
||||
|
||||
// HTTPPathNormalizer optionally maps a request to its "path" label (e.g. the
|
||||
// matched route template of your router). Return "" to fall back to the
|
||||
// default behaviour (ServeMux pattern, then generic ID normalization).
|
||||
HTTPPathNormalizer func(*http.Request) string `mapstructure:"-"`
|
||||
|
||||
// PushgatewayURL is the URL of the Prometheus Pushgateway (optional)
|
||||
// If set, metrics will be pushed to this gateway instead of only being scraped
|
||||
// Example: "http://pushgateway:9091"
|
||||
@@ -32,6 +45,34 @@ type Config struct {
|
||||
// Only used if PushgatewayURL is set. If 0, automatic pushing is disabled.
|
||||
// Default: 0 (no automatic pushing)
|
||||
PushgatewayInterval int `mapstructure:"pushgateway_interval"`
|
||||
|
||||
// PushEndpointURL is a custom HTTP endpoint that metrics are POSTed to
|
||||
// (independent of Pushgateway). Example: "https://collector.example.com/metrics"
|
||||
PushEndpointURL string `mapstructure:"push_endpoint_url"`
|
||||
|
||||
// PushEndpointFormat is the request body format: "text" (Prometheus text
|
||||
// exposition, Content-Type text/plain; version=0.0.4) or "json".
|
||||
// Default: "text"
|
||||
PushEndpointFormat string `mapstructure:"push_endpoint_format"`
|
||||
|
||||
// PushEndpointHeaders are extra headers sent with each POST (e.g. Authorization).
|
||||
PushEndpointHeaders map[string]string `mapstructure:"push_endpoint_headers"`
|
||||
|
||||
// PushEndpointInterval is the interval in seconds for automatic POSTs.
|
||||
// If 0, automatic posting is disabled (PushToEndpoint can still be called manually).
|
||||
PushEndpointInterval int `mapstructure:"push_endpoint_interval"`
|
||||
|
||||
// PushEndpointTimeout is the per-request timeout in seconds. Default: 10
|
||||
PushEndpointTimeout int `mapstructure:"push_endpoint_timeout"`
|
||||
|
||||
// PushEndpointResetOnSuccess clears local counters and histograms after the
|
||||
// endpoint answers with a 2xx status. Default: false.
|
||||
PushEndpointResetOnSuccess bool `mapstructure:"push_endpoint_reset_on_success"`
|
||||
|
||||
// PushgatewayResetOnPush clears the local counters and histograms after each
|
||||
// successful push (automatic or via PushAndReset), so each push carries only
|
||||
// the activity since the previous one. Default: false.
|
||||
PushgatewayResetOnPush bool `mapstructure:"pushgateway_reset_on_push"`
|
||||
}
|
||||
|
||||
// DefaultConfig returns a Config with sensible defaults
|
||||
@@ -43,6 +84,7 @@ func DefaultConfig() *Config {
|
||||
HTTPRequestBuckets: []float64{0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5, 10},
|
||||
// DB queries are usually faster
|
||||
DBQueryBuckets: []float64{0.001, 0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5},
|
||||
HTTPMaxPaths: defaultHTTPMaxPaths,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -57,6 +99,17 @@ func (c *Config) ApplyDefaults() {
|
||||
if len(c.DBQueryBuckets) == 0 {
|
||||
c.DBQueryBuckets = []float64{0.001, 0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5}
|
||||
}
|
||||
if c.PushEndpointURL != "" {
|
||||
if c.PushEndpointFormat == "" {
|
||||
c.PushEndpointFormat = "text"
|
||||
}
|
||||
if c.PushEndpointTimeout <= 0 {
|
||||
c.PushEndpointTimeout = 10
|
||||
}
|
||||
}
|
||||
if c.HTTPMaxPaths == 0 {
|
||||
c.HTTPMaxPaths = defaultHTTPMaxPaths
|
||||
}
|
||||
// Set default job name if pushgateway is configured but job name is empty
|
||||
if c.PushgatewayURL != "" && c.PushgatewayJobName == "" {
|
||||
c.PushgatewayJobName = "resolvespec"
|
||||
|
||||
@@ -0,0 +1,189 @@
|
||||
package metrics
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/prometheus/client_golang/prometheus"
|
||||
dto "github.com/prometheus/client_model/go"
|
||||
"github.com/prometheus/common/expfmt"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
)
|
||||
|
||||
const textContentType = "text/plain; version=0.0.4; charset=utf-8"
|
||||
|
||||
// endpointPusher POSTs gathered metrics to a user-configured HTTP endpoint.
|
||||
type endpointPusher struct {
|
||||
url string
|
||||
format string
|
||||
headers map[string]string
|
||||
client *http.Client
|
||||
resetOnOK bool
|
||||
provider *PrometheusProvider
|
||||
gatherer prometheus.Gatherer
|
||||
stopOnce sync.Once
|
||||
stopCh chan struct{}
|
||||
startedMu sync.Mutex
|
||||
started bool
|
||||
}
|
||||
|
||||
func newEndpointPusher(cfg *Config, p *PrometheusProvider) *endpointPusher {
|
||||
return &endpointPusher{
|
||||
url: cfg.PushEndpointURL,
|
||||
format: cfg.PushEndpointFormat,
|
||||
headers: cfg.PushEndpointHeaders,
|
||||
client: &http.Client{Timeout: time.Duration(cfg.PushEndpointTimeout) * time.Second},
|
||||
resetOnOK: cfg.PushEndpointResetOnSuccess,
|
||||
provider: p,
|
||||
gatherer: prometheus.DefaultGatherer,
|
||||
stopCh: make(chan struct{}),
|
||||
}
|
||||
}
|
||||
|
||||
func (e *endpointPusher) start(interval time.Duration) {
|
||||
e.startedMu.Lock()
|
||||
defer e.startedMu.Unlock()
|
||||
if e.started {
|
||||
return
|
||||
}
|
||||
e.started = true
|
||||
go func() {
|
||||
t := time.NewTicker(interval)
|
||||
defer t.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-t.C:
|
||||
if err := e.push(context.Background()); err != nil {
|
||||
logger.Warn("Failed to push metrics to endpoint %s: %v", e.url, err)
|
||||
}
|
||||
case <-e.stopCh:
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
func (e *endpointPusher) stop() {
|
||||
e.stopOnce.Do(func() { close(e.stopCh) })
|
||||
}
|
||||
|
||||
func (e *endpointPusher) push(ctx context.Context) error {
|
||||
mfs, err := e.gatherer.Gather()
|
||||
if err != nil && len(mfs) == 0 {
|
||||
return fmt.Errorf("gather metrics: %w", err)
|
||||
}
|
||||
|
||||
body, contentType, err := encodeMetrics(mfs, e.format)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, e.url, bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req.Header.Set("Content-Type", contentType)
|
||||
for k, v := range e.headers {
|
||||
req.Header.Set(k, v)
|
||||
}
|
||||
|
||||
resp, err := e.client.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
snippet, _ := io.ReadAll(io.LimitReader(resp.Body, 512))
|
||||
if resp.StatusCode < 200 || resp.StatusCode > 299 {
|
||||
return fmt.Errorf("endpoint returned %s: %s", resp.Status, bytes.TrimSpace(snippet))
|
||||
}
|
||||
|
||||
if e.resetOnOK {
|
||||
e.provider.Reset()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func encodeMetrics(mfs []*dto.MetricFamily, format string) (body []byte, contentType string, err error) {
|
||||
switch format {
|
||||
case "json":
|
||||
b, err := json.Marshal(toJSONFamilies(mfs))
|
||||
return b, "application/json", err
|
||||
case "", "text":
|
||||
var buf bytes.Buffer
|
||||
enc := expfmt.NewEncoder(&buf, expfmt.NewFormat(expfmt.TypeTextPlain))
|
||||
for _, mf := range mfs {
|
||||
if err := enc.Encode(mf); err != nil {
|
||||
return nil, "", err
|
||||
}
|
||||
}
|
||||
return buf.Bytes(), textContentType, nil
|
||||
default:
|
||||
return nil, "", fmt.Errorf("unsupported push endpoint format %q", format)
|
||||
}
|
||||
}
|
||||
|
||||
type jsonFamily struct {
|
||||
Name string `json:"name"`
|
||||
Help string `json:"help,omitempty"`
|
||||
Type string `json:"type"`
|
||||
Metrics []jsonMetric `json:"metrics"`
|
||||
}
|
||||
|
||||
type jsonMetric struct {
|
||||
Labels map[string]string `json:"labels,omitempty"`
|
||||
Value *float64 `json:"value,omitempty"`
|
||||
Count *uint64 `json:"count,omitempty"`
|
||||
Sum *float64 `json:"sum,omitempty"`
|
||||
Buckets []jsonBucket `json:"buckets,omitempty"`
|
||||
}
|
||||
|
||||
type jsonBucket struct {
|
||||
UpperBound float64 `json:"le"`
|
||||
Count uint64 `json:"count"`
|
||||
}
|
||||
|
||||
func toJSONFamilies(mfs []*dto.MetricFamily) []jsonFamily {
|
||||
out := make([]jsonFamily, 0, len(mfs))
|
||||
for _, mf := range mfs {
|
||||
f := jsonFamily{Name: mf.GetName(), Help: mf.GetHelp(), Type: mf.GetType().String()}
|
||||
for _, m := range mf.GetMetric() {
|
||||
jm := jsonMetric{}
|
||||
if len(m.GetLabel()) > 0 {
|
||||
jm.Labels = make(map[string]string, len(m.GetLabel()))
|
||||
for _, l := range m.GetLabel() {
|
||||
jm.Labels[l.GetName()] = l.GetValue()
|
||||
}
|
||||
}
|
||||
switch {
|
||||
case m.Counter != nil:
|
||||
v := m.Counter.GetValue()
|
||||
jm.Value = &v
|
||||
case m.Gauge != nil:
|
||||
v := m.Gauge.GetValue()
|
||||
jm.Value = &v
|
||||
case m.Untyped != nil:
|
||||
v := m.Untyped.GetValue()
|
||||
jm.Value = &v
|
||||
case m.Histogram != nil:
|
||||
c, s := m.Histogram.GetSampleCount(), m.Histogram.GetSampleSum()
|
||||
jm.Count, jm.Sum = &c, &s
|
||||
for _, b := range m.Histogram.GetBucket() {
|
||||
jm.Buckets = append(jm.Buckets, jsonBucket{UpperBound: b.GetUpperBound(), Count: b.GetCumulativeCount()})
|
||||
}
|
||||
case m.Summary != nil:
|
||||
c, s := m.Summary.GetSampleCount(), m.Summary.GetSampleSum()
|
||||
jm.Count, jm.Sum = &c, &s
|
||||
}
|
||||
f.Metrics = append(f.Metrics, jm)
|
||||
}
|
||||
out = append(out, f)
|
||||
}
|
||||
return out
|
||||
}
|
||||
@@ -47,6 +47,21 @@ type Provider interface {
|
||||
Handler() http.Handler
|
||||
}
|
||||
|
||||
// Resetter is optionally implemented by providers that can clear their recorded stats.
|
||||
type Resetter interface {
|
||||
Reset()
|
||||
}
|
||||
|
||||
// Reset clears the current provider's stats if it supports resetting.
|
||||
// It returns false if the provider does not implement Resetter.
|
||||
func Reset() bool {
|
||||
if r, ok := GetProvider().(Resetter); ok {
|
||||
r.Reset()
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// globalProvider is the global metrics provider, protected by globalProviderMu.
|
||||
var (
|
||||
globalProviderMu sync.RWMutex
|
||||
|
||||
@@ -0,0 +1,175 @@
|
||||
package metrics
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
)
|
||||
|
||||
const (
|
||||
// defaultHTTPMaxPaths is the default cap on distinct values of the "path" label.
|
||||
defaultHTTPMaxPaths = 1024
|
||||
|
||||
// overflowPathLabel is used once the cap on distinct path labels is reached.
|
||||
overflowPathLabel = "other"
|
||||
)
|
||||
|
||||
// routeLabel returns the low-cardinality path label for a request, preferring
|
||||
// (in order): the custom normalizer, the matched ServeMux pattern, and finally
|
||||
// the generic normalization of the raw URL path.
|
||||
func routeLabel(r *http.Request, custom func(*http.Request) string) string {
|
||||
if custom != nil {
|
||||
if p := custom(r); p != "" {
|
||||
return p
|
||||
}
|
||||
}
|
||||
if r.Pattern != "" {
|
||||
return stripPatternMethod(r.Pattern)
|
||||
}
|
||||
return NormalizePath(r.URL.Path)
|
||||
}
|
||||
|
||||
// stripPatternMethod removes the optional "METHOD " prefix (and host) from a
|
||||
// Go 1.22+ ServeMux pattern, e.g. "GET /users/{id}" -> "/users/{id}".
|
||||
func stripPatternMethod(pattern string) string {
|
||||
if i := strings.IndexByte(pattern, ' '); i >= 0 {
|
||||
pattern = strings.TrimLeft(pattern[i+1:], " ")
|
||||
}
|
||||
if i := strings.IndexByte(pattern, '/'); i > 0 {
|
||||
pattern = pattern[i:] // drop host part
|
||||
}
|
||||
return pattern
|
||||
}
|
||||
|
||||
// NormalizePath replaces dynamic-looking path segments (numeric IDs, UUIDs,
|
||||
// long hex strings and other long opaque tokens) with ":id" so that
|
||||
// /users/123 and /users/456 share one label value.
|
||||
func NormalizePath(path string) string {
|
||||
if path == "" {
|
||||
return "/"
|
||||
}
|
||||
if !strings.Contains(path, "/") {
|
||||
return path
|
||||
}
|
||||
segs := strings.Split(path, "/")
|
||||
for i, s := range segs {
|
||||
if isDynamicSegment(s) {
|
||||
segs[i] = ":id"
|
||||
}
|
||||
}
|
||||
return strings.Join(segs, "/")
|
||||
}
|
||||
|
||||
func isDynamicSegment(s string) bool {
|
||||
if s == "" {
|
||||
return false
|
||||
}
|
||||
if allDigits(s) {
|
||||
return true
|
||||
}
|
||||
if isUUID(s) {
|
||||
return true
|
||||
}
|
||||
// Long hex strings (hashes, object IDs)
|
||||
if len(s) >= 16 && allHex(s) {
|
||||
return true
|
||||
}
|
||||
// Long opaque tokens containing digits (base64/ULID-like)
|
||||
if len(s) >= 24 && hasDigit(s) && !strings.ContainsAny(s, ".") {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func allDigits(s string) bool {
|
||||
for i := 0; i < len(s); i++ {
|
||||
if s[i] < '0' || s[i] > '9' {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func hasDigit(s string) bool {
|
||||
for i := 0; i < len(s); i++ {
|
||||
if s[i] >= '0' && s[i] <= '9' {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func allHex(s string) bool {
|
||||
for i := 0; i < len(s); i++ {
|
||||
c := s[i]
|
||||
if !isHexByte(c) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func isUUID(s string) bool {
|
||||
if len(s) != 36 {
|
||||
return false
|
||||
}
|
||||
for i := 0; i < len(s); i++ {
|
||||
c := s[i]
|
||||
switch i {
|
||||
case 8, 13, 18, 23:
|
||||
if c != '-' {
|
||||
return false
|
||||
}
|
||||
default:
|
||||
if !isHexByte(c) {
|
||||
return false
|
||||
}
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// pathLimiter bounds the number of distinct path label values. Once the cap is
|
||||
// reached, unseen paths are reported as "other".
|
||||
type pathLimiter struct {
|
||||
mu sync.RWMutex
|
||||
max int // <= 0 disables the cap
|
||||
seen map[string]struct{}
|
||||
}
|
||||
|
||||
func newPathLimiter(limit int) *pathLimiter {
|
||||
return &pathLimiter{max: limit, seen: make(map[string]struct{})}
|
||||
}
|
||||
|
||||
func (l *pathLimiter) label(path string) string {
|
||||
if l.max <= 0 {
|
||||
return path
|
||||
}
|
||||
l.mu.RLock()
|
||||
_, ok := l.seen[path]
|
||||
l.mu.RUnlock()
|
||||
if ok {
|
||||
return path
|
||||
}
|
||||
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
if _, ok := l.seen[path]; ok {
|
||||
return path
|
||||
}
|
||||
if len(l.seen) >= l.max {
|
||||
return overflowPathLabel
|
||||
}
|
||||
l.seen[path] = struct{}{}
|
||||
return path
|
||||
}
|
||||
|
||||
func (l *pathLimiter) reset() {
|
||||
l.mu.Lock()
|
||||
l.seen = make(map[string]struct{})
|
||||
l.mu.Unlock()
|
||||
}
|
||||
|
||||
func isHexByte(c byte) bool {
|
||||
return c >= '0' && c <= '9' || c >= 'a' && c <= 'f' || c >= 'A' && c <= 'F'
|
||||
}
|
||||
@@ -0,0 +1,237 @@
|
||||
package metrics
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"github.com/prometheus/client_golang/prometheus"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestNormalizePath(t *testing.T) {
|
||||
cases := map[string]string{
|
||||
"": "/",
|
||||
"/": "/",
|
||||
"/users": "/users",
|
||||
"/users/123": "/users/:id",
|
||||
"/users/123/orders/9": "/users/:id/orders/:id",
|
||||
"/x/550e8400-e29b-41d4-a716-446655440000": "/x/:id",
|
||||
"/x/507f1f77bcf86cd799439011": "/x/:id",
|
||||
"/api/public/users": "/api/public/users",
|
||||
"/files/report.v2": "/files/report.v2",
|
||||
}
|
||||
for in, want := range cases {
|
||||
if got := NormalizePath(in); got != want {
|
||||
t.Errorf("NormalizePath(%q) = %q, want %q", in, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRouteLabel(t *testing.T) {
|
||||
r := httptest.NewRequest("GET", "/users/42", nil)
|
||||
if got := routeLabel(r, nil); got != "/users/:id" {
|
||||
t.Errorf("fallback = %q", got)
|
||||
}
|
||||
r.Pattern = "GET /users/{id}"
|
||||
if got := routeLabel(r, nil); got != "/users/{id}" {
|
||||
t.Errorf("pattern = %q", got)
|
||||
}
|
||||
got := routeLabel(r, func(*http.Request) string { return "/custom" })
|
||||
if got != "/custom" {
|
||||
t.Errorf("custom = %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPathLimiter(t *testing.T) {
|
||||
l := newPathLimiter(2)
|
||||
for _, p := range []string{"/a", "/b", "/a"} {
|
||||
if got := l.label(p); got != p {
|
||||
t.Errorf("label(%q) = %q", p, got)
|
||||
}
|
||||
}
|
||||
if got := l.label("/c"); got != overflowPathLabel {
|
||||
t.Errorf("overflow = %q", got)
|
||||
}
|
||||
if got := newPathLimiter(-1).label("/z"); got != "/z" {
|
||||
t.Errorf("disabled = %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMiddlewareUsesPattern(t *testing.T) {
|
||||
p := NewPrometheusProvider(&Config{Enabled: true, Namespace: "pathtest"})
|
||||
mux := http.NewServeMux()
|
||||
mux.HandleFunc("GET /users/{id}", func(w http.ResponseWriter, r *http.Request) {})
|
||||
h := p.Middleware(mux)
|
||||
for _, id := range []string{"1", "2", "abc"} {
|
||||
h.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest("GET", "/users/"+id, nil))
|
||||
}
|
||||
if n := len(p.pathLimiter.seen); n != 1 {
|
||||
t.Errorf("distinct paths = %d, want 1", n)
|
||||
}
|
||||
if _, ok := p.pathLimiter.seen["/users/{id}"]; !ok {
|
||||
t.Errorf("seen = %v", p.pathLimiter.seen)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResetAndHandler(t *testing.T) {
|
||||
p := NewPrometheusProvider(&Config{Enabled: true, Namespace: "resettest"})
|
||||
p.RecordHTTPRequest("GET", "/a/1", "200", 0)
|
||||
p.RecordDBQuery("SELECT", "s", "e", "t", 0, nil)
|
||||
p.IncRequestsInFlight()
|
||||
|
||||
count := func() int {
|
||||
mfs, _ := prometheus.DefaultGatherer.Gather()
|
||||
n := 0
|
||||
for _, mf := range mfs {
|
||||
if strings.HasPrefix(mf.GetName(), "resettest_") && mf.GetName() != "resettest_http_requests_in_flight" && mf.GetName() != "resettest_event_queue_size" {
|
||||
n += len(mf.GetMetric())
|
||||
}
|
||||
}
|
||||
return n
|
||||
}
|
||||
if count() == 0 {
|
||||
t.Fatal("expected recorded series")
|
||||
}
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
p.ResetHandler().ServeHTTP(rec, httptest.NewRequest("GET", "/reset", nil))
|
||||
if rec.Code != http.StatusMethodNotAllowed || count() == 0 {
|
||||
t.Fatalf("GET should be rejected, code=%d", rec.Code)
|
||||
}
|
||||
|
||||
rec = httptest.NewRecorder()
|
||||
p.ResetHandler().ServeHTTP(rec, httptest.NewRequest("POST", "/reset", nil))
|
||||
if rec.Code != http.StatusNoContent || count() != 0 {
|
||||
t.Fatalf("reset failed, code=%d series=%d", rec.Code, count())
|
||||
}
|
||||
if len(p.pathLimiter.seen) != 0 {
|
||||
t.Error("path limiter not reset")
|
||||
}
|
||||
|
||||
// push=true without a pushgateway must fail and not be silent
|
||||
rec = httptest.NewRecorder()
|
||||
p.ResetHandler().ServeHTTP(rec, httptest.NewRequest("POST", "/reset?push=true", nil))
|
||||
if rec.Code != http.StatusBadGateway {
|
||||
t.Errorf("push without gateway code=%d", rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPushAndResetKeepsStatsOnFailure(t *testing.T) {
|
||||
p := NewPrometheusProvider(&Config{Enabled: true, Namespace: "pushfail", PushgatewayURL: "http://127.0.0.1:1"})
|
||||
p.RecordHTTPRequest("GET", "/a", "200", 0)
|
||||
if err := p.PushAndReset(); err == nil {
|
||||
t.Fatal("expected push error")
|
||||
}
|
||||
if len(p.pathLimiter.seen) != 1 {
|
||||
t.Error("stats were reset despite failed push")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPushToEndpoint(t *testing.T) {
|
||||
for _, format := range []string{"text", "json"} {
|
||||
var gotCT, gotAuth string
|
||||
var gotBody []byte
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
t.Errorf("method = %s", r.Method)
|
||||
}
|
||||
gotCT, gotAuth = r.Header.Get("Content-Type"), r.Header.Get("Authorization")
|
||||
gotBody, _ = io.ReadAll(r.Body)
|
||||
}))
|
||||
|
||||
ns := "ep" + format
|
||||
p := NewPrometheusProvider(&Config{
|
||||
Enabled: true,
|
||||
Namespace: ns,
|
||||
PushEndpointURL: srv.URL,
|
||||
PushEndpointFormat: format,
|
||||
PushEndpointHeaders: map[string]string{"Authorization": "Bearer x"},
|
||||
PushEndpointResetOnSuccess: true,
|
||||
})
|
||||
p.RecordHTTPRequest("GET", "/a", "200", time.Millisecond)
|
||||
|
||||
if err := p.PushToEndpoint(context.Background()); err != nil {
|
||||
t.Fatalf("%s: %v", format, err)
|
||||
}
|
||||
srv.Close()
|
||||
if gotAuth != "Bearer x" || !strings.Contains(string(gotBody), ns+"_http_requests_total") {
|
||||
t.Errorf("%s: auth=%q body=%.200s", format, gotAuth, gotBody)
|
||||
}
|
||||
if format == "json" && gotCT != "application/json" || format == "text" && !strings.HasPrefix(gotCT, "text/plain") {
|
||||
t.Errorf("%s: content-type %q", format, gotCT)
|
||||
}
|
||||
if len(p.pathLimiter.seen) != 0 {
|
||||
t.Errorf("%s: stats not reset after success", format)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPushToEndpointFailureKeepsStats(t *testing.T) {
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
http.Error(w, "nope", http.StatusInternalServerError)
|
||||
}))
|
||||
defer srv.Close()
|
||||
p := NewPrometheusProvider(&Config{Enabled: true, Namespace: "epfail", PushEndpointURL: srv.URL, PushEndpointResetOnSuccess: true})
|
||||
p.RecordHTTPRequest("GET", "/a", "200", 0)
|
||||
if err := p.PushToEndpoint(context.Background()); err == nil {
|
||||
t.Fatal("expected error on 500")
|
||||
}
|
||||
if len(p.pathLimiter.seen) != 1 {
|
||||
t.Error("stats reset despite failure")
|
||||
}
|
||||
if err := NewPrometheusProvider(&Config{Enabled: true, Namespace: "epnone"}).PushToEndpoint(context.Background()); err == nil {
|
||||
t.Error("expected error without endpoint")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDisabledProvider(t *testing.T) {
|
||||
p := NewPrometheusProvider(&Config{Namespace: "disabled", PushEndpointURL: "http://127.0.0.1:1", PushEndpointInterval: 1})
|
||||
p.RecordHTTPRequest("GET", "/a", "200", 0)
|
||||
if len(p.pathLimiter.seen) != 0 {
|
||||
t.Error("disabled provider recorded")
|
||||
}
|
||||
for name, h := range map[string]http.Handler{"handler": p.Handler(), "json": p.JSONHandler()} {
|
||||
rec := httptest.NewRecorder()
|
||||
h.ServeHTTP(rec, httptest.NewRequest("GET", "/m", nil))
|
||||
if rec.Code != http.StatusNotFound {
|
||||
t.Errorf("%s code=%d", name, rec.Code)
|
||||
}
|
||||
}
|
||||
if p.endpoint != nil || p.PushToEndpoint(context.Background()) == nil || p.Push() == nil {
|
||||
t.Error("disabled provider must not push")
|
||||
}
|
||||
}
|
||||
|
||||
func TestJSONHandler(t *testing.T) {
|
||||
p := NewPrometheusProvider(&Config{Enabled: true, Namespace: "jsonpull"})
|
||||
p.RecordHTTPRequest("GET", "/a", "200", time.Millisecond)
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
p.JSONHandler().ServeHTTP(rec, httptest.NewRequest("GET", "/m", nil))
|
||||
if rec.Code != 200 || rec.Header().Get("Content-Type") != "application/json" {
|
||||
t.Fatalf("code=%d ct=%q", rec.Code, rec.Header().Get("Content-Type"))
|
||||
}
|
||||
var fams []map[string]any
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &fams); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
found := false
|
||||
for _, f := range fams {
|
||||
if f["name"] == "jsonpull_http_requests_total" {
|
||||
found = true
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Error("metric family missing from JSON")
|
||||
}
|
||||
|
||||
rec = httptest.NewRecorder()
|
||||
p.JSONHandler().ServeHTTP(rec, httptest.NewRequest("POST", "/m", nil))
|
||||
if rec.Code != http.StatusMethodNotAllowed {
|
||||
t.Errorf("POST code=%d", rec.Code)
|
||||
}
|
||||
}
|
||||
+200
-6
@@ -1,6 +1,8 @@
|
||||
package metrics
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"time"
|
||||
@@ -9,8 +11,12 @@ import (
|
||||
"github.com/prometheus/client_golang/prometheus/promauto"
|
||||
"github.com/prometheus/client_golang/prometheus/promhttp"
|
||||
"github.com/prometheus/client_golang/prometheus/push"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
)
|
||||
|
||||
var errMetricsDisabled = errors.New("metrics: disabled")
|
||||
|
||||
// PrometheusProvider implements the Provider interface using Prometheus
|
||||
type PrometheusProvider struct {
|
||||
requestDuration *prometheus.HistogramVec
|
||||
@@ -27,9 +33,16 @@ type PrometheusProvider struct {
|
||||
eventQueueSize prometheus.Gauge
|
||||
panicsTotal *prometheus.CounterVec
|
||||
|
||||
pathLimiter *pathLimiter
|
||||
pathNormalizer func(*http.Request) string
|
||||
|
||||
enabled bool
|
||||
endpoint *endpointPusher
|
||||
|
||||
// Pushgateway fields (optional)
|
||||
pushgatewayURL string
|
||||
pushgatewayJobName string
|
||||
resetOnPush bool
|
||||
pusher *push.Pusher
|
||||
pushTicker *time.Ticker
|
||||
pushStop chan bool
|
||||
@@ -55,6 +68,7 @@ func NewPrometheusProvider(cfg *Config) *PrometheusProvider {
|
||||
}
|
||||
|
||||
p := &PrometheusProvider{
|
||||
enabled: cfg.Enabled,
|
||||
requestDuration: promauto.NewHistogramVec(
|
||||
prometheus.HistogramOpts{
|
||||
Name: metricName("http_request_duration_seconds"),
|
||||
@@ -149,12 +163,17 @@ func NewPrometheusProvider(cfg *Config) *PrometheusProvider {
|
||||
[]string{"method"},
|
||||
),
|
||||
|
||||
pathLimiter: newPathLimiter(cfg.HTTPMaxPaths),
|
||||
pathNormalizer: cfg.HTTPPathNormalizer,
|
||||
|
||||
pushgatewayURL: cfg.PushgatewayURL,
|
||||
pushgatewayJobName: cfg.PushgatewayJobName,
|
||||
resetOnPush: cfg.PushgatewayResetOnPush,
|
||||
}
|
||||
|
||||
// Initialize pushgateway if configured
|
||||
if cfg.PushgatewayURL != "" {
|
||||
// Pushing is never started for a disabled provider
|
||||
if cfg.PushgatewayURL != "" && cfg.Enabled {
|
||||
p.pusher = push.New(cfg.PushgatewayURL, cfg.PushgatewayJobName).
|
||||
Gatherer(prometheus.DefaultGatherer)
|
||||
|
||||
@@ -166,6 +185,13 @@ func NewPrometheusProvider(cfg *Config) *PrometheusProvider {
|
||||
}
|
||||
}
|
||||
|
||||
if cfg.PushEndpointURL != "" && cfg.Enabled {
|
||||
p.endpoint = newEndpointPusher(cfg, p)
|
||||
if cfg.PushEndpointInterval > 0 {
|
||||
p.endpoint.start(time.Duration(cfg.PushEndpointInterval) * time.Second)
|
||||
}
|
||||
}
|
||||
|
||||
return p
|
||||
}
|
||||
|
||||
@@ -188,23 +214,37 @@ func (rw *ResponseWriter) WriteHeader(code int) {
|
||||
}
|
||||
|
||||
// RecordHTTPRequest implements Provider interface
|
||||
// The path is normalized and capped to keep label cardinality bounded.
|
||||
func (p *PrometheusProvider) RecordHTTPRequest(method, path, status string, duration time.Duration) {
|
||||
if !p.enabled {
|
||||
return
|
||||
}
|
||||
path = p.pathLimiter.label(NormalizePath(path))
|
||||
p.requestDuration.WithLabelValues(method, path, status).Observe(duration.Seconds())
|
||||
p.requestTotal.WithLabelValues(method, path, status).Inc()
|
||||
}
|
||||
|
||||
// IncRequestsInFlight implements Provider interface
|
||||
func (p *PrometheusProvider) IncRequestsInFlight() {
|
||||
if !p.enabled {
|
||||
return
|
||||
}
|
||||
p.requestsInFlight.Inc()
|
||||
}
|
||||
|
||||
// DecRequestsInFlight implements Provider interface
|
||||
func (p *PrometheusProvider) DecRequestsInFlight() {
|
||||
if !p.enabled {
|
||||
return
|
||||
}
|
||||
p.requestsInFlight.Dec()
|
||||
}
|
||||
|
||||
// RecordDBQuery implements Provider interface
|
||||
func (p *PrometheusProvider) RecordDBQuery(operation, schema, entity, table string, duration time.Duration, err error) {
|
||||
if !p.enabled {
|
||||
return
|
||||
}
|
||||
status := "success"
|
||||
if err != nil {
|
||||
status = "error"
|
||||
@@ -215,47 +255,115 @@ func (p *PrometheusProvider) RecordDBQuery(operation, schema, entity, table stri
|
||||
|
||||
// RecordCacheHit implements Provider interface
|
||||
func (p *PrometheusProvider) RecordCacheHit(provider string) {
|
||||
if !p.enabled {
|
||||
return
|
||||
}
|
||||
p.cacheHits.WithLabelValues(provider).Inc()
|
||||
}
|
||||
|
||||
// RecordCacheMiss implements Provider interface
|
||||
func (p *PrometheusProvider) RecordCacheMiss(provider string) {
|
||||
if !p.enabled {
|
||||
return
|
||||
}
|
||||
p.cacheMisses.WithLabelValues(provider).Inc()
|
||||
}
|
||||
|
||||
// UpdateCacheSize implements Provider interface
|
||||
func (p *PrometheusProvider) UpdateCacheSize(provider string, size int64) {
|
||||
if !p.enabled {
|
||||
return
|
||||
}
|
||||
p.cacheSize.WithLabelValues(provider).Set(float64(size))
|
||||
}
|
||||
|
||||
// RecordEventPublished implements Provider interface
|
||||
func (p *PrometheusProvider) RecordEventPublished(source, eventType string) {
|
||||
if !p.enabled {
|
||||
return
|
||||
}
|
||||
p.eventPublished.WithLabelValues(source, eventType).Inc()
|
||||
}
|
||||
|
||||
// RecordEventProcessed implements Provider interface
|
||||
func (p *PrometheusProvider) RecordEventProcessed(source, eventType, status string, duration time.Duration) {
|
||||
if !p.enabled {
|
||||
return
|
||||
}
|
||||
p.eventProcessed.WithLabelValues(source, eventType, status).Inc()
|
||||
p.eventDuration.WithLabelValues(source, eventType).Observe(duration.Seconds())
|
||||
}
|
||||
|
||||
// UpdateEventQueueSize implements Provider interface
|
||||
func (p *PrometheusProvider) UpdateEventQueueSize(size int64) {
|
||||
if !p.enabled {
|
||||
return
|
||||
}
|
||||
p.eventQueueSize.Set(float64(size))
|
||||
}
|
||||
|
||||
// RecordPanic implements the Provider interface
|
||||
func (p *PrometheusProvider) RecordPanic(methodName string) {
|
||||
if !p.enabled {
|
||||
return
|
||||
}
|
||||
p.panicsTotal.WithLabelValues(methodName).Inc()
|
||||
}
|
||||
|
||||
// Handler implements Provider interface
|
||||
// It responds 404 when metrics are disabled.
|
||||
func (p *PrometheusProvider) Handler() http.Handler {
|
||||
if !p.enabled {
|
||||
return disabledHandler()
|
||||
}
|
||||
return promhttp.Handler()
|
||||
}
|
||||
|
||||
// JSONHandler returns an HTTP handler serving the current metrics as JSON
|
||||
// (same shape as the "json" push endpoint format). Only GET and HEAD are
|
||||
// accepted, and it responds 404 when metrics are disabled. It performs no
|
||||
// authentication; mount it on an internal/protected route.
|
||||
func (p *PrometheusProvider) JSONHandler() http.Handler {
|
||||
if !p.enabled {
|
||||
return disabledHandler()
|
||||
}
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet && r.Method != http.MethodHead {
|
||||
w.Header().Set("Allow", "GET, HEAD")
|
||||
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
mfs, err := prometheus.DefaultGatherer.Gather()
|
||||
if err != nil && len(mfs) == 0 {
|
||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
body, contentType, err := encodeMetrics(mfs, "json")
|
||||
if err != nil {
|
||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", contentType)
|
||||
if r.Method == http.MethodGet {
|
||||
if _, err := w.Write(body); err != nil {
|
||||
logger.Warn("Failed to write metrics JSON: %v", err)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func disabledHandler() http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
http.Error(w, "metrics disabled", http.StatusNotFound)
|
||||
})
|
||||
}
|
||||
|
||||
// Middleware returns an HTTP middleware that collects metrics
|
||||
// When metrics are disabled it returns next unchanged.
|
||||
func (p *PrometheusProvider) Middleware(next http.Handler) http.Handler {
|
||||
if !p.enabled {
|
||||
return next
|
||||
}
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
start := time.Now()
|
||||
|
||||
@@ -273,13 +381,17 @@ func (p *PrometheusProvider) Middleware(next http.Handler) http.Handler {
|
||||
duration := time.Since(start)
|
||||
status := strconv.Itoa(rw.statusCode)
|
||||
|
||||
p.RecordHTTPRequest(r.Method, r.URL.Path, status, duration)
|
||||
// Read the label after next has run so the router has set r.Pattern.
|
||||
p.RecordHTTPRequest(r.Method, routeLabel(r, p.pathNormalizer), status, duration)
|
||||
})
|
||||
}
|
||||
|
||||
// Push manually pushes metrics to the configured Pushgateway
|
||||
// Returns an error if pushing fails or if Pushgateway is not configured
|
||||
func (p *PrometheusProvider) Push() error {
|
||||
if !p.enabled {
|
||||
return errMetricsDisabled
|
||||
}
|
||||
if p.pusher == nil {
|
||||
return nil // Pushgateway not configured, silently skip
|
||||
}
|
||||
@@ -291,10 +403,15 @@ func (p *PrometheusProvider) startAutoPush() {
|
||||
for {
|
||||
select {
|
||||
case <-p.pushTicker.C:
|
||||
if err := p.Push(); err != nil {
|
||||
// Log error but continue pushing
|
||||
// Note: In production, you might want to use a proper logger
|
||||
_ = err
|
||||
var err error
|
||||
if p.resetOnPush {
|
||||
err = p.PushAndReset()
|
||||
} else {
|
||||
err = p.Push()
|
||||
}
|
||||
if err != nil {
|
||||
// Log and keep going; the next tick retries (and nothing was reset)
|
||||
logger.Warn("Failed to push metrics to Pushgateway: %v", err)
|
||||
}
|
||||
case <-p.pushStop:
|
||||
p.pushTicker.Stop()
|
||||
@@ -303,10 +420,87 @@ func (p *PrometheusProvider) startAutoPush() {
|
||||
}
|
||||
}
|
||||
|
||||
// Reset clears all recorded counters, histograms and labelled gauges (cache size)
|
||||
// and forgets the tracked HTTP path labels. Live gauges (requests in flight,
|
||||
// event queue size) are left untouched since they reflect current state.
|
||||
// Prometheus treats the drop in counters as a counter reset, so rate() and
|
||||
// increase() keep working on the scraper side.
|
||||
func (p *PrometheusProvider) Reset() {
|
||||
p.requestDuration.Reset()
|
||||
p.requestTotal.Reset()
|
||||
p.dbQueryDuration.Reset()
|
||||
p.dbQueryTotal.Reset()
|
||||
p.cacheHits.Reset()
|
||||
p.cacheMisses.Reset()
|
||||
p.cacheSize.Reset()
|
||||
p.eventPublished.Reset()
|
||||
p.eventProcessed.Reset()
|
||||
p.eventDuration.Reset()
|
||||
p.panicsTotal.Reset()
|
||||
p.pathLimiter.reset()
|
||||
}
|
||||
|
||||
// PushAndReset pushes metrics to the Pushgateway and, only if the push
|
||||
// succeeded, clears the local stats. Returns an error if Pushgateway is not
|
||||
// configured, so stats are never discarded without being delivered. Observations
|
||||
// recorded between the push and the reset are lost.
|
||||
func (p *PrometheusProvider) PushAndReset() error {
|
||||
if !p.enabled {
|
||||
return errMetricsDisabled
|
||||
}
|
||||
if p.pusher == nil {
|
||||
return errors.New("metrics: pushgateway not configured, refusing to reset")
|
||||
}
|
||||
if err := p.pusher.Push(); err != nil {
|
||||
return err
|
||||
}
|
||||
p.Reset()
|
||||
return nil
|
||||
}
|
||||
|
||||
// PushToEndpoint POSTs the current metrics to the configured PushEndpointURL.
|
||||
// If PushEndpointResetOnSuccess is set, local stats are cleared after a 2xx reply.
|
||||
// Returns an error if no endpoint is configured.
|
||||
func (p *PrometheusProvider) PushToEndpoint(ctx context.Context) error {
|
||||
if !p.enabled {
|
||||
return errMetricsDisabled
|
||||
}
|
||||
if p.endpoint == nil {
|
||||
return errors.New("metrics: push endpoint not configured")
|
||||
}
|
||||
return p.endpoint.push(ctx)
|
||||
}
|
||||
|
||||
// ResetHandler returns an HTTP handler that clears local stats on POST.
|
||||
// With ?push=true it first pushes to the Pushgateway and only resets on success.
|
||||
// The handler performs no authentication; mount it on an internal/protected route.
|
||||
func (p *PrometheusProvider) ResetHandler() http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
w.Header().Set("Allow", http.MethodPost)
|
||||
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
if r.URL.Query().Get("push") == "true" {
|
||||
if err := p.PushAndReset(); err != nil {
|
||||
http.Error(w, err.Error(), http.StatusBadGateway)
|
||||
return
|
||||
}
|
||||
} else {
|
||||
p.Reset()
|
||||
}
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
})
|
||||
}
|
||||
|
||||
// StopAutoPush stops the automatic push goroutine
|
||||
// This should be called when shutting down the application
|
||||
func (p *PrometheusProvider) StopAutoPush() {
|
||||
if p.pushStop != nil {
|
||||
close(p.pushStop)
|
||||
p.pushStop = nil
|
||||
}
|
||||
if p.endpoint != nil {
|
||||
p.endpoint.stop()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -497,13 +497,13 @@ func TestIsColumnWritableWithEmbedded(t *testing.T) {
|
||||
|
||||
// Test models with relations for GetSQLModelColumns
|
||||
type User struct {
|
||||
ID int `bun:"id,pk" json:"id"`
|
||||
Name string `bun:"name" json:"name"`
|
||||
Email string `bun:"email" json:"email"`
|
||||
ProfileData string `json:"profile_data"` // No bun/gorm tag
|
||||
Posts []Post `bun:"rel:has-many,join:id=user_id" json:"posts"`
|
||||
Profile *Profile `bun:"rel:has-one,join:id=user_id" json:"profile"`
|
||||
RowNumber int64 `bun:",scanonly" json:"_rownumber"`
|
||||
ID int `bun:"id,pk" json:"id"`
|
||||
Name string `bun:"name" json:"name"`
|
||||
Email string `bun:"email" json:"email"`
|
||||
ProfileData string `json:"profile_data"` // No bun/gorm tag
|
||||
Posts []Post `bun:"rel:has-many,join:id=user_id" json:"posts"`
|
||||
Profile *Profile `bun:"rel:has-one,join:id=user_id" json:"profile"`
|
||||
RowNumber int64 `bun:",scanonly" json:"_rownumber"`
|
||||
}
|
||||
|
||||
type Post struct {
|
||||
@@ -528,8 +528,8 @@ type Tag struct {
|
||||
|
||||
// Model with scan-only embedded struct
|
||||
type EntityWithScanOnlyEmbedded struct {
|
||||
ID int `bun:"id,pk" json:"id"`
|
||||
Name string `bun:"name" json:"name"`
|
||||
ID int `bun:"id,pk" json:"id"`
|
||||
Name string `bun:"name" json:"name"`
|
||||
AdhocBuffer `bun:",scanonly"` // Entire embedded struct is scan-only
|
||||
}
|
||||
|
||||
@@ -1086,17 +1086,17 @@ func TestGetColumnTypeFromModel_SqlNullWrapper(t *testing.T) {
|
||||
|
||||
// Models for relation testing
|
||||
type Author struct {
|
||||
ID int `bun:"id,pk" json:"id"`
|
||||
Name string `bun:"name" json:"name"`
|
||||
Books []Book `bun:"rel:has-many,join:id=author_id" json:"books"`
|
||||
ID int `bun:"id,pk" json:"id"`
|
||||
Name string `bun:"name" json:"name"`
|
||||
Books []Book `bun:"rel:has-many,join:id=author_id" json:"books"`
|
||||
}
|
||||
|
||||
type Book struct {
|
||||
ID int `bun:"id,pk" json:"id"`
|
||||
Title string `bun:"title" json:"title"`
|
||||
AuthorID int `bun:"author_id" json:"author_id"`
|
||||
Author *Author `bun:"rel:belongs-to,join:author_id=id" json:"author"`
|
||||
Publisher *Publisher `bun:"rel:has-one,join:id=book_id" json:"publisher"`
|
||||
ID int `bun:"id,pk" json:"id"`
|
||||
Title string `bun:"title" json:"title"`
|
||||
AuthorID int `bun:"author_id" json:"author_id"`
|
||||
Author *Author `bun:"rel:belongs-to,join:author_id=id" json:"author"`
|
||||
Publisher *Publisher `bun:"rel:has-one,join:id=book_id" json:"publisher"`
|
||||
}
|
||||
|
||||
type Publisher struct {
|
||||
@@ -1106,9 +1106,9 @@ type Publisher struct {
|
||||
}
|
||||
|
||||
type Student struct {
|
||||
ID int `gorm:"column:id;primaryKey" json:"id"`
|
||||
Name string `gorm:"column:name" json:"name"`
|
||||
Courses []Course `gorm:"many2many:student_courses" json:"courses"`
|
||||
ID int `gorm:"column:id;primaryKey" json:"id"`
|
||||
Name string `gorm:"column:name" json:"name"`
|
||||
Courses []Course `gorm:"many2many:student_courses" json:"courses"`
|
||||
}
|
||||
|
||||
type Course struct {
|
||||
@@ -1119,11 +1119,11 @@ type Course struct {
|
||||
|
||||
// Recursive relation model
|
||||
type Category struct {
|
||||
ID int `bun:"id,pk" json:"id"`
|
||||
Name string `bun:"name" json:"name"`
|
||||
ParentID *int `bun:"parent_id" json:"parent_id"`
|
||||
Parent *Category `bun:"rel:belongs-to,join:parent_id=id" json:"parent"`
|
||||
Children []Category `bun:"rel:has-many,join:id=parent_id" json:"children"`
|
||||
ID int `bun:"id,pk" json:"id"`
|
||||
Name string `bun:"name" json:"name"`
|
||||
ParentID *int `bun:"parent_id" json:"parent_id"`
|
||||
Parent *Category `bun:"rel:belongs-to,join:parent_id=id" json:"parent"`
|
||||
Children []Category `bun:"rel:has-many,join:id=parent_id" json:"children"`
|
||||
}
|
||||
|
||||
func TestGetRelationType(t *testing.T) {
|
||||
@@ -1299,7 +1299,7 @@ func TestGetPrimaryKeyValue_EdgeCases(t *testing.T) {
|
||||
expected: nil,
|
||||
},
|
||||
{
|
||||
name: "model without primary key tags - fallback to ID field",
|
||||
name: "model without primary key tags - fallback to ID field",
|
||||
model: struct {
|
||||
ID int
|
||||
Name string
|
||||
@@ -1307,7 +1307,7 @@ func TestGetPrimaryKeyValue_EdgeCases(t *testing.T) {
|
||||
expected: 99,
|
||||
},
|
||||
{
|
||||
name: "model without ID field",
|
||||
name: "model without ID field",
|
||||
model: struct {
|
||||
Name string
|
||||
}{Name: "Test"},
|
||||
@@ -1508,10 +1508,10 @@ func TestGetSQLModelColumns_EdgeCases(t *testing.T) {
|
||||
|
||||
// Test models with table:, rel:, join: tags for ExtractColumnFromBunTag
|
||||
type BunSpecialTagsModel struct {
|
||||
Table string `bun:"table:users"`
|
||||
Relation []Post `bun:"rel:has-many"`
|
||||
Join string `bun:"join:id=user_id"`
|
||||
NormalCol string `bun:"normal_col"`
|
||||
Table string `bun:"table:users"`
|
||||
Relation []Post `bun:"rel:has-many"`
|
||||
Join string `bun:"join:id=user_id"`
|
||||
NormalCol string `bun:"normal_col"`
|
||||
}
|
||||
|
||||
func TestExtractColumnFromBunTag_SpecialTags(t *testing.T) {
|
||||
@@ -1592,8 +1592,8 @@ func TestGetRelationType_GORMFallback(t *testing.T) {
|
||||
func TestGetRelationType_AdditionalCases(t *testing.T) {
|
||||
// Test model with GORM has-one (pointer without foreignKey or with references)
|
||||
type Address struct {
|
||||
ID int `gorm:"column:id;primaryKey"`
|
||||
UserID int `gorm:"column:user_id"`
|
||||
ID int `gorm:"column:id;primaryKey"`
|
||||
UserID int `gorm:"column:user_id"`
|
||||
}
|
||||
|
||||
type UserWithAddress struct {
|
||||
@@ -1609,7 +1609,7 @@ func TestGetRelationType_AdditionalCases(t *testing.T) {
|
||||
|
||||
type Employee struct {
|
||||
ID int
|
||||
Company Company // Single struct (not pointer, not slice) - belongs-to
|
||||
Company Company // Single struct (not pointer, not slice) - belongs-to
|
||||
Coworkers []Employee // Slice without bun/gorm tags - has-many
|
||||
}
|
||||
|
||||
@@ -1963,3 +1963,31 @@ func TestNonWritableColumns(t *testing.T) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestNonWritableColumns_EmbeddedScanOnlyBuffer(t *testing.T) {
|
||||
type buffer struct {
|
||||
CQL1 string `json:"cql1,omitempty" gorm:"->" bun:",scanonly"`
|
||||
RowNumber int64 `json:"_rownumber,omitempty" gorm:"-" bun:",scanonly"`
|
||||
}
|
||||
type m struct {
|
||||
ID int `json:"id" bun:"id,pk"`
|
||||
Note string `json:"note" bun:"note,type:citext,"`
|
||||
buffer `json:",omitempty" bun:",scanonly"`
|
||||
}
|
||||
got := NonWritableColumns(&m{})
|
||||
has := map[string]bool{}
|
||||
for _, c := range got {
|
||||
has[c] = true
|
||||
}
|
||||
if !has["cql1"] {
|
||||
t.Errorf("cql1 should be non-writable, got %v", got)
|
||||
}
|
||||
if has["id"] || has["note"] {
|
||||
t.Errorf("writable columns reported as non-writable: %v", got)
|
||||
}
|
||||
vals := map[string]interface{}{"id": 1, "note": "x", "cql1": "y"}
|
||||
RemoveNonWritableColumns(&m{}, vals)
|
||||
if _, ok := vals["cql1"]; ok || len(vals) != 2 {
|
||||
t.Errorf("unexpected values: %v", vals)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,61 @@
|
||||
package resolvespec
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/uptrace/bunrouter"
|
||||
)
|
||||
|
||||
type wrapCtxKey struct{}
|
||||
|
||||
// The auth wrapper must hand the handler the middleware-enriched request
|
||||
// without dropping the bunrouter route params.
|
||||
func TestWrapBunRouterHandler_PreservesRouteParams(t *testing.T) {
|
||||
var gotSchema, gotEntity, gotID string
|
||||
var gotCtxVal any
|
||||
|
||||
handler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
gotSchema = req.Param("schema")
|
||||
gotEntity = req.Param("entity")
|
||||
gotID = req.Param("id")
|
||||
gotCtxVal = req.Context().Value(wrapCtxKey{})
|
||||
return nil
|
||||
}
|
||||
auth := func(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
next.ServeHTTP(w, r.WithContext(context.WithValue(r.Context(), wrapCtxKey{}, "enriched")))
|
||||
})
|
||||
}
|
||||
|
||||
router := bunrouter.New()
|
||||
router.GET("/:schema/:entity/:id", wrapBunRouterHandler(handler, auth))
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
router.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/public/users/42", nil))
|
||||
|
||||
if gotSchema != "public" || gotEntity != "users" || gotID != "42" {
|
||||
t.Errorf("route params lost: schema=%q entity=%q id=%q", gotSchema, gotEntity, gotID)
|
||||
}
|
||||
if gotCtxVal != "enriched" {
|
||||
t.Errorf("handler did not see middleware-enriched context, got %v", gotCtxVal)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWrapBunRouterHandler_NilAuthPassesThrough(t *testing.T) {
|
||||
var gotID string
|
||||
handler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
gotID = req.Param("id")
|
||||
return nil
|
||||
}
|
||||
|
||||
router := bunrouter.New()
|
||||
router.GET("/:schema/:entity/:id", wrapBunRouterHandler(handler, nil))
|
||||
router.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/public/users/7", nil))
|
||||
|
||||
if gotID != "7" {
|
||||
t.Errorf("id = %q, want 7", gotID)
|
||||
}
|
||||
}
|
||||
@@ -417,6 +417,14 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
||||
|
||||
if id == "" {
|
||||
options.SingleRecordAsObject = false
|
||||
} else {
|
||||
// The primary key is already filtered, so never return more than one
|
||||
// record regardless of limit/offset/cursor headers or joins.
|
||||
one := 1
|
||||
options.Limit = &one
|
||||
options.Offset = nil
|
||||
options.CursorForward = ""
|
||||
options.CursorBackward = ""
|
||||
}
|
||||
|
||||
// Validate and unwrap model type to get base struct
|
||||
@@ -726,7 +734,7 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
||||
sanitizedOr = common.EnsureOuterParentheses(sanitizedOr)
|
||||
}
|
||||
|
||||
if grouper, ok := query.(common.WhereGrouper); ok && sanitizedOr != "" && common.Hardening().SQLStrict {
|
||||
if grouper, ok := query.(common.WhereGrouper); ok && sanitizedOr != "" {
|
||||
query = grouper.WhereGroup(func(q common.SelectQuery) common.SelectQuery {
|
||||
return applyUserConds(q).WhereOr(sanitizedOr)
|
||||
})
|
||||
|
||||
@@ -0,0 +1,157 @@
|
||||
package restheadspec
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strconv"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/DATA-DOG/go-sqlmock"
|
||||
"github.com/uptrace/bun"
|
||||
"github.com/uptrace/bun/dialect/pgdialect"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||
)
|
||||
|
||||
// readCapturingSQL runs handleRead and returns every SELECT it issued.
|
||||
func readCapturingSQL(t *testing.T, id string, options ExtendedRequestOptions) []string {
|
||||
queries, _ := readCapturingSQLAndBody(t, id, options)
|
||||
return queries
|
||||
}
|
||||
|
||||
// readCapturingSQLAndBody is readCapturingSQL that also returns the response body.
|
||||
// The mocked row carries the requested id so the body can be checked against it.
|
||||
func readCapturingSQLAndBody(t *testing.T, id string, options ExtendedRequestOptions) ([]string, string) {
|
||||
t.Helper()
|
||||
resetTotalCache(t)
|
||||
var queries []string
|
||||
matcher := sqlmock.QueryMatcherFunc(func(_, actual string) error {
|
||||
queries = append(queries, actual)
|
||||
return nil
|
||||
})
|
||||
sqlDB, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(matcher))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
sqlDB.SetMaxOpenConns(1)
|
||||
t.Cleanup(func() { _ = sqlDB.Close() })
|
||||
h := NewHandler(database.NewBunAdapter(bun.NewDB(sqlDB, pgdialect.New())), modelregistry.NewModelRegistry())
|
||||
|
||||
rowID, err := strconv.Atoi(id)
|
||||
if err != nil {
|
||||
rowID = 7
|
||||
}
|
||||
mock.ExpectBegin()
|
||||
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"count"}).AddRow(1))
|
||||
mock.ExpectQuery(`SELECT`).WillReturnRows(sqlmock.NewRows([]string{"id", "name"}).AddRow(rowID, "a"))
|
||||
mock.ExpectCommit()
|
||||
mock.ExpectBegin()
|
||||
mock.ExpectCommit()
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
w, _ := common.WrapHTTPRequest(rec, httptest.NewRequest(http.MethodGet, "/", nil))
|
||||
h.handleRead(itemCtx(t), w, id, options)
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("status %d body %s", rec.Code, rec.Body)
|
||||
}
|
||||
return queries, rec.Body.String()
|
||||
}
|
||||
|
||||
func TestReadByIDIgnoresLimitOffsetAndCursor(t *testing.T) {
|
||||
limit, offset := 50, 10
|
||||
queries := readCapturingSQL(t, "7", ExtendedRequestOptions{
|
||||
RequestOptions: common.RequestOptions{
|
||||
Limit: &limit,
|
||||
Offset: &offset,
|
||||
},
|
||||
})
|
||||
last := queries[len(queries)-1]
|
||||
if !strings.Contains(last, "LIMIT 1") || strings.Contains(last, "OFFSET") {
|
||||
t.Fatalf("read by id must be LIMIT 1 with no OFFSET: %s", last)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadWithoutIDKeepsRequestedLimit(t *testing.T) {
|
||||
limit := 50
|
||||
queries := readCapturingSQL(t, "", ExtendedRequestOptions{
|
||||
RequestOptions: common.RequestOptions{Limit: &limit},
|
||||
})
|
||||
if last := queries[len(queries)-1]; !strings.Contains(last, "LIMIT 50") {
|
||||
t.Fatalf("list read must keep its limit: %s", last)
|
||||
}
|
||||
}
|
||||
|
||||
// topLevelOr reports whether the WHERE clause has an OR outside any parentheses,
|
||||
// i.e. one that would let rows bypass the AND-ed primary key condition.
|
||||
func topLevelOr(sql string) bool {
|
||||
where := sql[strings.Index(sql, "WHERE")+len("WHERE"):]
|
||||
depth, inStr := 0, false
|
||||
for i := 0; i < len(where); i++ {
|
||||
switch c := where[i]; {
|
||||
case c == '\'':
|
||||
inStr = !inStr
|
||||
case inStr:
|
||||
case c == '(':
|
||||
depth++
|
||||
case c == ')':
|
||||
depth--
|
||||
case depth == 0 && strings.HasPrefix(where[i:], " OR "):
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func TestReadByIDCustomSQLOrCannotEscapePrimaryKey(t *testing.T) {
|
||||
queries := readCapturingSQL(t, "7", ExtendedRequestOptions{
|
||||
RequestOptions: common.RequestOptions{
|
||||
Filters: []common.FilterOption{{Column: "name", Operator: "eq", Value: "a"}},
|
||||
},
|
||||
CustomSQLOr: "name = 'x'",
|
||||
})
|
||||
last := queries[len(queries)-1]
|
||||
if !strings.Contains(last, `"id" = '7'`) && !strings.Contains(last, `"id" = 7`) {
|
||||
t.Fatalf("primary key filter missing: %s", last)
|
||||
}
|
||||
if topLevelOr(last) {
|
||||
t.Fatalf("OR escapes the primary key filter: %s", last)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadByIDFiltersAndReturnsRequestedRecord(t *testing.T) {
|
||||
queries, body := readCapturingSQLAndBody(t, "42", ExtendedRequestOptions{})
|
||||
last := queries[len(queries)-1]
|
||||
if !strings.Contains(last, `"items"."id" = '42'`) && !strings.Contains(last, `"items"."id" = 42`) {
|
||||
t.Fatalf("query must filter the primary key to 42: %s", last)
|
||||
}
|
||||
if strings.Contains(last, "= 7") || strings.Contains(last, "= '7'") {
|
||||
t.Fatalf("query filters a different id: %s", last)
|
||||
}
|
||||
// every query that touches rows (count and select) must carry the id filter
|
||||
for _, q := range queries {
|
||||
if strings.Contains(q, "FROM") && !strings.Contains(q, "42") {
|
||||
t.Fatalf("query without the id filter: %s", q)
|
||||
}
|
||||
}
|
||||
var rows []struct {
|
||||
ID int `json:"id"`
|
||||
}
|
||||
data := body
|
||||
if i := strings.Index(body, `"data"`); i >= 0 {
|
||||
data = body[i+len(`"data"`):]
|
||||
}
|
||||
if i := strings.Index(data, "["); i >= 0 {
|
||||
data = data[i:]
|
||||
}
|
||||
dec := json.NewDecoder(strings.NewReader(data))
|
||||
if err := dec.Decode(&rows); err != nil {
|
||||
t.Fatalf("decode %q: %v", body, err)
|
||||
}
|
||||
if len(rows) != 1 || rows[0].ID != 42 {
|
||||
t.Fatalf("response must contain exactly the record with id 42: %s", body)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,61 @@
|
||||
package restheadspec
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/uptrace/bunrouter"
|
||||
)
|
||||
|
||||
type wrapCtxKey struct{}
|
||||
|
||||
// The auth wrapper must hand the handler the middleware-enriched request
|
||||
// without dropping the bunrouter route params.
|
||||
func TestWrapBunRouterHandler_PreservesRouteParams(t *testing.T) {
|
||||
var gotSchema, gotEntity, gotID string
|
||||
var gotCtxVal any
|
||||
|
||||
handler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
gotSchema = req.Param("schema")
|
||||
gotEntity = req.Param("entity")
|
||||
gotID = req.Param("id")
|
||||
gotCtxVal = req.Context().Value(wrapCtxKey{})
|
||||
return nil
|
||||
}
|
||||
auth := func(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
next.ServeHTTP(w, r.WithContext(context.WithValue(r.Context(), wrapCtxKey{}, "enriched")))
|
||||
})
|
||||
}
|
||||
|
||||
router := bunrouter.New()
|
||||
router.GET("/:schema/:entity/:id", wrapBunRouterHandler(handler, auth))
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
router.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/public/users/42", nil))
|
||||
|
||||
if gotSchema != "public" || gotEntity != "users" || gotID != "42" {
|
||||
t.Errorf("route params lost: schema=%q entity=%q id=%q", gotSchema, gotEntity, gotID)
|
||||
}
|
||||
if gotCtxVal != "enriched" {
|
||||
t.Errorf("handler did not see middleware-enriched context, got %v", gotCtxVal)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWrapBunRouterHandler_NilAuthPassesThrough(t *testing.T) {
|
||||
var gotID string
|
||||
handler := func(w http.ResponseWriter, req bunrouter.Request) error {
|
||||
gotID = req.Param("id")
|
||||
return nil
|
||||
}
|
||||
|
||||
router := bunrouter.New()
|
||||
router.GET("/:schema/:entity/:id", wrapBunRouterHandler(handler, nil))
|
||||
router.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest(http.MethodGet, "/public/users/7", nil))
|
||||
|
||||
if gotID != "7" {
|
||||
t.Errorf("id = %q, want 7", gotID)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user