Compare commits

..
4 Commits
Author SHA1 Message Date
Hein 5a3a1df3c8 feat(resolvemcp)!: make the server read-only by default
Tests / Integration Tests (push) Skipped
Build , Vet Test, and Lint / Build (push) Successful in 1m33s
Tests / Unit Tests (push) Successful in 2m1s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 2m36s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 2m38s
Build , Vet Test, and Lint / Lint Code (push) Successful in 2m49s
Tests / Race Detector (push) Successful in 4m31s
BREAKING CHANGE: Config.ReadOnly is now a *bool and unset means read-only.
Use ReadOnly: resolvemcp.Bool(false) to enable insert/update/delete,
annotations and function calls. Adds the Bool helper and updates docs.
2026-10-07 14:17:56 +02:00
Hein 4ed9506ad2 feat(resolvemcp): add read-only mode and function allowlist
- Config.ReadOnly disables insert/update/delete/annotation tools, reports only
  select in list_tables/describe_table and tells the agent it cannot write
- Config.AllowFunctionCalls keeps function tools on a read-only server
- Config.AllowedFunctions limits list_functions/call_function to named
  functions (empty allows all); others are reported as unknown
- reflect read-only mode in the usage guide and exported catalogue
2026-10-07 14:15:14 +02:00
Hein 431b674162 feat(resolvemcp): add model descriptions, usage guide and catalogue export
- modelregistry: ModelInfo (description, purpose, tags, column docs) with
  external JSON loader, Describer fallback and gorm/bun/comment tag support
- resolvemcp: surface descriptions in list_tables and describe_table, send a
  usage guide as MCP server instructions, add package docs
- add BuildCatalog/ExportCatalog to write a JSON or Markdown API catalogue
- document the descriptions map and catalogue in the README
2026-10-07 14:09:37 +02:00
Hein 234aac9770 feat(metrics): bound HTTP path labels, add reset, custom push endpoint and JSON pull
- normalize the HTTP path label (ServeMux pattern, custom normalizer, ID
  collapsing) and cap distinct values via HTTPMaxPaths (default 1024)
- add Reset, PushAndReset, ResetHandler and reset-on-push options
- add POST push to a custom endpoint (text or json) with optional reset
- add JSONHandler for JSON pull
- honour Config.Enabled; log Pushgateway push failures
2026-10-07 12:09:11 +02:00
24 changed files with 2043 additions and 41 deletions
+2 -2
View File
@@ -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
+50 -3
View File
@@ -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
+53
View File
@@ -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"
+189
View File
@@ -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
}
+15
View File
@@ -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
+175
View File
@@ -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'
}
+237
View File
@@ -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
View File
@@ -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()
}
}
+32
View File
@@ -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
+172
View File
@@ -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
}
+68
View File
@@ -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")
}
}
+1
View 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
}
+82
View File
@@ -16,6 +16,8 @@ import (
handler := resolvemcp.NewHandlerWithGORM(db, resolvemcp.Config{
BaseURL: "http://localhost:8080",
BasePath: "/mcp",
// Read-only by default; uncomment to allow writes:
// ReadOnly: resolvemcp.Bool(false),
})
securityList, _ := security.NewSecurityList(provider)
@@ -392,6 +394,85 @@ 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).
## Read-only mode
The server is **read-only unless you enable writes**: `Config.ReadOnly` is a `*bool` and an unset (nil) value means on. To allow inserts, updates and deletes:
```go
handler := resolvemcp.NewHandlerWithGORM(db, resolvemcp.Config{ReadOnly: resolvemcp.Bool(false)})
```
While read-only is on:
- The insert, update, delete and annotation tools are not registered, so the agent never sees them. `list_functions`/`call_function` are off too, because a registered function may change data, unless you set `AllowFunctionCalls` (below).
- `list_tables` and `describe_table` report only `select`; `describe_table` also sets `read_only: true` and lists no writable columns.
- The MCP server instructions (and the exported catalogue) say the server is read-only and tell the agent not to attempt writes.
- A write that reaches a handler anyway is refused with a `forbidden` error ("this server is read-only: writes are disabled").
### Function calls and the allowlist
```go
// Read-only server that may still run two named functions
resolvemcp.Config{
// ReadOnly is on by default
AllowFunctionCalls: true, // keep list_functions / call_function on a read-only server
AllowedFunctions: []string{"report_totals", "search_customers"},
}
```
- `AllowFunctionCalls` only matters while read-only is on; with writes enabled (`ReadOnly: resolvemcp.Bool(false)`), functions are always available. Set it only for functions that do not change data.
- `AllowedFunctions` works in either mode. When empty, every registered function is allowed. When set, only the named functions are listed and callable; any other is reported as `unknown function`, so its existence is not revealed. Per-function `Authorize` still applies on top.
## 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).
@@ -652,6 +733,7 @@ The handler resolves table names in priority order:
## Breaking changes
- The server is read-only by default. Writes (insert/update/delete), annotations and function calls need `Config{ReadOnly: resolvemcp.Bool(false)}` (function calls can also be kept on a read-only server with `AllowFunctionCalls`).
- Per-model tools (`read_/create_/update_/delete_{schema}_{entity}`) and per-model resources are gone; use the meta tools.
- `Setup*` / `NewSSEServer` / `NewStreamableHTTPHandler` take a `*security.SecurityList` and require authentication. `OptionalAuth*` helpers were removed; `*Unauthenticated` variants exist for explicit opt-out.
- `resolvespec_annotate` is opt-in via `Config.EnableAnnotations`.
+329
View File
@@ -0,0 +1,329 @@
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.`
// readOnlyGuide replaces usageGuide on a read-only server.
const readOnlyGuide = `This server exposes database tables through a fixed set of tools. It is READ-ONLY: you cannot insert, update or delete data or write annotations, and no tool for that exists. Do not attempt a write; tell the user it is not possible through this server.
1. Call list_tables to see the tables you may read and what they hold.
2. Call describe_table for a table before using it: columns, types, primary key, relations (preloadable) 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.
Read the error message when a call fails: it says which argument was wrong.`
// readOnlyFunctionsGuide is the extra step of a read-only server that still allows functions.
const readOnlyFunctionsGuide = `
4. Use list_functions / call_function for the registered functions. Only call functions that fit a read-only server; the server decides what is allowed.`
// guideFor returns the usage guide for the server mode.
func guideFor(readOnly, functions bool) string {
if !readOnly {
return usageGuide
}
if functions {
return readOnlyGuide + readOnlyFunctionsGuide
}
return readOnlyGuide
}
// 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"`
ReadOnly bool `json:"read_only"`
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,
ReadOnly: h.config.readOnly,
Guide: guideFor(h.config.readOnly, h.config.AllowFunctionCalls),
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 := h.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 !h.config.readOnly && 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))
if c.ReadOnly {
sb.WriteString("**This server is read-only.**\n\n")
}
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)
}
+126
View File
@@ -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"])
}
}
+50
View File
@@ -0,0 +1,50 @@
// 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.
//
// The server is read-only by default (Config.ReadOnly nil means on); set
// ReadOnly: resolvemcp.Bool(false) to enable the write tools.
//
// # 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
+13 -1
View File
@@ -108,6 +108,15 @@ func (h *Handler) function(name string) (Function, bool) {
return f, ok
}
// functionAllowed reports whether Config.AllowedFunctions lets the function through.
func (h *Handler) functionAllowed(name string) bool {
if h.allowedFns == nil {
return true
}
_, ok := h.allowedFns[name]
return ok
}
// visibleFunctions returns the functions the caller may call, sorted by name.
func (h *Handler) visibleFunctions(ctx context.Context) []Function {
h.functions.mu.RLock()
@@ -119,6 +128,9 @@ func (h *Handler) visibleFunctions(ctx context.Context) []Function {
sort.Slice(out, func(i, j int) bool { return out[i].Name < out[j].Name })
visible := out[:0]
for _, f := range out {
if !h.functionAllowed(f.Name) {
continue
}
if f.Authorize == nil || f.Authorize(ctx) == nil {
visible = append(visible, f)
}
@@ -236,7 +248,7 @@ func (h *Handler) executeCall(ctx context.Context, name string, rawArgs map[stri
defer cancel()
f, ok := h.function(name)
if !ok {
if !ok || !h.functionAllowed(name) {
return nil, invalidArg("unknown function %q", truncate(name))
}
hookCtx := &HookContext{Context: ctx, Handler: h, Entity: name, Operation: "call_function", Tx: h.db}
+9 -2
View File
@@ -24,6 +24,7 @@ import (
// Handler exposes registered database models as MCP tools and resources.
type Handler struct {
allowedFns map[string]struct{} // nil: every function is allowed
db common.Database
registry common.ModelRegistry
hooks *HookRegistry
@@ -43,14 +44,20 @@ 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(guideFor(cfg.withDefaults().readOnly, cfg.AllowFunctionCalls))),
config: cfg.withDefaults(),
confirms: newConfirmStore(),
name: "resolvemcp",
version: "1.0.0",
}
if len(cfg.AllowedFunctions) > 0 {
h.allowedFns = make(map[string]struct{}, len(cfg.AllowedFunctions))
for _, n := range cfg.AllowedFunctions {
h.allowedFns[n] = struct{}{}
}
}
registerMetaTools(h)
if cfg.EnableAnnotations {
if cfg.EnableAnnotations && !h.config.readOnly {
registerAnnotationTool(h)
}
return h
+37 -10
View File
@@ -56,6 +56,16 @@ func registerMetaTools(h *Handler) {
mcp.WithBoolean("include_count", mcp.Description("Also return the total number of matching rows (slower on large tables).")),
), h.handleSelect)
if !h.config.readOnly {
registerWriteTools(h, tableArg, idArg, filtersArg, dryRunArg, confirmArg)
}
if !h.config.readOnly || h.config.AllowFunctionCalls {
registerFunctionTools(h, readOnly)
}
}
// registerWriteTools adds the tools that change table rows.
func registerWriteTools(h *Handler, tableArg, idArg, filtersArg, dryRunArg, confirmArg mcp.ToolOption) {
h.mcpServer.AddTool(mcp.NewTool("insert_into_table",
mcp.WithDescription("Insert one row (object) or several rows (array, one transaction, capped). Unknown or read-only fields are rejected."),
tableArg, mcp.WithObject("data", mcp.Required(), mcp.Description("A row object or an array of row objects.")),
@@ -73,7 +83,10 @@ func registerMetaTools(h *Handler) {
mcp.WithDestructiveHintAnnotation(true),
tableArg, idArg, filtersArg, dryRunArg, confirmArg,
), h.handleDelete)
}
// registerFunctionTools adds list_functions and call_function.
func registerFunctionTools(h *Handler, readOnly mcp.ToolOption) {
h.mcpServer.AddTool(mcp.NewTool("list_functions", readOnly,
mcp.WithDescription("List the functions you can call with call_function, with their parameters.")),
h.handleListFunctions)
@@ -111,11 +124,15 @@ func (h *Handler) modelRules(schema, entity string) modelregistry.ModelRules {
return modelregistry.DefaultModelRules()
}
func opsFor(r modelregistry.ModelRules) []string {
// opsFor lists the operations the rules allow. A read-only server allows select only.
func (h *Handler) opsFor(r modelregistry.ModelRules) []string {
var ops []string
if r.CanRead {
ops = append(ops, opSelect)
}
if h.config.readOnly {
return ops
}
if r.CanCreate {
ops = append(ops, opInsert)
}
@@ -140,9 +157,12 @@ func (h *Handler) resolveTable(args map[string]any, op string) (schema, entity s
if _, err := h.registry.GetModelByEntity(schema, entity); err != nil {
return "", "", invalidArg("unknown table %q; see list_tables", truncate(table))
}
if op != "" && op != opSelect && h.config.readOnly {
return "", "", NewClientError(CodeForbidden, "this server is read-only: writes are disabled")
}
if op != "" {
allowed := false
for _, o := range opsFor(h.modelRules(schema, entity)) {
for _, o := range h.opsFor(h.modelRules(schema, entity)) {
if o == op {
allowed = true
}
@@ -156,14 +176,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})
if ops := h.opsFor(h.modelRules(schema, entity)); len(ops) > 0 {
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 })
@@ -180,17 +201,18 @@ func (h *Handler) handleDescribeTable(_ context.Context, req mcp.CallToolRequest
return toolError("describe_table", invalidArg("unknown table")), nil
}
rules := h.modelRules(schema, entity)
if len(opsFor(rules)) == 0 {
if len(h.opsFor(rules)) == 0 {
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) {
modelType = modelType.Elem()
}
writable := map[string]bool{}
if modelType != nil && modelType.Kind() == reflect.Struct {
if !h.config.readOnly && modelType != nil && modelType.Kind() == reflect.Struct {
for jsonKey := range reflection.BuildJSONToDBColumnMap(modelType) {
writable[jsonKey] = true
}
@@ -203,6 +225,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 +235,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,11 +243,15 @@ 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,
"writable_columns": writableNames,
"operations": opsFor(rules),
"operations": h.opsFor(rules),
"read_only": h.config.readOnly,
"filter_operators": filterOperators,
"limits": map[string]any{
"default_limit": h.config.DefaultLimit,
+168
View File
@@ -0,0 +1,168 @@
package resolvemcp
import (
"context"
"strings"
"testing"
"github.com/bitechdev/ResolveSpec/pkg/common"
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
"github.com/bitechdev/ResolveSpec/pkg/modelregistry"
)
func newReadOnlyHandler(t *testing.T) *Handler {
t.Helper()
h := NewHandler(database.NewPgSQLAdapter(nil), modelregistry.NewModelRegistry(),
Config{EnableAnnotations: true})
if err := h.RegisterModel("public", "items", &docItem{}); err != nil {
t.Fatal(err)
}
return h
}
func TestReadOnlyToolSet(t *testing.T) {
h := newReadOnlyHandler(t)
tools := h.mcpServer.ListTools()
for _, name := range []string{"list_tables", "describe_table", "select_table"} {
if tools[name] == nil {
t.Errorf("read tool %s missing", name)
}
}
for _, name := range []string{"insert_into_table", "update_table", "delete_from_table", "call_function", "list_functions", annotationToolName} {
if tools[name] != nil {
t.Errorf("tool %s must not be registered on a read-only server", name)
}
}
}
func TestReadOnlyRefusesWritesAndReportsIt(t *testing.T) {
h := newReadOnlyHandler(t)
ctx := context.Background()
args := map[string]any{"table": "public.items", "data": map[string]any{"name": "x"}, "id": 1}
for name, fn := range map[string]func() map[string]any{
"insert": func() map[string]any { r, _ := h.handleInsert(ctx, callReq(args)); return payload(t, r) },
"update": func() map[string]any { r, _ := h.handleUpdate(ctx, callReq(args)); return payload(t, r) },
"delete": func() map[string]any { r, _ := h.handleDelete(ctx, callReq(args)); return payload(t, r) },
} {
e, _ := fn()["error"].(map[string]any)
if e["code"] != CodeForbidden || !strings.Contains(e["message"].(string), "read-only") {
t.Errorf("%s: error = %v", name, e)
}
}
res, _ := h.handleListTables(ctx, callReq(nil))
tb := payload(t, res)["tables"].([]any)[0].(map[string]any)
if ops := tb["operations"].([]any); len(ops) != 1 || ops[0] != opSelect {
t.Errorf("list_tables operations = %v", ops)
}
res, _ = h.handleDescribeTable(ctx, callReq(map[string]any{"table": "public.items"}))
p := payload(t, res)
if p["read_only"] != true {
t.Errorf("describe_table read_only = %v", p["read_only"])
}
if w, _ := p["writable_columns"].([]any); len(w) != 0 {
t.Errorf("writable_columns = %v", w)
}
cat := h.BuildCatalog()
if !cat.ReadOnly || !strings.Contains(cat.Guide, "READ-ONLY") || !strings.Contains(cat.Markdown(), "read-only") {
t.Error("catalogue must say the server is read-only")
}
for _, c := range cat.Tables[0].Columns {
if c.Writable {
t.Errorf("column %s marked writable", c.Name)
}
}
if !strings.Contains(guideFor(true, false), "READ-ONLY") || strings.Contains(guideFor(false, false), "READ-ONLY") {
t.Error("guideFor")
}
}
func newFnHandler(t *testing.T, cfg Config) *Handler {
t.Helper()
h := NewHandler(database.NewPgSQLAdapter(nil), modelregistry.NewModelRegistry(), cfg)
for _, name := range []string{"alpha", "beta"} {
name := name
err := h.RegisterFunction(Function{Name: name, Handler: func(context.Context, common.Database, map[string]any) (any, error) {
return name, nil
}})
if err != nil {
t.Fatal(err)
}
}
return h
}
func TestReadOnlyAllowFunctionCalls(t *testing.T) {
h := newFnHandler(t, Config{AllowFunctionCalls: true})
tools := h.mcpServer.ListTools()
if tools["list_functions"] == nil || tools["call_function"] == nil {
t.Error("function tools must be registered")
}
if tools["insert_into_table"] != nil || tools["update_table"] != nil {
t.Error("write tools must stay off")
}
if g := guideFor(true, true); !strings.Contains(g, "READ-ONLY") || !strings.Contains(g, "call_function") {
t.Error("guide must mention functions")
}
if h.mcpServer.ListTools()["call_function"] == nil {
t.Error("call_function missing")
}
}
func TestAllowedFunctions(t *testing.T) {
ctx := context.Background()
for name, tc := range map[string]struct {
allowed []string
visible []string
}{
"empty allows all": {nil, []string{"alpha", "beta"}},
"only listed": {[]string{"beta"}, []string{"beta"}},
"unknown name": {[]string{"zzz"}, nil},
} {
h := newFnHandler(t, Config{AllowedFunctions: tc.allowed})
var got []string
for _, f := range h.visibleFunctions(ctx) {
got = append(got, f.Name)
}
if strings.Join(got, ",") != strings.Join(tc.visible, ",") {
t.Errorf("%s: visible = %v, want %v", name, got, tc.visible)
}
for _, fn := range []string{"alpha", "beta"} {
listed := false
for _, v := range tc.visible {
listed = listed || v == fn
}
if h.functionAllowed(fn) != listed {
t.Errorf("%s: functionAllowed(%s) = %v, want %v", name, fn, !listed, listed)
}
if !listed {
// refused before any database work, and indistinguishable from a missing function
if _, err := h.executeCall(ctx, fn, nil); err == nil || !strings.Contains(err.Error(), "unknown function") {
t.Errorf("%s: %s must be reported unknown, err=%v", name, fn, err)
}
}
}
}
}
func TestReadOnlyDefaultsOnAndCanBeDisabled(t *testing.T) {
if !(Config{}).withDefaults().readOnly {
t.Error("ReadOnly must default to on")
}
if !(Config{ReadOnly: Bool(true)}).withDefaults().readOnly {
t.Error("explicit true")
}
if (Config{ReadOnly: Bool(false)}).withDefaults().readOnly {
t.Error("Bool(false) must enable writes")
}
h := NewHandler(database.NewPgSQLAdapter(nil), modelregistry.NewModelRegistry(), Config{ReadOnly: Bool(false)})
if h.mcpServer.ListTools()["insert_into_table"] == nil {
t.Error("write tools must register when ReadOnly is Bool(false)")
}
if strings.Contains(h.BuildCatalog().Guide, "READ-ONLY") {
t.Error("guide must not claim read-only")
}
}
+26 -15
View File
@@ -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 (
@@ -66,6 +51,28 @@ type Config struct {
// host, with at most 32 distinct base URLs cached; prefer setting BaseURL.
AllowedHosts []string
// ReadOnly disables every write and is ON when left nil: set it to Bool(false) to allow
// writes. When on, the insert, update, delete and annotation tools are not registered,
// list_tables and describe_table report only the select operation (no writable columns),
// a write attempted anyway is refused with a "forbidden" error, and the server
// instructions tell the agent it cannot write. list_functions/call_function are also
// off, because a registered function may change data, unless AllowFunctionCalls is set.
ReadOnly *bool
// readOnly is ReadOnly after defaults (nil means true).
readOnly bool
// AllowFunctionCalls keeps list_functions and call_function available on a ReadOnly
// server. Only set it for functions that do not change data; pair it with
// AllowedFunctions to name them. It has no effect when writes are enabled (ReadOnly set to Bool(false)) (functions are
// always available then).
AllowFunctionCalls bool
// AllowedFunctions restricts list_functions and call_function to the named functions.
// Empty allows every registered function. A function outside the list is reported as
// unknown, so its existence is not revealed.
AllowedFunctions []string
// EnableAnnotations registers the resolvespec_annotate tool. Off by default: annotations
// are free text that agents read back, so enabling the tool opens a write channel into
// agent-visible text. When on, every call runs the BeforeHandle hooks (operation
@@ -73,6 +80,9 @@ type Config struct {
EnableAnnotations bool
}
// Bool returns a pointer to v, for the optional boolean fields of Config.
func Bool(v bool) *bool { return &v }
// withDefaults fills the zero limit fields.
func (c Config) withDefaults() Config {
def := func(v *int, d int) {
@@ -89,6 +99,7 @@ func (c Config) withDefaults() Config {
if c.DefaultLimit > c.MaxLimit {
c.DefaultLimit = c.MaxLimit
}
c.readOnly = c.ReadOnly == nil || *c.ReadOnly
if c.QueryTimeout <= 0 {
c.QueryTimeout = 30 * time.Second
}
+1 -1
View File
@@ -191,7 +191,7 @@ func TestAnnotationToolIsOptIn(t *testing.T) {
if h.mcpServer.GetTool(annotationToolName) != nil {
t.Fatal("annotation tool must be off by default")
}
on := NewHandler(h.db, modelregistry.NewModelRegistry(), Config{EnableAnnotations: true})
on := NewHandler(h.db, modelregistry.NewModelRegistry(), Config{EnableAnnotations: true, ReadOnly: Bool(false)})
if on.mcpServer.GetTool(annotationToolName) == nil {
t.Fatal("annotation tool missing when enabled")
}
+7
View File
@@ -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,
+1 -1
View File
@@ -29,7 +29,7 @@ func newTxHarness(t *testing.T) (*Handler, sqlmock.Sqlmock, context.Context) {
// connection and fails on the context timeout.
db.SetMaxOpenConns(1)
t.Cleanup(func() { _ = db.Close() })
h := NewHandler(database.NewPgSQLAdapter(db), modelregistry.NewModelRegistry(), Config{})
h := NewHandler(database.NewPgSQLAdapter(db), modelregistry.NewModelRegistry(), Config{ReadOnly: Bool(false)})
if err := h.RegisterModel("public", "items", &txItem{}); err != nil {
t.Fatal(err)
}