From 63cb0d22ebf1502a68ba2d4cb48f66558cc5b9a0 Mon Sep 17 00:00:00 2001 From: Hein Date: Wed, 7 Oct 2026 23:39:40 +0200 Subject: [PATCH] feat(aiproxy): add authenticated proxy for OpenAI-compatible APIs and MCP servers Hides upstream URLs and keys behind the security layer. Per-upstream roles, model/tool allowlists, per-user rate limits, Before/AfterProxy hooks, audit sink, Prometheus metrics (metrics.AIProxyRecorder) and upstreams loaded from the resolvespec_ai_proxies stored procedure. --- pkg/aiproxy/README.md | 82 ++++++ pkg/aiproxy/aiproxy.go | 191 ++++++++++++++ pkg/aiproxy/aiproxy_test.go | 329 +++++++++++++++++++++++++ pkg/aiproxy/handler.go | 315 +++++++++++++++++++++++ pkg/aiproxy/hooks.go | 96 ++++++++ pkg/aiproxy/observe.go | 142 +++++++++++ pkg/aiproxy/ratelimit.go | 80 ++++++ pkg/aiproxy/resolvespec_ai_proxies.sql | 42 ++++ pkg/aiproxy/store.go | 230 +++++++++++++++++ pkg/aiproxy/store_test.go | 269 ++++++++++++++++++++ pkg/aiproxy/upstream.go | 143 +++++++++++ pkg/aiproxy/usage.go | 118 +++++++++ pkg/metrics/interfaces.go | 8 + pkg/metrics/prometheus.go | 43 ++++ 14 files changed, 2088 insertions(+) create mode 100644 pkg/aiproxy/README.md create mode 100644 pkg/aiproxy/aiproxy.go create mode 100644 pkg/aiproxy/aiproxy_test.go create mode 100644 pkg/aiproxy/handler.go create mode 100644 pkg/aiproxy/hooks.go create mode 100644 pkg/aiproxy/observe.go create mode 100644 pkg/aiproxy/ratelimit.go create mode 100644 pkg/aiproxy/resolvespec_ai_proxies.sql create mode 100644 pkg/aiproxy/store.go create mode 100644 pkg/aiproxy/store_test.go create mode 100644 pkg/aiproxy/upstream.go create mode 100644 pkg/aiproxy/usage.go diff --git a/pkg/aiproxy/README.md b/pkg/aiproxy/README.md new file mode 100644 index 0000000..780d841 --- /dev/null +++ b/pkg/aiproxy/README.md @@ -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`. diff --git a/pkg/aiproxy/aiproxy.go b/pkg/aiproxy/aiproxy.go new file mode 100644 index 0000000..47afb31 --- /dev/null +++ b/pkg/aiproxy/aiproxy.go @@ -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, + } +} diff --git a/pkg/aiproxy/aiproxy_test.go b/pkg/aiproxy/aiproxy_test.go new file mode 100644 index 0000000..90d5e28 --- /dev/null +++ b/pkg/aiproxy/aiproxy_test.go @@ -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") + } +} diff --git a/pkg/aiproxy/handler.go b/pkg/aiproxy/handler.go new file mode 100644 index 0000000..0208c1f --- /dev/null +++ b/pkg/aiproxy/handler.go @@ -0,0 +1,315 @@ +package aiproxy + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "io" + "mime" + "net/http" + "net/http/httputil" + "path" + "strconv" + "strings" + "sync" + "time" + + "github.com/bitechdev/ResolveSpec/pkg/logger" + "github.com/bitechdev/ResolveSpec/pkg/security" + "github.com/tidwall/gjson" +) + +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) (int, string, 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)) +} diff --git a/pkg/aiproxy/hooks.go b/pkg/aiproxy/hooks.go new file mode 100644 index 0000000..d13a02b --- /dev/null +++ b/pkg/aiproxy/hooks.go @@ -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) + } +} diff --git a/pkg/aiproxy/observe.go b/pkg/aiproxy/observe.go new file mode 100644 index 0000000..6093589 --- /dev/null +++ b/pkg/aiproxy/observe.go @@ -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" +} diff --git a/pkg/aiproxy/ratelimit.go b/pkg/aiproxy/ratelimit.go new file mode 100644 index 0000000..234dc1a --- /dev/null +++ b/pkg/aiproxy/ratelimit.go @@ -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) + } + } +} diff --git a/pkg/aiproxy/resolvespec_ai_proxies.sql b/pkg/aiproxy/resolvespec_ai_proxies.sql new file mode 100644 index 0000000..93e4cf3 --- /dev/null +++ b/pkg/aiproxy/resolvespec_ai_proxies.sql @@ -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; diff --git a/pkg/aiproxy/store.go b/pkg/aiproxy/store.go new file mode 100644 index 0000000..f64aa5a --- /dev/null +++ b/pkg/aiproxy/store.go @@ -0,0 +1,230 @@ +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 _, d := range defs { + 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 +} diff --git a/pkg/aiproxy/store_test.go b/pkg/aiproxy/store_test.go new file mode 100644 index 0000000..0176bb1 --- /dev/null +++ b/pkg/aiproxy/store_test.go @@ -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", "")) +} diff --git a/pkg/aiproxy/upstream.go b/pkg/aiproxy/upstream.go new file mode 100644 index 0000000..5f773a5 --- /dev/null +++ b/pkg/aiproxy/upstream.go @@ -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 +} diff --git a/pkg/aiproxy/usage.go b/pkg/aiproxy/usage.go new file mode 100644 index 0000000..3fd79b1 --- /dev/null +++ b/pkg/aiproxy/usage.go @@ -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 +} diff --git a/pkg/metrics/interfaces.go b/pkg/metrics/interfaces.go index 5c0cdb1..e95b43a 100644 --- a/pkg/metrics/interfaces.go +++ b/pkg/metrics/interfaces.go @@ -47,6 +47,14 @@ 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() diff --git a/pkg/metrics/prometheus.go b/pkg/metrics/prometheus.go index 6c54db9..c543316 100644 --- a/pkg/metrics/prometheus.go +++ b/pkg/metrics/prometheus.go @@ -32,6 +32,9 @@ 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 @@ -162,6 +165,28 @@ 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, @@ -310,6 +335,21 @@ func (p *PrometheusProvider) RecordPanic(methodName string) { 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 { @@ -437,6 +477,9 @@ func (p *PrometheusProvider) Reset() { p.eventProcessed.Reset() p.eventDuration.Reset() p.panicsTotal.Reset() + p.aiRequests.Reset() + p.aiDuration.Reset() + p.aiTokens.Reset() p.pathLimiter.reset() }