Compare commits

...
3 Commits
Author SHA1 Message Date
warkanum 5b6406ef38 test(restheadspec): skip xfiles response test, fixture is not in the repo
Tests / Integration Tests (push) Skipped
Build , Vet Test, and Lint / Build (push) Successful in 1m33s
Tests / Unit Tests (push) Successful in 2m6s
Build , Vet Test, and Lint / Lint Code (push) Successful in 2m15s
Build , Vet Test, and Lint / Run Vet Tests (1.23.x) (push) Successful in 2m26s
Build , Vet Test, and Lint / Run Vet Tests (1.24.x) (push) Successful in 2m26s
Tests / Race Detector (push) Successful in 4m35s
2026-10-08 00:01:16 +02:00
warkanum d2f33b8f7d fix(aiproxy): handle response body closure and inspect return types 2026-10-07 23:41:28 +02:00
warkanum 63cb0d22eb 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.
2026-10-07 23:39:40 +02:00
15 changed files with 2092 additions and 0 deletions
+82
View File
@@ -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`.
+191
View File
@@ -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,
}
}
+329
View File
@@ -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")
}
}
+316
View File
@@ -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))
}
+96
View File
@@ -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)
}
}
+142
View File
@@ -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"
}
+80
View File
@@ -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)
}
}
}
+42
View File
@@ -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;
+231
View File
@@ -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
}
+269
View File
@@ -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", ""))
}
+143
View File
@@ -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
}
+118
View File
@@ -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
}
+8
View File
@@ -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()
+43
View File
@@ -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()
}
@@ -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)