mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-07 13:56:29 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
431b674162 | ||
|
|
234aac9770 | ||
|
|
8cff3bde85 | ||
|
|
3e6224698c | ||
|
|
aec87a81e7 | ||
|
|
9235292586 | ||
|
|
23f10387c5 |
@@ -129,9 +129,7 @@ jobs:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Rust
|
||||
run: |
|
||||
rustup toolchain install stable --profile minimal
|
||||
rustup default stable
|
||||
uses: dtolnay/rust-toolchain@stable
|
||||
|
||||
- name: Test
|
||||
run: cargo test
|
||||
@@ -252,6 +250,12 @@ jobs:
|
||||
run: |
|
||||
sed -i -E "s/^version: .*/version: ${VERSION}/" pubspec.yaml
|
||||
sed -i -E "s#^publish_to: .*#publish_to: ${SERVER_URL}/api/packages/${OWNER}/pub#" pubspec.yaml
|
||||
if ! grep -q "^## ${VERSION}\$" CHANGELOG.md; then
|
||||
{ head -n 1 CHANGELOG.md; printf '\n## %s\n\n- Release %s.\n' "$VERSION" "$VERSION"; tail -n +2 CHANGELOG.md; } > CHANGELOG.tmp
|
||||
mv CHANGELOG.tmp CHANGELOG.md
|
||||
fi
|
||||
# pub warns about a dirty git tree; commit the stamped files locally (never pushed)
|
||||
git -c user.name=ci -c user.email=ci@localhost commit -q -am "ci: stamp dart version ${VERSION}"
|
||||
|
||||
- name: Dry run
|
||||
if: ${{ env.PUBLISH != 'true' }}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
name: resolvespec
|
||||
description: Client for ResolveSpec (JSON body) and FunctionSpec endpoints.
|
||||
version: 0.1.0
|
||||
repository: https://git.warky.dev/wdevs/ResolveSpec
|
||||
publish_to: none
|
||||
|
||||
environment:
|
||||
|
||||
@@ -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,6 +1508,35 @@ 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 columns = bunWritableExcludes(b.query.GetModel(), columns); len(columns) > 0 {
|
||||
b.query = b.query.ExcludeColumn(columns...)
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
func (b *BunInsertQuery) Returning(columns ...string) common.InsertQuery {
|
||||
if len(columns) > 0 {
|
||||
b.query = b.query.Returning(strings.Join(columns, ", "))
|
||||
@@ -1619,6 +1649,13 @@ func (b *BunUpdateQuery) SetMap(values map[string]interface{}) common.UpdateQuer
|
||||
return b
|
||||
}
|
||||
|
||||
func (b *BunUpdateQuery) ExcludeColumn(columns ...string) common.UpdateQuery {
|
||||
if columns = bunWritableExcludes(b.query.GetModel(), columns); len(columns) > 0 {
|
||||
b.query = b.query.ExcludeColumn(columns...)
|
||||
}
|
||||
return b
|
||||
}
|
||||
|
||||
func (b *BunUpdateQuery) Where(query string, args ...interface{}) common.UpdateQuery {
|
||||
b.query = b.query.Where(query, args...)
|
||||
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
|
||||
}
|
||||
@@ -751,6 +751,13 @@ func (g *GormInsertQuery) OnConflict(action string) common.InsertQuery {
|
||||
return g
|
||||
}
|
||||
|
||||
func (g *GormInsertQuery) ExcludeColumn(columns ...string) common.InsertQuery {
|
||||
if len(columns) > 0 {
|
||||
g.db = g.db.Omit(columns...)
|
||||
}
|
||||
return g
|
||||
}
|
||||
|
||||
func (g *GormInsertQuery) Returning(columns ...string) common.InsertQuery {
|
||||
g.returningColumns = columns
|
||||
return g
|
||||
@@ -930,6 +937,13 @@ func (g *GormUpdateQuery) SetMap(values map[string]interface{}) common.UpdateQue
|
||||
return g
|
||||
}
|
||||
|
||||
func (g *GormUpdateQuery) ExcludeColumn(columns ...string) common.UpdateQuery {
|
||||
if len(columns) > 0 {
|
||||
g.db = g.db.Omit(columns...)
|
||||
}
|
||||
return g
|
||||
}
|
||||
|
||||
func (g *GormUpdateQuery) Where(query string, args ...interface{}) common.UpdateQuery {
|
||||
g.db = g.db.Where(query, args...)
|
||||
return g
|
||||
|
||||
@@ -691,6 +691,13 @@ func (p *PgSQLInsertQuery) OnConflict(action string) common.InsertQuery {
|
||||
return p
|
||||
}
|
||||
|
||||
func (p *PgSQLInsertQuery) ExcludeColumn(columns ...string) common.InsertQuery {
|
||||
for _, col := range columns {
|
||||
delete(p.values, col)
|
||||
}
|
||||
return p
|
||||
}
|
||||
|
||||
func (p *PgSQLInsertQuery) Returning(columns ...string) common.InsertQuery {
|
||||
p.returning = columns
|
||||
return p
|
||||
@@ -850,6 +857,13 @@ func (p *PgSQLUpdateQuery) Set(column string, value interface{}) common.UpdateQu
|
||||
return p
|
||||
}
|
||||
|
||||
func (p *PgSQLUpdateQuery) ExcludeColumn(columns ...string) common.UpdateQuery {
|
||||
for _, col := range columns {
|
||||
delete(p.sets, col)
|
||||
}
|
||||
return p
|
||||
}
|
||||
|
||||
func (p *PgSQLUpdateQuery) SetMap(values map[string]interface{}) common.UpdateQuery {
|
||||
pkName := ""
|
||||
if p.model != nil {
|
||||
|
||||
@@ -81,6 +81,8 @@ type InsertQuery interface {
|
||||
Table(table string) InsertQuery
|
||||
Value(column string, value interface{}) InsertQuery
|
||||
OnConflict(action string) InsertQuery
|
||||
// ExcludeColumn omits columns from a Model()-based INSERT (e.g. generated columns).
|
||||
ExcludeColumn(columns ...string) InsertQuery
|
||||
Returning(columns ...string) InsertQuery
|
||||
|
||||
// Execution
|
||||
@@ -94,6 +96,8 @@ type UpdateQuery interface {
|
||||
Table(table string) UpdateQuery
|
||||
Set(column string, value interface{}) UpdateQuery
|
||||
SetMap(values map[string]interface{}) UpdateQuery
|
||||
// ExcludeColumn omits columns from a Model()-based UPDATE (e.g. generated columns).
|
||||
ExcludeColumn(columns ...string) UpdateQuery
|
||||
Where(query string, args ...interface{}) UpdateQuery
|
||||
Returning(columns ...string) UpdateQuery
|
||||
|
||||
|
||||
@@ -116,7 +116,7 @@ func (p *NestedCUDProcessor) ProcessNestedCUD(
|
||||
case "insert", "create", "add":
|
||||
// Only perform insert if we have data to insert
|
||||
if hasData {
|
||||
id, err := p.processInsert(ctx, regularData, tableName)
|
||||
id, err := p.processInsert(ctx, regularData, model, tableName)
|
||||
if err != nil {
|
||||
logger.Error("Insert failed for table=%s, data=%+v, error=%v", tableName, regularData, err)
|
||||
return nil, fmt.Errorf("insert failed: %w", err)
|
||||
@@ -148,7 +148,7 @@ func (p *NestedCUDProcessor) ProcessNestedCUD(
|
||||
return result, nil
|
||||
}
|
||||
if hasData {
|
||||
rows, err := p.processUpdate(ctx, regularData, tableName, data[pkName])
|
||||
rows, err := p.processUpdate(ctx, regularData, model, tableName, data[pkName])
|
||||
if err != nil {
|
||||
logger.Error("Update failed for table=%s, id=%v, data=%+v, error=%v", tableName, data[pkName], regularData, err)
|
||||
return nil, fmt.Errorf("update failed: %w", err)
|
||||
@@ -295,10 +295,12 @@ func (p *NestedCUDProcessor) injectForeignKeys(data map[string]interface{}, mode
|
||||
func (p *NestedCUDProcessor) processInsert(
|
||||
ctx context.Context,
|
||||
data map[string]interface{},
|
||||
model interface{},
|
||||
tableName string,
|
||||
) (interface{}, error) {
|
||||
logger.Debug("Inserting into %s with data: %+v", tableName, data)
|
||||
|
||||
reflection.RemoveNonWritableColumns(model, data)
|
||||
query := p.db.NewInsert().Table(tableName)
|
||||
|
||||
for key, value := range data {
|
||||
@@ -335,6 +337,7 @@ func (p *NestedCUDProcessor) processSelect(ctx context.Context, tableName string
|
||||
func (p *NestedCUDProcessor) processUpdate(
|
||||
ctx context.Context,
|
||||
data map[string]interface{},
|
||||
model interface{},
|
||||
tableName string,
|
||||
id interface{},
|
||||
) (int64, error) {
|
||||
@@ -345,6 +348,7 @@ func (p *NestedCUDProcessor) processUpdate(
|
||||
|
||||
logger.Debug("Updating %s with ID %v, data: %+v", tableName, id, data)
|
||||
|
||||
reflection.RemoveNonWritableColumns(model, data)
|
||||
query := p.db.NewUpdate().Table(tableName).SetMap(data).Where(fmt.Sprintf("%s = ?", QuoteIdent(reflection.GetPrimaryKeyName(tableName))), id)
|
||||
|
||||
result, err := query.Exec(ctx)
|
||||
|
||||
@@ -99,6 +99,7 @@ func (m *mockInsertQuery) Value(column string, value interface{}) InsertQuery {
|
||||
return m
|
||||
}
|
||||
func (m *mockInsertQuery) OnConflict(action string) InsertQuery { return m }
|
||||
func (m *mockInsertQuery) ExcludeColumn(columns ...string) InsertQuery { return m }
|
||||
func (m *mockInsertQuery) Returning(columns ...string) InsertQuery { return m }
|
||||
func (m *mockInsertQuery) Exec(ctx context.Context) (Result, error) {
|
||||
m.db.insertCalls = append(m.db.insertCalls, m.values)
|
||||
@@ -131,6 +132,7 @@ func (m *mockUpdateQuery) SetMap(values map[string]interface{}) UpdateQuery {
|
||||
return m
|
||||
}
|
||||
func (m *mockUpdateQuery) Where(condition string, args ...interface{}) UpdateQuery { return m }
|
||||
func (m *mockUpdateQuery) ExcludeColumn(columns ...string) UpdateQuery { return m }
|
||||
func (m *mockUpdateQuery) Returning(columns ...string) UpdateQuery { return m }
|
||||
func (m *mockUpdateQuery) Exec(ctx context.Context) (Result, error) {
|
||||
// Record the update call
|
||||
|
||||
+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()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,32 @@
|
||||
// Package modelregistry is the shared catalogue of the Go model structs that the
|
||||
// ResolveSpec front ends (resolvespec, restheadspec, websocketspec, mqttspec,
|
||||
// resolvemcp, ...) expose as database entities.
|
||||
//
|
||||
// A registry maps a model name ("schema.entity") to a struct type and holds:
|
||||
// - ModelRules: which operations (read/create/update/delete, public or not)
|
||||
// are allowed, and whether security checks are disabled.
|
||||
// - ModelInfo: optional documentation (description, purpose, tags, per-column
|
||||
// descriptions) meant for humans and AI agents. It never affects queries or
|
||||
// permissions.
|
||||
//
|
||||
// Register models on a registry created with NewModelRegistry, or through the
|
||||
// package-level functions that use the default registry:
|
||||
//
|
||||
// reg := modelregistry.NewModelRegistry()
|
||||
// _ = reg.RegisterModelWithRules("public.users", User{}, modelregistry.DefaultModelRules())
|
||||
// reg.SetModelInfo("public.users", modelregistry.ModelInfo{
|
||||
// Description: "Application accounts",
|
||||
// Purpose: "Look up who a person is; never store credentials here",
|
||||
// Columns: map[string]string{"email": "Login address, unique"},
|
||||
// })
|
||||
//
|
||||
// Descriptions come from, in priority order:
|
||||
// 1. ModelInfo set with SetModelInfo or loaded from an external JSON map with
|
||||
// LoadModelInfoFile (the map can be maintained outside the Go code).
|
||||
// 2. The model's Describer (ModelDescription() string) for the description.
|
||||
// 3. Struct tags, read per column by FieldComment: comment, note, desc or
|
||||
// description tags, then "comment:" inside the gorm or bun tag.
|
||||
//
|
||||
// Models must be non-pointer structs; pointers, slices and arrays of structs are
|
||||
// unwrapped on registration. All registry methods are safe for concurrent use.
|
||||
package modelregistry
|
||||
@@ -0,0 +1,172 @@
|
||||
package modelregistry
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// ModelInfo is human/AI-facing documentation for a registered model: what it
|
||||
// is for and what its columns mean. It is optional and has no effect on
|
||||
// permissions or queries.
|
||||
type ModelInfo struct {
|
||||
// Description says what the model/table holds.
|
||||
Description string `json:"description,omitempty"`
|
||||
// Purpose says why it exists / when an agent should use it.
|
||||
Purpose string `json:"purpose,omitempty"`
|
||||
// Tags are free-form labels (e.g. "billing", "pii").
|
||||
Tags []string `json:"tags,omitempty"`
|
||||
// Columns maps a JSON column name to its description.
|
||||
Columns map[string]string `json:"columns,omitempty"`
|
||||
}
|
||||
|
||||
// IsZero reports whether the info carries no documentation.
|
||||
func (i ModelInfo) IsZero() bool {
|
||||
return i.Description == "" && i.Purpose == "" && len(i.Tags) == 0 && len(i.Columns) == 0
|
||||
}
|
||||
|
||||
// Describer can be implemented by a model to document itself. It is the
|
||||
// fallback used when no ModelInfo description was registered or loaded.
|
||||
// (The method is not called Description so models may keep a Description field.)
|
||||
type Describer interface {
|
||||
ModelDescription() string
|
||||
}
|
||||
|
||||
// commentTagKeys are the standalone struct tags read as a column description,
|
||||
// in priority order.
|
||||
var commentTagKeys = []string{"comment", "note", "desc", "description"}
|
||||
|
||||
// FieldComment returns the description of a struct field from its tags. Order:
|
||||
// standalone comment/note/desc/description tags, then a "comment:" entry inside
|
||||
// the gorm tag (semicolon separated), then inside the bun tag (comma separated).
|
||||
// It returns "" when the field carries none.
|
||||
func FieldComment(sf reflect.StructField) string {
|
||||
for _, key := range commentTagKeys {
|
||||
if v := strings.TrimSpace(sf.Tag.Get(key)); v != "" {
|
||||
return v
|
||||
}
|
||||
}
|
||||
if v := tagOption(sf.Tag.Get("gorm"), ';', "comment:"); v != "" {
|
||||
return v
|
||||
}
|
||||
return tagOption(sf.Tag.Get("bun"), ',', "comment:")
|
||||
}
|
||||
|
||||
func tagOption(tag string, sep byte, key string) string {
|
||||
for _, part := range strings.Split(tag, string(sep)) {
|
||||
part = strings.TrimSpace(part)
|
||||
if len(part) >= len(key) && strings.EqualFold(part[:len(key)], key) {
|
||||
return strings.Trim(strings.TrimSpace(part[len(key):]), `'"`)
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// SetModelInfo stores documentation for a model name ("schema.entity"). The
|
||||
// model does not have to be registered yet, so descriptions can be loaded
|
||||
// before or after registration. Any previous info for the name is replaced.
|
||||
func (r *DefaultModelRegistry) SetModelInfo(name string, info ModelInfo) {
|
||||
r.mutex.Lock()
|
||||
defer r.mutex.Unlock()
|
||||
if r.info == nil {
|
||||
r.info = make(map[string]ModelInfo)
|
||||
}
|
||||
r.info[name] = cloneInfo(info)
|
||||
}
|
||||
|
||||
// GetModelInfo returns the documentation stored with SetModelInfo (or loaded
|
||||
// from a descriptions file), without any fallback.
|
||||
func (r *DefaultModelRegistry) GetModelInfo(name string) (ModelInfo, bool) {
|
||||
r.mutex.RLock()
|
||||
defer r.mutex.RUnlock()
|
||||
info, ok := r.info[name]
|
||||
return cloneInfo(info), ok
|
||||
}
|
||||
|
||||
// RegisterModelWithInfo registers a model together with its documentation.
|
||||
func (r *DefaultModelRegistry) RegisterModelWithInfo(name string, model interface{}, info ModelInfo) error {
|
||||
if err := r.RegisterModel(name, model); err != nil {
|
||||
return err
|
||||
}
|
||||
r.SetModelInfo(name, info)
|
||||
return nil
|
||||
}
|
||||
|
||||
// ResolveModelInfo returns the effective documentation for a registered model.
|
||||
// Stored/loaded info wins; an empty Description falls back to the model's
|
||||
// Describer. Column descriptions are not resolved here: use the stored map and
|
||||
// fall back to FieldComment per field.
|
||||
func (r *DefaultModelRegistry) ResolveModelInfo(name string) ModelInfo {
|
||||
info, _ := r.GetModelInfo(name)
|
||||
if info.Description == "" {
|
||||
if model, err := r.GetModel(name); err == nil {
|
||||
info.Description = describerText(model)
|
||||
}
|
||||
}
|
||||
return info
|
||||
}
|
||||
|
||||
func describerText(model interface{}) (text string) {
|
||||
defer func() {
|
||||
if recover() != nil {
|
||||
text = ""
|
||||
}
|
||||
}()
|
||||
if d, ok := model.(Describer); ok {
|
||||
return strings.TrimSpace(d.ModelDescription())
|
||||
}
|
||||
if t := reflect.TypeOf(model); t != nil && t.Kind() != reflect.Pointer {
|
||||
if d, ok := reflect.New(t).Interface().(Describer); ok {
|
||||
return strings.TrimSpace(d.ModelDescription())
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// LoadModelInfo reads an external descriptions map from r and applies it. The
|
||||
// JSON is an object keyed by model name:
|
||||
//
|
||||
// {"public.users": {"description": "...", "purpose": "...", "tags": ["x"],
|
||||
// "columns": {"email": "Login address"}}}
|
||||
//
|
||||
// Entries replace any existing info for the same name and take precedence over
|
||||
// the model's Describer and its struct-tag comments. It returns the number of
|
||||
// models loaded.
|
||||
func (r *DefaultModelRegistry) LoadModelInfo(src io.Reader) (int, error) {
|
||||
var m map[string]ModelInfo
|
||||
dec := json.NewDecoder(src)
|
||||
dec.DisallowUnknownFields()
|
||||
if err := dec.Decode(&m); err != nil {
|
||||
return 0, fmt.Errorf("modelregistry: decode model info: %w", err)
|
||||
}
|
||||
for name, info := range m {
|
||||
r.SetModelInfo(name, info)
|
||||
}
|
||||
return len(m), nil
|
||||
}
|
||||
|
||||
// LoadModelInfoFile is LoadModelInfo reading from a JSON file.
|
||||
func (r *DefaultModelRegistry) LoadModelInfoFile(path string) (int, error) {
|
||||
f, err := os.Open(filepath.Clean(path)) //nolint:gosec // operator-supplied descriptions file
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("modelregistry: %w", err)
|
||||
}
|
||||
defer f.Close()
|
||||
return r.LoadModelInfo(f)
|
||||
}
|
||||
|
||||
func cloneInfo(in ModelInfo) ModelInfo {
|
||||
out := in
|
||||
out.Tags = append([]string(nil), in.Tags...)
|
||||
if in.Columns != nil {
|
||||
out.Columns = make(map[string]string, len(in.Columns))
|
||||
for k, v := range in.Columns {
|
||||
out.Columns[k] = v
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
@@ -0,0 +1,68 @@
|
||||
package modelregistry
|
||||
|
||||
import (
|
||||
"reflect"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
type infoModel struct {
|
||||
ID int `json:"id" gorm:"primaryKey;comment:Row id"`
|
||||
Email string `json:"email" bun:"email,comment:Login address"`
|
||||
Name string `json:"name" note:"Display name"`
|
||||
Plain string `json:"plain"`
|
||||
}
|
||||
|
||||
func (infoModel) ModelDescription() string { return " From the model " }
|
||||
|
||||
func TestFieldComment(t *testing.T) {
|
||||
typ := reflect.TypeOf(infoModel{})
|
||||
want := map[string]string{"ID": "Row id", "Email": "Login address", "Name": "Display name", "Plain": ""}
|
||||
for field, exp := range want {
|
||||
sf, _ := typ.FieldByName(field)
|
||||
if got := FieldComment(sf); got != exp {
|
||||
t.Errorf("%s = %q, want %q", field, got, exp)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestModelInfoPrecedence(t *testing.T) {
|
||||
r := NewModelRegistry()
|
||||
if err := r.RegisterModel("public.items", infoModel{}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got := r.ResolveModelInfo("public.items").Description; got != "From the model" {
|
||||
t.Errorf("describer fallback = %q", got)
|
||||
}
|
||||
|
||||
n, err := r.LoadModelInfo(strings.NewReader(
|
||||
`{"public.items":{"description":"From file","tags":["a"],"columns":{"email":"Mail"}},"public.later":{"purpose":"p"}}`))
|
||||
if err != nil || n != 2 {
|
||||
t.Fatalf("load n=%d err=%v", n, err)
|
||||
}
|
||||
info := r.ResolveModelInfo("public.items")
|
||||
if info.Description != "From file" || info.Columns["email"] != "Mail" || len(info.Tags) != 1 {
|
||||
t.Errorf("file info = %+v", info)
|
||||
}
|
||||
if _, ok := r.GetModelInfo("public.later"); !ok {
|
||||
t.Error("info for a not-yet-registered model must be kept")
|
||||
}
|
||||
|
||||
// returned info is a copy
|
||||
info.Columns["email"] = "changed"
|
||||
if got, _ := r.GetModelInfo("public.items"); got.Columns["email"] != "Mail" {
|
||||
t.Error("GetModelInfo leaked internal map")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoadModelInfoRejectsBadInput(t *testing.T) {
|
||||
r := NewModelRegistry()
|
||||
for _, in := range []string{`not json`, `{"a":{"descripton":"typo"}}`} {
|
||||
if _, err := r.LoadModelInfo(strings.NewReader(in)); err == nil {
|
||||
t.Errorf("expected error for %q", in)
|
||||
}
|
||||
}
|
||||
if _, err := r.LoadModelInfoFile("/nonexistent/x.json"); err == nil {
|
||||
t.Error("expected error for missing file")
|
||||
}
|
||||
}
|
||||
@@ -41,6 +41,7 @@ func DefaultModelRules() ModelRules {
|
||||
type DefaultModelRegistry struct {
|
||||
models map[string]interface{}
|
||||
rules map[string]ModelRules
|
||||
info map[string]ModelInfo
|
||||
mutex sync.RWMutex
|
||||
}
|
||||
|
||||
|
||||
@@ -895,6 +895,9 @@ func (h *Handler) create(hookCtx *HookContext) (interface{}, error) {
|
||||
|
||||
// Insert record
|
||||
query := hookCtx.Tx.NewInsert().Model(hookCtx.ModelPtr).Table(hookCtx.TableName)
|
||||
if generated := reflection.NonWritableColumns(hookCtx.Model); len(generated) > 0 {
|
||||
query = query.ExcludeColumn(generated...)
|
||||
}
|
||||
if _, err := query.Exec(hookCtx.Context); err != nil {
|
||||
return nil, fmt.Errorf("failed to create record: %w", err)
|
||||
}
|
||||
@@ -924,6 +927,8 @@ func (h *Handler) update(hookCtx *HookContext) error {
|
||||
// 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)
|
||||
|
||||
reflection.RemoveNonWritableColumns(hookCtx.Model, values)
|
||||
|
||||
if len(values) > 0 {
|
||||
query := hookCtx.Tx.NewUpdate().Table(hookCtx.TableName).SetMap(values).
|
||||
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), hookCtx.ID)
|
||||
|
||||
@@ -656,7 +656,7 @@ func isColumnWritableInType(typ reflect.Type, columnName string) (found bool, wr
|
||||
// Check bun tag for scanonly
|
||||
bunTag := field.Tag.Get("bun")
|
||||
if bunTag != "" {
|
||||
if isBunFieldScanOnly(bunTag) {
|
||||
if isBunFieldScanOnly(bunTag) || isBunFieldGenerated(bunTag) {
|
||||
return true, false
|
||||
}
|
||||
}
|
||||
@@ -689,6 +689,70 @@ func isBunFieldScanOnly(tag string) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// isBunFieldGenerated checks if a bun tag marks the column as database-generated
|
||||
// (GENERATED ALWAYS AS ... STORED), which can be read but never written.
|
||||
// Example: "email_normalized,generated" -> true
|
||||
func isBunFieldGenerated(tag string) bool {
|
||||
for _, part := range strings.Split(tag, ",") {
|
||||
if strings.TrimSpace(part) == "generated" {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// RemoveNonWritableColumns deletes from values every key that maps to a
|
||||
// non-writable model column (bun scanonly/generated, gorm read-only). Used
|
||||
// before writing a read-merged record back with UPDATE ... SET.
|
||||
func RemoveNonWritableColumns(model any, values map[string]interface{}) {
|
||||
for key := range values {
|
||||
if !IsColumnWritable(model, key) {
|
||||
delete(values, key)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// NonWritableColumns returns the column names of the model that cannot be
|
||||
// written (bun scanonly/generated, gorm read-only), including embedded structs.
|
||||
func NonWritableColumns(model any) []string {
|
||||
t := reflect.TypeOf(model)
|
||||
for t != nil && (t.Kind() == reflect.Pointer || t.Kind() == reflect.Slice || t.Kind() == reflect.Array) {
|
||||
t = t.Elem()
|
||||
}
|
||||
if t == nil || t.Kind() != reflect.Struct {
|
||||
return nil
|
||||
}
|
||||
var cols []string
|
||||
collectNonWritable(t, &cols)
|
||||
return cols
|
||||
}
|
||||
|
||||
func collectNonWritable(typ reflect.Type, cols *[]string) {
|
||||
for i := 0; i < typ.NumField(); i++ {
|
||||
field := typ.Field(i)
|
||||
if field.Anonymous {
|
||||
ft := field.Type
|
||||
if ft.Kind() == reflect.Pointer {
|
||||
ft = ft.Elem()
|
||||
}
|
||||
if ft.Kind() == reflect.Struct {
|
||||
collectNonWritable(ft, cols)
|
||||
continue
|
||||
}
|
||||
}
|
||||
bunTag, gormTag := field.Tag.Get("bun"), field.Tag.Get("gorm")
|
||||
if bunTag == "-" || gormTag == "-" {
|
||||
continue
|
||||
}
|
||||
if (bunTag != "" && (isBunFieldScanOnly(bunTag) || isBunFieldGenerated(bunTag))) ||
|
||||
(gormTag != "" && isGormFieldReadOnly(gormTag)) {
|
||||
if name := getColumnNameFromField(field); name != "" {
|
||||
*cols = append(*cols, name)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// isGormFieldReadOnly checks if a gorm tag indicates the field is read-only
|
||||
// Examples:
|
||||
// - "<-:false" -> true (no writes allowed)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -1920,3 +1920,74 @@ func TestMapToStruct_Errors(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRemoveNonWritableColumns_Generated(t *testing.T) {
|
||||
type m struct {
|
||||
ID int `bun:"id,pk"`
|
||||
Email string `bun:"email"`
|
||||
Norm string `bun:"email_normalized,generated"`
|
||||
Scan string `bun:"scan_col,scanonly"`
|
||||
}
|
||||
vals := map[string]interface{}{"id": 1, "email": "A", "email_normalized": "a", "scan_col": "x", "dynamic": 1}
|
||||
RemoveNonWritableColumns(&m{}, vals)
|
||||
if _, ok := vals["email_normalized"]; ok {
|
||||
t.Error("generated column not removed")
|
||||
}
|
||||
if _, ok := vals["scan_col"]; ok {
|
||||
t.Error("scanonly column not removed")
|
||||
}
|
||||
if len(vals) != 3 {
|
||||
t.Errorf("unexpected keys: %v", vals)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNonWritableColumns(t *testing.T) {
|
||||
type base struct {
|
||||
Created string `bun:"created_at,scanonly"`
|
||||
}
|
||||
type m struct {
|
||||
base
|
||||
ID int `bun:"id,pk"`
|
||||
Email string `bun:"email"`
|
||||
Norm string `bun:"email_normalized,generated"`
|
||||
Ro string `gorm:"column:ro;->"`
|
||||
}
|
||||
got := NonWritableColumns(&m{})
|
||||
want := map[string]bool{"created_at": true, "email_normalized": true, "ro": true}
|
||||
if len(got) != len(want) {
|
||||
t.Fatalf("got %v", got)
|
||||
}
|
||||
for _, c := range got {
|
||||
if !want[c] {
|
||||
t.Errorf("unexpected %s", c)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -392,6 +392,56 @@ handler.SetModelRules("public", "users", modelregistry.ModelRules{
|
||||
|
||||
---
|
||||
|
||||
## Describing the API for agents
|
||||
|
||||
Give agents context about what each table is for:
|
||||
|
||||
```go
|
||||
// 1. Explicitly, in code
|
||||
handler.SetModelDescription("public", "users", modelregistry.ModelInfo{
|
||||
Description: "Application accounts",
|
||||
Purpose: "Look up who a person is",
|
||||
Tags: []string{"identity"},
|
||||
Columns: map[string]string{"email": "Login address, unique"},
|
||||
})
|
||||
|
||||
// 2. From an external JSON map (keyed by "schema.entity"); entries here win
|
||||
n, err := handler.LoadModelDescriptions("docs/model-descriptions.json")
|
||||
```
|
||||
|
||||
Example `docs/model-descriptions.json` (every key is optional; unknown keys are rejected):
|
||||
|
||||
```json
|
||||
{
|
||||
"public.users": {
|
||||
"description": "Application accounts, one row per person who can sign in.",
|
||||
"purpose": "Look up who someone is. Use public.orders for what they bought.",
|
||||
"tags": ["identity", "pii"],
|
||||
"columns": {
|
||||
"id": "Internal account id",
|
||||
"email": "Login address, unique and lower-cased",
|
||||
"created_at": "When the account was created (UTC)"
|
||||
}
|
||||
},
|
||||
"public.orders": {
|
||||
"description": "Customer orders.",
|
||||
"columns": {
|
||||
"status": "One of: pending, paid, shipped, cancelled"
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
Keys are `schema.entity` names as registered with `RegisterModel`. Column keys are the JSON column names shown by `describe_table`. An entry replaces any info set earlier for that table, and columns it leaves out still fall back to field tags.
|
||||
|
||||
Fallbacks when nothing is set for a table or column, in order: the model's `ModelDescription() string` method (table), then field tags (column): `comment`, `note`, `desc` or `description` tags, then `comment:` inside the `gorm` or `bun` tag.
|
||||
|
||||
The text appears in `list_tables` and `describe_table`. The server also sends a short usage guide as MCP `instructions` on connect.
|
||||
|
||||
### Catalogue file
|
||||
|
||||
`handler.ExportCatalog(path)` writes the usage guide, tools, limits and every table (columns, types, keys, relations, allowed operations, descriptions) to disk, JSON for a `.json` path and Markdown otherwise. The file is replaced atomically. It lists every table with at least one allowed operation, regardless of caller, so keep it out of public directories. Call it after registering models (for example at startup, or from a `go generate` step).
|
||||
|
||||
## MCP Tools
|
||||
|
||||
Fixed set, independent of the models. `table` is `schema.entity`. Errors return `{"success":false,"error":{"code","message"}}` with codes `invalid_argument`, `not_found`, `forbidden`, `limit_exceeded`, `internal` (internal details are logged, the client gets a reference id).
|
||||
|
||||
@@ -0,0 +1,302 @@
|
||||
package resolvemcp
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"reflect"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||
)
|
||||
|
||||
// usageGuide is the short agent-facing guide sent as the MCP server instructions and
|
||||
// embedded in the exported catalogue. It is generic: it never mentions concrete models.
|
||||
const usageGuide = `This server exposes database tables through a fixed set of tools.
|
||||
1. Call list_tables to see the tables you may use, what they hold and the operations allowed.
|
||||
2. Call describe_table for a table before using it: columns, types, primary key, relations (preloadable), writable columns and limits.
|
||||
3. Read with select_table (filters, sort, columns, preloads). Results are paged; use limit/offset or cursors, and include_count only when you need a total.
|
||||
4. Write with insert_into_table, update_table, delete_from_table. Address one row by id, or several by filters. A filter-based write first returns a preview; repeat the call with the confirm_token to apply it (dry_run only previews).
|
||||
5. Use list_functions / call_function for registered functions.
|
||||
Read the error message when a call fails: it says which argument was wrong.`
|
||||
|
||||
// Catalog is a snapshot of what the server offers: the usage guide, the tools, the limits
|
||||
// and every table with its columns, relations, allowed operations and descriptions.
|
||||
type Catalog struct {
|
||||
GeneratedAt time.Time `json:"generated_at"`
|
||||
Server string `json:"server"`
|
||||
Version string `json:"version"`
|
||||
Guide string `json:"guide"`
|
||||
Limits CatalogLimits `json:"limits"`
|
||||
Tools []CatalogTool `json:"tools"`
|
||||
Tables []CatalogTable `json:"tables"`
|
||||
}
|
||||
|
||||
// CatalogLimits mirrors the configured server limits.
|
||||
type CatalogLimits struct {
|
||||
DefaultLimit int `json:"default_limit"`
|
||||
MaxLimit int `json:"max_limit"`
|
||||
MaxOffset int `json:"max_offset"`
|
||||
MaxBatch int `json:"max_batch"`
|
||||
MaxPreloadDepth int `json:"max_preload_depth"`
|
||||
MaxWriteRows int `json:"max_write_rows"`
|
||||
}
|
||||
|
||||
// CatalogTool is one MCP tool.
|
||||
type CatalogTool struct {
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
}
|
||||
|
||||
// CatalogTable is one table in the catalogue.
|
||||
type CatalogTable struct {
|
||||
Table string `json:"table"`
|
||||
Description string `json:"description,omitempty"`
|
||||
Purpose string `json:"purpose,omitempty"`
|
||||
Tags []string `json:"tags,omitempty"`
|
||||
Operations []string `json:"operations"`
|
||||
PrimaryKey string `json:"primary_key,omitempty"`
|
||||
Columns []CatalogColumn `json:"columns"`
|
||||
Relations []string `json:"relations,omitempty"`
|
||||
}
|
||||
|
||||
// CatalogColumn is one column of a table.
|
||||
type CatalogColumn struct {
|
||||
Name string `json:"name"`
|
||||
Type string `json:"type,omitempty"`
|
||||
Nullable bool `json:"nullable"`
|
||||
PrimaryKey bool `json:"primary_key,omitempty"`
|
||||
Unique bool `json:"unique,omitempty"`
|
||||
Writable bool `json:"writable"`
|
||||
Description string `json:"description,omitempty"`
|
||||
}
|
||||
|
||||
// modelDocs returns the effective documentation of a table. Registry info (including a
|
||||
// loaded descriptions file) wins, then the model's ModelDescription().
|
||||
func (h *Handler) modelDocs(schema, entity string) modelregistry.ModelInfo {
|
||||
if reg, ok := h.registry.(*modelregistry.DefaultModelRegistry); ok {
|
||||
return reg.ResolveModelInfo(buildModelName(schema, entity))
|
||||
}
|
||||
return modelregistry.ModelInfo{}
|
||||
}
|
||||
|
||||
// columnDescription picks a column's description: the registry/file map first, then the
|
||||
// struct-tag comment.
|
||||
func columnDescription(docs modelregistry.ModelInfo, c columnInfo) string {
|
||||
if d := docs.Columns[c.jsonName]; d != "" {
|
||||
return d
|
||||
}
|
||||
return c.comment
|
||||
}
|
||||
|
||||
// SetModelDescription stores documentation for a registered or soon-to-be-registered table.
|
||||
// It returns an error when the handler's registry does not keep model info.
|
||||
func (h *Handler) SetModelDescription(schema, entity string, info modelregistry.ModelInfo) error {
|
||||
reg, ok := h.registry.(*modelregistry.DefaultModelRegistry)
|
||||
if !ok {
|
||||
return fmt.Errorf("resolvemcp: registry does not support model descriptions (use NewHandlerWithGORM/Bun/DB)")
|
||||
}
|
||||
reg.SetModelInfo(buildModelName(schema, entity), info)
|
||||
return nil
|
||||
}
|
||||
|
||||
// LoadModelDescriptions loads an external JSON map of descriptions keyed by "schema.entity"
|
||||
// (see modelregistry.LoadModelInfo for the format) and returns how many tables it covered.
|
||||
// Loaded entries override the model's own comments.
|
||||
func (h *Handler) LoadModelDescriptions(path string) (int, error) {
|
||||
reg, ok := h.registry.(*modelregistry.DefaultModelRegistry)
|
||||
if !ok {
|
||||
return 0, fmt.Errorf("resolvemcp: registry does not support model descriptions (use NewHandlerWithGORM/Bun/DB)")
|
||||
}
|
||||
return reg.LoadModelInfoFile(path)
|
||||
}
|
||||
|
||||
// BuildCatalog snapshots the server's tools and tables. Tables with no allowed operation are
|
||||
// left out, exactly as list_tables does.
|
||||
func (h *Handler) BuildCatalog() Catalog {
|
||||
cat := Catalog{
|
||||
GeneratedAt: time.Now().UTC(),
|
||||
Server: h.name,
|
||||
Version: h.version,
|
||||
Guide: usageGuide,
|
||||
Limits: CatalogLimits{
|
||||
DefaultLimit: h.config.DefaultLimit,
|
||||
MaxLimit: h.config.MaxLimit,
|
||||
MaxOffset: h.config.MaxOffset,
|
||||
MaxBatch: h.config.MaxBatch,
|
||||
MaxPreloadDepth: h.config.MaxPreloadDepth,
|
||||
MaxWriteRows: h.config.MaxWriteRows,
|
||||
},
|
||||
Tools: []CatalogTool{},
|
||||
Tables: []CatalogTable{},
|
||||
}
|
||||
|
||||
for name, tool := range h.mcpServer.ListTools() {
|
||||
cat.Tools = append(cat.Tools, CatalogTool{Name: name, Description: tool.Tool.Description})
|
||||
}
|
||||
sort.Slice(cat.Tools, func(i, j int) bool { return cat.Tools[i].Name < cat.Tools[j].Name })
|
||||
|
||||
for name, model := range h.registry.GetAllModels() {
|
||||
schema, entity, _ := splitTable(name)
|
||||
rules := h.modelRules(schema, entity)
|
||||
ops := opsFor(rules)
|
||||
if len(ops) == 0 {
|
||||
continue
|
||||
}
|
||||
info := buildModelInfo(schema, entity, model)
|
||||
docs := h.modelDocs(schema, entity)
|
||||
|
||||
writable := map[string]bool{}
|
||||
mt := reflect.TypeOf(model)
|
||||
for mt != nil && (mt.Kind() == reflect.Pointer || mt.Kind() == reflect.Slice) {
|
||||
mt = mt.Elem()
|
||||
}
|
||||
if mt != nil && mt.Kind() == reflect.Struct {
|
||||
for k := range reflectionJSONColumns(mt) {
|
||||
writable[k] = true
|
||||
}
|
||||
}
|
||||
|
||||
t := CatalogTable{
|
||||
Table: info.fullName,
|
||||
Description: docs.Description,
|
||||
Purpose: docs.Purpose,
|
||||
Tags: docs.Tags,
|
||||
Operations: ops,
|
||||
PrimaryKey: info.pkName,
|
||||
Relations: info.relationNames,
|
||||
Columns: make([]CatalogColumn, 0, len(info.columns)),
|
||||
}
|
||||
for _, c := range info.columns {
|
||||
typ := c.sqlType
|
||||
if typ == "" {
|
||||
typ = c.goType
|
||||
}
|
||||
t.Columns = append(t.Columns, CatalogColumn{
|
||||
Name: c.jsonName, Type: typ, Nullable: c.nullable, PrimaryKey: c.isPrimary,
|
||||
Unique: c.isUnique, Writable: writable[c.jsonName], Description: columnDescription(docs, c),
|
||||
})
|
||||
}
|
||||
cat.Tables = append(cat.Tables, t)
|
||||
}
|
||||
sort.Slice(cat.Tables, func(i, j int) bool { return cat.Tables[i].Table < cat.Tables[j].Table })
|
||||
return cat
|
||||
}
|
||||
|
||||
// ExportCatalog writes the catalogue to path. A ".json" extension writes JSON; anything
|
||||
// else writes Markdown. The file is replaced atomically (written to a temp file in the same
|
||||
// directory, then renamed) and created with mode 0600. It lists every table the registry
|
||||
// allows any operation on, regardless of caller, so keep it out of public directories.
|
||||
func (h *Handler) ExportCatalog(path string) error {
|
||||
cat := h.BuildCatalog()
|
||||
|
||||
var data []byte
|
||||
if strings.EqualFold(filepath.Ext(path), ".json") {
|
||||
b, err := json.MarshalIndent(cat, "", " ")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
data = b
|
||||
data = append(data, '\n')
|
||||
} else {
|
||||
data = []byte(cat.Markdown())
|
||||
}
|
||||
|
||||
dir := filepath.Dir(path)
|
||||
if err := os.MkdirAll(dir, 0o750); err != nil {
|
||||
return fmt.Errorf("resolvemcp: export catalog: %w", err)
|
||||
}
|
||||
tmp, err := os.CreateTemp(dir, ".catalog-*")
|
||||
if err != nil {
|
||||
return fmt.Errorf("resolvemcp: export catalog: %w", err)
|
||||
}
|
||||
tmpName := tmp.Name()
|
||||
_, werr := tmp.Write(data)
|
||||
cerr := tmp.Close()
|
||||
if werr == nil {
|
||||
werr = cerr
|
||||
}
|
||||
if werr == nil {
|
||||
werr = os.Rename(tmpName, path)
|
||||
}
|
||||
if werr != nil {
|
||||
_ = os.Remove(tmpName)
|
||||
return fmt.Errorf("resolvemcp: export catalog: %w", werr)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Markdown renders the catalogue as a Markdown document.
|
||||
func (c Catalog) Markdown() string {
|
||||
var sb strings.Builder
|
||||
fmt.Fprintf(&sb, "# %s API catalogue\n\nGenerated %s.\n\n", c.Server, c.GeneratedAt.Format(time.RFC3339))
|
||||
sb.WriteString("## How to use\n\n" + c.Guide + "\n\n")
|
||||
fmt.Fprintf(&sb, "## Limits\n\ndefault limit %d, max limit %d, max offset %d, max batch %d, max preload depth %d, max rows per filter write %d.\n\n",
|
||||
c.Limits.DefaultLimit, c.Limits.MaxLimit, c.Limits.MaxOffset, c.Limits.MaxBatch, c.Limits.MaxPreloadDepth, c.Limits.MaxWriteRows)
|
||||
|
||||
sb.WriteString("## Tools\n\n")
|
||||
for _, t := range c.Tools {
|
||||
fmt.Fprintf(&sb, "- `%s`: %s\n", t.Name, oneLine(t.Description))
|
||||
}
|
||||
|
||||
sb.WriteString("\n## Tables\n\n")
|
||||
if len(c.Tables) == 0 {
|
||||
sb.WriteString("No tables are registered.\n")
|
||||
}
|
||||
for i := range c.Tables {
|
||||
t := &c.Tables[i]
|
||||
fmt.Fprintf(&sb, "### %s\n\n", t.Table)
|
||||
if t.Description != "" {
|
||||
sb.WriteString(t.Description + "\n\n")
|
||||
}
|
||||
if t.Purpose != "" {
|
||||
sb.WriteString("Purpose: " + t.Purpose + "\n\n")
|
||||
}
|
||||
if len(t.Tags) > 0 {
|
||||
sb.WriteString("Tags: " + strings.Join(t.Tags, ", ") + "\n\n")
|
||||
}
|
||||
fmt.Fprintf(&sb, "Operations: %s", strings.Join(t.Operations, ", "))
|
||||
if t.PrimaryKey != "" {
|
||||
fmt.Fprintf(&sb, " · Primary key: `%s`", t.PrimaryKey)
|
||||
}
|
||||
sb.WriteString("\n\n| Column | Type | Flags | Description |\n|---|---|---|---|\n")
|
||||
for _, col := range t.Columns {
|
||||
var flags []string
|
||||
if col.PrimaryKey {
|
||||
flags = append(flags, "pk")
|
||||
}
|
||||
if col.Unique {
|
||||
flags = append(flags, "unique")
|
||||
}
|
||||
if col.Nullable {
|
||||
flags = append(flags, "nullable")
|
||||
}
|
||||
if !col.Writable {
|
||||
flags = append(flags, "read-only")
|
||||
}
|
||||
fmt.Fprintf(&sb, "| `%s` | %s | %s | %s |\n", col.Name, mdCell(col.Type), strings.Join(flags, ", "), mdCell(col.Description))
|
||||
}
|
||||
if len(t.Relations) > 0 {
|
||||
sb.WriteString("\nRelations (preloadable): " + strings.Join(t.Relations, ", ") + "\n")
|
||||
}
|
||||
sb.WriteString("\n")
|
||||
}
|
||||
return sb.String()
|
||||
}
|
||||
|
||||
func oneLine(s string) string {
|
||||
return strings.Join(strings.Fields(s), " ")
|
||||
}
|
||||
|
||||
func mdCell(s string) string {
|
||||
return strings.ReplaceAll(oneLine(s), "|", `\|`)
|
||||
}
|
||||
|
||||
// reflectionJSONColumns returns the JSON names of the columns a write may set.
|
||||
func reflectionJSONColumns(t reflect.Type) map[string]string {
|
||||
return reflection.BuildJSONToDBColumnMap(t)
|
||||
}
|
||||
@@ -0,0 +1,126 @@
|
||||
package resolvemcp
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||
)
|
||||
|
||||
type docItem struct {
|
||||
ID int `json:"id" bun:"id,pk"`
|
||||
Email string `json:"email" bun:"email,comment:Tag comment"`
|
||||
Name string `json:"name" bun:"name" note:"Name from tag"`
|
||||
}
|
||||
|
||||
func (docItem) ModelDescription() string { return "Model-level fallback" }
|
||||
|
||||
func newDocHandler(t *testing.T) *Handler {
|
||||
t.Helper()
|
||||
h := NewHandler(database.NewPgSQLAdapter(nil), modelregistry.NewModelRegistry(), Config{})
|
||||
if err := h.RegisterModel("public", "items", &docItem{}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
hidden := modelregistry.ModelRules{} // no operation allowed
|
||||
if err := h.RegisterModelWithRules("public", "secret", &docItem{}, hidden); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return h
|
||||
}
|
||||
|
||||
func TestCatalogDescriptionsPrecedence(t *testing.T) {
|
||||
h := newDocHandler(t)
|
||||
cat := h.BuildCatalog()
|
||||
if len(cat.Tables) != 1 || cat.Tables[0].Table != "public.items" {
|
||||
t.Fatalf("tables = %+v (hidden table must be left out)", cat.Tables)
|
||||
}
|
||||
tb := cat.Tables[0]
|
||||
if tb.Description != "Model-level fallback" {
|
||||
t.Errorf("fallback description = %q", tb.Description)
|
||||
}
|
||||
col := map[string]string{}
|
||||
for _, c := range tb.Columns {
|
||||
col[c.Name] = c.Description
|
||||
}
|
||||
if col["email"] != "Tag comment" || col["name"] != "Name from tag" {
|
||||
t.Errorf("tag comments = %v", col)
|
||||
}
|
||||
|
||||
path := filepath.Join(t.TempDir(), "desc.json")
|
||||
if err := os.WriteFile(path, []byte(`{"public.items":{"description":"From file","columns":{"email":"File email"}}}`), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if n, err := h.LoadModelDescriptions(path); err != nil || n != 1 {
|
||||
t.Fatalf("load n=%d err=%v", n, err)
|
||||
}
|
||||
tb = h.BuildCatalog().Tables[0]
|
||||
if tb.Description != "From file" {
|
||||
t.Errorf("file must win, got %q", tb.Description)
|
||||
}
|
||||
for _, c := range tb.Columns {
|
||||
switch c.Name {
|
||||
case "email":
|
||||
if c.Description != "File email" {
|
||||
t.Errorf("email = %q", c.Description)
|
||||
}
|
||||
case "name":
|
||||
if c.Description != "Name from tag" {
|
||||
t.Errorf("name must fall back to tag, got %q", c.Description)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestExportCatalogFiles(t *testing.T) {
|
||||
h := newDocHandler(t)
|
||||
dir := t.TempDir()
|
||||
|
||||
md := filepath.Join(dir, "sub", "catalog.md")
|
||||
if err := h.ExportCatalog(md); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
b, _ := os.ReadFile(md)
|
||||
for _, want := range []string{"# resolvemcp API catalogue", "### public.items", "Model-level fallback", "`list_tables`", "Tag comment"} {
|
||||
if !strings.Contains(string(b), want) {
|
||||
t.Errorf("markdown missing %q", want)
|
||||
}
|
||||
}
|
||||
if strings.Contains(string(b), "public.secret") {
|
||||
t.Error("table without operations leaked into the catalogue")
|
||||
}
|
||||
|
||||
js := filepath.Join(dir, "catalog.json")
|
||||
if err := h.ExportCatalog(js); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
var cat Catalog
|
||||
b, _ = os.ReadFile(js)
|
||||
if err := json.Unmarshal(b, &cat); err != nil || len(cat.Tables) != 1 || len(cat.Tools) == 0 {
|
||||
t.Fatalf("json catalog bad: err=%v %+v", err, cat)
|
||||
}
|
||||
|
||||
entries, _ := os.ReadDir(dir)
|
||||
for _, e := range entries {
|
||||
if strings.HasPrefix(e.Name(), ".catalog-") {
|
||||
t.Errorf("temp file left behind: %s", e.Name())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDescribeAndListIncludeDescriptions(t *testing.T) {
|
||||
h := newDocHandler(t)
|
||||
res, _ := h.handleListTables(nil, callReq(nil))
|
||||
tables, _ := payload(t, res)["tables"].([]any)
|
||||
if len(tables) != 1 || tables[0].(map[string]any)["description"] != "Model-level fallback" {
|
||||
t.Errorf("list_tables = %v", tables)
|
||||
}
|
||||
res, _ = h.handleDescribeTable(nil, callReq(map[string]any{"table": "public.items"}))
|
||||
p := payload(t, res)
|
||||
if p["description"] != "Model-level fallback" {
|
||||
t.Errorf("describe_table description = %v", p["description"])
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
// Package resolvemcp exposes registered database models as Model Context Protocol (MCP)
|
||||
// tools over HTTP (SSE and streamable HTTP), so an AI agent can discover, read and
|
||||
// change data without a tool per table.
|
||||
//
|
||||
// # How an agent uses it
|
||||
//
|
||||
// The tool set is fixed and does not grow with the models:
|
||||
//
|
||||
// list_tables tables the caller may use, their description and allowed operations
|
||||
// describe_table columns, types, keys, relations, writable fields, limits
|
||||
// select_table read rows: filters, sort, columns, preloads, paging, cursors
|
||||
// insert_into_table / update_table / delete_from_table
|
||||
// writes; filter-based writes are previewed (dry_run) and need the
|
||||
// confirm_token from the preview
|
||||
// list_functions / call_function registered stored functions
|
||||
// resolvespec_annotate optional free-text notes (Config.EnableAnnotations)
|
||||
//
|
||||
// The same guide is sent to MCP clients as the server instructions.
|
||||
//
|
||||
// # Setting it up
|
||||
//
|
||||
// handler := resolvemcp.NewHandlerWithGORM(db, resolvemcp.Config{BaseURL: "http://localhost:8080"})
|
||||
// handler.RegisterModel("public", "users", &User{})
|
||||
//
|
||||
// r := mux.NewRouter()
|
||||
// resolvemcp.SetupMuxRoutes(r, handler, securityList) // requires an authenticated caller
|
||||
//
|
||||
// # Describing the API for agents
|
||||
//
|
||||
// Models are documented through the model registry (see package modelregistry):
|
||||
// SetModelDescription / LoadModelDescriptions on the handler, a ModelDescription()
|
||||
// method on the model, or comment tags on its fields (gorm/bun "comment:" or
|
||||
// comment/note/desc tags). The descriptions show up in list_tables and
|
||||
// describe_table.
|
||||
//
|
||||
// ExportCatalog writes the whole picture (guide, tools, limits, tables with columns,
|
||||
// relations, operations and descriptions) to a JSON or Markdown file on disk, so
|
||||
// agents and developers can learn the API without connecting:
|
||||
//
|
||||
// handler.ExportCatalog("docs/mcp-catalog.md")
|
||||
//
|
||||
// # Security
|
||||
//
|
||||
// Routes must be mounted behind Guard(securityList); the *Unauthenticated setup
|
||||
// functions exist only for use behind another trusted layer. Per-entity rules come from
|
||||
// modelregistry.ModelRules, and BeforeHandle/AfterHandle hooks can veto or audit any call.
|
||||
package resolvemcp
|
||||
@@ -43,7 +43,7 @@ func NewHandler(db common.Database, registry common.ModelRegistry, cfg Config) *
|
||||
db: db,
|
||||
registry: registry,
|
||||
hooks: NewHookRegistry(),
|
||||
mcpServer: server.NewMCPServer("resolvemcp", "1.0.0"),
|
||||
mcpServer: server.NewMCPServer("resolvemcp", "1.0.0", server.WithInstructions(usageGuide)),
|
||||
config: cfg.withDefaults(),
|
||||
confirms: newConfirmStore(),
|
||||
name: "resolvemcp",
|
||||
@@ -559,6 +559,7 @@ func (h *Handler) executeCreate(ctx context.Context, schema, entity string, data
|
||||
if len(cols) == 0 {
|
||||
return invalidArg("no writable fields in data")
|
||||
}
|
||||
reflection.RemoveNonWritableColumns(model, cols)
|
||||
q := tx.NewInsert().Table(tableName)
|
||||
for key, value := range cols {
|
||||
q = q.Value(key, value)
|
||||
@@ -726,6 +727,7 @@ func (h *Handler) executeUpdate(ctx context.Context, schema, entity, id string,
|
||||
existingMap[key] = v
|
||||
}
|
||||
|
||||
reflection.RemoveNonWritableColumns(model, setCols)
|
||||
q := tx.NewUpdate().Table(tableName).SetMap(setCols).
|
||||
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), id)
|
||||
res, err := q.Exec(ctx)
|
||||
|
||||
+10
-4
@@ -156,14 +156,15 @@ func (h *Handler) resolveTable(args map[string]any, op string) (schema, entity s
|
||||
|
||||
func (h *Handler) handleListTables(ctx context.Context, _ mcp.CallToolRequest) (*mcp.CallToolResult, error) {
|
||||
type table struct {
|
||||
Table string `json:"table"`
|
||||
Operations []string `json:"operations"`
|
||||
Table string `json:"table"`
|
||||
Description string `json:"description,omitempty"`
|
||||
Operations []string `json:"operations"`
|
||||
}
|
||||
var tables []table
|
||||
for name := range h.registry.GetAllModels() {
|
||||
schema, entity, _ := splitTable(name)
|
||||
if ops := opsFor(h.modelRules(schema, entity)); len(ops) > 0 {
|
||||
tables = append(tables, table{Table: name, Operations: ops})
|
||||
tables = append(tables, table{Table: name, Description: h.modelDocs(schema, entity).Description, Operations: ops})
|
||||
}
|
||||
}
|
||||
sort.Slice(tables, func(i, j int) bool { return tables[i].Table < tables[j].Table })
|
||||
@@ -184,6 +185,7 @@ func (h *Handler) handleDescribeTable(_ context.Context, req mcp.CallToolRequest
|
||||
return toolError("describe_table", invalidArg("unknown table %q; see list_tables", buildModelName(schema, entity))), nil
|
||||
}
|
||||
info := buildModelInfo(schema, entity, model)
|
||||
docs := h.modelDocs(schema, entity)
|
||||
|
||||
modelType := reflect.TypeOf(model)
|
||||
for modelType != nil && (modelType.Kind() == reflect.Pointer || modelType.Kind() == reflect.Slice) {
|
||||
@@ -203,6 +205,7 @@ func (h *Handler) handleDescribeTable(_ context.Context, req mcp.CallToolRequest
|
||||
PrimaryKey bool `json:"primary_key,omitempty"`
|
||||
Unique bool `json:"unique,omitempty"`
|
||||
Writable bool `json:"writable"`
|
||||
Comment string `json:"description,omitempty"`
|
||||
}
|
||||
cols := make([]column, 0, len(info.columns))
|
||||
var writableNames []string
|
||||
@@ -212,7 +215,7 @@ func (h *Handler) handleDescribeTable(_ context.Context, req mcp.CallToolRequest
|
||||
typ = c.goType
|
||||
}
|
||||
w := writable[c.jsonName]
|
||||
cols = append(cols, column{Name: c.jsonName, Type: typ, Nullable: c.nullable, PrimaryKey: c.isPrimary, Unique: c.isUnique, Writable: w})
|
||||
cols = append(cols, column{Name: c.jsonName, Type: typ, Nullable: c.nullable, PrimaryKey: c.isPrimary, Unique: c.isUnique, Writable: w, Comment: columnDescription(docs, c)})
|
||||
if w && !c.isPrimary {
|
||||
writableNames = append(writableNames, c.jsonName)
|
||||
}
|
||||
@@ -220,6 +223,9 @@ func (h *Handler) handleDescribeTable(_ context.Context, req mcp.CallToolRequest
|
||||
return marshalResult(map[string]any{
|
||||
"success": true,
|
||||
"table": info.fullName,
|
||||
"description": docs.Description,
|
||||
"purpose": docs.Purpose,
|
||||
"tags": docs.Tags,
|
||||
"primary_key": info.pkName,
|
||||
"columns": cols,
|
||||
"relations": info.relationNames,
|
||||
|
||||
@@ -1,18 +1,3 @@
|
||||
// Package resolvemcp exposes registered database models as Model Context Protocol (MCP) tools
|
||||
// and resources over HTTP/SSE transport.
|
||||
//
|
||||
// It mirrors the resolvespec package patterns:
|
||||
// - Same model registration API
|
||||
// - Same filter, sort, cursor pagination, preload options
|
||||
// - Same lifecycle hook system
|
||||
//
|
||||
// Usage:
|
||||
//
|
||||
// handler := resolvemcp.NewHandlerWithGORM(db, resolvemcp.Config{BaseURL: "http://localhost:8080"})
|
||||
// handler.RegisterModel("public", "users", &User{})
|
||||
//
|
||||
// r := mux.NewRouter()
|
||||
// resolvemcp.SetupMuxRoutes(r, handler, securityList) // requires an authenticated caller
|
||||
package resolvemcp
|
||||
|
||||
import (
|
||||
|
||||
@@ -9,6 +9,7 @@ import (
|
||||
"github.com/mark3labs/mcp-go/mcp"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||
)
|
||||
|
||||
@@ -30,6 +31,7 @@ type columnInfo struct {
|
||||
isUnique bool
|
||||
isFK bool
|
||||
nullable bool
|
||||
comment string // from struct tags (gorm/bun comment:, comment/note/desc tags)
|
||||
}
|
||||
|
||||
// buildModelInfo extracts column metadata and pre-builds the schema documentation string.
|
||||
@@ -100,7 +102,12 @@ func buildModelInfo(schema, entity string, model interface{}) modelInfo {
|
||||
isPrimary := d.SQLKey == "primary_key" ||
|
||||
(info.pkName != "" && (sqlName == info.pkName || jsonName == info.pkName))
|
||||
|
||||
comment := ""
|
||||
if found {
|
||||
comment = modelregistry.FieldComment(fieldType)
|
||||
}
|
||||
ci := columnInfo{
|
||||
comment: comment,
|
||||
jsonName: jsonName,
|
||||
sqlName: sqlName,
|
||||
goType: goType,
|
||||
|
||||
@@ -194,6 +194,7 @@ func (h *Handler) executeWhere(ctx context.Context, req whereRequest) (_ *whereR
|
||||
cond := fmt.Sprintf("%s IN (%s)", common.QuoteIdent(pkName), strings.Join(inList, ", "))
|
||||
var affected int64
|
||||
if req.op == "update" {
|
||||
reflection.RemoveNonWritableColumns(model, setCols)
|
||||
r, err := tx.NewUpdate().Table(tableName).SetMap(setCols).Where(cond, ids...).Exec(ctx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("error updating records: %w", err)
|
||||
|
||||
@@ -824,6 +824,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
|
||||
}
|
||||
responseData = v
|
||||
|
||||
reflection.RemoveNonWritableColumns(model, v)
|
||||
query := tx.NewInsert().Table(tableName)
|
||||
for key, value := range v {
|
||||
query = query.Value(key, common.ConvertSliceForBun(value))
|
||||
@@ -971,6 +972,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
|
||||
item = modifiedData
|
||||
}
|
||||
|
||||
reflection.RemoveNonWritableColumns(model, item)
|
||||
txQuery := tx.NewInsert().Table(tableName)
|
||||
for key, value := range item {
|
||||
txQuery = txQuery.Value(key, common.ConvertSliceForBun(value))
|
||||
@@ -1127,6 +1129,7 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
|
||||
itemMap = modifiedData
|
||||
}
|
||||
|
||||
reflection.RemoveNonWritableColumns(model, itemMap)
|
||||
txQuery := tx.NewInsert().Table(tableName)
|
||||
for key, value := range itemMap {
|
||||
txQuery = txQuery.Value(key, common.ConvertSliceForBun(value))
|
||||
@@ -1322,6 +1325,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
|
||||
|
||||
// Overwrite with every key present in the request (including "" and null unless disallowed)
|
||||
common.MergeUpdateValues(existingMap, updates, h.disallowNulls)
|
||||
reflection.RemoveNonWritableColumns(model, existingMap)
|
||||
|
||||
// Build update query with merged data
|
||||
query := tx.NewUpdate().Table(tableName).SetMap(existingMap)
|
||||
@@ -1507,6 +1511,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
|
||||
|
||||
// Overwrite with every key present in the request (including "" and null unless disallowed)
|
||||
common.MergeUpdateValues(existingMap, item, h.disallowNulls)
|
||||
reflection.RemoveNonWritableColumns(model, existingMap)
|
||||
|
||||
txQuery := tx.NewUpdate().Table(tableName).SetMap(existingMap).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID)
|
||||
if _, err := txQuery.Exec(ctx); err != nil {
|
||||
@@ -1662,6 +1667,7 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, url
|
||||
|
||||
// Overwrite with every key present in the request (including "" and null unless disallowed)
|
||||
common.MergeUpdateValues(existingMap, itemMap, h.disallowNulls)
|
||||
reflection.RemoveNonWritableColumns(model, existingMap)
|
||||
|
||||
txQuery := tx.NewUpdate().Table(tableName).SetMap(existingMap).Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), itemID)
|
||||
if _, err := txQuery.Exec(ctx); err != nil {
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
@@ -1410,6 +1418,9 @@ func (h *Handler) handleCreate(ctx context.Context, w common.ResponseWriter, dat
|
||||
if provider, ok := modelValue.(common.TableNameProvider); !ok || provider.TableName() == "" {
|
||||
query = query.Table(tableName)
|
||||
}
|
||||
if generated := reflection.NonWritableColumns(model); len(generated) > 0 {
|
||||
query = query.ExcludeColumn(generated...)
|
||||
}
|
||||
fields := reflection.GetSQLModelColumns(model)
|
||||
query = query.Returning(fields...)
|
||||
|
||||
@@ -1657,6 +1668,9 @@ func (h *Handler) handleUpdate(ctx context.Context, w common.ResponseWriter, id
|
||||
|
||||
// Create update query using Model() to preserve custom types and driver.Valuer interfaces
|
||||
query := tx.NewUpdate().Model(modelInstance)
|
||||
if generated := reflection.NonWritableColumns(model); len(generated) > 0 {
|
||||
query = query.ExcludeColumn(generated...)
|
||||
}
|
||||
query = query.Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), targetID)
|
||||
|
||||
// Execute BeforeScan hooks - pass query chain so hooks can modify it
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -758,6 +758,9 @@ func (h *Handler) create(hookCtx *HookContext) (interface{}, error) {
|
||||
|
||||
// Insert record
|
||||
query := hookCtx.Tx.NewInsert().Model(hookCtx.ModelPtr).Table(hookCtx.TableName)
|
||||
if generated := reflection.NonWritableColumns(hookCtx.Model); len(generated) > 0 {
|
||||
query = query.ExcludeColumn(generated...)
|
||||
}
|
||||
if _, err := query.Exec(hookCtx.Context); err != nil {
|
||||
return nil, fmt.Errorf("failed to create record: %w", err)
|
||||
}
|
||||
@@ -786,6 +789,8 @@ func (h *Handler) update(hookCtx *HookContext) error {
|
||||
// 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)
|
||||
|
||||
reflection.RemoveNonWritableColumns(hookCtx.Model, values)
|
||||
|
||||
if len(values) > 0 {
|
||||
query := hookCtx.Tx.NewUpdate().Table(hookCtx.TableName).SetMap(values).
|
||||
Where(fmt.Sprintf("%s = ?", common.QuoteIdent(pkName)), hookCtx.ID)
|
||||
|
||||
@@ -226,6 +226,11 @@ func (m *MockInsertQuery) OnConflict(action string) common.InsertQuery {
|
||||
return args.Get(0).(common.InsertQuery)
|
||||
}
|
||||
|
||||
func (m *MockInsertQuery) ExcludeColumn(columns ...string) common.InsertQuery {
|
||||
args := m.Called(columns)
|
||||
return args.Get(0).(common.InsertQuery)
|
||||
}
|
||||
|
||||
func (m *MockInsertQuery) Returning(columns ...string) common.InsertQuery {
|
||||
args := m.Called(columns)
|
||||
return args.Get(0).(common.InsertQuery)
|
||||
@@ -254,6 +259,11 @@ func (m *MockUpdateQuery) Model(model interface{}) common.UpdateQuery {
|
||||
return args.Get(0).(common.UpdateQuery)
|
||||
}
|
||||
|
||||
func (m *MockUpdateQuery) ExcludeColumn(columns ...string) common.UpdateQuery {
|
||||
args := m.Called(columns)
|
||||
return args.Get(0).(common.UpdateQuery)
|
||||
}
|
||||
|
||||
func (m *MockUpdateQuery) Table(table string) common.UpdateQuery {
|
||||
args := m.Called(table)
|
||||
return args.Get(0).(common.UpdateQuery)
|
||||
|
||||
Reference in New Issue
Block a user