mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-07 22:06:28 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
5b6406ef38 | ||
|
|
d2f33b8f7d | ||
|
|
63cb0d22eb | ||
|
|
e4c4315f4b | ||
|
|
3efd539e0f | ||
|
|
5a3a1df3c8 | ||
|
|
4ed9506ad2 | ||
|
|
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
|
||||
|
||||
@@ -0,0 +1,82 @@
|
||||
# aiproxy
|
||||
|
||||
Authenticated reverse proxy for OpenAI-compatible APIs and MCP servers (streamable HTTP).
|
||||
Upstream URLs and keys stay server-side; clients only use the proxy with normal ResolveSpec auth.
|
||||
|
||||
## Setup
|
||||
```go
|
||||
p := aiproxy.New(aiproxy.Config{Prefix: "/ai"})
|
||||
p.RegisterOpenAI(aiproxy.OpenAIUpstream{Upstream: aiproxy.Upstream{Name: "gpt", BaseURL: "https://api.openai.com/v1", APIKey: key}})
|
||||
p.RegisterMCP(aiproxy.MCPUpstream{Upstream: aiproxy.Upstream{Name: "tools", BaseURL: "http://localhost:3000/mcp", APIKey: key}})
|
||||
mux.Handle("/ai/", p.Handler(security.NewAuthMiddleware(securityList)))
|
||||
```
|
||||
|
||||
## Routes
|
||||
| Kind | Client URL | Upstream URL |
|
||||
|---|---|---|
|
||||
| OpenAI | `{prefix}/{name}/v1/*` | `BaseURL` + path after `/v1` (BaseURL includes the version) |
|
||||
| MCP | `{prefix}/{name}` | `BaseURL` as is |
|
||||
|
||||
Client SDK base URL: `https://host/ai/gpt/v1` (OpenAI), `https://host/ai/tools` (MCP).
|
||||
Query string, `Mcp-Session-Id`, SSE streaming all pass through.
|
||||
|
||||
## Upstream fields
|
||||
| Field | Notes |
|
||||
|---|---|
|
||||
| `Name` | route segment, `[A-Za-z0-9_.-]`, unique across kinds |
|
||||
| `BaseURL` | http/https |
|
||||
| `APIKey` | optional; never logged/serialized (redacted in `String`/JSON) |
|
||||
| `AuthHeader` | default `Authorization` |
|
||||
| `AuthFormat` | one `%s`; default `Bearer %s` for `Authorization`, else raw key |
|
||||
| `Headers` | static extra upstream headers |
|
||||
| `AllowedRoles` | any-of; empty = any authenticated user |
|
||||
| `RateLimit` | `{PerSecond, Burst}` per user per upstream; 429 + `Retry-After` |
|
||||
| `Timeout` | response-header timeout, default 120s (streams not cut) |
|
||||
| `AllowedModels` (OpenAI) | JSON body `model` must match; non-JSON writes rejected |
|
||||
| `AllowedTools` (MCP) | `tools/call` name must match (batches checked); `tools/list` not filtered |
|
||||
|
||||
## Security behavior
|
||||
- Auth required; guests rejected (401). Roles checked (403).
|
||||
- Client `Authorization`, `Cookie`, `Proxy-Authorization`, `X-Api-Key` (+ `Config.StripHeaders`) removed; key injected.
|
||||
- Response `Set-Cookie`, `Www-Authenticate`, redirect `Location` removed.
|
||||
- Upstream 401 becomes 502 (proxy credentials problem, not the client's); upstream errors hide URL.
|
||||
- JSON bodies inspected up to `Config.MaxBodyBytes` (10 MB), else 413.
|
||||
- Fails closed if `Handler(nil)`.
|
||||
|
||||
## Hooks (`p.Hooks()`)
|
||||
| Hook | When | Context |
|
||||
|---|---|---|
|
||||
| `BeforeProxy` | after auth/limit/inspection | `Request` (headers editable), `UserContext`, `Upstream`, `Kind`, `Model`, `Tools`; abort via `Abort`/`AbortMessage`/`AbortCode` (default 403) |
|
||||
| `AfterProxy` | response finished/failed | `StatusCode`, `Duration`, `Usage`, `Error` |
|
||||
|
||||
`Usage` (OpenAI): from `usage` / `response.usage`; streams need `stream_options.include_usage`.
|
||||
|
||||
## Audit
|
||||
`Config.Audit` (`AuditSink`) gets one `AuditRecord` per handled request: user, upstream, method, client path (no query),
|
||||
model/tools, status, outcome (`ok|upstream_error|denied|rate_limited`), reason, duration, usage. No bodies, no keys.
|
||||
Default `LogAuditSink` (pkg/logger); `NopAuditSink{}` disables. Sinks run on the request path: keep them fast.
|
||||
|
||||
## Metrics
|
||||
Used automatically when the global `metrics.Provider` implements `metrics.AIProxyRecorder` (`PrometheusProvider` does):
|
||||
`aiproxy_requests_total{upstream,kind,model,status,outcome}`, `aiproxy_request_duration_seconds{upstream,kind}`,
|
||||
`aiproxy_tokens_total{upstream,model,type}`. `model` is client-supplied: bounded to 400 distinct pairs, then `other`.
|
||||
|
||||
## Dynamic upstreams
|
||||
| Call | Effect |
|
||||
|---|---|
|
||||
| `Register(Definition)` / `Unregister(name)` | runtime add/remove |
|
||||
| `Reload(ctx, store)` | sync store-managed upstreams: add, replace changed, drop missing; unchanged keep rate-limit state |
|
||||
| `AutoReload(ctx, store, interval)` | `Reload` now, then periodically |
|
||||
| `NewProcStore(*sql.DB, proc)` / `NewProcStoreFromDatabase(common.Database, proc)` | calls stored procedure (PostgreSQL), default `resolvespec_ai_proxies` |
|
||||
|
||||
- Code-registered upstreams are never replaced/removed by `Reload`; a same-name stored row is skipped and reported.
|
||||
- Invalid rows are skipped (valid ones applied; last good definition kept); store failure changes nothing.
|
||||
- Procedure contract (like the security procedures): `SELECT p_success, p_error, p_data FROM resolvespec_ai_proxies()`;
|
||||
`p_data` = JSON array of `{name, kind, base_url, api_key, auth_header, auth_format, headers, allowed_roles, allowed, rate_per_second, rate_burst, timeout_seconds, enabled}`.
|
||||
Reference table + function: `resolvespec_ai_proxies.sql` (replace the body to source from anywhere).
|
||||
- Entries with `enabled:false` or undecodable JSON are skipped.
|
||||
- API keys come back in plaintext from the procedure: restrict execute rights on it.
|
||||
|
||||
## Notes
|
||||
- Legacy MCP SSE transport is not supported.
|
||||
- `GET /v1/models` is not filtered by `AllowedModels`.
|
||||
@@ -0,0 +1,191 @@
|
||||
// Package aiproxy is an authenticated reverse proxy for OpenAI-compatible APIs and
|
||||
// MCP servers (streamable HTTP). Upstream URLs and keys stay on the server; clients
|
||||
// authenticate with the normal security layer and call the proxy instead.
|
||||
//
|
||||
// p := aiproxy.New(aiproxy.Config{Prefix: "/ai"})
|
||||
// p.RegisterOpenAI(aiproxy.OpenAIUpstream{Upstream: aiproxy.Upstream{Name: "gpt", BaseURL: "https://api.openai.com/v1", APIKey: key}})
|
||||
// p.RegisterMCP(aiproxy.MCPUpstream{Upstream: aiproxy.Upstream{Name: "tools", BaseURL: "http://localhost:3000/mcp", APIKey: key}})
|
||||
// mux.Handle("/ai/", p.Handler(security.NewAuthMiddleware(securityList)))
|
||||
//
|
||||
// Client routes: {prefix}/{name}/v1/... (OpenAI) and {prefix}/{name} (MCP).
|
||||
package aiproxy
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httputil"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultMaxBody = 10 << 20
|
||||
defaultTimeout = 120 * time.Second
|
||||
)
|
||||
|
||||
// Config configures a Proxy.
|
||||
type Config struct {
|
||||
// Prefix the handler is mounted under, e.g. "/ai". Empty when mounted with http.StripPrefix.
|
||||
Prefix string
|
||||
// MaxBodyBytes caps JSON request bodies that are inspected. Default 10 MB.
|
||||
MaxBodyBytes int64
|
||||
// StripHeaders are removed from client requests in addition to the defaults
|
||||
// (Authorization, Cookie, Proxy-Authorization, X-Api-Key).
|
||||
StripHeaders []string
|
||||
// Audit receives one record per request. Default: LogAuditSink. Use NopAuditSink to disable.
|
||||
Audit AuditSink
|
||||
}
|
||||
|
||||
// Definition describes one upstream independent of how it was registered.
|
||||
type Definition struct {
|
||||
Kind Kind
|
||||
Upstream Upstream
|
||||
// Allowed holds AllowedModels (OpenAI) or AllowedTools (MCP).
|
||||
Allowed []string
|
||||
}
|
||||
|
||||
// target is a registered upstream, ready to serve.
|
||||
type target struct {
|
||||
def Definition
|
||||
kind Kind
|
||||
up Upstream
|
||||
base *url.URL
|
||||
roles map[string]struct{}
|
||||
models map[string]struct{}
|
||||
tools map[string]struct{}
|
||||
proxy *httputil.ReverseProxy
|
||||
managed bool // loaded by Reload; code-registered upstreams are never touched by it
|
||||
|
||||
authName string
|
||||
authVal string
|
||||
}
|
||||
|
||||
// Proxy routes authenticated requests to registered upstreams.
|
||||
type Proxy struct {
|
||||
cfg Config
|
||||
hooks *HookRegistry
|
||||
limits *limiters
|
||||
labels modelLabels
|
||||
mu sync.RWMutex
|
||||
upstream map[string]*target
|
||||
}
|
||||
|
||||
// New creates a Proxy.
|
||||
func New(cfg Config) *Proxy {
|
||||
cfg.Prefix = strings.TrimRight(cfg.Prefix, "/")
|
||||
if cfg.Audit == nil {
|
||||
cfg.Audit = LogAuditSink{}
|
||||
}
|
||||
if cfg.MaxBodyBytes <= 0 {
|
||||
cfg.MaxBodyBytes = defaultMaxBody
|
||||
}
|
||||
return &Proxy{
|
||||
cfg: cfg,
|
||||
hooks: NewHookRegistry(),
|
||||
limits: newLimiters(),
|
||||
upstream: make(map[string]*target),
|
||||
}
|
||||
}
|
||||
|
||||
// Hooks returns the hook registry.
|
||||
func (p *Proxy) Hooks() *HookRegistry { return p.hooks }
|
||||
|
||||
// RegisterOpenAI registers an OpenAI-compatible upstream.
|
||||
func (p *Proxy) RegisterOpenAI(u OpenAIUpstream) error {
|
||||
return p.Register(Definition{Kind: KindOpenAI, Upstream: u.Upstream, Allowed: u.AllowedModels})
|
||||
}
|
||||
|
||||
// RegisterMCP registers an MCP server (streamable HTTP).
|
||||
func (p *Proxy) RegisterMCP(u MCPUpstream) error {
|
||||
return p.Register(Definition{Kind: KindMCP, Upstream: u.Upstream, Allowed: u.AllowedTools})
|
||||
}
|
||||
|
||||
// Register registers an upstream from a Definition. The name must be unused.
|
||||
func (p *Proxy) Register(d Definition) error {
|
||||
t, err := p.build(d)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
if _, dup := p.upstream[t.up.Name]; dup {
|
||||
return fmt.Errorf("aiproxy: upstream %q already registered", t.up.Name)
|
||||
}
|
||||
p.upstream[t.up.Name] = t
|
||||
return nil
|
||||
}
|
||||
|
||||
// Unregister removes an upstream (code-registered or loaded). It reports whether it existed.
|
||||
// In-flight requests finish normally.
|
||||
func (p *Proxy) Unregister(name string) bool {
|
||||
p.mu.Lock()
|
||||
t, ok := p.upstream[name]
|
||||
delete(p.upstream, name)
|
||||
p.mu.Unlock()
|
||||
if ok {
|
||||
p.retire(t)
|
||||
}
|
||||
return ok
|
||||
}
|
||||
|
||||
// retire releases what a removed or replaced target held.
|
||||
func (p *Proxy) retire(t *target) {
|
||||
p.limits.forget(t.up.Name)
|
||||
if tr, ok := t.proxy.Transport.(*http.Transport); ok {
|
||||
tr.CloseIdleConnections()
|
||||
}
|
||||
}
|
||||
|
||||
func (p *Proxy) build(d Definition) (*target, error) {
|
||||
if d.Kind != KindOpenAI && d.Kind != KindMCP {
|
||||
return nil, fmt.Errorf("aiproxy: upstream %q has unknown kind %q", d.Upstream.Name, d.Kind)
|
||||
}
|
||||
base, err := d.Upstream.validate()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
t := &target{def: d, kind: d.Kind, up: d.Upstream, base: base, roles: toSet(d.Upstream.AllowedRoles)}
|
||||
if d.Kind == KindOpenAI {
|
||||
t.models = toSet(d.Allowed)
|
||||
} else {
|
||||
t.tools = toSet(d.Allowed)
|
||||
}
|
||||
t.authName, t.authVal = d.Upstream.authValue()
|
||||
t.proxy = p.buildProxy(t)
|
||||
return t, nil
|
||||
}
|
||||
|
||||
func (p *Proxy) get(name string) *target {
|
||||
p.mu.RLock()
|
||||
defer p.mu.RUnlock()
|
||||
return p.upstream[name]
|
||||
}
|
||||
|
||||
// Names lists registered upstream names (no URLs or keys).
|
||||
func (p *Proxy) Names() []string {
|
||||
p.mu.RLock()
|
||||
defer p.mu.RUnlock()
|
||||
out := make([]string, 0, len(p.upstream))
|
||||
for n := range p.upstream {
|
||||
out = append(out, n)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func newTransport(timeout time.Duration) *http.Transport {
|
||||
if timeout <= 0 {
|
||||
timeout = defaultTimeout
|
||||
}
|
||||
return &http.Transport{
|
||||
Proxy: http.ProxyFromEnvironment,
|
||||
DialContext: (&net.Dialer{Timeout: 10 * time.Second, KeepAlive: 30 * time.Second}).DialContext,
|
||||
ForceAttemptHTTP2: true,
|
||||
MaxIdleConnsPerHost: 32,
|
||||
IdleConnTimeout: 90 * time.Second,
|
||||
TLSHandshakeTimeout: 10 * time.Second,
|
||||
ResponseHeaderTimeout: timeout,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,329 @@
|
||||
package aiproxy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// testAuth authenticates by the X-Test-User header: "name:role1,role2". Empty means no user.
|
||||
func testAuth(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
h := r.Header.Get("X-Test-User")
|
||||
if h == "" {
|
||||
http.Error(w, "unauthenticated", http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
name, roles, _ := strings.Cut(h, ":")
|
||||
uc := &security.UserContext{UserID: len(name), UserName: name, Roles: strings.Split(roles, ",")}
|
||||
next.ServeHTTP(w, r.WithContext(context.WithValue(r.Context(), security.UserContextKey, uc)))
|
||||
})
|
||||
}
|
||||
|
||||
type captured struct {
|
||||
mu sync.Mutex
|
||||
req *http.Request
|
||||
body string
|
||||
}
|
||||
|
||||
func (c *captured) set(r *http.Request) {
|
||||
b, _ := io.ReadAll(r.Body)
|
||||
c.mu.Lock()
|
||||
c.req, c.body = r.Clone(context.Background()), string(b)
|
||||
c.mu.Unlock()
|
||||
}
|
||||
|
||||
func newProxy(t *testing.T, upstream http.Handler, mod func(*Proxy, string)) (*Proxy, *httptest.Server) {
|
||||
t.Helper()
|
||||
up := httptest.NewServer(upstream)
|
||||
t.Cleanup(up.Close)
|
||||
p := New(Config{Prefix: "/ai"})
|
||||
mod(p, up.URL)
|
||||
front := httptest.NewServer(p.Handler(testAuth))
|
||||
t.Cleanup(front.Close)
|
||||
return p, front
|
||||
}
|
||||
|
||||
func do(t *testing.T, method, url, body, user string, hdr map[string]string) (*http.Response, string) {
|
||||
t.Helper()
|
||||
req, err := http.NewRequest(method, url, strings.NewReader(body))
|
||||
require.NoError(t, err)
|
||||
if body != "" {
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
}
|
||||
if user != "" {
|
||||
req.Header.Set("X-Test-User", user)
|
||||
}
|
||||
for k, v := range hdr {
|
||||
req.Header.Set(k, v)
|
||||
}
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
require.NoError(t, err)
|
||||
defer resp.Body.Close()
|
||||
b, _ := io.ReadAll(resp.Body)
|
||||
return resp, string(b)
|
||||
}
|
||||
|
||||
func TestOpenAIKeyInjectionAndHiding(t *testing.T) {
|
||||
var c captured
|
||||
_, front := newProxy(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
c.set(r)
|
||||
w.Header().Set("Set-Cookie", "s=1")
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
io.WriteString(w, `{"ok":true}`)
|
||||
}), func(p *Proxy, u string) {
|
||||
require.NoError(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "gpt", BaseURL: u + "/v1", APIKey: "sk-secret"}}))
|
||||
})
|
||||
|
||||
resp, body := do(t, "POST", front.URL+"/ai/gpt/v1/chat/completions?x=1", `{"model":"m"}`, "bob:user",
|
||||
map[string]string{"Authorization": "Bearer client-token", "Cookie": "a=b"})
|
||||
assert.Equal(t, 200, resp.StatusCode)
|
||||
assert.Equal(t, `{"ok":true}`, body)
|
||||
assert.Empty(t, resp.Header.Get("Set-Cookie"))
|
||||
assert.Equal(t, "/v1/chat/completions", c.req.URL.Path)
|
||||
assert.Equal(t, "x=1", c.req.URL.RawQuery)
|
||||
assert.Equal(t, "Bearer sk-secret", c.req.Header.Get("Authorization"))
|
||||
assert.Empty(t, c.req.Header.Get("Cookie"))
|
||||
assert.Equal(t, `{"model":"m"}`, c.body)
|
||||
assert.NotContains(t, body, "sk-secret")
|
||||
}
|
||||
|
||||
func TestCustomAuthHeaderAndFormat(t *testing.T) {
|
||||
var c captured
|
||||
_, front := newProxy(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { c.set(r); io.WriteString(w, "{}") }),
|
||||
func(p *Proxy, u string) {
|
||||
require.NoError(t, p.RegisterMCP(MCPUpstream{Upstream: Upstream{Name: "a", BaseURL: u + "/mcp", APIKey: "k1", AuthHeader: "X-Token", Headers: map[string]string{"X-Tenant": "t1"}}}))
|
||||
require.NoError(t, p.RegisterMCP(MCPUpstream{Upstream: Upstream{Name: "b", BaseURL: u + "/mcp", APIKey: "k2", AuthFormat: "Token %s"}}))
|
||||
})
|
||||
|
||||
do(t, "POST", front.URL+"/ai/a", `{"jsonrpc":"2.0","id":1,"method":"ping"}`, "bob:user", map[string]string{"X-Token": "client"})
|
||||
assert.Equal(t, "k1", c.req.Header.Get("X-Token"))
|
||||
assert.Equal(t, "t1", c.req.Header.Get("X-Tenant"))
|
||||
assert.Equal(t, "/mcp", c.req.URL.Path)
|
||||
|
||||
do(t, "POST", front.URL+"/ai/b", `{"jsonrpc":"2.0","id":1,"method":"ping"}`, "bob:user", nil)
|
||||
assert.Equal(t, "Token k2", c.req.Header.Get("Authorization"))
|
||||
}
|
||||
|
||||
func TestAuthAndRoles(t *testing.T) {
|
||||
_, front := newProxy(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { io.WriteString(w, "{}") }),
|
||||
func(p *Proxy, u string) {
|
||||
require.NoError(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "gpt", BaseURL: u + "/v1", AllowedRoles: []string{"ai"}}}))
|
||||
})
|
||||
url := front.URL + "/ai/gpt/v1/models"
|
||||
|
||||
resp, _ := do(t, "GET", url, "", "", nil)
|
||||
assert.Equal(t, 401, resp.StatusCode)
|
||||
resp, _ = do(t, "GET", url, "", "bob:user", nil)
|
||||
assert.Equal(t, 403, resp.StatusCode)
|
||||
resp, _ = do(t, "GET", url, "", "bob:guest", nil)
|
||||
assert.Equal(t, 401, resp.StatusCode)
|
||||
resp, _ = do(t, "GET", url, "", "bob:user,ai", nil)
|
||||
assert.Equal(t, 200, resp.StatusCode)
|
||||
|
||||
resp, _ = do(t, "GET", front.URL+"/ai/nope/v1/models", "", "bob:ai", nil)
|
||||
assert.Equal(t, 404, resp.StatusCode)
|
||||
resp, _ = do(t, "GET", front.URL+"/ai/gpt/other", "", "bob:ai", nil)
|
||||
assert.Equal(t, 404, resp.StatusCode)
|
||||
}
|
||||
|
||||
func TestNilAuthRefuses(t *testing.T) {
|
||||
rec := httptest.NewRecorder()
|
||||
New(Config{}).Handler(nil).ServeHTTP(rec, httptest.NewRequest("GET", "/", nil))
|
||||
assert.Equal(t, 500, rec.Code)
|
||||
}
|
||||
|
||||
func TestModelAllowlist(t *testing.T) {
|
||||
_, front := newProxy(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { io.WriteString(w, "{}") }),
|
||||
func(p *Proxy, u string) {
|
||||
require.NoError(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "gpt", BaseURL: u + "/v1"}, AllowedModels: []string{"good"}}))
|
||||
})
|
||||
url := front.URL + "/ai/gpt/v1/chat/completions"
|
||||
resp, _ := do(t, "POST", url, `{"model":"good"}`, "bob:u", nil)
|
||||
assert.Equal(t, 200, resp.StatusCode)
|
||||
resp, _ = do(t, "POST", url, `{"model":"bad"}`, "bob:u", nil)
|
||||
assert.Equal(t, 403, resp.StatusCode)
|
||||
resp, _ = do(t, "POST", url, `{}`, "bob:u", nil)
|
||||
assert.Equal(t, 403, resp.StatusCode)
|
||||
resp, _ = do(t, "POST", url, `x`, "bob:u", map[string]string{"Content-Type": "text/plain"})
|
||||
assert.Equal(t, 400, resp.StatusCode)
|
||||
resp, _ = do(t, "GET", front.URL+"/ai/gpt/v1/models", "", "bob:u", nil)
|
||||
assert.Equal(t, 200, resp.StatusCode)
|
||||
}
|
||||
|
||||
func TestBodyTooLarge(t *testing.T) {
|
||||
up := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {}))
|
||||
defer up.Close()
|
||||
p := New(Config{Prefix: "/ai", MaxBodyBytes: 10})
|
||||
require.NoError(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "g", BaseURL: up.URL + "/v1"}}))
|
||||
front := httptest.NewServer(p.Handler(testAuth))
|
||||
defer front.Close()
|
||||
resp, _ := do(t, "POST", front.URL+"/ai/g/v1/x", `{"model":"aaaaaaaaaaaaaaaa"}`, "bob:u", nil)
|
||||
assert.Equal(t, 413, resp.StatusCode)
|
||||
}
|
||||
|
||||
func TestMCPToolAllowlist(t *testing.T) {
|
||||
var c captured
|
||||
_, front := newProxy(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
c.set(r)
|
||||
w.Header().Set("Mcp-Session-Id", "sess1")
|
||||
io.WriteString(w, "{}")
|
||||
}), func(p *Proxy, u string) {
|
||||
require.NoError(t, p.RegisterMCP(MCPUpstream{Upstream: Upstream{Name: "t", BaseURL: u + "/mcp"}, AllowedTools: []string{"ok"}}))
|
||||
})
|
||||
url := front.URL + "/ai/t"
|
||||
call := func(name string) string {
|
||||
return `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"` + name + `"}}`
|
||||
}
|
||||
|
||||
resp, _ := do(t, "POST", url, call("ok"), "bob:u", nil)
|
||||
assert.Equal(t, 200, resp.StatusCode)
|
||||
assert.Equal(t, "sess1", resp.Header.Get("Mcp-Session-Id"))
|
||||
resp, _ = do(t, "POST", url, call("evil"), "bob:u", nil)
|
||||
assert.Equal(t, 403, resp.StatusCode)
|
||||
resp, _ = do(t, "POST", url, `[`+call("ok")+`,`+call("evil")+`]`, "bob:u", nil)
|
||||
assert.Equal(t, 403, resp.StatusCode)
|
||||
resp, _ = do(t, "POST", url, `{"jsonrpc":"2.0","id":2,"method":"tools/list"}`, "bob:u", nil)
|
||||
assert.Equal(t, 200, resp.StatusCode)
|
||||
resp, _ = do(t, "DELETE", url, "", "bob:u", map[string]string{"Mcp-Session-Id": "sess1"})
|
||||
assert.Equal(t, 200, resp.StatusCode)
|
||||
assert.Equal(t, "sess1", c.req.Header.Get("Mcp-Session-Id"))
|
||||
}
|
||||
|
||||
func TestRateLimit(t *testing.T) {
|
||||
_, front := newProxy(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { io.WriteString(w, "{}") }),
|
||||
func(p *Proxy, u string) {
|
||||
require.NoError(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "g", BaseURL: u + "/v1", RateLimit: &RateLimit{PerSecond: 0.1, Burst: 2}}}))
|
||||
})
|
||||
url := front.URL + "/ai/g/v1/models"
|
||||
for i := 0; i < 2; i++ {
|
||||
resp, _ := do(t, "GET", url, "", "bob:u", nil)
|
||||
assert.Equal(t, 200, resp.StatusCode)
|
||||
}
|
||||
resp, _ := do(t, "GET", url, "", "bob:u", nil)
|
||||
assert.Equal(t, 429, resp.StatusCode)
|
||||
assert.NotEmpty(t, resp.Header.Get("Retry-After"))
|
||||
// other user has own bucket
|
||||
resp, _ = do(t, "GET", url, "", "alice:u", nil)
|
||||
assert.Equal(t, 200, resp.StatusCode)
|
||||
}
|
||||
|
||||
func TestHooksAndUsage(t *testing.T) {
|
||||
var p *Proxy
|
||||
_, front := newProxy(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Header.Get("X-From-Hook") == "" {
|
||||
t.Error("hook header missing")
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
io.WriteString(w, `{"usage":{"prompt_tokens":3,"completion_tokens":4,"total_tokens":7}}`)
|
||||
}), func(pp *Proxy, u string) {
|
||||
p = pp
|
||||
require.NoError(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "g", BaseURL: u + "/v1"}}))
|
||||
})
|
||||
done := make(chan *HookContext, 1)
|
||||
p.Hooks().Register(BeforeProxy, func(h *HookContext) error {
|
||||
h.Request.Header.Set("X-From-Hook", "1")
|
||||
if h.Model == "blocked" {
|
||||
h.Abort, h.AbortMessage, h.AbortCode = true, "nope", 418
|
||||
}
|
||||
return nil
|
||||
})
|
||||
p.Hooks().Register(AfterProxy, func(h *HookContext) error { done <- h; return nil })
|
||||
|
||||
resp, _ := do(t, "POST", front.URL+"/ai/g/v1/chat/completions", `{"model":"blocked"}`, "bob:u", nil)
|
||||
assert.Equal(t, 418, resp.StatusCode)
|
||||
|
||||
resp, _ = do(t, "POST", front.URL+"/ai/g/v1/chat/completions", `{"model":"m"}`, "bob:u", nil)
|
||||
assert.Equal(t, 200, resp.StatusCode)
|
||||
select {
|
||||
case h := <-done:
|
||||
assert.Equal(t, "m", h.Model)
|
||||
assert.Equal(t, "bob", h.UserContext.UserName)
|
||||
assert.Equal(t, 200, h.StatusCode)
|
||||
assert.Equal(t, Usage{3, 4, 7}, h.Usage)
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Fatal("AfterProxy not called")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSSEUsageAndStreaming(t *testing.T) {
|
||||
var p *Proxy
|
||||
_, front := newProxy(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "text/event-stream")
|
||||
f := w.(http.Flusher)
|
||||
io.WriteString(w, "data: {\"choices\":[{}]}\n\n")
|
||||
f.Flush()
|
||||
io.WriteString(w, "data: {\"choices\":[],\"usage\":{\"prompt_tokens\":1,\"completion_tokens\":2,\"total_tokens\":3}}\n\ndata: [DONE]\n\n")
|
||||
f.Flush()
|
||||
}), func(pp *Proxy, u string) {
|
||||
p = pp
|
||||
require.NoError(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "g", BaseURL: u + "/v1"}}))
|
||||
})
|
||||
done := make(chan *HookContext, 1)
|
||||
p.Hooks().Register(AfterProxy, func(h *HookContext) error { done <- h; return nil })
|
||||
|
||||
resp, body := do(t, "POST", front.URL+"/ai/g/v1/chat/completions", `{"model":"m","stream":true}`, "bob:u", nil)
|
||||
assert.Equal(t, 200, resp.StatusCode)
|
||||
assert.Contains(t, body, "[DONE]")
|
||||
h := <-done
|
||||
assert.Equal(t, Usage{1, 2, 3}, h.Usage)
|
||||
}
|
||||
|
||||
func TestResponsesAPIUsage(t *testing.T) {
|
||||
u, ok := parseUsage([]byte(`{"type":"response.completed","response":{"usage":{"input_tokens":5,"output_tokens":6}}}`))
|
||||
assert.True(t, ok)
|
||||
assert.Equal(t, Usage{5, 6, 11}, u)
|
||||
}
|
||||
|
||||
func TestUpstreamUnauthorizedIsMasked(t *testing.T) {
|
||||
_, front := newProxy(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Www-Authenticate", "Bearer realm=x")
|
||||
w.WriteHeader(401)
|
||||
}), func(p *Proxy, u string) {
|
||||
require.NoError(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "g", BaseURL: u + "/v1", APIKey: "bad"}}))
|
||||
})
|
||||
resp, body := do(t, "GET", front.URL+"/ai/g/v1/models", "", "bob:u", nil)
|
||||
assert.Equal(t, 502, resp.StatusCode)
|
||||
assert.Empty(t, resp.Header.Get("Www-Authenticate"))
|
||||
assert.Contains(t, body, "bad_gateway")
|
||||
}
|
||||
|
||||
func TestUpstreamDownHidesURL(t *testing.T) {
|
||||
up := httptest.NewServer(http.NotFoundHandler())
|
||||
url := up.URL
|
||||
up.Close()
|
||||
p := New(Config{Prefix: "/ai"})
|
||||
require.NoError(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "g", BaseURL: url + "/v1"}}))
|
||||
front := httptest.NewServer(p.Handler(testAuth))
|
||||
defer front.Close()
|
||||
resp, body := do(t, "GET", front.URL+"/ai/g/v1/models", "", "bob:u", nil)
|
||||
assert.Equal(t, 502, resp.StatusCode)
|
||||
assert.NotContains(t, body, url)
|
||||
}
|
||||
|
||||
func TestRegisterValidationAndRedaction(t *testing.T) {
|
||||
p := New(Config{})
|
||||
assert.Error(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "bad name", BaseURL: "http://x"}}))
|
||||
assert.Error(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "a", BaseURL: "ftp://x"}}))
|
||||
assert.Error(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "a", BaseURL: "http://x", AuthFormat: "%d"}}))
|
||||
assert.Error(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "a", BaseURL: "http://x", RateLimit: &RateLimit{}}}))
|
||||
require.NoError(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "a", BaseURL: "http://x", APIKey: "sk-secret"}}))
|
||||
assert.Error(t, p.RegisterMCP(MCPUpstream{Upstream: Upstream{Name: "a", BaseURL: "http://x"}}))
|
||||
|
||||
u := Upstream{Name: "a", BaseURL: "http://x", APIKey: "sk-secret"}
|
||||
js, _ := json.Marshal(u)
|
||||
for _, s := range []string{u.String(), string(js), strings.TrimSpace(errors.New(u.String()).Error())} {
|
||||
assert.NotContains(t, s, "sk-secret")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,316 @@
|
||||
package aiproxy
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"mime"
|
||||
"net/http"
|
||||
"net/http/httputil"
|
||||
"path"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/tidwall/gjson"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||
)
|
||||
|
||||
type ctxKey struct{}
|
||||
|
||||
// call is the per-request state shared between the handler and the proxy callbacks.
|
||||
type call struct {
|
||||
p *Proxy
|
||||
t *target
|
||||
hc *HookContext
|
||||
start time.Time
|
||||
once sync.Once
|
||||
}
|
||||
|
||||
func (c *call) finish(status int, usage Usage, err error) {
|
||||
c.once.Do(func() {
|
||||
c.hc.StatusCode = status
|
||||
c.hc.Usage = usage
|
||||
c.hc.Error = err
|
||||
c.hc.Duration = time.Since(c.start)
|
||||
c.p.hooks.executeAfter(c.hc)
|
||||
outcome := OutcomeOK
|
||||
if err != nil || status >= 400 {
|
||||
outcome = OutcomeUpstreamError
|
||||
}
|
||||
c.p.record(c.t, c.hc, status, outcome, "", err)
|
||||
})
|
||||
}
|
||||
|
||||
var defaultStrip = []string{"Authorization", "Cookie", "Proxy-Authorization", "X-Api-Key"}
|
||||
|
||||
// Handler returns the proxy handler. auth is the authentication middleware, normally
|
||||
// security.NewAuthMiddleware(securityList). It is required: without it every request is refused.
|
||||
func (p *Proxy) Handler(auth func(http.Handler) http.Handler) http.Handler {
|
||||
if auth == nil {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
writeError(w, http.StatusInternalServerError, "server_error", "authentication is not configured")
|
||||
})
|
||||
}
|
||||
return auth(http.HandlerFunc(p.serve))
|
||||
}
|
||||
|
||||
func (p *Proxy) serve(w http.ResponseWriter, r *http.Request) {
|
||||
// route: {prefix}/{name}/{rest}
|
||||
rel := strings.TrimPrefix(r.URL.Path, p.cfg.Prefix)
|
||||
if rel == r.URL.Path && p.cfg.Prefix != "" {
|
||||
writeError(w, http.StatusNotFound, "not_found", "not found")
|
||||
return
|
||||
}
|
||||
rel = strings.TrimPrefix(rel, "/")
|
||||
name, rest, _ := strings.Cut(rel, "/")
|
||||
t := p.get(name)
|
||||
if t == nil {
|
||||
writeError(w, http.StatusNotFound, "not_found", "unknown upstream")
|
||||
return
|
||||
}
|
||||
rest = "/" + rest
|
||||
if rest != "/" {
|
||||
rest = path.Clean(rest)
|
||||
}
|
||||
if t.kind == KindOpenAI {
|
||||
switch {
|
||||
case rest == "/v1":
|
||||
rest = ""
|
||||
case strings.HasPrefix(rest, "/v1/"):
|
||||
rest = strings.TrimPrefix(rest, "/v1")
|
||||
default:
|
||||
writeError(w, http.StatusNotFound, "not_found", "unknown path")
|
||||
return
|
||||
}
|
||||
} else if rest == "/" {
|
||||
rest = ""
|
||||
}
|
||||
|
||||
start := time.Now()
|
||||
user, _ := security.GetUserContext(r.Context())
|
||||
hc := &HookContext{
|
||||
Context: r.Context(),
|
||||
Request: r,
|
||||
UserContext: user,
|
||||
Upstream: t.up.Name,
|
||||
Kind: t.kind,
|
||||
}
|
||||
deny := func(code int, outcome Outcome, typ, msg string) {
|
||||
hc.Duration = time.Since(start)
|
||||
p.record(t, hc, code, outcome, msg, nil)
|
||||
writeError(w, code, typ, msg)
|
||||
}
|
||||
|
||||
// authenticated user, no guests
|
||||
if user == nil || hasRole(user.Roles, "guest") {
|
||||
deny(http.StatusUnauthorized, OutcomeDenied, "invalid_api_key", "authentication required")
|
||||
return
|
||||
}
|
||||
if len(t.roles) > 0 && !anyRole(user.Roles, t.roles) {
|
||||
deny(http.StatusForbidden, OutcomeDenied, "forbidden", "not allowed to use this upstream")
|
||||
return
|
||||
}
|
||||
if rl := t.up.RateLimit; rl != nil {
|
||||
key := t.up.Name + "|" + strconv.Itoa(user.UserID) + "|" + user.UserName
|
||||
if allowed, wait := p.limits.allow(key, rl); !allowed {
|
||||
w.Header().Set("Retry-After", strconv.Itoa(int(wait.Seconds())+1))
|
||||
deny(http.StatusTooManyRequests, OutcomeRateLimited, "rate_limit_exceeded", "rate limit exceeded")
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if code, typ, msg := p.inspect(t, r, hc); code != 0 {
|
||||
deny(code, OutcomeDenied, typ, msg)
|
||||
return
|
||||
}
|
||||
|
||||
if err := p.hooks.Execute(BeforeProxy, hc); err != nil {
|
||||
code := hc.AbortCode
|
||||
if code == 0 {
|
||||
code = http.StatusForbidden
|
||||
}
|
||||
msg := hc.AbortMessage
|
||||
if msg == "" {
|
||||
msg = "request rejected"
|
||||
}
|
||||
if !hc.Abort {
|
||||
logger.Error("aiproxy: %v", err)
|
||||
code, msg = http.StatusInternalServerError, "request failed"
|
||||
}
|
||||
deny(code, OutcomeDenied, "rejected", msg)
|
||||
return
|
||||
}
|
||||
|
||||
c := &call{p: p, t: t, hc: hc, start: start}
|
||||
ctx := context.WithValue(r.Context(), ctxKey{}, &routed{call: c, rest: rest})
|
||||
t.proxy.ServeHTTP(w, hc.Request.WithContext(ctx))
|
||||
}
|
||||
|
||||
// routed is stored on the request context for the ReverseProxy callbacks.
|
||||
type routed struct {
|
||||
call *call
|
||||
rest string
|
||||
}
|
||||
|
||||
func (p *Proxy) buildProxy(t *target) *httputil.ReverseProxy {
|
||||
strip := append(append([]string{}, defaultStrip...), p.cfg.StripHeaders...)
|
||||
if t.authName != "" {
|
||||
strip = append(strip, t.authName)
|
||||
}
|
||||
return &httputil.ReverseProxy{
|
||||
Transport: newTransport(t.up.Timeout),
|
||||
FlushInterval: -1,
|
||||
Rewrite: func(pr *httputil.ProxyRequest) {
|
||||
rt, _ := pr.In.Context().Value(ctxKey{}).(*routed)
|
||||
out := pr.Out
|
||||
u := *t.base
|
||||
u.Path = t.base.Path
|
||||
if rt != nil {
|
||||
u.Path += rt.rest
|
||||
}
|
||||
u.RawPath = ""
|
||||
u.RawQuery = pr.In.URL.RawQuery
|
||||
out.URL = &u
|
||||
out.Host = u.Host
|
||||
for _, h := range strip {
|
||||
out.Header.Del(h)
|
||||
}
|
||||
out.Header.Del("Accept-Encoding") // let the transport decompress so usage can be read
|
||||
for k, v := range t.up.Headers {
|
||||
out.Header.Set(k, v)
|
||||
}
|
||||
if t.authName != "" {
|
||||
out.Header.Set(t.authName, t.authVal)
|
||||
}
|
||||
},
|
||||
ModifyResponse: func(resp *http.Response) error {
|
||||
rt, _ := resp.Request.Context().Value(ctxKey{}).(*routed)
|
||||
h := resp.Header
|
||||
h.Del("Set-Cookie")
|
||||
h.Del("Www-Authenticate")
|
||||
if resp.StatusCode >= 300 && resp.StatusCode < 400 {
|
||||
h.Del("Location")
|
||||
}
|
||||
if resp.StatusCode == http.StatusUnauthorized {
|
||||
logger.Warn("aiproxy: upstream %q rejected its credentials", t.up.Name)
|
||||
_ = resp.Body.Close()
|
||||
body, _ := json.Marshal(errorBody("bad_gateway", "upstream rejected the proxy credentials"))
|
||||
resp.StatusCode, resp.Status = http.StatusBadGateway, "502 Bad Gateway"
|
||||
resp.Body = io.NopCloser(bytes.NewReader(body))
|
||||
resp.ContentLength = int64(len(body))
|
||||
h.Set("Content-Type", "application/json")
|
||||
h.Set("Content-Length", strconv.Itoa(len(body)))
|
||||
}
|
||||
if rt != nil {
|
||||
sse := strings.HasPrefix(h.Get("Content-Type"), "text/event-stream")
|
||||
resp.Body = newWatchBody(resp.Body, t.kind == KindOpenAI, sse, resp.StatusCode, rt.call)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
ErrorHandler: func(w http.ResponseWriter, r *http.Request, err error) {
|
||||
if rt, _ := r.Context().Value(ctxKey{}).(*routed); rt != nil {
|
||||
rt.call.finish(http.StatusBadGateway, Usage{}, err)
|
||||
}
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
logger.Warn("aiproxy: upstream %q error: %v", t.up.Name, err)
|
||||
}
|
||||
writeError(w, http.StatusBadGateway, "bad_gateway", "upstream request failed")
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// inspect reads JSON request bodies to extract/enforce model and tools.
|
||||
// It returns a non-zero status when the request must be rejected.
|
||||
func (p *Proxy) inspect(t *target, r *http.Request, hc *HookContext) (status int, errType, msg string) {
|
||||
switch r.Method {
|
||||
case http.MethodPost, http.MethodPut, http.MethodPatch:
|
||||
default:
|
||||
return 0, "", ""
|
||||
}
|
||||
if r.Body == nil || r.Body == http.NoBody {
|
||||
return 0, "", ""
|
||||
}
|
||||
mt, _, _ := mime.ParseMediaType(r.Header.Get("Content-Type"))
|
||||
isJSON := mt == "application/json" || strings.HasSuffix(mt, "+json")
|
||||
|
||||
if !isJSON {
|
||||
if t.kind == KindOpenAI && len(t.models) > 0 {
|
||||
return http.StatusBadRequest, "invalid_request_error", "model cannot be verified for this content type"
|
||||
}
|
||||
return 0, "", ""
|
||||
}
|
||||
|
||||
body, err := io.ReadAll(io.LimitReader(r.Body, p.cfg.MaxBodyBytes+1))
|
||||
_ = r.Body.Close()
|
||||
if err != nil {
|
||||
return http.StatusBadRequest, "invalid_request_error", "could not read request body"
|
||||
}
|
||||
if int64(len(body)) > p.cfg.MaxBodyBytes {
|
||||
return http.StatusRequestEntityTooLarge, "invalid_request_error", "request body too large"
|
||||
}
|
||||
r.Body = io.NopCloser(bytes.NewReader(body))
|
||||
r.ContentLength = int64(len(body))
|
||||
r.Header.Set("Content-Length", strconv.Itoa(len(body)))
|
||||
|
||||
switch t.kind {
|
||||
case KindOpenAI:
|
||||
hc.Model = gjson.GetBytes(body, "model").String()
|
||||
if len(t.models) > 0 {
|
||||
if _, ok := t.models[hc.Model]; !ok {
|
||||
return http.StatusForbidden, "model_not_allowed", "model not allowed"
|
||||
}
|
||||
}
|
||||
case KindMCP:
|
||||
msgs := []gjson.Result{gjson.ParseBytes(body)}
|
||||
if msgs[0].IsArray() {
|
||||
msgs = msgs[0].Array()
|
||||
}
|
||||
for _, m := range msgs {
|
||||
if m.Get("method").String() != "tools/call" {
|
||||
continue
|
||||
}
|
||||
tool := m.Get("params.name").String()
|
||||
if len(t.tools) > 0 {
|
||||
if _, ok := t.tools[tool]; !ok {
|
||||
return http.StatusForbidden, "tool_not_allowed", "tool not allowed"
|
||||
}
|
||||
}
|
||||
hc.Tools = append(hc.Tools, tool)
|
||||
}
|
||||
}
|
||||
return 0, "", ""
|
||||
}
|
||||
|
||||
func hasRole(roles []string, want string) bool {
|
||||
for _, r := range roles {
|
||||
if r == want {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func anyRole(roles []string, allowed map[string]struct{}) bool {
|
||||
for _, r := range roles {
|
||||
if _, ok := allowed[r]; ok {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func errorBody(typ, msg string) map[string]any {
|
||||
return map[string]any{"error": map[string]any{"message": msg, "type": typ, "code": typ}}
|
||||
}
|
||||
|
||||
func writeError(w http.ResponseWriter, status int, typ, msg string) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(status)
|
||||
_ = json.NewEncoder(w).Encode(errorBody(typ, msg))
|
||||
}
|
||||
@@ -0,0 +1,96 @@
|
||||
package aiproxy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||
)
|
||||
|
||||
// HookType defines when a hook runs.
|
||||
type HookType string
|
||||
|
||||
const (
|
||||
// BeforeProxy runs after auth, rate limit and request inspection, before the upstream call.
|
||||
BeforeProxy HookType = "before_proxy"
|
||||
// AfterProxy runs once the response finished (or failed). Errors are logged only.
|
||||
AfterProxy HookType = "after_proxy"
|
||||
)
|
||||
|
||||
// Usage is the token usage reported by an OpenAI upstream.
|
||||
type Usage struct {
|
||||
PromptTokens int64
|
||||
CompletionTokens int64
|
||||
TotalTokens int64
|
||||
}
|
||||
|
||||
// HookContext is passed to hooks.
|
||||
type HookContext struct {
|
||||
Context context.Context
|
||||
Request *http.Request // BeforeProxy may modify headers/query; the key is injected later
|
||||
UserContext *security.UserContext
|
||||
|
||||
Upstream string
|
||||
Kind Kind
|
||||
Model string // OpenAI: requested model
|
||||
Tools []string // MCP: tools named by tools/call
|
||||
|
||||
// AfterProxy only
|
||||
StatusCode int
|
||||
Duration time.Duration
|
||||
Usage Usage // OpenAI only; needs stream_options.include_usage for streams
|
||||
Error error
|
||||
|
||||
// BeforeProxy only
|
||||
Abort bool
|
||||
AbortMessage string
|
||||
AbortCode int // default 403
|
||||
}
|
||||
|
||||
// HookFunc is a hook. A returned error aborts a BeforeProxy request.
|
||||
type HookFunc func(*HookContext) error
|
||||
|
||||
// HookRegistry holds hooks per type.
|
||||
type HookRegistry struct {
|
||||
mu sync.RWMutex
|
||||
hooks map[HookType][]HookFunc
|
||||
}
|
||||
|
||||
// NewHookRegistry creates an empty registry.
|
||||
func NewHookRegistry() *HookRegistry {
|
||||
return &HookRegistry{hooks: make(map[HookType][]HookFunc)}
|
||||
}
|
||||
|
||||
// Register adds a hook.
|
||||
func (r *HookRegistry) Register(t HookType, h HookFunc) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.hooks[t] = append(r.hooks[t], h)
|
||||
}
|
||||
|
||||
// Execute runs the hooks in order and stops at the first error or abort.
|
||||
func (r *HookRegistry) Execute(t HookType, ctx *HookContext) error {
|
||||
r.mu.RLock()
|
||||
list := append([]HookFunc(nil), r.hooks[t]...)
|
||||
r.mu.RUnlock()
|
||||
|
||||
for i, h := range list {
|
||||
if err := h(ctx); err != nil {
|
||||
return fmt.Errorf("aiproxy hook %d for %s failed: %w", i+1, t, err)
|
||||
}
|
||||
if ctx.Abort {
|
||||
return fmt.Errorf("aborted by hook: %s", ctx.AbortMessage)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *HookRegistry) executeAfter(ctx *HookContext) {
|
||||
if err := r.Execute(AfterProxy, ctx); err != nil {
|
||||
logger.Warn("aiproxy: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,142 @@
|
||||
package aiproxy
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/metrics"
|
||||
)
|
||||
|
||||
// Outcome classifies a handled request.
|
||||
type Outcome string
|
||||
|
||||
const (
|
||||
OutcomeOK Outcome = "ok" // upstream answered < 400
|
||||
OutcomeUpstreamError Outcome = "upstream_error" // upstream answered >= 400 or failed
|
||||
OutcomeDenied Outcome = "denied" // refused by auth, roles, allowlist, size or a hook
|
||||
OutcomeRateLimited Outcome = "rate_limited"
|
||||
)
|
||||
|
||||
// AuditRecord is written once per handled request. It holds no bodies and no keys.
|
||||
type AuditRecord struct {
|
||||
Time time.Time
|
||||
UserID int
|
||||
User string
|
||||
RemoteID string
|
||||
Upstream string
|
||||
Kind Kind
|
||||
Method string
|
||||
Path string // client path, no query string
|
||||
Model string
|
||||
Tools []string
|
||||
Status int
|
||||
Outcome Outcome
|
||||
Reason string // why a request was refused
|
||||
Duration time.Duration
|
||||
Usage Usage
|
||||
Error string
|
||||
}
|
||||
|
||||
// AuditSink receives audit records. It runs on the request path, so keep it fast
|
||||
// (queue slow work such as database writes).
|
||||
type AuditSink interface {
|
||||
Record(AuditRecord)
|
||||
}
|
||||
|
||||
// LogAuditSink writes records through pkg/logger. It is the default sink.
|
||||
type LogAuditSink struct{}
|
||||
|
||||
// Record implements AuditSink.
|
||||
func (LogAuditSink) Record(r AuditRecord) {
|
||||
logger.Info("aiproxy audit: user=%s(%d) upstream=%s kind=%s %s %s model=%q tools=%v status=%d outcome=%s reason=%q duration=%s tokens=%d/%d/%d err=%q",
|
||||
r.User, r.UserID, r.Upstream, r.Kind, r.Method, r.Path, r.Model, r.Tools, r.Status, r.Outcome, r.Reason, r.Duration,
|
||||
r.Usage.PromptTokens, r.Usage.CompletionTokens, r.Usage.TotalTokens, r.Error)
|
||||
}
|
||||
|
||||
// NopAuditSink discards records (use it to switch the default logging off).
|
||||
type NopAuditSink struct{}
|
||||
|
||||
// Record implements AuditSink.
|
||||
func (NopAuditSink) Record(AuditRecord) {}
|
||||
|
||||
// maxModelLabels bounds the distinct (upstream, model) label pairs, since the
|
||||
// model comes from the client.
|
||||
const maxModelLabels = 400
|
||||
|
||||
type modelLabels struct {
|
||||
mu sync.Mutex
|
||||
seen map[string]struct{}
|
||||
}
|
||||
|
||||
func (m *modelLabels) label(upstream, model string) string {
|
||||
if model == "" {
|
||||
return ""
|
||||
}
|
||||
key := upstream + "|" + model
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
if _, ok := m.seen[key]; ok {
|
||||
return model
|
||||
}
|
||||
if m.seen == nil {
|
||||
m.seen = make(map[string]struct{})
|
||||
}
|
||||
if len(m.seen) >= maxModelLabels {
|
||||
return "other"
|
||||
}
|
||||
m.seen[key] = struct{}{}
|
||||
return model
|
||||
}
|
||||
|
||||
// record emits the audit record and the metrics for a finished or refused request.
|
||||
func (p *Proxy) record(t *target, hc *HookContext, status int, outcome Outcome, reason string, err error) {
|
||||
rec := AuditRecord{
|
||||
Time: time.Now(),
|
||||
Upstream: hc.Upstream,
|
||||
Kind: hc.Kind,
|
||||
Model: hc.Model,
|
||||
Tools: hc.Tools,
|
||||
Status: status,
|
||||
Outcome: outcome,
|
||||
Reason: reason,
|
||||
Duration: hc.Duration,
|
||||
Usage: hc.Usage,
|
||||
}
|
||||
if hc.Request != nil {
|
||||
rec.Method = hc.Request.Method
|
||||
rec.Path = hc.Request.URL.Path
|
||||
}
|
||||
if u := hc.UserContext; u != nil {
|
||||
rec.UserID, rec.User, rec.RemoteID = u.UserID, u.UserName, u.RemoteID
|
||||
}
|
||||
if err != nil {
|
||||
rec.Error = err.Error()
|
||||
}
|
||||
|
||||
func() {
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
logger.Error("aiproxy: audit sink panic: %v", r)
|
||||
}
|
||||
}()
|
||||
p.cfg.Audit.Record(rec)
|
||||
}()
|
||||
|
||||
if m, ok := metrics.GetProvider().(metrics.AIProxyRecorder); ok && m != nil {
|
||||
model := ""
|
||||
if outcome == OutcomeOK || outcome == OutcomeUpstreamError {
|
||||
model = p.labels.label(hc.Upstream, hc.Model)
|
||||
}
|
||||
m.RecordAIProxy(hc.Upstream, string(hc.Kind), model, statusClass(status), string(outcome), hc.Duration,
|
||||
hc.Usage.PromptTokens, hc.Usage.CompletionTokens)
|
||||
}
|
||||
}
|
||||
|
||||
// statusClass maps a status code to a low-cardinality label (2xx, 4xx, ...).
|
||||
func statusClass(status int) string {
|
||||
if status < 100 || status > 599 {
|
||||
return "unknown"
|
||||
}
|
||||
return string(rune('0'+status/100)) + "xx"
|
||||
}
|
||||
@@ -0,0 +1,80 @@
|
||||
package aiproxy
|
||||
|
||||
import (
|
||||
"math"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"golang.org/x/time/rate"
|
||||
)
|
||||
|
||||
const (
|
||||
limiterSweepEvery = 5 * time.Minute
|
||||
limiterIdleAfter = 10 * time.Minute
|
||||
)
|
||||
|
||||
type limiterEntry struct {
|
||||
lim *rate.Limiter
|
||||
last time.Time
|
||||
}
|
||||
|
||||
// limiters keeps one token bucket per key and drops idle ones.
|
||||
type limiters struct {
|
||||
mu sync.Mutex
|
||||
entries map[string]*limiterEntry
|
||||
lastSweep time.Time
|
||||
}
|
||||
|
||||
func newLimiters() *limiters {
|
||||
return &limiters{entries: make(map[string]*limiterEntry), lastSweep: time.Now()}
|
||||
}
|
||||
|
||||
// allow takes a token. When denied it returns how long to wait.
|
||||
func (l *limiters) allow(key string, rl *RateLimit) (bool, time.Duration) {
|
||||
now := time.Now()
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
|
||||
if now.Sub(l.lastSweep) > limiterSweepEvery {
|
||||
for k, e := range l.entries {
|
||||
if now.Sub(e.last) > limiterIdleAfter {
|
||||
delete(l.entries, k)
|
||||
}
|
||||
}
|
||||
l.lastSweep = now
|
||||
}
|
||||
|
||||
e, ok := l.entries[key]
|
||||
if !ok {
|
||||
burst := rl.Burst
|
||||
if burst < 1 {
|
||||
burst = int(math.Ceil(rl.PerSecond))
|
||||
if burst < 1 {
|
||||
burst = 1
|
||||
}
|
||||
}
|
||||
e = &limiterEntry{lim: rate.NewLimiter(rate.Limit(rl.PerSecond), burst)}
|
||||
l.entries[key] = e
|
||||
}
|
||||
e.last = now
|
||||
|
||||
res := e.lim.ReserveN(now, 1)
|
||||
if d := res.DelayFrom(now); d > 0 {
|
||||
res.CancelAt(now)
|
||||
return false, d
|
||||
}
|
||||
return true, 0
|
||||
}
|
||||
|
||||
// forget drops the buckets of an upstream (after it was replaced or removed).
|
||||
func (l *limiters) forget(upstream string) {
|
||||
prefix := upstream + "|"
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
for k := range l.entries {
|
||||
if strings.HasPrefix(k, prefix) {
|
||||
delete(l.entries, k)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,42 @@
|
||||
-- aiproxy upstream definitions (PostgreSQL), same contract as the resolvespec security procedures.
|
||||
-- ProcStore calls: SELECT p_success, p_error, p_data FROM resolvespec_ai_proxies()
|
||||
-- Replace the function body to source the definitions from anywhere; keep the signature.
|
||||
--
|
||||
-- p_data: JSON array, one object per upstream
|
||||
-- name, kind ('openai'|'mcp'), base_url required
|
||||
-- api_key, auth_header, auth_format, headers (object) optional
|
||||
-- allowed_roles (array), allowed (array: models for openai, tools for mcp)
|
||||
-- rate_per_second, rate_burst, timeout_seconds, enabled (default true)
|
||||
|
||||
CREATE TABLE IF NOT EXISTS ai_proxies (
|
||||
name text PRIMARY KEY,
|
||||
kind text NOT NULL CHECK (kind IN ('openai', 'mcp')),
|
||||
base_url text NOT NULL,
|
||||
api_key text,
|
||||
auth_header text,
|
||||
auth_format text,
|
||||
headers jsonb,
|
||||
allowed_roles jsonb,
|
||||
allowed jsonb,
|
||||
rate_per_second double precision,
|
||||
rate_burst integer,
|
||||
timeout_seconds integer,
|
||||
enabled boolean NOT NULL DEFAULT true
|
||||
);
|
||||
|
||||
CREATE OR REPLACE FUNCTION resolvespec_ai_proxies()
|
||||
RETURNS TABLE(p_success boolean, p_error text, p_data jsonb) AS $$
|
||||
DECLARE
|
||||
v_data jsonb;
|
||||
BEGIN
|
||||
SELECT COALESCE(jsonb_agg(to_jsonb(p) - 'enabled' ORDER BY p.name), '[]'::jsonb)
|
||||
INTO v_data
|
||||
FROM ai_proxies p
|
||||
WHERE p.enabled = true;
|
||||
|
||||
RETURN QUERY SELECT true, NULL::text, v_data;
|
||||
EXCEPTION
|
||||
WHEN OTHERS THEN
|
||||
RETURN QUERY SELECT false, SQLERRM::text, '[]'::jsonb;
|
||||
END;
|
||||
$$ LANGUAGE plpgsql;
|
||||
@@ -0,0 +1,231 @@
|
||||
package aiproxy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"regexp"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
)
|
||||
|
||||
// UpstreamStore supplies upstream definitions for Reload.
|
||||
type UpstreamStore interface {
|
||||
List(ctx context.Context) ([]Definition, error)
|
||||
}
|
||||
|
||||
// Reload syncs the store-managed upstreams with the store.
|
||||
//
|
||||
// - New definitions are added, changed ones replaced, unchanged ones left alone
|
||||
// (their rate-limit state and connections survive), missing ones removed.
|
||||
// - Upstreams registered in code are never replaced or removed; a stored definition
|
||||
// with the same name is skipped and reported.
|
||||
// - Invalid definitions are skipped and reported in the returned error; valid ones
|
||||
// are still applied. If the store itself fails, nothing changes.
|
||||
func (p *Proxy) Reload(ctx context.Context, store UpstreamStore) error {
|
||||
defs, err := store.List(ctx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("aiproxy: load upstreams: %w", err)
|
||||
}
|
||||
|
||||
var errs []error
|
||||
next := make(map[string]*target, len(defs))
|
||||
var retired []*target
|
||||
|
||||
p.mu.Lock()
|
||||
for i := range defs {
|
||||
d := defs[i]
|
||||
name := d.Upstream.Name
|
||||
if _, dup := next[name]; dup {
|
||||
errs = append(errs, fmt.Errorf("aiproxy: duplicate stored upstream %q", name))
|
||||
continue
|
||||
}
|
||||
cur := p.upstream[name]
|
||||
if cur != nil && !cur.managed {
|
||||
errs = append(errs, fmt.Errorf("aiproxy: stored upstream %q conflicts with a code-registered one", name))
|
||||
continue
|
||||
}
|
||||
if cur != nil && reflect.DeepEqual(cur.def, d) {
|
||||
next[name] = cur
|
||||
continue
|
||||
}
|
||||
t, err := p.build(d)
|
||||
if err != nil {
|
||||
errs = append(errs, err)
|
||||
if cur != nil { // keep serving the last good definition
|
||||
next[name] = cur
|
||||
}
|
||||
continue
|
||||
}
|
||||
t.managed = true
|
||||
next[name] = t
|
||||
if cur != nil {
|
||||
retired = append(retired, cur)
|
||||
}
|
||||
}
|
||||
for name, cur := range p.upstream {
|
||||
if !cur.managed {
|
||||
continue
|
||||
}
|
||||
if _, keep := next[name]; !keep {
|
||||
delete(p.upstream, name)
|
||||
retired = append(retired, cur)
|
||||
}
|
||||
}
|
||||
for name, t := range next {
|
||||
p.upstream[name] = t
|
||||
}
|
||||
p.mu.Unlock()
|
||||
|
||||
for _, t := range retired {
|
||||
p.retire(t)
|
||||
}
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
// AutoReload calls Reload now and then every interval until ctx ends. The first error
|
||||
// is returned; later ones are logged.
|
||||
func (p *Proxy) AutoReload(ctx context.Context, store UpstreamStore, interval time.Duration) error {
|
||||
if interval <= 0 {
|
||||
return errors.New("aiproxy: reload interval must be > 0")
|
||||
}
|
||||
first := p.Reload(ctx, store)
|
||||
go func() {
|
||||
t := time.NewTicker(interval)
|
||||
defer t.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-t.C:
|
||||
if err := p.Reload(ctx, store); err != nil {
|
||||
logger.Warn("aiproxy: reload: %v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
}()
|
||||
return first
|
||||
}
|
||||
|
||||
// DefaultProc is the stored procedure (function) ProcStore calls when no name is given.
|
||||
const DefaultProc = "resolvespec_ai_proxies"
|
||||
|
||||
var procRe = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*(\.[A-Za-z_][A-Za-z0-9_]*)?$`)
|
||||
|
||||
// ProcStore loads upstream definitions from a stored procedure (PostgreSQL), using the
|
||||
// same contract as the security procedures:
|
||||
//
|
||||
// SELECT p_success, p_error, p_data FROM resolvespec_ai_proxies()
|
||||
//
|
||||
// p_data is a JSON array; each element has: name, kind ("openai"|"mcp"), base_url, and
|
||||
// optionally api_key, auth_header, auth_format, headers (object), allowed_roles (array),
|
||||
// allowed (array: models or tools), rate_per_second, rate_burst, timeout_seconds,
|
||||
// enabled (default true). See resolvespec_ai_proxies.sql.
|
||||
type ProcStore struct {
|
||||
db func() *sql.DB
|
||||
proc string
|
||||
}
|
||||
|
||||
// NewProcStore calls proc through db. An empty proc means DefaultProc.
|
||||
func NewProcStore(db *sql.DB, proc string) (*ProcStore, error) {
|
||||
if db == nil {
|
||||
return nil, errors.New("aiproxy: nil database")
|
||||
}
|
||||
return newProcStore(func() *sql.DB { return db }, proc)
|
||||
}
|
||||
|
||||
// NewProcStoreFromDatabase is NewProcStore for an application's common.Database. The
|
||||
// connection is fetched on every call, so adapter reconnects are followed.
|
||||
func NewProcStoreFromDatabase(db common.Database, proc string) (*ProcStore, error) {
|
||||
if db == nil {
|
||||
return nil, errors.New("aiproxy: nil database")
|
||||
}
|
||||
p, ok := db.(common.SQLDBProvider)
|
||||
if !ok {
|
||||
return nil, fmt.Errorf("aiproxy: %T does not expose a *sql.DB", db)
|
||||
}
|
||||
return newProcStore(p.SQLDB, proc)
|
||||
}
|
||||
|
||||
func newProcStore(db func() *sql.DB, proc string) (*ProcStore, error) {
|
||||
if proc == "" {
|
||||
proc = DefaultProc
|
||||
}
|
||||
if !procRe.MatchString(proc) {
|
||||
return nil, fmt.Errorf("aiproxy: invalid procedure name %q", proc)
|
||||
}
|
||||
return &ProcStore{db: db, proc: proc}, nil
|
||||
}
|
||||
|
||||
type procRow struct {
|
||||
Name string `json:"name"`
|
||||
Kind string `json:"kind"`
|
||||
BaseURL string `json:"base_url"`
|
||||
APIKey string `json:"api_key"`
|
||||
AuthHeader string `json:"auth_header"`
|
||||
AuthFormat string `json:"auth_format"`
|
||||
Headers map[string]string `json:"headers"`
|
||||
AllowedRoles []string `json:"allowed_roles"`
|
||||
Allowed []string `json:"allowed"`
|
||||
RatePerSecond float64 `json:"rate_per_second"`
|
||||
RateBurst int `json:"rate_burst"`
|
||||
TimeoutSeconds int `json:"timeout_seconds"`
|
||||
Enabled *bool `json:"enabled"`
|
||||
}
|
||||
|
||||
// List implements UpstreamStore. Disabled and undecodable entries are skipped.
|
||||
func (s *ProcStore) List(ctx context.Context) ([]Definition, error) {
|
||||
db := s.db()
|
||||
if db == nil {
|
||||
return nil, errors.New("database connection is nil")
|
||||
}
|
||||
var (
|
||||
success bool
|
||||
errMsg sql.NullString
|
||||
data []byte
|
||||
)
|
||||
if err := db.QueryRowContext(ctx, `SELECT p_success, p_error, p_data FROM `+s.proc+`()`).Scan(&success, &errMsg, &data); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !success {
|
||||
if errMsg.Valid {
|
||||
return nil, errors.New(errMsg.String)
|
||||
}
|
||||
return nil, errors.New("failed to load AI proxies")
|
||||
}
|
||||
if len(data) == 0 {
|
||||
return nil, nil
|
||||
}
|
||||
var raw []json.RawMessage
|
||||
if err := json.Unmarshal(data, &raw); err != nil {
|
||||
return nil, fmt.Errorf("parse AI proxies: %w", err)
|
||||
}
|
||||
|
||||
out := make([]Definition, 0, len(raw))
|
||||
for i, r := range raw {
|
||||
var row procRow
|
||||
if err := json.Unmarshal(r, &row); err != nil {
|
||||
logger.Warn("aiproxy: proxy entry %d: %v", i, err)
|
||||
continue
|
||||
}
|
||||
if row.Enabled != nil && !*row.Enabled {
|
||||
continue
|
||||
}
|
||||
d := Definition{Kind: Kind(row.Kind), Allowed: row.Allowed, Upstream: Upstream{
|
||||
Name: row.Name, BaseURL: row.BaseURL, APIKey: row.APIKey,
|
||||
AuthHeader: row.AuthHeader, AuthFormat: row.AuthFormat,
|
||||
Headers: row.Headers, AllowedRoles: row.AllowedRoles,
|
||||
Timeout: time.Duration(row.TimeoutSeconds) * time.Second,
|
||||
}}
|
||||
if row.RatePerSecond > 0 {
|
||||
d.Upstream.RateLimit = &RateLimit{PerSecond: row.RatePerSecond, Burst: row.RateBurst}
|
||||
}
|
||||
out = append(out, d)
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
@@ -0,0 +1,269 @@
|
||||
package aiproxy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/DATA-DOG/go-sqlmock"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/metrics"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type memStore struct {
|
||||
mu sync.Mutex
|
||||
defs []Definition
|
||||
err error
|
||||
}
|
||||
|
||||
func (m *memStore) List(context.Context) ([]Definition, error) {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
return append([]Definition(nil), m.defs...), m.err
|
||||
}
|
||||
|
||||
func def(name, url, key string) Definition {
|
||||
return Definition{Kind: KindOpenAI, Upstream: Upstream{Name: name, BaseURL: url + "/v1", APIKey: key}}
|
||||
}
|
||||
|
||||
func TestReloadAddChangeRemove(t *testing.T) {
|
||||
up := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
io.WriteString(w, r.Header.Get("Authorization"))
|
||||
}))
|
||||
defer up.Close()
|
||||
p := New(Config{Prefix: "/ai", Audit: NopAuditSink{}})
|
||||
require.NoError(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "code", BaseURL: up.URL + "/v1", APIKey: "c"}}))
|
||||
front := httptest.NewServer(p.Handler(testAuth))
|
||||
defer front.Close()
|
||||
get := func(name string) (int, string) {
|
||||
resp, body := do(t, "GET", front.URL+"/ai/"+name+"/v1/models", "", "bob:u", nil)
|
||||
return resp.StatusCode, body
|
||||
}
|
||||
|
||||
st := &memStore{defs: []Definition{def("a", up.URL, "k1"), def("code", up.URL, "stolen")}}
|
||||
err := p.Reload(context.Background(), st)
|
||||
assert.ErrorContains(t, err, "conflicts with a code-registered")
|
||||
code, body := get("a")
|
||||
assert.Equal(t, 200, code)
|
||||
assert.Equal(t, "Bearer k1", body)
|
||||
_, body = get("code")
|
||||
assert.Equal(t, "Bearer c", body) // code-registered untouched
|
||||
|
||||
st.defs = []Definition{def("a", up.URL, "k2"), def("b", up.URL, "kb")}
|
||||
require.NoError(t, p.Reload(context.Background(), st))
|
||||
_, body = get("a")
|
||||
assert.Equal(t, "Bearer k2", body)
|
||||
_, body = get("b")
|
||||
assert.Equal(t, "Bearer kb", body)
|
||||
|
||||
st.defs = []Definition{def("b", up.URL, "kb")}
|
||||
require.NoError(t, p.Reload(context.Background(), st))
|
||||
code, _ = get("a")
|
||||
assert.Equal(t, 404, code)
|
||||
code, _ = get("code")
|
||||
assert.Equal(t, 200, code)
|
||||
|
||||
// store failure changes nothing; invalid def keeps last good one
|
||||
st.err = errors.New("db down")
|
||||
assert.Error(t, p.Reload(context.Background(), st))
|
||||
code, _ = get("b")
|
||||
assert.Equal(t, 200, code)
|
||||
st.err = nil
|
||||
st.defs = []Definition{{Kind: KindOpenAI, Upstream: Upstream{Name: "b", BaseURL: "ftp://x"}}}
|
||||
assert.Error(t, p.Reload(context.Background(), st))
|
||||
_, body = get("b")
|
||||
assert.Equal(t, "Bearer kb", body)
|
||||
|
||||
assert.True(t, p.Unregister("code"))
|
||||
assert.False(t, p.Unregister("code"))
|
||||
}
|
||||
|
||||
func TestReloadKeepsRateLimitStateWhenUnchanged(t *testing.T) {
|
||||
up := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {}))
|
||||
defer up.Close()
|
||||
p := New(Config{Prefix: "/ai", Audit: NopAuditSink{}})
|
||||
d := def("a", up.URL, "")
|
||||
d.Upstream.RateLimit = &RateLimit{PerSecond: 0.01, Burst: 1}
|
||||
st := &memStore{defs: []Definition{d}}
|
||||
require.NoError(t, p.Reload(context.Background(), st))
|
||||
front := httptest.NewServer(p.Handler(testAuth))
|
||||
defer front.Close()
|
||||
|
||||
resp, _ := do(t, "GET", front.URL+"/ai/a/v1/x", "", "bob:u", nil)
|
||||
assert.Equal(t, 200, resp.StatusCode)
|
||||
require.NoError(t, p.Reload(context.Background(), st))
|
||||
resp, _ = do(t, "GET", front.URL+"/ai/a/v1/x", "", "bob:u", nil)
|
||||
assert.Equal(t, 429, resp.StatusCode)
|
||||
}
|
||||
|
||||
func TestAutoReload(t *testing.T) {
|
||||
p := New(Config{Audit: NopAuditSink{}})
|
||||
st := &memStore{}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
require.NoError(t, p.AutoReload(ctx, st, 20*time.Millisecond))
|
||||
assert.Empty(t, p.Names())
|
||||
st.mu.Lock()
|
||||
st.defs = []Definition{def("a", "http://x", "")}
|
||||
st.mu.Unlock()
|
||||
assert.Eventually(t, func() bool { return len(p.Names()) == 1 }, 2*time.Second, 10*time.Millisecond)
|
||||
assert.Error(t, p.AutoReload(ctx, st, 0))
|
||||
}
|
||||
|
||||
func procRows(data string) *sqlmock.Rows {
|
||||
return sqlmock.NewRows([]string{"p_success", "p_error", "p_data"}).AddRow(true, nil, []byte(data))
|
||||
}
|
||||
|
||||
func TestProcStore(t *testing.T) {
|
||||
db, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(sqlmock.QueryMatcherRegexp))
|
||||
require.NoError(t, err)
|
||||
defer db.Close()
|
||||
|
||||
_, err = NewProcStore(db, "x; drop table y")
|
||||
assert.Error(t, err)
|
||||
_, err = NewProcStore(nil, "")
|
||||
assert.Error(t, err)
|
||||
s, err := NewProcStore(db, "")
|
||||
require.NoError(t, err)
|
||||
|
||||
mock.ExpectQuery(`SELECT p_success, p_error, p_data FROM resolvespec_ai_proxies\(\)`).WillReturnRows(procRows(`[
|
||||
{"name":"gpt","kind":"openai","base_url":"https://api.openai.com/v1","api_key":"sk","auth_header":"api-key",
|
||||
"auth_format":"%s","headers":{"X-A":"1"},"allowed_roles":["ai"],"allowed":["m1","m2"],
|
||||
"rate_per_second":2.5,"rate_burst":5,"timeout_seconds":30},
|
||||
{"name":"tools","kind":"mcp","base_url":"http://localhost:3000/mcp","allowed":["t1"],"enabled":true},
|
||||
{"name":"off","kind":"mcp","base_url":"http://localhost:3001/mcp","enabled":false},
|
||||
{"name":5}
|
||||
]`))
|
||||
defs, err := s.List(context.Background())
|
||||
require.NoError(t, err)
|
||||
require.Len(t, defs, 2)
|
||||
|
||||
gpt := defs[0]
|
||||
assert.Equal(t, KindOpenAI, gpt.Kind)
|
||||
assert.Equal(t, "sk", gpt.Upstream.APIKey)
|
||||
assert.Equal(t, "api-key", gpt.Upstream.AuthHeader)
|
||||
assert.Equal(t, map[string]string{"X-A": "1"}, gpt.Upstream.Headers)
|
||||
assert.Equal(t, []string{"ai"}, gpt.Upstream.AllowedRoles)
|
||||
assert.Equal(t, []string{"m1", "m2"}, gpt.Allowed)
|
||||
assert.Equal(t, &RateLimit{PerSecond: 2.5, Burst: 5}, gpt.Upstream.RateLimit)
|
||||
assert.Equal(t, 30*time.Second, gpt.Upstream.Timeout)
|
||||
assert.Equal(t, KindMCP, defs[1].Kind)
|
||||
assert.Nil(t, defs[1].Upstream.RateLimit)
|
||||
|
||||
// feeds Reload
|
||||
mock.ExpectQuery(`FROM resolvespec_ai_proxies`).WillReturnRows(procRows(`[{"name":"a","kind":"mcp","base_url":"http://x/mcp"},{"name":"bad","kind":"nope","base_url":"http://x"}]`))
|
||||
p := New(Config{Audit: NopAuditSink{}})
|
||||
assert.Error(t, p.Reload(context.Background(), s)) // "bad" reported
|
||||
assert.Equal(t, []string{"a"}, p.Names())
|
||||
assert.NoError(t, mock.ExpectationsWereMet())
|
||||
}
|
||||
|
||||
func TestProcStoreFailures(t *testing.T) {
|
||||
db, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(sqlmock.QueryMatcherRegexp))
|
||||
require.NoError(t, err)
|
||||
defer db.Close()
|
||||
s, err := NewProcStore(db, "myschema.my_proxies")
|
||||
require.NoError(t, err)
|
||||
|
||||
mock.ExpectQuery(`FROM myschema.my_proxies\(\)`).WillReturnRows(
|
||||
sqlmock.NewRows([]string{"p_success", "p_error", "p_data"}).AddRow(false, "boom", []byte(`[]`)))
|
||||
_, err = s.List(context.Background())
|
||||
assert.EqualError(t, err, "boom")
|
||||
|
||||
mock.ExpectQuery(`FROM myschema.my_proxies`).WillReturnRows(procRows(`{not json`))
|
||||
_, err = s.List(context.Background())
|
||||
assert.ErrorContains(t, err, "parse AI proxies")
|
||||
|
||||
mock.ExpectQuery(`FROM myschema.my_proxies`).WillReturnError(errors.New("conn lost"))
|
||||
_, err = s.List(context.Background())
|
||||
assert.ErrorContains(t, err, "conn lost")
|
||||
|
||||
mock.ExpectQuery(`FROM myschema.my_proxies`).WillReturnRows(
|
||||
sqlmock.NewRows([]string{"p_success", "p_error", "p_data"}).AddRow(true, nil, nil))
|
||||
defs, err := s.List(context.Background())
|
||||
assert.NoError(t, err)
|
||||
assert.Empty(t, defs)
|
||||
}
|
||||
|
||||
type auditCollector struct {
|
||||
mu sync.Mutex
|
||||
recs []AuditRecord
|
||||
}
|
||||
|
||||
func (a *auditCollector) Record(r AuditRecord) {
|
||||
a.mu.Lock()
|
||||
a.recs = append(a.recs, r)
|
||||
a.mu.Unlock()
|
||||
}
|
||||
|
||||
type metricsCollector struct {
|
||||
metrics.Provider
|
||||
mu sync.Mutex
|
||||
calls []string
|
||||
toks int64
|
||||
}
|
||||
|
||||
func (m *metricsCollector) RecordAIProxy(up, kind, model, class, outcome string, d time.Duration, pt, ct int64) {
|
||||
m.mu.Lock()
|
||||
m.calls = append(m.calls, up+"|"+kind+"|"+model+"|"+class+"|"+outcome)
|
||||
m.toks += pt + ct
|
||||
m.mu.Unlock()
|
||||
}
|
||||
|
||||
func TestAuditAndMetrics(t *testing.T) {
|
||||
old := metrics.GetProvider()
|
||||
mc := &metricsCollector{Provider: old}
|
||||
metrics.SetProvider(mc)
|
||||
defer metrics.SetProvider(old)
|
||||
|
||||
ac := &auditCollector{}
|
||||
up := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
io.WriteString(w, `{"usage":{"prompt_tokens":1,"completion_tokens":2}}`)
|
||||
}))
|
||||
defer up.Close()
|
||||
p := New(Config{Prefix: "/ai", Audit: ac})
|
||||
require.NoError(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "g", BaseURL: up.URL + "/v1", APIKey: "sk-secret", AllowedRoles: []string{"ai"}}, AllowedModels: []string{"m"}}))
|
||||
front := httptest.NewServer(p.Handler(testAuth))
|
||||
defer front.Close()
|
||||
|
||||
do(t, "POST", front.URL+"/ai/g/v1/chat/completions?k=v", `{"model":"m"}`, "bob:ai", nil)
|
||||
do(t, "POST", front.URL+"/ai/g/v1/chat/completions", `{"model":"zzz"}`, "bob:ai", nil)
|
||||
do(t, "GET", front.URL+"/ai/g/v1/models", "", "bob:user", nil)
|
||||
|
||||
require.Eventually(t, func() bool { ac.mu.Lock(); defer ac.mu.Unlock(); return len(ac.recs) == 3 }, 2*time.Second, 10*time.Millisecond)
|
||||
ac.mu.Lock()
|
||||
r := ac.recs
|
||||
ac.mu.Unlock()
|
||||
assert.Equal(t, OutcomeOK, r[0].Outcome)
|
||||
assert.Equal(t, "bob", r[0].User)
|
||||
assert.Equal(t, "/ai/g/v1/chat/completions", r[0].Path) // no query
|
||||
assert.Equal(t, "m", r[0].Model)
|
||||
assert.Equal(t, int64(3), r[0].Usage.TotalTokens)
|
||||
assert.Equal(t, OutcomeDenied, r[1].Outcome)
|
||||
assert.Equal(t, 403, r[1].Status)
|
||||
assert.Equal(t, OutcomeDenied, r[2].Outcome)
|
||||
for _, rec := range r {
|
||||
assert.NotContains(t, rec.Reason+rec.Error+rec.Path, "sk-secret")
|
||||
}
|
||||
|
||||
mc.mu.Lock()
|
||||
defer mc.mu.Unlock()
|
||||
assert.Equal(t, []string{"g|openai|m|2xx|ok", "g|openai||4xx|denied", "g|openai||4xx|denied"}, mc.calls)
|
||||
assert.Equal(t, int64(3), mc.toks)
|
||||
}
|
||||
|
||||
func TestModelLabelsAreBounded(t *testing.T) {
|
||||
var m modelLabels
|
||||
for i := 0; i < maxModelLabels+50; i++ {
|
||||
m.label("u", "model-"+string(rune('a'+i%26))+string(rune('a'+(i/26)%26))+string(rune('a'+(i/676)%26)))
|
||||
}
|
||||
assert.Equal(t, "other", m.label("u", "brand-new"))
|
||||
assert.Equal(t, "", m.label("u", ""))
|
||||
}
|
||||
@@ -0,0 +1,143 @@
|
||||
package aiproxy
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"regexp"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Kind is the type of an upstream.
|
||||
type Kind string
|
||||
|
||||
const (
|
||||
KindOpenAI Kind = "openai"
|
||||
KindMCP Kind = "mcp"
|
||||
)
|
||||
|
||||
const redacted = "[redacted]"
|
||||
|
||||
var nameRe = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9_.-]*$`)
|
||||
|
||||
// RateLimit limits requests per user per upstream (token bucket).
|
||||
type RateLimit struct {
|
||||
PerSecond float64 // sustained requests per second
|
||||
Burst int // bucket size; defaults to ceil(PerSecond), min 1
|
||||
}
|
||||
|
||||
// Upstream holds the settings shared by every upstream kind.
|
||||
type Upstream struct {
|
||||
// Name is the route segment clients use: {prefix}/{Name}/...
|
||||
Name string
|
||||
// BaseURL of the upstream. OpenAI: include the version path (https://api.openai.com/v1).
|
||||
// MCP: the full streamable HTTP endpoint (http://localhost:3000/mcp).
|
||||
BaseURL string
|
||||
// APIKey is sent to the upstream. Never exposed to clients or logs.
|
||||
APIKey string
|
||||
// AuthHeader carries the key. Default "Authorization". Applies to OpenAI and MCP upstreams.
|
||||
AuthHeader string
|
||||
// AuthFormat renders the header value; %s is the key (e.g. "Bearer %s", "Token %s", "%s").
|
||||
// Default: "Bearer %s" for the Authorization header, "%s" (raw key) for any other header.
|
||||
AuthFormat string
|
||||
// Headers are extra static headers sent to the upstream (applied before the key header).
|
||||
// Per-request values can be set in a BeforeProxy hook.
|
||||
Headers map[string]string
|
||||
// AllowedRoles: caller needs at least one. Empty means any authenticated user.
|
||||
AllowedRoles []string
|
||||
// RateLimit per user for this upstream. Nil means unlimited.
|
||||
RateLimit *RateLimit
|
||||
// Timeout waiting for the upstream response headers. Default 120s. Streams are not cut off.
|
||||
Timeout time.Duration
|
||||
}
|
||||
|
||||
// OpenAIUpstream is an OpenAI-compatible API (OpenAI, Azure, Ollama, vLLM, ...).
|
||||
type OpenAIUpstream struct {
|
||||
Upstream
|
||||
// AllowedModels restricts the "model" a request may use. Empty means any.
|
||||
// When set, write requests must be JSON bodies carrying an allowed model.
|
||||
AllowedModels []string
|
||||
}
|
||||
|
||||
// MCPUpstream is an MCP server speaking streamable HTTP.
|
||||
type MCPUpstream struct {
|
||||
Upstream
|
||||
// AllowedTools restricts tools/call to these tool names. Empty means any.
|
||||
AllowedTools []string
|
||||
}
|
||||
|
||||
// String redacts the key.
|
||||
func (u Upstream) String() string {
|
||||
return fmt.Sprintf("Upstream{Name:%s BaseURL:%s APIKey:%s}", u.Name, u.BaseURL, keyMask(u.APIKey))
|
||||
}
|
||||
|
||||
// GoString redacts the key (%#v).
|
||||
func (u Upstream) GoString() string { return u.String() }
|
||||
|
||||
// MarshalJSON redacts the key.
|
||||
func (u Upstream) MarshalJSON() ([]byte, error) {
|
||||
return json.Marshal(struct {
|
||||
Name string `json:"name"`
|
||||
BaseURL string `json:"base_url"`
|
||||
APIKey string `json:"api_key,omitempty"`
|
||||
AllowedRoles []string `json:"allowed_roles,omitempty"`
|
||||
}{u.Name, u.BaseURL, keyMask(u.APIKey), u.AllowedRoles})
|
||||
}
|
||||
|
||||
func keyMask(k string) string {
|
||||
if k == "" {
|
||||
return ""
|
||||
}
|
||||
return redacted
|
||||
}
|
||||
|
||||
// validate checks the shared fields and returns the parsed base URL.
|
||||
func (u Upstream) validate() (*url.URL, error) {
|
||||
if !nameRe.MatchString(u.Name) {
|
||||
return nil, fmt.Errorf("aiproxy: invalid upstream name %q", u.Name)
|
||||
}
|
||||
base, err := url.Parse(strings.TrimRight(u.BaseURL, "/"))
|
||||
if err != nil || (base.Scheme != "http" && base.Scheme != "https") || base.Host == "" {
|
||||
return nil, fmt.Errorf("aiproxy: upstream %q has an invalid base URL", u.Name)
|
||||
}
|
||||
if u.AuthFormat != "" && (strings.Count(u.AuthFormat, "%s") != 1 || strings.Count(u.AuthFormat, "%") != 1) {
|
||||
return nil, fmt.Errorf("aiproxy: upstream %q auth format must contain exactly one %%s", u.Name)
|
||||
}
|
||||
if u.RateLimit != nil && u.RateLimit.PerSecond <= 0 {
|
||||
return nil, fmt.Errorf("aiproxy: upstream %q rate limit must be > 0", u.Name)
|
||||
}
|
||||
return base, nil
|
||||
}
|
||||
|
||||
// authValue returns the header name and value carrying the key, or "" when no key is set.
|
||||
func (u Upstream) authValue() (name, value string) {
|
||||
if u.APIKey == "" {
|
||||
return "", ""
|
||||
}
|
||||
name = u.AuthHeader
|
||||
if name == "" {
|
||||
name = "Authorization"
|
||||
}
|
||||
format := u.AuthFormat
|
||||
if format == "" {
|
||||
if http.CanonicalHeaderKey(name) == "Authorization" {
|
||||
format = "Bearer %s"
|
||||
} else {
|
||||
format = "%s"
|
||||
}
|
||||
}
|
||||
return name, fmt.Sprintf(format, u.APIKey)
|
||||
}
|
||||
|
||||
func toSet(items []string) map[string]struct{} {
|
||||
if len(items) == 0 {
|
||||
return nil
|
||||
}
|
||||
m := make(map[string]struct{}, len(items))
|
||||
for _, i := range items {
|
||||
m[i] = struct{}{}
|
||||
}
|
||||
return m
|
||||
}
|
||||
@@ -0,0 +1,118 @@
|
||||
package aiproxy
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"io"
|
||||
"sync"
|
||||
|
||||
"github.com/tidwall/gjson"
|
||||
)
|
||||
|
||||
const maxCapture = 1 << 20
|
||||
|
||||
// watchBody passes the response through while it looks for OpenAI usage, and fires
|
||||
// AfterProxy once the body ends or is closed.
|
||||
type watchBody struct {
|
||||
rc io.ReadCloser
|
||||
capture bool // read usage
|
||||
sse bool
|
||||
status int
|
||||
c *call
|
||||
|
||||
buf bytes.Buffer // JSON body (capped)
|
||||
overflow bool
|
||||
line []byte // partial SSE line
|
||||
usage Usage
|
||||
once sync.Once
|
||||
}
|
||||
|
||||
func newWatchBody(rc io.ReadCloser, openai, sse bool, status int, c *call) io.ReadCloser {
|
||||
return &watchBody{rc: rc, capture: openai, sse: sse, status: status, c: c}
|
||||
}
|
||||
|
||||
func (w *watchBody) Read(p []byte) (int, error) {
|
||||
n, err := w.rc.Read(p)
|
||||
if w.capture && n > 0 {
|
||||
w.feed(p[:n])
|
||||
}
|
||||
if err != nil {
|
||||
w.end(err)
|
||||
}
|
||||
return n, err
|
||||
}
|
||||
|
||||
func (w *watchBody) Close() error {
|
||||
err := w.rc.Close()
|
||||
w.end(nil)
|
||||
return err
|
||||
}
|
||||
|
||||
func (w *watchBody) end(err error) {
|
||||
w.once.Do(func() {
|
||||
if w.capture && !w.sse && !w.overflow {
|
||||
if u, ok := parseUsage(w.buf.Bytes()); ok {
|
||||
w.usage = u
|
||||
}
|
||||
}
|
||||
if err == io.EOF {
|
||||
err = nil
|
||||
}
|
||||
w.c.finish(w.status, w.usage, err)
|
||||
})
|
||||
}
|
||||
|
||||
func (w *watchBody) feed(b []byte) {
|
||||
if !w.sse {
|
||||
if w.overflow || w.buf.Len()+len(b) > maxCapture {
|
||||
w.overflow = true
|
||||
return
|
||||
}
|
||||
w.buf.Write(b)
|
||||
return
|
||||
}
|
||||
w.line = append(w.line, b...)
|
||||
for {
|
||||
i := bytes.IndexByte(w.line, '\n')
|
||||
if i < 0 {
|
||||
break
|
||||
}
|
||||
line := bytes.TrimSpace(w.line[:i])
|
||||
w.line = w.line[i+1:]
|
||||
if rest, ok := bytes.CutPrefix(line, []byte("data:")); ok && bytes.Contains(rest, []byte("usage")) {
|
||||
if u, ok := parseUsage(bytes.TrimSpace(rest)); ok {
|
||||
w.usage = u
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(w.line) > maxCapture {
|
||||
w.line = nil
|
||||
}
|
||||
}
|
||||
|
||||
// parseUsage reads chat/embeddings ("usage") and responses ("response.usage") shapes.
|
||||
func parseUsage(data []byte) (Usage, bool) {
|
||||
u := gjson.GetBytes(data, "usage")
|
||||
if !u.IsObject() {
|
||||
u = gjson.GetBytes(data, "response.usage")
|
||||
}
|
||||
if !u.IsObject() {
|
||||
return Usage{}, false
|
||||
}
|
||||
first := func(keys ...string) int64 {
|
||||
for _, k := range keys {
|
||||
if v := u.Get(k); v.Exists() {
|
||||
return v.Int()
|
||||
}
|
||||
}
|
||||
return 0
|
||||
}
|
||||
out := Usage{
|
||||
PromptTokens: first("prompt_tokens", "input_tokens"),
|
||||
CompletionTokens: first("completion_tokens", "output_tokens"),
|
||||
TotalTokens: first("total_tokens"),
|
||||
}
|
||||
if out.TotalTokens == 0 {
|
||||
out.TotalTokens = out.PromptTokens + out.CompletionTokens
|
||||
}
|
||||
return out, true
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
//go:build integration
|
||||
// +build integration
|
||||
|
||||
package database
|
||||
@@ -18,11 +19,11 @@ import (
|
||||
|
||||
// Integration test models
|
||||
type IntegrationUser struct {
|
||||
ID int `db:"id"`
|
||||
Name string `db:"name"`
|
||||
Email string `db:"email"`
|
||||
Age int `db:"age"`
|
||||
CreatedAt time.Time `db:"created_at"`
|
||||
ID int `db:"id"`
|
||||
Name string `db:"name"`
|
||||
Email string `db:"email"`
|
||||
Age int `db:"age"`
|
||||
CreatedAt time.Time `db:"created_at"`
|
||||
Posts []*IntegrationPost `bun:"rel:has-many,join:id=user_id"`
|
||||
}
|
||||
|
||||
@@ -46,10 +47,10 @@ func (p IntegrationPost) TableName() string {
|
||||
}
|
||||
|
||||
type IntegrationComment struct {
|
||||
ID int `db:"id"`
|
||||
Content string `db:"content"`
|
||||
PostID int `db:"post_id"`
|
||||
CreatedAt time.Time `db:"created_at"`
|
||||
ID int `db:"id"`
|
||||
Content string `db:"content"`
|
||||
PostID int `db:"post_id"`
|
||||
CreatedAt time.Time `db:"created_at"`
|
||||
Post *IntegrationPost `bun:"rel:belongs-to,join:post_id=id"`
|
||||
}
|
||||
|
||||
|
||||
@@ -26,11 +26,11 @@ func (u TestUser) TableName() string {
|
||||
}
|
||||
|
||||
type TestPost struct {
|
||||
ID int `db:"id"`
|
||||
Title string `db:"title"`
|
||||
Content string `db:"content"`
|
||||
UserID int `db:"user_id"`
|
||||
User *TestUser `bun:"rel:belongs-to,join:user_id=id"`
|
||||
ID int `db:"id"`
|
||||
Title string `db:"title"`
|
||||
Content string `db:"content"`
|
||||
UserID int `db:"user_id"`
|
||||
User *TestUser `bun:"rel:belongs-to,join:user_id=id"`
|
||||
Comments []TestComment `bun:"rel:has-many,join:id=post_id"`
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -0,0 +1,77 @@
|
||||
package common
|
||||
|
||||
import (
|
||||
"regexp"
|
||||
"strings"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||
)
|
||||
|
||||
var rePlainIdent = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*$`)
|
||||
|
||||
// MainTableAlias returns the alias the main table gets in the SELECT:
|
||||
// the model's TableAlias() when provided, otherwise the bare table name.
|
||||
func MainTableAlias(model interface{}, tableName string) string {
|
||||
if p, ok := model.(TableAliasProvider); ok {
|
||||
if a := p.TableAlias(); a != "" {
|
||||
return a
|
||||
}
|
||||
}
|
||||
return reflection.ExtractTableNameOnly(tableName)
|
||||
}
|
||||
|
||||
func isModelSQLColumn(model interface{}, column string) bool {
|
||||
for _, c := range reflection.GetSQLModelColumns(model) {
|
||||
if strings.EqualFold(c, column) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// QualifyModelColumn returns "alias"."column" when column is a plain identifier
|
||||
// that exists on the model, so it stays unambiguous once joins are added.
|
||||
// Anything else (expressions, JSON paths, other tables' columns) is returned unchanged.
|
||||
func QualifyModelColumn(model interface{}, alias, column string) string {
|
||||
if alias == "" || model == nil || !rePlainIdent.MatchString(column) || !isModelSQLColumn(model, column) {
|
||||
return column
|
||||
}
|
||||
return QuoteIdent(alias) + "." + QuoteIdent(column)
|
||||
}
|
||||
|
||||
// StripMainTablePrefix turns "<prefix>.<column>" into "<column>" when prefix is
|
||||
// one of the main table's names/aliases and column exists on the model.
|
||||
// Any other input is returned unchanged.
|
||||
func StripMainTablePrefix(model interface{}, column string, prefixes ...string) string {
|
||||
idx := strings.Index(column, ".")
|
||||
if idx <= 0 || model == nil {
|
||||
return column
|
||||
}
|
||||
prefix := strings.Trim(column[:idx], `"`)
|
||||
col := strings.Trim(column[idx+1:], `"`)
|
||||
if !rePlainIdent.MatchString(prefix) || !rePlainIdent.MatchString(col) {
|
||||
return column
|
||||
}
|
||||
for _, p := range prefixes {
|
||||
if p != "" && strings.EqualFold(p, prefix) && isModelSQLColumn(model, col) {
|
||||
return col
|
||||
}
|
||||
}
|
||||
return column
|
||||
}
|
||||
|
||||
// StripMainTablePrefixFromFilters applies StripMainTablePrefix to every filter column.
|
||||
func StripMainTablePrefixFromFilters(model interface{}, filters []FilterOption, prefixes ...string) {
|
||||
for i := range filters {
|
||||
filters[i].Column = StripMainTablePrefix(model, filters[i].Column, prefixes...)
|
||||
}
|
||||
}
|
||||
|
||||
// NormalizeMainTableFilters rewrites "<main table or alias>.<column>" filter
|
||||
// columns on opts to the bare model column (on a copy, never the caller's
|
||||
// slice) so the column validator keeps them.
|
||||
func NormalizeMainTableFilters(model interface{}, tableName string, opts *RequestOptions) {
|
||||
opts.Filters = append([]FilterOption(nil), opts.Filters...)
|
||||
StripMainTablePrefixFromFilters(model, opts.Filters,
|
||||
MainTableAlias(model, tableName), reflection.ExtractTableNameOnly(tableName))
|
||||
}
|
||||
@@ -0,0 +1,55 @@
|
||||
package common
|
||||
|
||||
import "testing"
|
||||
|
||||
type qualifyModel struct {
|
||||
ID int64 `bun:"id,pk"`
|
||||
Name string `bun:"name"`
|
||||
}
|
||||
|
||||
func TestQualifyModelColumn(t *testing.T) {
|
||||
m := qualifyModel{}
|
||||
cases := []struct{ alias, col, want string }{
|
||||
{"t", "name", `"t"."name"`},
|
||||
{"t", "NAME", `"t"."NAME"`},
|
||||
{"t", "missing", "missing"},
|
||||
{"t", "data->>'x'", "data->>'x'"},
|
||||
{"t", "rel.name", "rel.name"},
|
||||
{"", "name", "name"},
|
||||
}
|
||||
for _, c := range cases {
|
||||
if got := QualifyModelColumn(m, c.alias, c.col); got != c.want {
|
||||
t.Errorf("Qualify(%q,%q) = %q, want %q", c.alias, c.col, got, c.want)
|
||||
}
|
||||
}
|
||||
if got := QualifyModelColumn(nil, "t", "name"); got != "name" {
|
||||
t.Errorf("nil model: %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStripMainTablePrefix(t *testing.T) {
|
||||
m := qualifyModel{}
|
||||
cases := []struct{ col, want string }{
|
||||
{"province_state.name", "name"},
|
||||
{`"province_state"."name"`, "name"},
|
||||
{"PROVINCE_STATE.name", "name"},
|
||||
{"rel_rid_country.name", "rel_rid_country.name"},
|
||||
{"province_state.missing", "province_state.missing"},
|
||||
{"name", "name"},
|
||||
}
|
||||
for _, c := range cases {
|
||||
if got := StripMainTablePrefix(m, c.col, "province_state"); got != c.want {
|
||||
t.Errorf("Strip(%q) = %q, want %q", c.col, got, c.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestStripMainTablePrefix_ValidatorKeepsFilter(t *testing.T) {
|
||||
m := qualifyModel{}
|
||||
filters := []FilterOption{{Column: "province_state.name", Operator: "eq", Value: "x"}}
|
||||
StripMainTablePrefixFromFilters(m, filters, "province_state")
|
||||
out := NewColumnValidator(m).FilterRequestOptions(RequestOptions{Filters: filters})
|
||||
if len(out.Filters) != 1 || out.Filters[0].Column != "name" {
|
||||
t.Fatalf("filters = %+v", out.Filters)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -25,10 +25,10 @@ func newMockDatabase() *mockDatabase {
|
||||
}
|
||||
}
|
||||
|
||||
func (m *mockDatabase) NewSelect() SelectQuery { return &mockSelectQuery{} }
|
||||
func (m *mockDatabase) NewInsert() InsertQuery { return &mockInsertQuery{db: m} }
|
||||
func (m *mockDatabase) NewUpdate() UpdateQuery { return &mockUpdateQuery{db: m} }
|
||||
func (m *mockDatabase) NewDelete() DeleteQuery { return &mockDeleteQuery{db: m} }
|
||||
func (m *mockDatabase) NewSelect() SelectQuery { return &mockSelectQuery{} }
|
||||
func (m *mockDatabase) NewInsert() InsertQuery { return &mockInsertQuery{db: m} }
|
||||
func (m *mockDatabase) NewUpdate() UpdateQuery { return &mockUpdateQuery{db: m} }
|
||||
func (m *mockDatabase) NewDelete() DeleteQuery { return &mockDeleteQuery{db: m} }
|
||||
func (m *mockDatabase) RunInTransaction(ctx context.Context, fn func(Database) error) error {
|
||||
return fn(m)
|
||||
}
|
||||
@@ -57,27 +57,31 @@ func (m *mockDatabase) DriverName() string {
|
||||
// Mock SelectQuery
|
||||
type mockSelectQuery struct{}
|
||||
|
||||
func (m *mockSelectQuery) Model(model interface{}) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Table(name string) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Column(columns ...string) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) ColumnExpr(query string, args ...interface{}) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Where(condition string, args ...interface{}) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) WhereOr(query string, args ...interface{}) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Join(query string, args ...interface{}) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) LeftJoin(query string, args ...interface{}) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Model(model interface{}) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Table(name string) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Column(columns ...string) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) ColumnExpr(query string, args ...interface{}) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Where(condition string, args ...interface{}) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) WhereOr(query string, args ...interface{}) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Join(query string, args ...interface{}) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) LeftJoin(query string, args ...interface{}) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Preload(relation string, conditions ...interface{}) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) PreloadRelation(relation string, apply ...func(SelectQuery) SelectQuery) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) JoinRelation(relation string, apply ...func(SelectQuery) SelectQuery) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Order(order string) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) OrderExpr(order string, args ...interface{}) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Limit(n int) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Offset(n int) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Group(group string) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) PreloadRelation(relation string, apply ...func(SelectQuery) SelectQuery) SelectQuery {
|
||||
return m
|
||||
}
|
||||
func (m *mockSelectQuery) JoinRelation(relation string, apply ...func(SelectQuery) SelectQuery) SelectQuery {
|
||||
return m
|
||||
}
|
||||
func (m *mockSelectQuery) Order(order string) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) OrderExpr(order string, args ...interface{}) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Limit(n int) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Offset(n int) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Group(group string) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Having(condition string, args ...interface{}) SelectQuery { return m }
|
||||
func (m *mockSelectQuery) Scan(ctx context.Context, dest interface{}) error { return nil }
|
||||
func (m *mockSelectQuery) ScanModel(ctx context.Context) error { return nil }
|
||||
func (m *mockSelectQuery) Count(ctx context.Context) (int, error) { return 0, nil }
|
||||
func (m *mockSelectQuery) Exists(ctx context.Context) (bool, error) { return false, nil }
|
||||
func (m *mockSelectQuery) Scan(ctx context.Context, dest interface{}) error { return nil }
|
||||
func (m *mockSelectQuery) ScanModel(ctx context.Context) error { return nil }
|
||||
func (m *mockSelectQuery) Count(ctx context.Context) (int, error) { return 0, nil }
|
||||
func (m *mockSelectQuery) Exists(ctx context.Context) (bool, error) { return false, nil }
|
||||
|
||||
// Mock InsertQuery
|
||||
type mockInsertQuery struct {
|
||||
@@ -98,8 +102,9 @@ func (m *mockInsertQuery) Value(column string, value interface{}) InsertQuery {
|
||||
m.values[column] = value
|
||||
return m
|
||||
}
|
||||
func (m *mockInsertQuery) OnConflict(action string) InsertQuery { return m }
|
||||
func (m *mockInsertQuery) Returning(columns ...string) 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)
|
||||
m.db.lastID++
|
||||
@@ -131,7 +136,8 @@ 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) Returning(columns ...string) 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
|
||||
m.db.updateCalls = append(m.db.updateCalls, m.setValues)
|
||||
@@ -169,9 +175,13 @@ func (m *mockResult) RowsAffected() int64 { return m.rowsAffected }
|
||||
type mockModelRegistry struct{}
|
||||
|
||||
func (m *mockModelRegistry) GetModel(name string) (interface{}, error) { return nil, nil }
|
||||
func (m *mockModelRegistry) GetModelByEntity(schema, entity string) (interface{}, error) { return nil, nil }
|
||||
func (m *mockModelRegistry) GetModelByEntity(schema, entity string) (interface{}, error) {
|
||||
return nil, nil
|
||||
}
|
||||
func (m *mockModelRegistry) RegisterModel(name string, model interface{}) error { return nil }
|
||||
func (m *mockModelRegistry) GetAllModels() map[string]interface{} { return make(map[string]interface{}) }
|
||||
func (m *mockModelRegistry) GetAllModels() map[string]interface{} {
|
||||
return make(map[string]interface{})
|
||||
}
|
||||
|
||||
// Mock RelationshipInfoProvider
|
||||
type mockRelationshipProvider struct {
|
||||
@@ -196,9 +206,9 @@ func (m *mockRelationshipProvider) RegisterRelation(modelTypeName, relationName
|
||||
|
||||
// Test Models
|
||||
type Department struct {
|
||||
ID int64 `json:"id" bun:"id,pk"`
|
||||
Name string `json:"name"`
|
||||
Employees []*Employee `json:"employees,omitempty"`
|
||||
ID int64 `json:"id" bun:"id,pk"`
|
||||
Name string `json:"name"`
|
||||
Employees []*Employee `json:"employees,omitempty"`
|
||||
}
|
||||
|
||||
func (d Department) TableName() string { return "departments" }
|
||||
@@ -225,9 +235,9 @@ func (t Task) TableName() string { return "tasks" }
|
||||
func (t Task) GetIDName() string { return "ID" }
|
||||
|
||||
type Comment struct {
|
||||
ID int64 `json:"id" bun:"id,pk"`
|
||||
Text string `json:"text"`
|
||||
TaskID int64 `json:"task_id"`
|
||||
ID int64 `json:"id" bun:"id,pk"`
|
||||
Text string `json:"text"`
|
||||
TaskID int64 `json:"task_id"`
|
||||
}
|
||||
|
||||
func (c Comment) TableName() string { return "comments" }
|
||||
|
||||
@@ -0,0 +1,92 @@
|
||||
package funcspec
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/uptrace/bun/driver/sqliteshim"
|
||||
)
|
||||
|
||||
const joinedBaseSQL = "SELECT p.id, p.name, c.name AS country_name FROM province_state p LEFT JOIN country c ON c.id = p.rid_country"
|
||||
|
||||
// funcspec filters are appended to author-written SQL with no model, so columns
|
||||
// are used verbatim: clients disambiguate by sending "alias.column", and the
|
||||
// dot must survive ValidSQL and every filter path.
|
||||
func TestApplyFilters_QualifiedColumnsPreservedWithJoin(t *testing.T) {
|
||||
h := NewHandler(&MockDatabase{})
|
||||
|
||||
cases := []struct {
|
||||
name string
|
||||
params *RequestParameters
|
||||
want string
|
||||
}{
|
||||
{"field filter", &RequestParameters{FieldFilters: map[string]string{"p.name": "abc"}}, "p.name = abc"},
|
||||
{"search filter", &RequestParameters{SearchFilters: map[string]string{"p.name": "abc"}}, "CAST(p.name AS TEXT) ILIKE '%abc%'"},
|
||||
{"search op eq", &RequestParameters{SearchOps: map[string]FilterOperator{"p.name": {Operator: "eq", Value: "abc", Logic: "AND"}}}, "p.name = 'abc'"},
|
||||
{"search op contains", &RequestParameters{SearchOps: map[string]FilterOperator{"p.name": {Operator: "contains", Value: "abc", Logic: "AND"}}}, "CAST(p.name AS TEXT) ILIKE '%abc%'"},
|
||||
}
|
||||
for _, c := range cases {
|
||||
t.Run(c.name, func(t *testing.T) {
|
||||
got := h.ApplyFilters(joinedBaseSQL, c.params)
|
||||
if !strings.Contains(got, c.want) {
|
||||
t.Fatalf("SQL %q does not contain %q", got, c.want)
|
||||
}
|
||||
if !strings.Contains(got, "LEFT JOIN country") || !strings.Contains(got, " WHERE ") {
|
||||
t.Fatalf("join/where lost: %q", got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyFilters_WithAndWithoutJoin_RealQuery(t *testing.T) {
|
||||
sqldb, err := sql.Open(sqliteshim.ShimName, "file:funcspecjoin?mode=memory&cache=private")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer sqldb.Close()
|
||||
for _, stmt := range []string{
|
||||
"CREATE TABLE country (id INTEGER PRIMARY KEY, name TEXT)",
|
||||
"CREATE TABLE province_state (id INTEGER PRIMARY KEY, name TEXT, rid_country INTEGER)",
|
||||
"INSERT INTO country VALUES (1,'Abcland'),(2,'Other')",
|
||||
"INSERT INTO province_state VALUES (1,'abc one',1),(2,'xyz two',1),(3,'nope',2)",
|
||||
} {
|
||||
if _, err := sqldb.Exec(stmt); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
h := NewHandler(&MockDatabase{})
|
||||
count := func(base string, params *RequestParameters) (int, error) {
|
||||
rows, err := sqldb.Query(h.ApplyFilters(base, params))
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
defer rows.Close()
|
||||
n := 0
|
||||
for rows.Next() {
|
||||
n++
|
||||
}
|
||||
return n, rows.Err()
|
||||
}
|
||||
|
||||
const noJoin = "SELECT p.id, p.name FROM province_state p"
|
||||
params := func(col string) *RequestParameters {
|
||||
return &RequestParameters{SearchOps: map[string]FilterOperator{col: {Operator: "eq", Value: "nope", Logic: "AND"}}}
|
||||
}
|
||||
|
||||
if n, err := count(noJoin, params("name")); err != nil || n != 1 {
|
||||
t.Fatalf("no join, unqualified: n=%d err=%v", n, err)
|
||||
}
|
||||
if n, err := count(joinedBaseSQL, params("p.name")); err != nil || n != 1 {
|
||||
t.Fatalf("join, qualified: n=%d err=%v", n, err)
|
||||
}
|
||||
if n, err := count(joinedBaseSQL, params("c.name")); err != nil || n != 0 {
|
||||
t.Fatalf("join, joined-table column: n=%d err=%v", n, err)
|
||||
}
|
||||
// Documented limit: with a join in the author's SQL, an unqualified shared
|
||||
// column is ambiguous and the client must send "alias.column".
|
||||
if _, err := count(joinedBaseSQL, params("name")); err == nil || !strings.Contains(strings.ToLower(err.Error()), "ambiguous") {
|
||||
t.Fatalf("expected ambiguous error for unqualified shared column, got %v", err)
|
||||
}
|
||||
}
|
||||
+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,29 @@ type Provider interface {
|
||||
Handler() http.Handler
|
||||
}
|
||||
|
||||
// AIProxyRecorder is optionally implemented by providers that record aiproxy traffic.
|
||||
// pkg/aiproxy uses it when the global provider implements it.
|
||||
type AIProxyRecorder interface {
|
||||
// RecordAIProxy records one proxied (or refused) request. statusClass is "2xx".."5xx",
|
||||
// outcome is ok, upstream_error, denied or rate_limited, model may be empty.
|
||||
RecordAIProxy(upstream, kind, model, statusClass, outcome string, duration time.Duration, promptTokens, completionTokens int64)
|
||||
}
|
||||
|
||||
// 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)
|
||||
}
|
||||
}
|
||||
+243
-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
|
||||
@@ -26,10 +32,20 @@ type PrometheusProvider struct {
|
||||
eventDuration *prometheus.HistogramVec
|
||||
eventQueueSize prometheus.Gauge
|
||||
panicsTotal *prometheus.CounterVec
|
||||
aiRequests *prometheus.CounterVec
|
||||
aiDuration *prometheus.HistogramVec
|
||||
aiTokens *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 +71,7 @@ func NewPrometheusProvider(cfg *Config) *PrometheusProvider {
|
||||
}
|
||||
|
||||
p := &PrometheusProvider{
|
||||
enabled: cfg.Enabled,
|
||||
requestDuration: promauto.NewHistogramVec(
|
||||
prometheus.HistogramOpts{
|
||||
Name: metricName("http_request_duration_seconds"),
|
||||
@@ -148,13 +165,40 @@ func NewPrometheusProvider(cfg *Config) *PrometheusProvider {
|
||||
},
|
||||
[]string{"method"},
|
||||
),
|
||||
aiRequests: promauto.NewCounterVec(
|
||||
prometheus.CounterOpts{
|
||||
Name: metricName("aiproxy_requests_total"),
|
||||
Help: "Total number of requests handled by the AI proxy",
|
||||
},
|
||||
[]string{"upstream", "kind", "model", "status", "outcome"},
|
||||
),
|
||||
aiDuration: promauto.NewHistogramVec(
|
||||
prometheus.HistogramOpts{
|
||||
Name: metricName("aiproxy_request_duration_seconds"),
|
||||
Help: "AI proxy request duration in seconds",
|
||||
Buckets: cfg.HTTPRequestBuckets,
|
||||
},
|
||||
[]string{"upstream", "kind"},
|
||||
),
|
||||
aiTokens: promauto.NewCounterVec(
|
||||
prometheus.CounterOpts{
|
||||
Name: metricName("aiproxy_tokens_total"),
|
||||
Help: "Tokens reported by AI proxy upstreams",
|
||||
},
|
||||
[]string{"upstream", "model", "type"},
|
||||
),
|
||||
|
||||
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 +210,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 +239,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 +280,130 @@ 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()
|
||||
}
|
||||
|
||||
// RecordAIProxy implements AIProxyRecorder
|
||||
func (p *PrometheusProvider) RecordAIProxy(upstream, kind, model, statusClass, outcome string, duration time.Duration, promptTokens, completionTokens int64) {
|
||||
if !p.enabled {
|
||||
return
|
||||
}
|
||||
p.aiRequests.WithLabelValues(upstream, kind, model, statusClass, outcome).Inc()
|
||||
p.aiDuration.WithLabelValues(upstream, kind).Observe(duration.Seconds())
|
||||
if promptTokens > 0 {
|
||||
p.aiTokens.WithLabelValues(upstream, model, "prompt").Add(float64(promptTokens))
|
||||
}
|
||||
if completionTokens > 0 {
|
||||
p.aiTokens.WithLabelValues(upstream, model, "completion").Add(float64(completionTokens))
|
||||
}
|
||||
}
|
||||
|
||||
// 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 +421,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 +443,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 +460,90 @@ 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.aiRequests.Reset()
|
||||
p.aiDuration.Reset()
|
||||
p.aiTokens.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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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`.
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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,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
|
||||
@@ -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}
|
||||
|
||||
@@ -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
|
||||
@@ -559,6 +566,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 +734,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)
|
||||
|
||||
+37
-10
@@ -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,
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -0,0 +1,183 @@
|
||||
package resolvespec
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/uptrace/bun"
|
||||
"github.com/uptrace/bun/dialect/sqlitedialect"
|
||||
"github.com/uptrace/bun/driver/sqliteshim"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
|
||||
)
|
||||
|
||||
type joinCountry struct {
|
||||
bun.BaseModel `bun:"table:country,alias:country"`
|
||||
ID int64 `bun:"id,pk"`
|
||||
Name string `bun:"name"`
|
||||
}
|
||||
|
||||
type joinProvince struct {
|
||||
bun.BaseModel `bun:"table:province_state,alias:province_state"`
|
||||
ID int64 `bun:"id,pk"`
|
||||
Name string `bun:"name"`
|
||||
Abbreviation string `bun:"abbreviation"`
|
||||
RidCountry int64 `bun:"rid_country"`
|
||||
}
|
||||
|
||||
// SQLite has no ILIKE; the ilike operator is covered by SQL-string tests, and
|
||||
// LIKE here proves the CAST(... AS TEXT) wrapping executes with qualified columns.
|
||||
func setupJoinDB(t *testing.T) *bun.DB {
|
||||
t.Helper()
|
||||
sqldb, err := sql.Open(sqliteshim.ShimName, "file:filterjoin?mode=memory&cache=private")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
db := bun.NewDB(sqldb, sqlitedialect.New())
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
ctx := context.Background()
|
||||
for _, m := range []interface{}{(*joinCountry)(nil), (*joinProvince)(nil)} {
|
||||
if _, err := db.NewCreateTable().Model(m).IfNotExists().Exec(ctx); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if _, err := db.NewInsert().Model(&[]joinCountry{{ID: 1, Name: "Abcland"}, {ID: 2, Name: "Other"}}).Exec(ctx); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := db.NewInsert().Model(&[]joinProvince{
|
||||
{ID: 1, Name: "abc one", Abbreviation: "A1", RidCountry: 1},
|
||||
{ID: 2, Name: "xyz two", Abbreviation: "ABC", RidCountry: 1},
|
||||
{ID: 3, Name: "nope", Abbreviation: "N3", RidCountry: 2},
|
||||
}).Exec(ctx); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return db
|
||||
}
|
||||
|
||||
const joinSQL = "LEFT JOIN country AS rel_rid_country ON rel_rid_country.id = province_state.rid_country"
|
||||
|
||||
func countWith(t *testing.T, db *bun.DB, withJoin bool, filters []common.FilterOption, alias string) (int, error) {
|
||||
t.Helper()
|
||||
h := &Handler{}
|
||||
var q common.SelectQuery = database.NewBunAdapter(db).NewSelect().Model(&[]*joinProvince{})
|
||||
if withJoin {
|
||||
q = q.Join(joinSQL)
|
||||
}
|
||||
q = h.applyFilters(q, filters, &joinProvince{}, alias)
|
||||
return q.Count(context.Background())
|
||||
}
|
||||
|
||||
func TestFilters_AmbiguousColumnWithJoin_RealQuery(t *testing.T) {
|
||||
db := setupJoinDB(t)
|
||||
|
||||
filters := []common.FilterOption{
|
||||
{Column: "name", Operator: "like", Value: "%abc%"},
|
||||
{Column: "abbreviation", Operator: "like", Value: "%abc%", LogicOperator: "OR"},
|
||||
}
|
||||
|
||||
t.Run("unqualified with join is ambiguous (the bug)", func(t *testing.T) {
|
||||
_, err := countWith(t, db, true, filters, "")
|
||||
if err == nil || !strings.Contains(strings.ToLower(err.Error()), "ambiguous") {
|
||||
t.Fatalf("expected ambiguous column error, got %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("qualified with join", func(t *testing.T) {
|
||||
n, err := countWith(t, db, true, filters, "province_state")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if n != 2 {
|
||||
t.Fatalf("count = %d, want 2", n)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("qualified without join", func(t *testing.T) {
|
||||
n, err := countWith(t, db, false, filters, "province_state")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if n != 2 {
|
||||
t.Fatalf("count = %d, want 2", n)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("unqualified without join still works", func(t *testing.T) {
|
||||
n, err := countWith(t, db, false, filters, "")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if n != 2 {
|
||||
t.Fatalf("count = %d, want 2", n)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestFilters_AllOperatorsWithJoin_RealQuery(t *testing.T) {
|
||||
db := setupJoinDB(t)
|
||||
cases := []struct {
|
||||
name string
|
||||
filter common.FilterOption
|
||||
want int
|
||||
}{
|
||||
{"eq", common.FilterOption{Column: "name", Operator: "eq", Value: "nope"}, 1},
|
||||
{"neq", common.FilterOption{Column: "name", Operator: "neq", Value: "nope"}, 2},
|
||||
{"like", common.FilterOption{Column: "name", Operator: "like", Value: "abc%"}, 1},
|
||||
{"in", common.FilterOption{Column: "name", Operator: "in", Value: []string{"nope", "xyz two"}}, 2},
|
||||
{"gt", common.FilterOption{Column: "id", Operator: "gt", Value: 1}, 2},
|
||||
{"gte", common.FilterOption{Column: "id", Operator: "gte", Value: 2}, 2},
|
||||
{"lt", common.FilterOption{Column: "id", Operator: "lt", Value: 3}, 2},
|
||||
{"lte", common.FilterOption{Column: "id", Operator: "lte", Value: 1}, 1},
|
||||
}
|
||||
for _, c := range cases {
|
||||
t.Run(c.name, func(t *testing.T) {
|
||||
n, err := countWith(t, db, true, []common.FilterOption{c.filter}, "province_state")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if n != c.want {
|
||||
t.Fatalf("count = %d, want %d", n, c.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// A filter on the joined table's own column, sent already qualified, must pass through untouched.
|
||||
func TestFilters_JoinedTableColumnPassesThrough_RealQuery(t *testing.T) {
|
||||
db := setupJoinDB(t)
|
||||
n, err := countWith(t, db, true, []common.FilterOption{
|
||||
{Column: "rel_rid_country.name", Operator: "eq", Value: "Abcland"},
|
||||
}, "province_state")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if n != 2 {
|
||||
t.Fatalf("count = %d, want 2", n)
|
||||
}
|
||||
}
|
||||
|
||||
// Client sends "province_state.name": it must survive validation and be applied.
|
||||
func TestFilters_ClientQualifiedColumn_NotDropped(t *testing.T) {
|
||||
db := setupJoinDB(t)
|
||||
model := &joinProvince{}
|
||||
|
||||
opts := common.RequestOptions{Filters: []common.FilterOption{
|
||||
{Column: "province_state.name", Operator: "like", Value: "%abc%"},
|
||||
}}
|
||||
common.NormalizeMainTableFilters(model, "public.province_state", &opts)
|
||||
opts = common.NewColumnValidator(model).FilterRequestOptions(opts)
|
||||
if len(opts.Filters) != 1 {
|
||||
t.Fatalf("filter was dropped: %+v", opts.Filters)
|
||||
}
|
||||
|
||||
n, err := countWith(t, db, true, opts.Filters, "province_state")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if n != 1 {
|
||||
t.Fatalf("count = %d, want 1 (filter must narrow the result)", n)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,65 @@
|
||||
package resolvespec
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
)
|
||||
|
||||
func TestBuildFilterConditionAlias_QualifiesModelColumns(t *testing.T) {
|
||||
h := &Handler{}
|
||||
model := jsonColModel{}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
filter common.FilterOption
|
||||
want string
|
||||
}{
|
||||
{"eq", common.FilterOption{Column: "name", Operator: "eq", Value: "x"}, `"province_state"."name" = ?`},
|
||||
{"ilike", common.FilterOption{Column: "name", Operator: "ilike", Value: "%a%"}, `CAST("province_state"."name" AS TEXT) ILIKE ?`},
|
||||
{"like", common.FilterOption{Column: "name", Operator: "like", Value: "%a%"}, `CAST("province_state"."name" AS TEXT) LIKE ?`},
|
||||
{"in", common.FilterOption{Column: "name", Operator: "in", Value: []string{"a", "b"}}, `"province_state"."name" IN (?,?)`},
|
||||
{"non-model column untouched", common.FilterOption{Column: "other", Operator: "eq", Value: 1}, `other = ?`},
|
||||
{"already qualified untouched", common.FilterOption{Column: "rel.name", Operator: "eq", Value: 1}, `rel.name = ?`},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got, _ := h.buildFilterConditionAlias(tt.filter, model, "province_state")
|
||||
if got != tt.want {
|
||||
t.Fatalf("got %q, want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildFilterConditionAlias_NoAliasUnchanged(t *testing.T) {
|
||||
h := &Handler{}
|
||||
got, _ := h.buildFilterConditionAlias(common.FilterOption{Column: "name", Operator: "ilike", Value: "%a%"}, jsonColModel{}, "")
|
||||
if got != "CAST(name AS TEXT) ILIKE ?" {
|
||||
t.Fatalf("got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyFilter_QualifiesWithAlias(t *testing.T) {
|
||||
h := &Handler{}
|
||||
q := &jsonCapQuery{}
|
||||
h.applyFilter(q, common.FilterOption{Column: "name", Operator: "ilike", Value: "%a%"}, jsonColModel{}, "province_state")
|
||||
c := q.only(t)
|
||||
if c.query != `CAST("province_state"."name" AS TEXT) ILIKE ?` {
|
||||
t.Fatalf("query = %q", c.query)
|
||||
}
|
||||
}
|
||||
|
||||
func TestApplyFilters_OrGroupQualified(t *testing.T) {
|
||||
h := &Handler{}
|
||||
q := &jsonCapQuery{}
|
||||
h.applyFilters(q, []common.FilterOption{
|
||||
{Column: "name", Operator: "ilike", Value: "%a%"},
|
||||
{Column: "id", Operator: "eq", Value: 1, LogicOperator: "OR"},
|
||||
}, jsonColModel{}, "t")
|
||||
c := q.only(t)
|
||||
want := `(CAST("t"."name" AS TEXT) ILIKE ? OR "t"."id" = ?)`
|
||||
if c.query != want {
|
||||
t.Fatalf("query = %q, want %q", c.query, want)
|
||||
}
|
||||
}
|
||||
+36
-15
@@ -172,6 +172,9 @@ func (h *Handler) Handle(w common.ResponseWriter, r common.Request, params map[s
|
||||
// Add request-scoped data to context
|
||||
ctx = WithRequestData(ctx, schema, entity, tableName, model, modelPtr)
|
||||
|
||||
// Accept "<main table or alias>.<column>" for model columns before validation drops them
|
||||
common.NormalizeMainTableFilters(model, tableName, &req.Options)
|
||||
|
||||
// Validate and filter columns in options (log warnings for invalid columns)
|
||||
validator := common.NewColumnValidator(model)
|
||||
req.Options = validator.FilterRequestOptions(req.Options)
|
||||
@@ -411,7 +414,7 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
||||
}
|
||||
|
||||
// Apply filters with proper grouping for OR logic
|
||||
query = h.applyFilters(query, options.Filters, model)
|
||||
query = h.applyFilters(query, options.Filters, model, common.MainTableAlias(model, tableName))
|
||||
|
||||
// Apply custom operators
|
||||
for _, customOp := range options.CustomOperators {
|
||||
@@ -558,7 +561,7 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
||||
|
||||
// Apply the same filters as the main query
|
||||
for _, filter := range options.Filters {
|
||||
rowNumQuery = h.applyFilter(rowNumQuery, filter, model)
|
||||
rowNumQuery = h.applyFilter(rowNumQuery, filter, model, common.MainTableAlias(model, tableName))
|
||||
}
|
||||
|
||||
// Apply custom operators
|
||||
@@ -824,6 +827,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 +975,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 +1132,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 +1328,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 +1514,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 +1670,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 {
|
||||
@@ -1926,7 +1935,7 @@ func (h *Handler) executeDelete(ctx context.Context, tx common.Database, hookCtx
|
||||
// applyFilters applies all filters with proper grouping for OR logic
|
||||
// Groups consecutive OR filters together to ensure proper query precedence
|
||||
// Example: [A, B(OR), C(OR), D(AND)] => WHERE (A OR B OR C) AND D
|
||||
func (h *Handler) applyFilters(query common.SelectQuery, filters []common.FilterOption, model interface{}) common.SelectQuery {
|
||||
func (h *Handler) applyFilters(query common.SelectQuery, filters []common.FilterOption, model interface{}, alias string) common.SelectQuery {
|
||||
if len(filters) == 0 {
|
||||
return query
|
||||
}
|
||||
@@ -1946,11 +1955,11 @@ func (h *Handler) applyFilters(query common.SelectQuery, filters []common.Filter
|
||||
}
|
||||
|
||||
// Apply the OR group as a single grouped WHERE clause
|
||||
query = h.applyFilterGroup(query, orGroup, model)
|
||||
query = h.applyFilterGroup(query, orGroup, model, alias)
|
||||
i = j
|
||||
} else {
|
||||
// Single filter with AND logic (or first filter)
|
||||
condition, args := h.buildFilterCondition(filters[i], model)
|
||||
condition, args := h.buildFilterConditionAlias(filters[i], model, alias)
|
||||
if condition != "" {
|
||||
query = query.Where(condition, args...)
|
||||
}
|
||||
@@ -1963,7 +1972,7 @@ func (h *Handler) applyFilters(query common.SelectQuery, filters []common.Filter
|
||||
|
||||
// applyFilterGroup applies a group of filters that should be OR'd together
|
||||
// Always wraps them in parentheses and applies as a single WHERE clause
|
||||
func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.FilterOption, model interface{}) common.SelectQuery {
|
||||
func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.FilterOption, model interface{}, alias string) common.SelectQuery {
|
||||
if len(filters) == 0 {
|
||||
return query
|
||||
}
|
||||
@@ -1973,7 +1982,7 @@ func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.Fi
|
||||
var args []interface{}
|
||||
|
||||
for _, filter := range filters {
|
||||
condition, filterArgs := h.buildFilterCondition(filter, model)
|
||||
condition, filterArgs := h.buildFilterConditionAlias(filter, model, alias)
|
||||
if condition != "" {
|
||||
conditions = append(conditions, condition)
|
||||
args = append(args, filterArgs...)
|
||||
@@ -1999,6 +2008,12 @@ func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.Fi
|
||||
// or the dotted data.x shorthand for a JSON column) resolve to a safe,
|
||||
// parameterised expression before the ordinary operator handling below.
|
||||
func (h *Handler) buildFilterCondition(filter common.FilterOption, model interface{}) (conditionString string, conditionArgs []interface{}) {
|
||||
return h.buildFilterConditionAlias(filter, model, "")
|
||||
}
|
||||
|
||||
// buildFilterConditionAlias is buildFilterCondition with plain model columns
|
||||
// qualified by the main table alias, so joins from preloads can't make them ambiguous.
|
||||
func (h *Handler) buildFilterConditionAlias(filter common.FilterOption, model interface{}, alias string) (conditionString string, conditionArgs []interface{}) {
|
||||
var condition string
|
||||
var args []interface{}
|
||||
|
||||
@@ -2006,6 +2021,9 @@ func (h *Handler) buildFilterCondition(filter common.FilterOption, model interfa
|
||||
return cond, jargs
|
||||
}
|
||||
|
||||
rawColumn := filter.Column
|
||||
filter.Column = common.QualifyModelColumn(model, alias, filter.Column)
|
||||
|
||||
switch filter.Operator {
|
||||
case "eq", "=":
|
||||
condition = fmt.Sprintf("%s = ?", filter.Column)
|
||||
@@ -2026,10 +2044,10 @@ func (h *Handler) buildFilterCondition(filter common.FilterOption, model interfa
|
||||
condition = fmt.Sprintf("%s <= ?", filter.Column)
|
||||
args = []interface{}{filter.Value}
|
||||
case "like":
|
||||
condition = fmt.Sprintf("%s LIKE ?", likeColumn(filter.Column, model))
|
||||
condition = fmt.Sprintf("%s LIKE ?", likeColumn(filter.Column, rawColumn, model))
|
||||
args = []interface{}{filter.Value}
|
||||
case "ilike":
|
||||
condition = fmt.Sprintf("%s ILIKE ?", likeColumn(filter.Column, model))
|
||||
condition = fmt.Sprintf("%s ILIKE ?", likeColumn(filter.Column, rawColumn, model))
|
||||
args = []interface{}{filter.Value}
|
||||
case "in":
|
||||
condition, args = common.BuildInCondition(filter.Column, filter.Value)
|
||||
@@ -2067,14 +2085,14 @@ func (h *Handler) buildFilterCondition(filter common.FilterOption, model interfa
|
||||
// CAST(... AS TEXT) would switch to case-sensitive matching and defeat a
|
||||
// citext index. Every other column is cast to TEXT so LIKE/ILIKE also works
|
||||
// against date/time/timestamp and numeric columns.
|
||||
func likeColumn(column string, model interface{}) string {
|
||||
if reflection.IsCitextColumn(model, column) {
|
||||
func likeColumn(column, rawColumn string, model interface{}) string {
|
||||
if reflection.IsCitextColumn(model, rawColumn) {
|
||||
return column
|
||||
}
|
||||
return fmt.Sprintf("CAST(%s AS TEXT)", column)
|
||||
}
|
||||
|
||||
func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOption, model interface{}) common.SelectQuery {
|
||||
func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOption, model interface{}, alias string) common.SelectQuery {
|
||||
// Determine which method to use based on LogicOperator
|
||||
useOrLogic := strings.EqualFold(filter.LogicOperator, "OR")
|
||||
|
||||
@@ -2088,6 +2106,9 @@ func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOpti
|
||||
return query.Where(cond, jargs...)
|
||||
}
|
||||
|
||||
rawColumn := filter.Column
|
||||
filter.Column = common.QualifyModelColumn(model, alias, filter.Column)
|
||||
|
||||
switch filter.Operator {
|
||||
case "eq", "=":
|
||||
condition = fmt.Sprintf("%s = ?", filter.Column)
|
||||
@@ -2108,10 +2129,10 @@ func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOpti
|
||||
condition = fmt.Sprintf("%s <= ?", filter.Column)
|
||||
args = []interface{}{filter.Value}
|
||||
case "like":
|
||||
condition = fmt.Sprintf("%s LIKE ?", likeColumn(filter.Column, model))
|
||||
condition = fmt.Sprintf("%s LIKE ?", likeColumn(filter.Column, rawColumn, model))
|
||||
args = []interface{}{filter.Value}
|
||||
case "ilike":
|
||||
condition = fmt.Sprintf("%s ILIKE ?", likeColumn(filter.Column, model))
|
||||
condition = fmt.Sprintf("%s ILIKE ?", likeColumn(filter.Column, rawColumn, model))
|
||||
args = []interface{}{filter.Value}
|
||||
case "in":
|
||||
condition, args = common.BuildInCondition(filter.Column, filter.Value)
|
||||
@@ -2517,7 +2538,7 @@ func (h *Handler) applyPreloads(model interface{}, query common.SelectQuery, pre
|
||||
|
||||
if len(preload.Filters) > 0 {
|
||||
for _, filter := range preload.Filters {
|
||||
sq = h.applyFilter(sq, filter, nil)
|
||||
sq = h.applyFilter(sq, filter, nil, "")
|
||||
}
|
||||
}
|
||||
if len(preload.Sort) > 0 {
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
//go:build integration
|
||||
// +build integration
|
||||
|
||||
package resolvespec
|
||||
@@ -22,12 +23,12 @@ import (
|
||||
|
||||
// Test models
|
||||
type TestUser struct {
|
||||
ID uint `gorm:"primaryKey" json:"id"`
|
||||
Name string `gorm:"not null" json:"name"`
|
||||
Email string `gorm:"uniqueIndex;not null" json:"email"`
|
||||
Age int `json:"age"`
|
||||
Active bool `gorm:"default:true" json:"active"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
ID uint `gorm:"primaryKey" json:"id"`
|
||||
Name string `gorm:"not null" json:"name"`
|
||||
Email string `gorm:"uniqueIndex;not null" json:"email"`
|
||||
Age int `json:"age"`
|
||||
Active bool `gorm:"default:true" json:"active"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
Posts []TestPost `gorm:"foreignKey:UserID" json:"posts,omitempty"`
|
||||
}
|
||||
|
||||
@@ -36,13 +37,13 @@ func (TestUser) TableName() string {
|
||||
}
|
||||
|
||||
type TestPost struct {
|
||||
ID uint `gorm:"primaryKey" json:"id"`
|
||||
UserID uint `gorm:"not null" json:"user_id"`
|
||||
Title string `gorm:"not null" json:"title"`
|
||||
Content string `json:"content"`
|
||||
Published bool `gorm:"default:false" json:"published"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
User *TestUser `gorm:"foreignKey:UserID" json:"user,omitempty"`
|
||||
ID uint `gorm:"primaryKey" json:"id"`
|
||||
UserID uint `gorm:"not null" json:"user_id"`
|
||||
Title string `gorm:"not null" json:"title"`
|
||||
Content string `json:"content"`
|
||||
Published bool `gorm:"default:false" json:"published"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
User *TestUser `gorm:"foreignKey:UserID" json:"user,omitempty"`
|
||||
Comments []TestComment `gorm:"foreignKey:PostID" json:"comments,omitempty"`
|
||||
}
|
||||
|
||||
@@ -55,7 +56,7 @@ type TestComment struct {
|
||||
PostID uint `gorm:"not null" json:"post_id"`
|
||||
Content string `gorm:"not null" json:"content"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
Post *TestPost `gorm:"foreignKey:PostID" json:"post,omitempty"`
|
||||
Post *TestPost `gorm:"foreignKey:PostID" json:"post,omitempty"`
|
||||
}
|
||||
|
||||
func (TestComment) TableName() string {
|
||||
|
||||
@@ -142,7 +142,7 @@ func TestApplyFilter_JSONColumn(t *testing.T) {
|
||||
q := &jsonCapQuery{}
|
||||
h.applyFilter(q, common.FilterOption{
|
||||
Column: "data->>'tier'", Operator: "in", Value: []string{"a", "b"}, LogicOperator: "OR",
|
||||
}, model)
|
||||
}, model, "")
|
||||
c := q.only(t)
|
||||
if c.method != "WhereOr" || c.query != `("data" #>> ?::text[]) IN (?,?)` {
|
||||
t.Fatalf("call = %+v", c)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,180 @@
|
||||
package restheadspec
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"testing"
|
||||
|
||||
"github.com/uptrace/bun"
|
||||
"github.com/uptrace/bun/dialect/sqlitedialect"
|
||||
"github.com/uptrace/bun/driver/sqliteshim"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
|
||||
)
|
||||
|
||||
type joinCountry struct {
|
||||
bun.BaseModel `bun:"table:country,alias:country"`
|
||||
ID int64 `bun:"id,pk"`
|
||||
Name string `bun:"name"`
|
||||
}
|
||||
|
||||
type joinProvince struct {
|
||||
bun.BaseModel `bun:"table:province_state,alias:province_state"`
|
||||
ID int64 `bun:"id,pk"`
|
||||
Name string `bun:"name"`
|
||||
Abbreviation string `bun:"abbreviation"`
|
||||
RidCountry int64 `bun:"rid_country"`
|
||||
}
|
||||
|
||||
func setupJoinDB(t *testing.T) *bun.DB {
|
||||
t.Helper()
|
||||
sqldb, err := sql.Open(sqliteshim.ShimName, "file:rhsfilterjoin?mode=memory&cache=private")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
db := bun.NewDB(sqldb, sqlitedialect.New())
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
ctx := context.Background()
|
||||
for _, m := range []interface{}{(*joinCountry)(nil), (*joinProvince)(nil)} {
|
||||
if _, err := db.NewCreateTable().Model(m).IfNotExists().Exec(ctx); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
if _, err := db.NewInsert().Model(&[]joinCountry{{ID: 1, Name: "Abcland"}, {ID: 2, Name: "Other"}}).Exec(ctx); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := db.NewInsert().Model(&[]joinProvince{
|
||||
{ID: 1, Name: "abc one", Abbreviation: "A1", RidCountry: 1},
|
||||
{ID: 2, Name: "xyz two", Abbreviation: "ABC", RidCountry: 1},
|
||||
{ID: 3, Name: "nope", Abbreviation: "N3", RidCountry: 2},
|
||||
}).Exec(ctx); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return db
|
||||
}
|
||||
|
||||
const joinSQL = "LEFT JOIN country AS rel_rid_country ON rel_rid_country.id = province_state.rid_country"
|
||||
|
||||
// countFilters applies filters the way handleRead does (single AND filters via
|
||||
// applyFilter, consecutive OR filters via applyOrFilterGroup) and counts.
|
||||
func countFilters(t *testing.T, db *bun.DB, withJoin bool, filters []common.FilterOption) (int, error) {
|
||||
t.Helper()
|
||||
h := &Handler{}
|
||||
model := &joinProvince{}
|
||||
var q common.SelectQuery = database.NewBunAdapter(db).NewSelect().Model(&[]*joinProvince{})
|
||||
if withJoin {
|
||||
q = q.Join(joinSQL)
|
||||
}
|
||||
for i := 0; i < len(filters); {
|
||||
f := filters[i]
|
||||
castInfo := h.ValidateAndAdjustFilterForColumnType(&f, model)
|
||||
if f.LogicOperator == "OR" {
|
||||
group := []*common.FilterOption{&f}
|
||||
info := []ColumnCastInfo{castInfo}
|
||||
j := i + 1
|
||||
for j < len(filters) && filters[j].LogicOperator == "OR" {
|
||||
g := filters[j]
|
||||
info = append(info, h.ValidateAndAdjustFilterForColumnType(&g, model))
|
||||
group = append(group, &g)
|
||||
j++
|
||||
}
|
||||
q = h.applyOrFilterGroup(q, group, info, "public.province_state", model)
|
||||
i = j
|
||||
continue
|
||||
}
|
||||
q = h.applyFilter(q, f, "public.province_state", castInfo.NeedsCast, "AND", model)
|
||||
i++
|
||||
}
|
||||
return q.Count(context.Background())
|
||||
}
|
||||
|
||||
func TestRHSFilters_WithAndWithoutJoin_RealQuery(t *testing.T) {
|
||||
db := setupJoinDB(t)
|
||||
cases := []struct {
|
||||
name string
|
||||
filters []common.FilterOption
|
||||
want int
|
||||
}{
|
||||
{"eq", []common.FilterOption{{Column: "name", Operator: "eq", Value: "nope"}}, 1},
|
||||
{"neq", []common.FilterOption{{Column: "name", Operator: "neq", Value: "nope"}}, 2},
|
||||
{"like", []common.FilterOption{{Column: "name", Operator: "like", Value: "%abc%"}}, 1},
|
||||
{"in", []common.FilterOption{{Column: "name", Operator: "in", Value: []string{"nope", "xyz two"}}}, 2},
|
||||
{"gt", []common.FilterOption{{Column: "id", Operator: "gt", Value: 1}}, 2},
|
||||
{"between", []common.FilterOption{{Column: "id", Operator: "between", Value: []interface{}{0, 3}}}, 2},
|
||||
{"between_inclusive", []common.FilterOption{{Column: "id", Operator: "between_inclusive", Value: []interface{}{1, 3}}}, 3},
|
||||
{"is_not_null", []common.FilterOption{{Column: "name", Operator: "is_not_null"}}, 3},
|
||||
{"or group", []common.FilterOption{
|
||||
{Column: "name", Operator: "like", Value: "%abc%", LogicOperator: "OR"},
|
||||
{Column: "abbreviation", Operator: "like", Value: "%abc%", LogicOperator: "OR"},
|
||||
}, 2},
|
||||
{"or group then and", []common.FilterOption{
|
||||
{Column: "name", Operator: "like", Value: "%abc%", LogicOperator: "OR"},
|
||||
{Column: "abbreviation", Operator: "like", Value: "%abc%", LogicOperator: "OR"},
|
||||
{Column: "id", Operator: "eq", Value: 2},
|
||||
}, 1},
|
||||
}
|
||||
for _, c := range cases {
|
||||
for _, withJoin := range []bool{false, true} {
|
||||
name := c.name + "/no_join"
|
||||
if withJoin {
|
||||
name = c.name + "/join"
|
||||
}
|
||||
t.Run(name, func(t *testing.T) {
|
||||
n, err := countFilters(t, db, withJoin, c.filters)
|
||||
if err != nil {
|
||||
t.Fatalf("query failed (ambiguous column?): %v", err)
|
||||
}
|
||||
if n != c.want {
|
||||
t.Fatalf("count = %d, want %d", n, c.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRHSFilters_JoinedTableColumnPassesThrough(t *testing.T) {
|
||||
db := setupJoinDB(t)
|
||||
n, err := countFilters(t, db, true, []common.FilterOption{
|
||||
{Column: "rel_rid_country.name", Operator: "eq", Value: "Abcland"},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if n != 2 {
|
||||
t.Fatalf("count = %d, want 2", n)
|
||||
}
|
||||
}
|
||||
|
||||
// Client sends "province_state.name": it must survive validation and be applied.
|
||||
func TestRHSFilters_ClientQualifiedColumn_NotDropped(t *testing.T) {
|
||||
db := setupJoinDB(t)
|
||||
h := &Handler{}
|
||||
model := &joinProvince{}
|
||||
|
||||
opts := ExtendedRequestOptions{}
|
||||
opts.Filters = []common.FilterOption{{Column: "province_state.name", Operator: "like", Value: "%abc%"}}
|
||||
common.NormalizeMainTableFilters(model, "public.province_state", &opts.RequestOptions)
|
||||
opts = h.filterExtendedOptions(common.NewColumnValidator(model), opts, model)
|
||||
if len(opts.Filters) != 1 || opts.Filters[0].Column != "name" {
|
||||
t.Fatalf("filter was dropped or not normalised: %+v", opts.Filters)
|
||||
}
|
||||
|
||||
n, err := countFilters(t, db, true, opts.Filters)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if n != 1 {
|
||||
t.Fatalf("count = %d, want 1", n)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRHSFilters_ILikeSQLQualified(t *testing.T) {
|
||||
h := &Handler{}
|
||||
q := &jsonCapQuery{}
|
||||
f := common.FilterOption{Column: "name", Operator: "ilike", Value: "%abc%"}
|
||||
h.applyFilter(q, f, "info.province_state", false, "AND", jsonColModel{})
|
||||
if got := q.only(t).query; got != "CAST(province_state.name AS TEXT) ILIKE ?" {
|
||||
t.Fatalf("query = %q", got)
|
||||
}
|
||||
}
|
||||
@@ -164,6 +164,9 @@ func (h *Handler) Handle(w common.ResponseWriter, r common.Request, params map[s
|
||||
// Parse options from headers - this now includes relation name resolution
|
||||
options := h.parseOptionsFromHeaders(r, model)
|
||||
|
||||
// Accept "<main table or alias>.<column>" for model columns before validation drops them
|
||||
common.NormalizeMainTableFilters(model, tableName, &options.RequestOptions)
|
||||
|
||||
// Validate and filter columns in options (log warnings for invalid columns)
|
||||
validator := common.NewColumnValidator(model)
|
||||
options = h.filterExtendedOptions(validator, options, model)
|
||||
@@ -417,6 +420,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 +737,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 +1421,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 +1671,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)
|
||||
}
|
||||
}
|
||||
@@ -389,6 +389,8 @@ func TestXFilesRecursivePreloadDepth(t *testing.T) {
|
||||
// TestXFilesResponseStructure validates the actual structure of the response
|
||||
// This test can be expanded when we have a full database integration test environment
|
||||
func TestXFilesResponseStructure(t *testing.T) {
|
||||
t.Skip("disabled: needs tests/data/xfiles.response.correct.json, which is gitignored and not in the repo")
|
||||
|
||||
// Load the expected correct response
|
||||
correctResponsePath := filepath.Join("..", "..", "tests", "data", "xfiles.response.correct.json")
|
||||
correctData, err := os.ReadFile(correctResponsePath)
|
||||
|
||||
@@ -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