mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-08 14:26:28 +00:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
5b6406ef38 | ||
|
|
d2f33b8f7d | ||
|
|
63cb0d22eb |
@@ -0,0 +1,82 @@
|
|||||||
|
# aiproxy
|
||||||
|
|
||||||
|
Authenticated reverse proxy for OpenAI-compatible APIs and MCP servers (streamable HTTP).
|
||||||
|
Upstream URLs and keys stay server-side; clients only use the proxy with normal ResolveSpec auth.
|
||||||
|
|
||||||
|
## Setup
|
||||||
|
```go
|
||||||
|
p := aiproxy.New(aiproxy.Config{Prefix: "/ai"})
|
||||||
|
p.RegisterOpenAI(aiproxy.OpenAIUpstream{Upstream: aiproxy.Upstream{Name: "gpt", BaseURL: "https://api.openai.com/v1", APIKey: key}})
|
||||||
|
p.RegisterMCP(aiproxy.MCPUpstream{Upstream: aiproxy.Upstream{Name: "tools", BaseURL: "http://localhost:3000/mcp", APIKey: key}})
|
||||||
|
mux.Handle("/ai/", p.Handler(security.NewAuthMiddleware(securityList)))
|
||||||
|
```
|
||||||
|
|
||||||
|
## Routes
|
||||||
|
| Kind | Client URL | Upstream URL |
|
||||||
|
|---|---|---|
|
||||||
|
| OpenAI | `{prefix}/{name}/v1/*` | `BaseURL` + path after `/v1` (BaseURL includes the version) |
|
||||||
|
| MCP | `{prefix}/{name}` | `BaseURL` as is |
|
||||||
|
|
||||||
|
Client SDK base URL: `https://host/ai/gpt/v1` (OpenAI), `https://host/ai/tools` (MCP).
|
||||||
|
Query string, `Mcp-Session-Id`, SSE streaming all pass through.
|
||||||
|
|
||||||
|
## Upstream fields
|
||||||
|
| Field | Notes |
|
||||||
|
|---|---|
|
||||||
|
| `Name` | route segment, `[A-Za-z0-9_.-]`, unique across kinds |
|
||||||
|
| `BaseURL` | http/https |
|
||||||
|
| `APIKey` | optional; never logged/serialized (redacted in `String`/JSON) |
|
||||||
|
| `AuthHeader` | default `Authorization` |
|
||||||
|
| `AuthFormat` | one `%s`; default `Bearer %s` for `Authorization`, else raw key |
|
||||||
|
| `Headers` | static extra upstream headers |
|
||||||
|
| `AllowedRoles` | any-of; empty = any authenticated user |
|
||||||
|
| `RateLimit` | `{PerSecond, Burst}` per user per upstream; 429 + `Retry-After` |
|
||||||
|
| `Timeout` | response-header timeout, default 120s (streams not cut) |
|
||||||
|
| `AllowedModels` (OpenAI) | JSON body `model` must match; non-JSON writes rejected |
|
||||||
|
| `AllowedTools` (MCP) | `tools/call` name must match (batches checked); `tools/list` not filtered |
|
||||||
|
|
||||||
|
## Security behavior
|
||||||
|
- Auth required; guests rejected (401). Roles checked (403).
|
||||||
|
- Client `Authorization`, `Cookie`, `Proxy-Authorization`, `X-Api-Key` (+ `Config.StripHeaders`) removed; key injected.
|
||||||
|
- Response `Set-Cookie`, `Www-Authenticate`, redirect `Location` removed.
|
||||||
|
- Upstream 401 becomes 502 (proxy credentials problem, not the client's); upstream errors hide URL.
|
||||||
|
- JSON bodies inspected up to `Config.MaxBodyBytes` (10 MB), else 413.
|
||||||
|
- Fails closed if `Handler(nil)`.
|
||||||
|
|
||||||
|
## Hooks (`p.Hooks()`)
|
||||||
|
| Hook | When | Context |
|
||||||
|
|---|---|---|
|
||||||
|
| `BeforeProxy` | after auth/limit/inspection | `Request` (headers editable), `UserContext`, `Upstream`, `Kind`, `Model`, `Tools`; abort via `Abort`/`AbortMessage`/`AbortCode` (default 403) |
|
||||||
|
| `AfterProxy` | response finished/failed | `StatusCode`, `Duration`, `Usage`, `Error` |
|
||||||
|
|
||||||
|
`Usage` (OpenAI): from `usage` / `response.usage`; streams need `stream_options.include_usage`.
|
||||||
|
|
||||||
|
## Audit
|
||||||
|
`Config.Audit` (`AuditSink`) gets one `AuditRecord` per handled request: user, upstream, method, client path (no query),
|
||||||
|
model/tools, status, outcome (`ok|upstream_error|denied|rate_limited`), reason, duration, usage. No bodies, no keys.
|
||||||
|
Default `LogAuditSink` (pkg/logger); `NopAuditSink{}` disables. Sinks run on the request path: keep them fast.
|
||||||
|
|
||||||
|
## Metrics
|
||||||
|
Used automatically when the global `metrics.Provider` implements `metrics.AIProxyRecorder` (`PrometheusProvider` does):
|
||||||
|
`aiproxy_requests_total{upstream,kind,model,status,outcome}`, `aiproxy_request_duration_seconds{upstream,kind}`,
|
||||||
|
`aiproxy_tokens_total{upstream,model,type}`. `model` is client-supplied: bounded to 400 distinct pairs, then `other`.
|
||||||
|
|
||||||
|
## Dynamic upstreams
|
||||||
|
| Call | Effect |
|
||||||
|
|---|---|
|
||||||
|
| `Register(Definition)` / `Unregister(name)` | runtime add/remove |
|
||||||
|
| `Reload(ctx, store)` | sync store-managed upstreams: add, replace changed, drop missing; unchanged keep rate-limit state |
|
||||||
|
| `AutoReload(ctx, store, interval)` | `Reload` now, then periodically |
|
||||||
|
| `NewProcStore(*sql.DB, proc)` / `NewProcStoreFromDatabase(common.Database, proc)` | calls stored procedure (PostgreSQL), default `resolvespec_ai_proxies` |
|
||||||
|
|
||||||
|
- Code-registered upstreams are never replaced/removed by `Reload`; a same-name stored row is skipped and reported.
|
||||||
|
- Invalid rows are skipped (valid ones applied; last good definition kept); store failure changes nothing.
|
||||||
|
- Procedure contract (like the security procedures): `SELECT p_success, p_error, p_data FROM resolvespec_ai_proxies()`;
|
||||||
|
`p_data` = JSON array of `{name, kind, base_url, api_key, auth_header, auth_format, headers, allowed_roles, allowed, rate_per_second, rate_burst, timeout_seconds, enabled}`.
|
||||||
|
Reference table + function: `resolvespec_ai_proxies.sql` (replace the body to source from anywhere).
|
||||||
|
- Entries with `enabled:false` or undecodable JSON are skipped.
|
||||||
|
- API keys come back in plaintext from the procedure: restrict execute rights on it.
|
||||||
|
|
||||||
|
## Notes
|
||||||
|
- Legacy MCP SSE transport is not supported.
|
||||||
|
- `GET /v1/models` is not filtered by `AllowedModels`.
|
||||||
@@ -0,0 +1,191 @@
|
|||||||
|
// Package aiproxy is an authenticated reverse proxy for OpenAI-compatible APIs and
|
||||||
|
// MCP servers (streamable HTTP). Upstream URLs and keys stay on the server; clients
|
||||||
|
// authenticate with the normal security layer and call the proxy instead.
|
||||||
|
//
|
||||||
|
// p := aiproxy.New(aiproxy.Config{Prefix: "/ai"})
|
||||||
|
// p.RegisterOpenAI(aiproxy.OpenAIUpstream{Upstream: aiproxy.Upstream{Name: "gpt", BaseURL: "https://api.openai.com/v1", APIKey: key}})
|
||||||
|
// p.RegisterMCP(aiproxy.MCPUpstream{Upstream: aiproxy.Upstream{Name: "tools", BaseURL: "http://localhost:3000/mcp", APIKey: key}})
|
||||||
|
// mux.Handle("/ai/", p.Handler(security.NewAuthMiddleware(securityList)))
|
||||||
|
//
|
||||||
|
// Client routes: {prefix}/{name}/v1/... (OpenAI) and {prefix}/{name} (MCP).
|
||||||
|
package aiproxy
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httputil"
|
||||||
|
"net/url"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
defaultMaxBody = 10 << 20
|
||||||
|
defaultTimeout = 120 * time.Second
|
||||||
|
)
|
||||||
|
|
||||||
|
// Config configures a Proxy.
|
||||||
|
type Config struct {
|
||||||
|
// Prefix the handler is mounted under, e.g. "/ai". Empty when mounted with http.StripPrefix.
|
||||||
|
Prefix string
|
||||||
|
// MaxBodyBytes caps JSON request bodies that are inspected. Default 10 MB.
|
||||||
|
MaxBodyBytes int64
|
||||||
|
// StripHeaders are removed from client requests in addition to the defaults
|
||||||
|
// (Authorization, Cookie, Proxy-Authorization, X-Api-Key).
|
||||||
|
StripHeaders []string
|
||||||
|
// Audit receives one record per request. Default: LogAuditSink. Use NopAuditSink to disable.
|
||||||
|
Audit AuditSink
|
||||||
|
}
|
||||||
|
|
||||||
|
// Definition describes one upstream independent of how it was registered.
|
||||||
|
type Definition struct {
|
||||||
|
Kind Kind
|
||||||
|
Upstream Upstream
|
||||||
|
// Allowed holds AllowedModels (OpenAI) or AllowedTools (MCP).
|
||||||
|
Allowed []string
|
||||||
|
}
|
||||||
|
|
||||||
|
// target is a registered upstream, ready to serve.
|
||||||
|
type target struct {
|
||||||
|
def Definition
|
||||||
|
kind Kind
|
||||||
|
up Upstream
|
||||||
|
base *url.URL
|
||||||
|
roles map[string]struct{}
|
||||||
|
models map[string]struct{}
|
||||||
|
tools map[string]struct{}
|
||||||
|
proxy *httputil.ReverseProxy
|
||||||
|
managed bool // loaded by Reload; code-registered upstreams are never touched by it
|
||||||
|
|
||||||
|
authName string
|
||||||
|
authVal string
|
||||||
|
}
|
||||||
|
|
||||||
|
// Proxy routes authenticated requests to registered upstreams.
|
||||||
|
type Proxy struct {
|
||||||
|
cfg Config
|
||||||
|
hooks *HookRegistry
|
||||||
|
limits *limiters
|
||||||
|
labels modelLabels
|
||||||
|
mu sync.RWMutex
|
||||||
|
upstream map[string]*target
|
||||||
|
}
|
||||||
|
|
||||||
|
// New creates a Proxy.
|
||||||
|
func New(cfg Config) *Proxy {
|
||||||
|
cfg.Prefix = strings.TrimRight(cfg.Prefix, "/")
|
||||||
|
if cfg.Audit == nil {
|
||||||
|
cfg.Audit = LogAuditSink{}
|
||||||
|
}
|
||||||
|
if cfg.MaxBodyBytes <= 0 {
|
||||||
|
cfg.MaxBodyBytes = defaultMaxBody
|
||||||
|
}
|
||||||
|
return &Proxy{
|
||||||
|
cfg: cfg,
|
||||||
|
hooks: NewHookRegistry(),
|
||||||
|
limits: newLimiters(),
|
||||||
|
upstream: make(map[string]*target),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Hooks returns the hook registry.
|
||||||
|
func (p *Proxy) Hooks() *HookRegistry { return p.hooks }
|
||||||
|
|
||||||
|
// RegisterOpenAI registers an OpenAI-compatible upstream.
|
||||||
|
func (p *Proxy) RegisterOpenAI(u OpenAIUpstream) error {
|
||||||
|
return p.Register(Definition{Kind: KindOpenAI, Upstream: u.Upstream, Allowed: u.AllowedModels})
|
||||||
|
}
|
||||||
|
|
||||||
|
// RegisterMCP registers an MCP server (streamable HTTP).
|
||||||
|
func (p *Proxy) RegisterMCP(u MCPUpstream) error {
|
||||||
|
return p.Register(Definition{Kind: KindMCP, Upstream: u.Upstream, Allowed: u.AllowedTools})
|
||||||
|
}
|
||||||
|
|
||||||
|
// Register registers an upstream from a Definition. The name must be unused.
|
||||||
|
func (p *Proxy) Register(d Definition) error {
|
||||||
|
t, err := p.build(d)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
p.mu.Lock()
|
||||||
|
defer p.mu.Unlock()
|
||||||
|
if _, dup := p.upstream[t.up.Name]; dup {
|
||||||
|
return fmt.Errorf("aiproxy: upstream %q already registered", t.up.Name)
|
||||||
|
}
|
||||||
|
p.upstream[t.up.Name] = t
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// Unregister removes an upstream (code-registered or loaded). It reports whether it existed.
|
||||||
|
// In-flight requests finish normally.
|
||||||
|
func (p *Proxy) Unregister(name string) bool {
|
||||||
|
p.mu.Lock()
|
||||||
|
t, ok := p.upstream[name]
|
||||||
|
delete(p.upstream, name)
|
||||||
|
p.mu.Unlock()
|
||||||
|
if ok {
|
||||||
|
p.retire(t)
|
||||||
|
}
|
||||||
|
return ok
|
||||||
|
}
|
||||||
|
|
||||||
|
// retire releases what a removed or replaced target held.
|
||||||
|
func (p *Proxy) retire(t *target) {
|
||||||
|
p.limits.forget(t.up.Name)
|
||||||
|
if tr, ok := t.proxy.Transport.(*http.Transport); ok {
|
||||||
|
tr.CloseIdleConnections()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Proxy) build(d Definition) (*target, error) {
|
||||||
|
if d.Kind != KindOpenAI && d.Kind != KindMCP {
|
||||||
|
return nil, fmt.Errorf("aiproxy: upstream %q has unknown kind %q", d.Upstream.Name, d.Kind)
|
||||||
|
}
|
||||||
|
base, err := d.Upstream.validate()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
t := &target{def: d, kind: d.Kind, up: d.Upstream, base: base, roles: toSet(d.Upstream.AllowedRoles)}
|
||||||
|
if d.Kind == KindOpenAI {
|
||||||
|
t.models = toSet(d.Allowed)
|
||||||
|
} else {
|
||||||
|
t.tools = toSet(d.Allowed)
|
||||||
|
}
|
||||||
|
t.authName, t.authVal = d.Upstream.authValue()
|
||||||
|
t.proxy = p.buildProxy(t)
|
||||||
|
return t, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Proxy) get(name string) *target {
|
||||||
|
p.mu.RLock()
|
||||||
|
defer p.mu.RUnlock()
|
||||||
|
return p.upstream[name]
|
||||||
|
}
|
||||||
|
|
||||||
|
// Names lists registered upstream names (no URLs or keys).
|
||||||
|
func (p *Proxy) Names() []string {
|
||||||
|
p.mu.RLock()
|
||||||
|
defer p.mu.RUnlock()
|
||||||
|
out := make([]string, 0, len(p.upstream))
|
||||||
|
for n := range p.upstream {
|
||||||
|
out = append(out, n)
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func newTransport(timeout time.Duration) *http.Transport {
|
||||||
|
if timeout <= 0 {
|
||||||
|
timeout = defaultTimeout
|
||||||
|
}
|
||||||
|
return &http.Transport{
|
||||||
|
Proxy: http.ProxyFromEnvironment,
|
||||||
|
DialContext: (&net.Dialer{Timeout: 10 * time.Second, KeepAlive: 30 * time.Second}).DialContext,
|
||||||
|
ForceAttemptHTTP2: true,
|
||||||
|
MaxIdleConnsPerHost: 32,
|
||||||
|
IdleConnTimeout: 90 * time.Second,
|
||||||
|
TLSHandshakeTimeout: 10 * time.Second,
|
||||||
|
ResponseHeaderTimeout: timeout,
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,329 @@
|
|||||||
|
package aiproxy
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
// testAuth authenticates by the X-Test-User header: "name:role1,role2". Empty means no user.
|
||||||
|
func testAuth(next http.Handler) http.Handler {
|
||||||
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
h := r.Header.Get("X-Test-User")
|
||||||
|
if h == "" {
|
||||||
|
http.Error(w, "unauthenticated", http.StatusUnauthorized)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
name, roles, _ := strings.Cut(h, ":")
|
||||||
|
uc := &security.UserContext{UserID: len(name), UserName: name, Roles: strings.Split(roles, ",")}
|
||||||
|
next.ServeHTTP(w, r.WithContext(context.WithValue(r.Context(), security.UserContextKey, uc)))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
type captured struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
req *http.Request
|
||||||
|
body string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *captured) set(r *http.Request) {
|
||||||
|
b, _ := io.ReadAll(r.Body)
|
||||||
|
c.mu.Lock()
|
||||||
|
c.req, c.body = r.Clone(context.Background()), string(b)
|
||||||
|
c.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
func newProxy(t *testing.T, upstream http.Handler, mod func(*Proxy, string)) (*Proxy, *httptest.Server) {
|
||||||
|
t.Helper()
|
||||||
|
up := httptest.NewServer(upstream)
|
||||||
|
t.Cleanup(up.Close)
|
||||||
|
p := New(Config{Prefix: "/ai"})
|
||||||
|
mod(p, up.URL)
|
||||||
|
front := httptest.NewServer(p.Handler(testAuth))
|
||||||
|
t.Cleanup(front.Close)
|
||||||
|
return p, front
|
||||||
|
}
|
||||||
|
|
||||||
|
func do(t *testing.T, method, url, body, user string, hdr map[string]string) (*http.Response, string) {
|
||||||
|
t.Helper()
|
||||||
|
req, err := http.NewRequest(method, url, strings.NewReader(body))
|
||||||
|
require.NoError(t, err)
|
||||||
|
if body != "" {
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
}
|
||||||
|
if user != "" {
|
||||||
|
req.Header.Set("X-Test-User", user)
|
||||||
|
}
|
||||||
|
for k, v := range hdr {
|
||||||
|
req.Header.Set(k, v)
|
||||||
|
}
|
||||||
|
resp, err := http.DefaultClient.Do(req)
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer resp.Body.Close()
|
||||||
|
b, _ := io.ReadAll(resp.Body)
|
||||||
|
return resp, string(b)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOpenAIKeyInjectionAndHiding(t *testing.T) {
|
||||||
|
var c captured
|
||||||
|
_, front := newProxy(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
c.set(r)
|
||||||
|
w.Header().Set("Set-Cookie", "s=1")
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
io.WriteString(w, `{"ok":true}`)
|
||||||
|
}), func(p *Proxy, u string) {
|
||||||
|
require.NoError(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "gpt", BaseURL: u + "/v1", APIKey: "sk-secret"}}))
|
||||||
|
})
|
||||||
|
|
||||||
|
resp, body := do(t, "POST", front.URL+"/ai/gpt/v1/chat/completions?x=1", `{"model":"m"}`, "bob:user",
|
||||||
|
map[string]string{"Authorization": "Bearer client-token", "Cookie": "a=b"})
|
||||||
|
assert.Equal(t, 200, resp.StatusCode)
|
||||||
|
assert.Equal(t, `{"ok":true}`, body)
|
||||||
|
assert.Empty(t, resp.Header.Get("Set-Cookie"))
|
||||||
|
assert.Equal(t, "/v1/chat/completions", c.req.URL.Path)
|
||||||
|
assert.Equal(t, "x=1", c.req.URL.RawQuery)
|
||||||
|
assert.Equal(t, "Bearer sk-secret", c.req.Header.Get("Authorization"))
|
||||||
|
assert.Empty(t, c.req.Header.Get("Cookie"))
|
||||||
|
assert.Equal(t, `{"model":"m"}`, c.body)
|
||||||
|
assert.NotContains(t, body, "sk-secret")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCustomAuthHeaderAndFormat(t *testing.T) {
|
||||||
|
var c captured
|
||||||
|
_, front := newProxy(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { c.set(r); io.WriteString(w, "{}") }),
|
||||||
|
func(p *Proxy, u string) {
|
||||||
|
require.NoError(t, p.RegisterMCP(MCPUpstream{Upstream: Upstream{Name: "a", BaseURL: u + "/mcp", APIKey: "k1", AuthHeader: "X-Token", Headers: map[string]string{"X-Tenant": "t1"}}}))
|
||||||
|
require.NoError(t, p.RegisterMCP(MCPUpstream{Upstream: Upstream{Name: "b", BaseURL: u + "/mcp", APIKey: "k2", AuthFormat: "Token %s"}}))
|
||||||
|
})
|
||||||
|
|
||||||
|
do(t, "POST", front.URL+"/ai/a", `{"jsonrpc":"2.0","id":1,"method":"ping"}`, "bob:user", map[string]string{"X-Token": "client"})
|
||||||
|
assert.Equal(t, "k1", c.req.Header.Get("X-Token"))
|
||||||
|
assert.Equal(t, "t1", c.req.Header.Get("X-Tenant"))
|
||||||
|
assert.Equal(t, "/mcp", c.req.URL.Path)
|
||||||
|
|
||||||
|
do(t, "POST", front.URL+"/ai/b", `{"jsonrpc":"2.0","id":1,"method":"ping"}`, "bob:user", nil)
|
||||||
|
assert.Equal(t, "Token k2", c.req.Header.Get("Authorization"))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAuthAndRoles(t *testing.T) {
|
||||||
|
_, front := newProxy(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { io.WriteString(w, "{}") }),
|
||||||
|
func(p *Proxy, u string) {
|
||||||
|
require.NoError(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "gpt", BaseURL: u + "/v1", AllowedRoles: []string{"ai"}}}))
|
||||||
|
})
|
||||||
|
url := front.URL + "/ai/gpt/v1/models"
|
||||||
|
|
||||||
|
resp, _ := do(t, "GET", url, "", "", nil)
|
||||||
|
assert.Equal(t, 401, resp.StatusCode)
|
||||||
|
resp, _ = do(t, "GET", url, "", "bob:user", nil)
|
||||||
|
assert.Equal(t, 403, resp.StatusCode)
|
||||||
|
resp, _ = do(t, "GET", url, "", "bob:guest", nil)
|
||||||
|
assert.Equal(t, 401, resp.StatusCode)
|
||||||
|
resp, _ = do(t, "GET", url, "", "bob:user,ai", nil)
|
||||||
|
assert.Equal(t, 200, resp.StatusCode)
|
||||||
|
|
||||||
|
resp, _ = do(t, "GET", front.URL+"/ai/nope/v1/models", "", "bob:ai", nil)
|
||||||
|
assert.Equal(t, 404, resp.StatusCode)
|
||||||
|
resp, _ = do(t, "GET", front.URL+"/ai/gpt/other", "", "bob:ai", nil)
|
||||||
|
assert.Equal(t, 404, resp.StatusCode)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNilAuthRefuses(t *testing.T) {
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
New(Config{}).Handler(nil).ServeHTTP(rec, httptest.NewRequest("GET", "/", nil))
|
||||||
|
assert.Equal(t, 500, rec.Code)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestModelAllowlist(t *testing.T) {
|
||||||
|
_, front := newProxy(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { io.WriteString(w, "{}") }),
|
||||||
|
func(p *Proxy, u string) {
|
||||||
|
require.NoError(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "gpt", BaseURL: u + "/v1"}, AllowedModels: []string{"good"}}))
|
||||||
|
})
|
||||||
|
url := front.URL + "/ai/gpt/v1/chat/completions"
|
||||||
|
resp, _ := do(t, "POST", url, `{"model":"good"}`, "bob:u", nil)
|
||||||
|
assert.Equal(t, 200, resp.StatusCode)
|
||||||
|
resp, _ = do(t, "POST", url, `{"model":"bad"}`, "bob:u", nil)
|
||||||
|
assert.Equal(t, 403, resp.StatusCode)
|
||||||
|
resp, _ = do(t, "POST", url, `{}`, "bob:u", nil)
|
||||||
|
assert.Equal(t, 403, resp.StatusCode)
|
||||||
|
resp, _ = do(t, "POST", url, `x`, "bob:u", map[string]string{"Content-Type": "text/plain"})
|
||||||
|
assert.Equal(t, 400, resp.StatusCode)
|
||||||
|
resp, _ = do(t, "GET", front.URL+"/ai/gpt/v1/models", "", "bob:u", nil)
|
||||||
|
assert.Equal(t, 200, resp.StatusCode)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBodyTooLarge(t *testing.T) {
|
||||||
|
up := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {}))
|
||||||
|
defer up.Close()
|
||||||
|
p := New(Config{Prefix: "/ai", MaxBodyBytes: 10})
|
||||||
|
require.NoError(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "g", BaseURL: up.URL + "/v1"}}))
|
||||||
|
front := httptest.NewServer(p.Handler(testAuth))
|
||||||
|
defer front.Close()
|
||||||
|
resp, _ := do(t, "POST", front.URL+"/ai/g/v1/x", `{"model":"aaaaaaaaaaaaaaaa"}`, "bob:u", nil)
|
||||||
|
assert.Equal(t, 413, resp.StatusCode)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMCPToolAllowlist(t *testing.T) {
|
||||||
|
var c captured
|
||||||
|
_, front := newProxy(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
c.set(r)
|
||||||
|
w.Header().Set("Mcp-Session-Id", "sess1")
|
||||||
|
io.WriteString(w, "{}")
|
||||||
|
}), func(p *Proxy, u string) {
|
||||||
|
require.NoError(t, p.RegisterMCP(MCPUpstream{Upstream: Upstream{Name: "t", BaseURL: u + "/mcp"}, AllowedTools: []string{"ok"}}))
|
||||||
|
})
|
||||||
|
url := front.URL + "/ai/t"
|
||||||
|
call := func(name string) string {
|
||||||
|
return `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"` + name + `"}}`
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, _ := do(t, "POST", url, call("ok"), "bob:u", nil)
|
||||||
|
assert.Equal(t, 200, resp.StatusCode)
|
||||||
|
assert.Equal(t, "sess1", resp.Header.Get("Mcp-Session-Id"))
|
||||||
|
resp, _ = do(t, "POST", url, call("evil"), "bob:u", nil)
|
||||||
|
assert.Equal(t, 403, resp.StatusCode)
|
||||||
|
resp, _ = do(t, "POST", url, `[`+call("ok")+`,`+call("evil")+`]`, "bob:u", nil)
|
||||||
|
assert.Equal(t, 403, resp.StatusCode)
|
||||||
|
resp, _ = do(t, "POST", url, `{"jsonrpc":"2.0","id":2,"method":"tools/list"}`, "bob:u", nil)
|
||||||
|
assert.Equal(t, 200, resp.StatusCode)
|
||||||
|
resp, _ = do(t, "DELETE", url, "", "bob:u", map[string]string{"Mcp-Session-Id": "sess1"})
|
||||||
|
assert.Equal(t, 200, resp.StatusCode)
|
||||||
|
assert.Equal(t, "sess1", c.req.Header.Get("Mcp-Session-Id"))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRateLimit(t *testing.T) {
|
||||||
|
_, front := newProxy(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { io.WriteString(w, "{}") }),
|
||||||
|
func(p *Proxy, u string) {
|
||||||
|
require.NoError(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "g", BaseURL: u + "/v1", RateLimit: &RateLimit{PerSecond: 0.1, Burst: 2}}}))
|
||||||
|
})
|
||||||
|
url := front.URL + "/ai/g/v1/models"
|
||||||
|
for i := 0; i < 2; i++ {
|
||||||
|
resp, _ := do(t, "GET", url, "", "bob:u", nil)
|
||||||
|
assert.Equal(t, 200, resp.StatusCode)
|
||||||
|
}
|
||||||
|
resp, _ := do(t, "GET", url, "", "bob:u", nil)
|
||||||
|
assert.Equal(t, 429, resp.StatusCode)
|
||||||
|
assert.NotEmpty(t, resp.Header.Get("Retry-After"))
|
||||||
|
// other user has own bucket
|
||||||
|
resp, _ = do(t, "GET", url, "", "alice:u", nil)
|
||||||
|
assert.Equal(t, 200, resp.StatusCode)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestHooksAndUsage(t *testing.T) {
|
||||||
|
var p *Proxy
|
||||||
|
_, front := newProxy(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
if r.Header.Get("X-From-Hook") == "" {
|
||||||
|
t.Error("hook header missing")
|
||||||
|
}
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
io.WriteString(w, `{"usage":{"prompt_tokens":3,"completion_tokens":4,"total_tokens":7}}`)
|
||||||
|
}), func(pp *Proxy, u string) {
|
||||||
|
p = pp
|
||||||
|
require.NoError(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "g", BaseURL: u + "/v1"}}))
|
||||||
|
})
|
||||||
|
done := make(chan *HookContext, 1)
|
||||||
|
p.Hooks().Register(BeforeProxy, func(h *HookContext) error {
|
||||||
|
h.Request.Header.Set("X-From-Hook", "1")
|
||||||
|
if h.Model == "blocked" {
|
||||||
|
h.Abort, h.AbortMessage, h.AbortCode = true, "nope", 418
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
})
|
||||||
|
p.Hooks().Register(AfterProxy, func(h *HookContext) error { done <- h; return nil })
|
||||||
|
|
||||||
|
resp, _ := do(t, "POST", front.URL+"/ai/g/v1/chat/completions", `{"model":"blocked"}`, "bob:u", nil)
|
||||||
|
assert.Equal(t, 418, resp.StatusCode)
|
||||||
|
|
||||||
|
resp, _ = do(t, "POST", front.URL+"/ai/g/v1/chat/completions", `{"model":"m"}`, "bob:u", nil)
|
||||||
|
assert.Equal(t, 200, resp.StatusCode)
|
||||||
|
select {
|
||||||
|
case h := <-done:
|
||||||
|
assert.Equal(t, "m", h.Model)
|
||||||
|
assert.Equal(t, "bob", h.UserContext.UserName)
|
||||||
|
assert.Equal(t, 200, h.StatusCode)
|
||||||
|
assert.Equal(t, Usage{3, 4, 7}, h.Usage)
|
||||||
|
case <-time.After(2 * time.Second):
|
||||||
|
t.Fatal("AfterProxy not called")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestSSEUsageAndStreaming(t *testing.T) {
|
||||||
|
var p *Proxy
|
||||||
|
_, front := newProxy(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
w.Header().Set("Content-Type", "text/event-stream")
|
||||||
|
f := w.(http.Flusher)
|
||||||
|
io.WriteString(w, "data: {\"choices\":[{}]}\n\n")
|
||||||
|
f.Flush()
|
||||||
|
io.WriteString(w, "data: {\"choices\":[],\"usage\":{\"prompt_tokens\":1,\"completion_tokens\":2,\"total_tokens\":3}}\n\ndata: [DONE]\n\n")
|
||||||
|
f.Flush()
|
||||||
|
}), func(pp *Proxy, u string) {
|
||||||
|
p = pp
|
||||||
|
require.NoError(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "g", BaseURL: u + "/v1"}}))
|
||||||
|
})
|
||||||
|
done := make(chan *HookContext, 1)
|
||||||
|
p.Hooks().Register(AfterProxy, func(h *HookContext) error { done <- h; return nil })
|
||||||
|
|
||||||
|
resp, body := do(t, "POST", front.URL+"/ai/g/v1/chat/completions", `{"model":"m","stream":true}`, "bob:u", nil)
|
||||||
|
assert.Equal(t, 200, resp.StatusCode)
|
||||||
|
assert.Contains(t, body, "[DONE]")
|
||||||
|
h := <-done
|
||||||
|
assert.Equal(t, Usage{1, 2, 3}, h.Usage)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResponsesAPIUsage(t *testing.T) {
|
||||||
|
u, ok := parseUsage([]byte(`{"type":"response.completed","response":{"usage":{"input_tokens":5,"output_tokens":6}}}`))
|
||||||
|
assert.True(t, ok)
|
||||||
|
assert.Equal(t, Usage{5, 6, 11}, u)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUpstreamUnauthorizedIsMasked(t *testing.T) {
|
||||||
|
_, front := newProxy(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
w.Header().Set("Www-Authenticate", "Bearer realm=x")
|
||||||
|
w.WriteHeader(401)
|
||||||
|
}), func(p *Proxy, u string) {
|
||||||
|
require.NoError(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "g", BaseURL: u + "/v1", APIKey: "bad"}}))
|
||||||
|
})
|
||||||
|
resp, body := do(t, "GET", front.URL+"/ai/g/v1/models", "", "bob:u", nil)
|
||||||
|
assert.Equal(t, 502, resp.StatusCode)
|
||||||
|
assert.Empty(t, resp.Header.Get("Www-Authenticate"))
|
||||||
|
assert.Contains(t, body, "bad_gateway")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestUpstreamDownHidesURL(t *testing.T) {
|
||||||
|
up := httptest.NewServer(http.NotFoundHandler())
|
||||||
|
url := up.URL
|
||||||
|
up.Close()
|
||||||
|
p := New(Config{Prefix: "/ai"})
|
||||||
|
require.NoError(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "g", BaseURL: url + "/v1"}}))
|
||||||
|
front := httptest.NewServer(p.Handler(testAuth))
|
||||||
|
defer front.Close()
|
||||||
|
resp, body := do(t, "GET", front.URL+"/ai/g/v1/models", "", "bob:u", nil)
|
||||||
|
assert.Equal(t, 502, resp.StatusCode)
|
||||||
|
assert.NotContains(t, body, url)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRegisterValidationAndRedaction(t *testing.T) {
|
||||||
|
p := New(Config{})
|
||||||
|
assert.Error(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "bad name", BaseURL: "http://x"}}))
|
||||||
|
assert.Error(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "a", BaseURL: "ftp://x"}}))
|
||||||
|
assert.Error(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "a", BaseURL: "http://x", AuthFormat: "%d"}}))
|
||||||
|
assert.Error(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "a", BaseURL: "http://x", RateLimit: &RateLimit{}}}))
|
||||||
|
require.NoError(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "a", BaseURL: "http://x", APIKey: "sk-secret"}}))
|
||||||
|
assert.Error(t, p.RegisterMCP(MCPUpstream{Upstream: Upstream{Name: "a", BaseURL: "http://x"}}))
|
||||||
|
|
||||||
|
u := Upstream{Name: "a", BaseURL: "http://x", APIKey: "sk-secret"}
|
||||||
|
js, _ := json.Marshal(u)
|
||||||
|
for _, s := range []string{u.String(), string(js), strings.TrimSpace(errors.New(u.String()).Error())} {
|
||||||
|
assert.NotContains(t, s, "sk-secret")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,316 @@
|
|||||||
|
package aiproxy
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
"mime"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httputil"
|
||||||
|
"path"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/tidwall/gjson"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||||
|
)
|
||||||
|
|
||||||
|
type ctxKey struct{}
|
||||||
|
|
||||||
|
// call is the per-request state shared between the handler and the proxy callbacks.
|
||||||
|
type call struct {
|
||||||
|
p *Proxy
|
||||||
|
t *target
|
||||||
|
hc *HookContext
|
||||||
|
start time.Time
|
||||||
|
once sync.Once
|
||||||
|
}
|
||||||
|
|
||||||
|
func (c *call) finish(status int, usage Usage, err error) {
|
||||||
|
c.once.Do(func() {
|
||||||
|
c.hc.StatusCode = status
|
||||||
|
c.hc.Usage = usage
|
||||||
|
c.hc.Error = err
|
||||||
|
c.hc.Duration = time.Since(c.start)
|
||||||
|
c.p.hooks.executeAfter(c.hc)
|
||||||
|
outcome := OutcomeOK
|
||||||
|
if err != nil || status >= 400 {
|
||||||
|
outcome = OutcomeUpstreamError
|
||||||
|
}
|
||||||
|
c.p.record(c.t, c.hc, status, outcome, "", err)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
var defaultStrip = []string{"Authorization", "Cookie", "Proxy-Authorization", "X-Api-Key"}
|
||||||
|
|
||||||
|
// Handler returns the proxy handler. auth is the authentication middleware, normally
|
||||||
|
// security.NewAuthMiddleware(securityList). It is required: without it every request is refused.
|
||||||
|
func (p *Proxy) Handler(auth func(http.Handler) http.Handler) http.Handler {
|
||||||
|
if auth == nil {
|
||||||
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
writeError(w, http.StatusInternalServerError, "server_error", "authentication is not configured")
|
||||||
|
})
|
||||||
|
}
|
||||||
|
return auth(http.HandlerFunc(p.serve))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Proxy) serve(w http.ResponseWriter, r *http.Request) {
|
||||||
|
// route: {prefix}/{name}/{rest}
|
||||||
|
rel := strings.TrimPrefix(r.URL.Path, p.cfg.Prefix)
|
||||||
|
if rel == r.URL.Path && p.cfg.Prefix != "" {
|
||||||
|
writeError(w, http.StatusNotFound, "not_found", "not found")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
rel = strings.TrimPrefix(rel, "/")
|
||||||
|
name, rest, _ := strings.Cut(rel, "/")
|
||||||
|
t := p.get(name)
|
||||||
|
if t == nil {
|
||||||
|
writeError(w, http.StatusNotFound, "not_found", "unknown upstream")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
rest = "/" + rest
|
||||||
|
if rest != "/" {
|
||||||
|
rest = path.Clean(rest)
|
||||||
|
}
|
||||||
|
if t.kind == KindOpenAI {
|
||||||
|
switch {
|
||||||
|
case rest == "/v1":
|
||||||
|
rest = ""
|
||||||
|
case strings.HasPrefix(rest, "/v1/"):
|
||||||
|
rest = strings.TrimPrefix(rest, "/v1")
|
||||||
|
default:
|
||||||
|
writeError(w, http.StatusNotFound, "not_found", "unknown path")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
} else if rest == "/" {
|
||||||
|
rest = ""
|
||||||
|
}
|
||||||
|
|
||||||
|
start := time.Now()
|
||||||
|
user, _ := security.GetUserContext(r.Context())
|
||||||
|
hc := &HookContext{
|
||||||
|
Context: r.Context(),
|
||||||
|
Request: r,
|
||||||
|
UserContext: user,
|
||||||
|
Upstream: t.up.Name,
|
||||||
|
Kind: t.kind,
|
||||||
|
}
|
||||||
|
deny := func(code int, outcome Outcome, typ, msg string) {
|
||||||
|
hc.Duration = time.Since(start)
|
||||||
|
p.record(t, hc, code, outcome, msg, nil)
|
||||||
|
writeError(w, code, typ, msg)
|
||||||
|
}
|
||||||
|
|
||||||
|
// authenticated user, no guests
|
||||||
|
if user == nil || hasRole(user.Roles, "guest") {
|
||||||
|
deny(http.StatusUnauthorized, OutcomeDenied, "invalid_api_key", "authentication required")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if len(t.roles) > 0 && !anyRole(user.Roles, t.roles) {
|
||||||
|
deny(http.StatusForbidden, OutcomeDenied, "forbidden", "not allowed to use this upstream")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if rl := t.up.RateLimit; rl != nil {
|
||||||
|
key := t.up.Name + "|" + strconv.Itoa(user.UserID) + "|" + user.UserName
|
||||||
|
if allowed, wait := p.limits.allow(key, rl); !allowed {
|
||||||
|
w.Header().Set("Retry-After", strconv.Itoa(int(wait.Seconds())+1))
|
||||||
|
deny(http.StatusTooManyRequests, OutcomeRateLimited, "rate_limit_exceeded", "rate limit exceeded")
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if code, typ, msg := p.inspect(t, r, hc); code != 0 {
|
||||||
|
deny(code, OutcomeDenied, typ, msg)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := p.hooks.Execute(BeforeProxy, hc); err != nil {
|
||||||
|
code := hc.AbortCode
|
||||||
|
if code == 0 {
|
||||||
|
code = http.StatusForbidden
|
||||||
|
}
|
||||||
|
msg := hc.AbortMessage
|
||||||
|
if msg == "" {
|
||||||
|
msg = "request rejected"
|
||||||
|
}
|
||||||
|
if !hc.Abort {
|
||||||
|
logger.Error("aiproxy: %v", err)
|
||||||
|
code, msg = http.StatusInternalServerError, "request failed"
|
||||||
|
}
|
||||||
|
deny(code, OutcomeDenied, "rejected", msg)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
c := &call{p: p, t: t, hc: hc, start: start}
|
||||||
|
ctx := context.WithValue(r.Context(), ctxKey{}, &routed{call: c, rest: rest})
|
||||||
|
t.proxy.ServeHTTP(w, hc.Request.WithContext(ctx))
|
||||||
|
}
|
||||||
|
|
||||||
|
// routed is stored on the request context for the ReverseProxy callbacks.
|
||||||
|
type routed struct {
|
||||||
|
call *call
|
||||||
|
rest string
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *Proxy) buildProxy(t *target) *httputil.ReverseProxy {
|
||||||
|
strip := append(append([]string{}, defaultStrip...), p.cfg.StripHeaders...)
|
||||||
|
if t.authName != "" {
|
||||||
|
strip = append(strip, t.authName)
|
||||||
|
}
|
||||||
|
return &httputil.ReverseProxy{
|
||||||
|
Transport: newTransport(t.up.Timeout),
|
||||||
|
FlushInterval: -1,
|
||||||
|
Rewrite: func(pr *httputil.ProxyRequest) {
|
||||||
|
rt, _ := pr.In.Context().Value(ctxKey{}).(*routed)
|
||||||
|
out := pr.Out
|
||||||
|
u := *t.base
|
||||||
|
u.Path = t.base.Path
|
||||||
|
if rt != nil {
|
||||||
|
u.Path += rt.rest
|
||||||
|
}
|
||||||
|
u.RawPath = ""
|
||||||
|
u.RawQuery = pr.In.URL.RawQuery
|
||||||
|
out.URL = &u
|
||||||
|
out.Host = u.Host
|
||||||
|
for _, h := range strip {
|
||||||
|
out.Header.Del(h)
|
||||||
|
}
|
||||||
|
out.Header.Del("Accept-Encoding") // let the transport decompress so usage can be read
|
||||||
|
for k, v := range t.up.Headers {
|
||||||
|
out.Header.Set(k, v)
|
||||||
|
}
|
||||||
|
if t.authName != "" {
|
||||||
|
out.Header.Set(t.authName, t.authVal)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
ModifyResponse: func(resp *http.Response) error {
|
||||||
|
rt, _ := resp.Request.Context().Value(ctxKey{}).(*routed)
|
||||||
|
h := resp.Header
|
||||||
|
h.Del("Set-Cookie")
|
||||||
|
h.Del("Www-Authenticate")
|
||||||
|
if resp.StatusCode >= 300 && resp.StatusCode < 400 {
|
||||||
|
h.Del("Location")
|
||||||
|
}
|
||||||
|
if resp.StatusCode == http.StatusUnauthorized {
|
||||||
|
logger.Warn("aiproxy: upstream %q rejected its credentials", t.up.Name)
|
||||||
|
_ = resp.Body.Close()
|
||||||
|
body, _ := json.Marshal(errorBody("bad_gateway", "upstream rejected the proxy credentials"))
|
||||||
|
resp.StatusCode, resp.Status = http.StatusBadGateway, "502 Bad Gateway"
|
||||||
|
resp.Body = io.NopCloser(bytes.NewReader(body))
|
||||||
|
resp.ContentLength = int64(len(body))
|
||||||
|
h.Set("Content-Type", "application/json")
|
||||||
|
h.Set("Content-Length", strconv.Itoa(len(body)))
|
||||||
|
}
|
||||||
|
if rt != nil {
|
||||||
|
sse := strings.HasPrefix(h.Get("Content-Type"), "text/event-stream")
|
||||||
|
resp.Body = newWatchBody(resp.Body, t.kind == KindOpenAI, sse, resp.StatusCode, rt.call)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
},
|
||||||
|
ErrorHandler: func(w http.ResponseWriter, r *http.Request, err error) {
|
||||||
|
if rt, _ := r.Context().Value(ctxKey{}).(*routed); rt != nil {
|
||||||
|
rt.call.finish(http.StatusBadGateway, Usage{}, err)
|
||||||
|
}
|
||||||
|
if !errors.Is(err, context.Canceled) {
|
||||||
|
logger.Warn("aiproxy: upstream %q error: %v", t.up.Name, err)
|
||||||
|
}
|
||||||
|
writeError(w, http.StatusBadGateway, "bad_gateway", "upstream request failed")
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// inspect reads JSON request bodies to extract/enforce model and tools.
|
||||||
|
// It returns a non-zero status when the request must be rejected.
|
||||||
|
func (p *Proxy) inspect(t *target, r *http.Request, hc *HookContext) (status int, errType, msg string) {
|
||||||
|
switch r.Method {
|
||||||
|
case http.MethodPost, http.MethodPut, http.MethodPatch:
|
||||||
|
default:
|
||||||
|
return 0, "", ""
|
||||||
|
}
|
||||||
|
if r.Body == nil || r.Body == http.NoBody {
|
||||||
|
return 0, "", ""
|
||||||
|
}
|
||||||
|
mt, _, _ := mime.ParseMediaType(r.Header.Get("Content-Type"))
|
||||||
|
isJSON := mt == "application/json" || strings.HasSuffix(mt, "+json")
|
||||||
|
|
||||||
|
if !isJSON {
|
||||||
|
if t.kind == KindOpenAI && len(t.models) > 0 {
|
||||||
|
return http.StatusBadRequest, "invalid_request_error", "model cannot be verified for this content type"
|
||||||
|
}
|
||||||
|
return 0, "", ""
|
||||||
|
}
|
||||||
|
|
||||||
|
body, err := io.ReadAll(io.LimitReader(r.Body, p.cfg.MaxBodyBytes+1))
|
||||||
|
_ = r.Body.Close()
|
||||||
|
if err != nil {
|
||||||
|
return http.StatusBadRequest, "invalid_request_error", "could not read request body"
|
||||||
|
}
|
||||||
|
if int64(len(body)) > p.cfg.MaxBodyBytes {
|
||||||
|
return http.StatusRequestEntityTooLarge, "invalid_request_error", "request body too large"
|
||||||
|
}
|
||||||
|
r.Body = io.NopCloser(bytes.NewReader(body))
|
||||||
|
r.ContentLength = int64(len(body))
|
||||||
|
r.Header.Set("Content-Length", strconv.Itoa(len(body)))
|
||||||
|
|
||||||
|
switch t.kind {
|
||||||
|
case KindOpenAI:
|
||||||
|
hc.Model = gjson.GetBytes(body, "model").String()
|
||||||
|
if len(t.models) > 0 {
|
||||||
|
if _, ok := t.models[hc.Model]; !ok {
|
||||||
|
return http.StatusForbidden, "model_not_allowed", "model not allowed"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
case KindMCP:
|
||||||
|
msgs := []gjson.Result{gjson.ParseBytes(body)}
|
||||||
|
if msgs[0].IsArray() {
|
||||||
|
msgs = msgs[0].Array()
|
||||||
|
}
|
||||||
|
for _, m := range msgs {
|
||||||
|
if m.Get("method").String() != "tools/call" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
tool := m.Get("params.name").String()
|
||||||
|
if len(t.tools) > 0 {
|
||||||
|
if _, ok := t.tools[tool]; !ok {
|
||||||
|
return http.StatusForbidden, "tool_not_allowed", "tool not allowed"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
hc.Tools = append(hc.Tools, tool)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return 0, "", ""
|
||||||
|
}
|
||||||
|
|
||||||
|
func hasRole(roles []string, want string) bool {
|
||||||
|
for _, r := range roles {
|
||||||
|
if r == want {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func anyRole(roles []string, allowed map[string]struct{}) bool {
|
||||||
|
for _, r := range roles {
|
||||||
|
if _, ok := allowed[r]; ok {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func errorBody(typ, msg string) map[string]any {
|
||||||
|
return map[string]any{"error": map[string]any{"message": msg, "type": typ, "code": typ}}
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeError(w http.ResponseWriter, status int, typ, msg string) {
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
w.WriteHeader(status)
|
||||||
|
_ = json.NewEncoder(w).Encode(errorBody(typ, msg))
|
||||||
|
}
|
||||||
@@ -0,0 +1,96 @@
|
|||||||
|
package aiproxy
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||||
|
)
|
||||||
|
|
||||||
|
// HookType defines when a hook runs.
|
||||||
|
type HookType string
|
||||||
|
|
||||||
|
const (
|
||||||
|
// BeforeProxy runs after auth, rate limit and request inspection, before the upstream call.
|
||||||
|
BeforeProxy HookType = "before_proxy"
|
||||||
|
// AfterProxy runs once the response finished (or failed). Errors are logged only.
|
||||||
|
AfterProxy HookType = "after_proxy"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Usage is the token usage reported by an OpenAI upstream.
|
||||||
|
type Usage struct {
|
||||||
|
PromptTokens int64
|
||||||
|
CompletionTokens int64
|
||||||
|
TotalTokens int64
|
||||||
|
}
|
||||||
|
|
||||||
|
// HookContext is passed to hooks.
|
||||||
|
type HookContext struct {
|
||||||
|
Context context.Context
|
||||||
|
Request *http.Request // BeforeProxy may modify headers/query; the key is injected later
|
||||||
|
UserContext *security.UserContext
|
||||||
|
|
||||||
|
Upstream string
|
||||||
|
Kind Kind
|
||||||
|
Model string // OpenAI: requested model
|
||||||
|
Tools []string // MCP: tools named by tools/call
|
||||||
|
|
||||||
|
// AfterProxy only
|
||||||
|
StatusCode int
|
||||||
|
Duration time.Duration
|
||||||
|
Usage Usage // OpenAI only; needs stream_options.include_usage for streams
|
||||||
|
Error error
|
||||||
|
|
||||||
|
// BeforeProxy only
|
||||||
|
Abort bool
|
||||||
|
AbortMessage string
|
||||||
|
AbortCode int // default 403
|
||||||
|
}
|
||||||
|
|
||||||
|
// HookFunc is a hook. A returned error aborts a BeforeProxy request.
|
||||||
|
type HookFunc func(*HookContext) error
|
||||||
|
|
||||||
|
// HookRegistry holds hooks per type.
|
||||||
|
type HookRegistry struct {
|
||||||
|
mu sync.RWMutex
|
||||||
|
hooks map[HookType][]HookFunc
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewHookRegistry creates an empty registry.
|
||||||
|
func NewHookRegistry() *HookRegistry {
|
||||||
|
return &HookRegistry{hooks: make(map[HookType][]HookFunc)}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Register adds a hook.
|
||||||
|
func (r *HookRegistry) Register(t HookType, h HookFunc) {
|
||||||
|
r.mu.Lock()
|
||||||
|
defer r.mu.Unlock()
|
||||||
|
r.hooks[t] = append(r.hooks[t], h)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Execute runs the hooks in order and stops at the first error or abort.
|
||||||
|
func (r *HookRegistry) Execute(t HookType, ctx *HookContext) error {
|
||||||
|
r.mu.RLock()
|
||||||
|
list := append([]HookFunc(nil), r.hooks[t]...)
|
||||||
|
r.mu.RUnlock()
|
||||||
|
|
||||||
|
for i, h := range list {
|
||||||
|
if err := h(ctx); err != nil {
|
||||||
|
return fmt.Errorf("aiproxy hook %d for %s failed: %w", i+1, t, err)
|
||||||
|
}
|
||||||
|
if ctx.Abort {
|
||||||
|
return fmt.Errorf("aborted by hook: %s", ctx.AbortMessage)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (r *HookRegistry) executeAfter(ctx *HookContext) {
|
||||||
|
if err := r.Execute(AfterProxy, ctx); err != nil {
|
||||||
|
logger.Warn("aiproxy: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,142 @@
|
|||||||
|
package aiproxy
|
||||||
|
|
||||||
|
import (
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/metrics"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Outcome classifies a handled request.
|
||||||
|
type Outcome string
|
||||||
|
|
||||||
|
const (
|
||||||
|
OutcomeOK Outcome = "ok" // upstream answered < 400
|
||||||
|
OutcomeUpstreamError Outcome = "upstream_error" // upstream answered >= 400 or failed
|
||||||
|
OutcomeDenied Outcome = "denied" // refused by auth, roles, allowlist, size or a hook
|
||||||
|
OutcomeRateLimited Outcome = "rate_limited"
|
||||||
|
)
|
||||||
|
|
||||||
|
// AuditRecord is written once per handled request. It holds no bodies and no keys.
|
||||||
|
type AuditRecord struct {
|
||||||
|
Time time.Time
|
||||||
|
UserID int
|
||||||
|
User string
|
||||||
|
RemoteID string
|
||||||
|
Upstream string
|
||||||
|
Kind Kind
|
||||||
|
Method string
|
||||||
|
Path string // client path, no query string
|
||||||
|
Model string
|
||||||
|
Tools []string
|
||||||
|
Status int
|
||||||
|
Outcome Outcome
|
||||||
|
Reason string // why a request was refused
|
||||||
|
Duration time.Duration
|
||||||
|
Usage Usage
|
||||||
|
Error string
|
||||||
|
}
|
||||||
|
|
||||||
|
// AuditSink receives audit records. It runs on the request path, so keep it fast
|
||||||
|
// (queue slow work such as database writes).
|
||||||
|
type AuditSink interface {
|
||||||
|
Record(AuditRecord)
|
||||||
|
}
|
||||||
|
|
||||||
|
// LogAuditSink writes records through pkg/logger. It is the default sink.
|
||||||
|
type LogAuditSink struct{}
|
||||||
|
|
||||||
|
// Record implements AuditSink.
|
||||||
|
func (LogAuditSink) Record(r AuditRecord) {
|
||||||
|
logger.Info("aiproxy audit: user=%s(%d) upstream=%s kind=%s %s %s model=%q tools=%v status=%d outcome=%s reason=%q duration=%s tokens=%d/%d/%d err=%q",
|
||||||
|
r.User, r.UserID, r.Upstream, r.Kind, r.Method, r.Path, r.Model, r.Tools, r.Status, r.Outcome, r.Reason, r.Duration,
|
||||||
|
r.Usage.PromptTokens, r.Usage.CompletionTokens, r.Usage.TotalTokens, r.Error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// NopAuditSink discards records (use it to switch the default logging off).
|
||||||
|
type NopAuditSink struct{}
|
||||||
|
|
||||||
|
// Record implements AuditSink.
|
||||||
|
func (NopAuditSink) Record(AuditRecord) {}
|
||||||
|
|
||||||
|
// maxModelLabels bounds the distinct (upstream, model) label pairs, since the
|
||||||
|
// model comes from the client.
|
||||||
|
const maxModelLabels = 400
|
||||||
|
|
||||||
|
type modelLabels struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
seen map[string]struct{}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *modelLabels) label(upstream, model string) string {
|
||||||
|
if model == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
key := upstream + "|" + model
|
||||||
|
m.mu.Lock()
|
||||||
|
defer m.mu.Unlock()
|
||||||
|
if _, ok := m.seen[key]; ok {
|
||||||
|
return model
|
||||||
|
}
|
||||||
|
if m.seen == nil {
|
||||||
|
m.seen = make(map[string]struct{})
|
||||||
|
}
|
||||||
|
if len(m.seen) >= maxModelLabels {
|
||||||
|
return "other"
|
||||||
|
}
|
||||||
|
m.seen[key] = struct{}{}
|
||||||
|
return model
|
||||||
|
}
|
||||||
|
|
||||||
|
// record emits the audit record and the metrics for a finished or refused request.
|
||||||
|
func (p *Proxy) record(t *target, hc *HookContext, status int, outcome Outcome, reason string, err error) {
|
||||||
|
rec := AuditRecord{
|
||||||
|
Time: time.Now(),
|
||||||
|
Upstream: hc.Upstream,
|
||||||
|
Kind: hc.Kind,
|
||||||
|
Model: hc.Model,
|
||||||
|
Tools: hc.Tools,
|
||||||
|
Status: status,
|
||||||
|
Outcome: outcome,
|
||||||
|
Reason: reason,
|
||||||
|
Duration: hc.Duration,
|
||||||
|
Usage: hc.Usage,
|
||||||
|
}
|
||||||
|
if hc.Request != nil {
|
||||||
|
rec.Method = hc.Request.Method
|
||||||
|
rec.Path = hc.Request.URL.Path
|
||||||
|
}
|
||||||
|
if u := hc.UserContext; u != nil {
|
||||||
|
rec.UserID, rec.User, rec.RemoteID = u.UserID, u.UserName, u.RemoteID
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
rec.Error = err.Error()
|
||||||
|
}
|
||||||
|
|
||||||
|
func() {
|
||||||
|
defer func() {
|
||||||
|
if r := recover(); r != nil {
|
||||||
|
logger.Error("aiproxy: audit sink panic: %v", r)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
p.cfg.Audit.Record(rec)
|
||||||
|
}()
|
||||||
|
|
||||||
|
if m, ok := metrics.GetProvider().(metrics.AIProxyRecorder); ok && m != nil {
|
||||||
|
model := ""
|
||||||
|
if outcome == OutcomeOK || outcome == OutcomeUpstreamError {
|
||||||
|
model = p.labels.label(hc.Upstream, hc.Model)
|
||||||
|
}
|
||||||
|
m.RecordAIProxy(hc.Upstream, string(hc.Kind), model, statusClass(status), string(outcome), hc.Duration,
|
||||||
|
hc.Usage.PromptTokens, hc.Usage.CompletionTokens)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// statusClass maps a status code to a low-cardinality label (2xx, 4xx, ...).
|
||||||
|
func statusClass(status int) string {
|
||||||
|
if status < 100 || status > 599 {
|
||||||
|
return "unknown"
|
||||||
|
}
|
||||||
|
return string(rune('0'+status/100)) + "xx"
|
||||||
|
}
|
||||||
@@ -0,0 +1,80 @@
|
|||||||
|
package aiproxy
|
||||||
|
|
||||||
|
import (
|
||||||
|
"math"
|
||||||
|
"strings"
|
||||||
|
"sync"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"golang.org/x/time/rate"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
limiterSweepEvery = 5 * time.Minute
|
||||||
|
limiterIdleAfter = 10 * time.Minute
|
||||||
|
)
|
||||||
|
|
||||||
|
type limiterEntry struct {
|
||||||
|
lim *rate.Limiter
|
||||||
|
last time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
// limiters keeps one token bucket per key and drops idle ones.
|
||||||
|
type limiters struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
entries map[string]*limiterEntry
|
||||||
|
lastSweep time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
func newLimiters() *limiters {
|
||||||
|
return &limiters{entries: make(map[string]*limiterEntry), lastSweep: time.Now()}
|
||||||
|
}
|
||||||
|
|
||||||
|
// allow takes a token. When denied it returns how long to wait.
|
||||||
|
func (l *limiters) allow(key string, rl *RateLimit) (bool, time.Duration) {
|
||||||
|
now := time.Now()
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
|
if now.Sub(l.lastSweep) > limiterSweepEvery {
|
||||||
|
for k, e := range l.entries {
|
||||||
|
if now.Sub(e.last) > limiterIdleAfter {
|
||||||
|
delete(l.entries, k)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
l.lastSweep = now
|
||||||
|
}
|
||||||
|
|
||||||
|
e, ok := l.entries[key]
|
||||||
|
if !ok {
|
||||||
|
burst := rl.Burst
|
||||||
|
if burst < 1 {
|
||||||
|
burst = int(math.Ceil(rl.PerSecond))
|
||||||
|
if burst < 1 {
|
||||||
|
burst = 1
|
||||||
|
}
|
||||||
|
}
|
||||||
|
e = &limiterEntry{lim: rate.NewLimiter(rate.Limit(rl.PerSecond), burst)}
|
||||||
|
l.entries[key] = e
|
||||||
|
}
|
||||||
|
e.last = now
|
||||||
|
|
||||||
|
res := e.lim.ReserveN(now, 1)
|
||||||
|
if d := res.DelayFrom(now); d > 0 {
|
||||||
|
res.CancelAt(now)
|
||||||
|
return false, d
|
||||||
|
}
|
||||||
|
return true, 0
|
||||||
|
}
|
||||||
|
|
||||||
|
// forget drops the buckets of an upstream (after it was replaced or removed).
|
||||||
|
func (l *limiters) forget(upstream string) {
|
||||||
|
prefix := upstream + "|"
|
||||||
|
l.mu.Lock()
|
||||||
|
defer l.mu.Unlock()
|
||||||
|
for k := range l.entries {
|
||||||
|
if strings.HasPrefix(k, prefix) {
|
||||||
|
delete(l.entries, k)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,42 @@
|
|||||||
|
-- aiproxy upstream definitions (PostgreSQL), same contract as the resolvespec security procedures.
|
||||||
|
-- ProcStore calls: SELECT p_success, p_error, p_data FROM resolvespec_ai_proxies()
|
||||||
|
-- Replace the function body to source the definitions from anywhere; keep the signature.
|
||||||
|
--
|
||||||
|
-- p_data: JSON array, one object per upstream
|
||||||
|
-- name, kind ('openai'|'mcp'), base_url required
|
||||||
|
-- api_key, auth_header, auth_format, headers (object) optional
|
||||||
|
-- allowed_roles (array), allowed (array: models for openai, tools for mcp)
|
||||||
|
-- rate_per_second, rate_burst, timeout_seconds, enabled (default true)
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS ai_proxies (
|
||||||
|
name text PRIMARY KEY,
|
||||||
|
kind text NOT NULL CHECK (kind IN ('openai', 'mcp')),
|
||||||
|
base_url text NOT NULL,
|
||||||
|
api_key text,
|
||||||
|
auth_header text,
|
||||||
|
auth_format text,
|
||||||
|
headers jsonb,
|
||||||
|
allowed_roles jsonb,
|
||||||
|
allowed jsonb,
|
||||||
|
rate_per_second double precision,
|
||||||
|
rate_burst integer,
|
||||||
|
timeout_seconds integer,
|
||||||
|
enabled boolean NOT NULL DEFAULT true
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE OR REPLACE FUNCTION resolvespec_ai_proxies()
|
||||||
|
RETURNS TABLE(p_success boolean, p_error text, p_data jsonb) AS $$
|
||||||
|
DECLARE
|
||||||
|
v_data jsonb;
|
||||||
|
BEGIN
|
||||||
|
SELECT COALESCE(jsonb_agg(to_jsonb(p) - 'enabled' ORDER BY p.name), '[]'::jsonb)
|
||||||
|
INTO v_data
|
||||||
|
FROM ai_proxies p
|
||||||
|
WHERE p.enabled = true;
|
||||||
|
|
||||||
|
RETURN QUERY SELECT true, NULL::text, v_data;
|
||||||
|
EXCEPTION
|
||||||
|
WHEN OTHERS THEN
|
||||||
|
RETURN QUERY SELECT false, SQLERRM::text, '[]'::jsonb;
|
||||||
|
END;
|
||||||
|
$$ LANGUAGE plpgsql;
|
||||||
@@ -0,0 +1,231 @@
|
|||||||
|
package aiproxy
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"reflect"
|
||||||
|
"regexp"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||||
|
)
|
||||||
|
|
||||||
|
// UpstreamStore supplies upstream definitions for Reload.
|
||||||
|
type UpstreamStore interface {
|
||||||
|
List(ctx context.Context) ([]Definition, error)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Reload syncs the store-managed upstreams with the store.
|
||||||
|
//
|
||||||
|
// - New definitions are added, changed ones replaced, unchanged ones left alone
|
||||||
|
// (their rate-limit state and connections survive), missing ones removed.
|
||||||
|
// - Upstreams registered in code are never replaced or removed; a stored definition
|
||||||
|
// with the same name is skipped and reported.
|
||||||
|
// - Invalid definitions are skipped and reported in the returned error; valid ones
|
||||||
|
// are still applied. If the store itself fails, nothing changes.
|
||||||
|
func (p *Proxy) Reload(ctx context.Context, store UpstreamStore) error {
|
||||||
|
defs, err := store.List(ctx)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("aiproxy: load upstreams: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
var errs []error
|
||||||
|
next := make(map[string]*target, len(defs))
|
||||||
|
var retired []*target
|
||||||
|
|
||||||
|
p.mu.Lock()
|
||||||
|
for i := range defs {
|
||||||
|
d := defs[i]
|
||||||
|
name := d.Upstream.Name
|
||||||
|
if _, dup := next[name]; dup {
|
||||||
|
errs = append(errs, fmt.Errorf("aiproxy: duplicate stored upstream %q", name))
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
cur := p.upstream[name]
|
||||||
|
if cur != nil && !cur.managed {
|
||||||
|
errs = append(errs, fmt.Errorf("aiproxy: stored upstream %q conflicts with a code-registered one", name))
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if cur != nil && reflect.DeepEqual(cur.def, d) {
|
||||||
|
next[name] = cur
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
t, err := p.build(d)
|
||||||
|
if err != nil {
|
||||||
|
errs = append(errs, err)
|
||||||
|
if cur != nil { // keep serving the last good definition
|
||||||
|
next[name] = cur
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
t.managed = true
|
||||||
|
next[name] = t
|
||||||
|
if cur != nil {
|
||||||
|
retired = append(retired, cur)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for name, cur := range p.upstream {
|
||||||
|
if !cur.managed {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if _, keep := next[name]; !keep {
|
||||||
|
delete(p.upstream, name)
|
||||||
|
retired = append(retired, cur)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for name, t := range next {
|
||||||
|
p.upstream[name] = t
|
||||||
|
}
|
||||||
|
p.mu.Unlock()
|
||||||
|
|
||||||
|
for _, t := range retired {
|
||||||
|
p.retire(t)
|
||||||
|
}
|
||||||
|
return errors.Join(errs...)
|
||||||
|
}
|
||||||
|
|
||||||
|
// AutoReload calls Reload now and then every interval until ctx ends. The first error
|
||||||
|
// is returned; later ones are logged.
|
||||||
|
func (p *Proxy) AutoReload(ctx context.Context, store UpstreamStore, interval time.Duration) error {
|
||||||
|
if interval <= 0 {
|
||||||
|
return errors.New("aiproxy: reload interval must be > 0")
|
||||||
|
}
|
||||||
|
first := p.Reload(ctx, store)
|
||||||
|
go func() {
|
||||||
|
t := time.NewTicker(interval)
|
||||||
|
defer t.Stop()
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
case <-t.C:
|
||||||
|
if err := p.Reload(ctx, store); err != nil {
|
||||||
|
logger.Warn("aiproxy: reload: %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
return first
|
||||||
|
}
|
||||||
|
|
||||||
|
// DefaultProc is the stored procedure (function) ProcStore calls when no name is given.
|
||||||
|
const DefaultProc = "resolvespec_ai_proxies"
|
||||||
|
|
||||||
|
var procRe = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*(\.[A-Za-z_][A-Za-z0-9_]*)?$`)
|
||||||
|
|
||||||
|
// ProcStore loads upstream definitions from a stored procedure (PostgreSQL), using the
|
||||||
|
// same contract as the security procedures:
|
||||||
|
//
|
||||||
|
// SELECT p_success, p_error, p_data FROM resolvespec_ai_proxies()
|
||||||
|
//
|
||||||
|
// p_data is a JSON array; each element has: name, kind ("openai"|"mcp"), base_url, and
|
||||||
|
// optionally api_key, auth_header, auth_format, headers (object), allowed_roles (array),
|
||||||
|
// allowed (array: models or tools), rate_per_second, rate_burst, timeout_seconds,
|
||||||
|
// enabled (default true). See resolvespec_ai_proxies.sql.
|
||||||
|
type ProcStore struct {
|
||||||
|
db func() *sql.DB
|
||||||
|
proc string
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewProcStore calls proc through db. An empty proc means DefaultProc.
|
||||||
|
func NewProcStore(db *sql.DB, proc string) (*ProcStore, error) {
|
||||||
|
if db == nil {
|
||||||
|
return nil, errors.New("aiproxy: nil database")
|
||||||
|
}
|
||||||
|
return newProcStore(func() *sql.DB { return db }, proc)
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewProcStoreFromDatabase is NewProcStore for an application's common.Database. The
|
||||||
|
// connection is fetched on every call, so adapter reconnects are followed.
|
||||||
|
func NewProcStoreFromDatabase(db common.Database, proc string) (*ProcStore, error) {
|
||||||
|
if db == nil {
|
||||||
|
return nil, errors.New("aiproxy: nil database")
|
||||||
|
}
|
||||||
|
p, ok := db.(common.SQLDBProvider)
|
||||||
|
if !ok {
|
||||||
|
return nil, fmt.Errorf("aiproxy: %T does not expose a *sql.DB", db)
|
||||||
|
}
|
||||||
|
return newProcStore(p.SQLDB, proc)
|
||||||
|
}
|
||||||
|
|
||||||
|
func newProcStore(db func() *sql.DB, proc string) (*ProcStore, error) {
|
||||||
|
if proc == "" {
|
||||||
|
proc = DefaultProc
|
||||||
|
}
|
||||||
|
if !procRe.MatchString(proc) {
|
||||||
|
return nil, fmt.Errorf("aiproxy: invalid procedure name %q", proc)
|
||||||
|
}
|
||||||
|
return &ProcStore{db: db, proc: proc}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type procRow struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Kind string `json:"kind"`
|
||||||
|
BaseURL string `json:"base_url"`
|
||||||
|
APIKey string `json:"api_key"`
|
||||||
|
AuthHeader string `json:"auth_header"`
|
||||||
|
AuthFormat string `json:"auth_format"`
|
||||||
|
Headers map[string]string `json:"headers"`
|
||||||
|
AllowedRoles []string `json:"allowed_roles"`
|
||||||
|
Allowed []string `json:"allowed"`
|
||||||
|
RatePerSecond float64 `json:"rate_per_second"`
|
||||||
|
RateBurst int `json:"rate_burst"`
|
||||||
|
TimeoutSeconds int `json:"timeout_seconds"`
|
||||||
|
Enabled *bool `json:"enabled"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// List implements UpstreamStore. Disabled and undecodable entries are skipped.
|
||||||
|
func (s *ProcStore) List(ctx context.Context) ([]Definition, error) {
|
||||||
|
db := s.db()
|
||||||
|
if db == nil {
|
||||||
|
return nil, errors.New("database connection is nil")
|
||||||
|
}
|
||||||
|
var (
|
||||||
|
success bool
|
||||||
|
errMsg sql.NullString
|
||||||
|
data []byte
|
||||||
|
)
|
||||||
|
if err := db.QueryRowContext(ctx, `SELECT p_success, p_error, p_data FROM `+s.proc+`()`).Scan(&success, &errMsg, &data); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if !success {
|
||||||
|
if errMsg.Valid {
|
||||||
|
return nil, errors.New(errMsg.String)
|
||||||
|
}
|
||||||
|
return nil, errors.New("failed to load AI proxies")
|
||||||
|
}
|
||||||
|
if len(data) == 0 {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
var raw []json.RawMessage
|
||||||
|
if err := json.Unmarshal(data, &raw); err != nil {
|
||||||
|
return nil, fmt.Errorf("parse AI proxies: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
out := make([]Definition, 0, len(raw))
|
||||||
|
for i, r := range raw {
|
||||||
|
var row procRow
|
||||||
|
if err := json.Unmarshal(r, &row); err != nil {
|
||||||
|
logger.Warn("aiproxy: proxy entry %d: %v", i, err)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if row.Enabled != nil && !*row.Enabled {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
d := Definition{Kind: Kind(row.Kind), Allowed: row.Allowed, Upstream: Upstream{
|
||||||
|
Name: row.Name, BaseURL: row.BaseURL, APIKey: row.APIKey,
|
||||||
|
AuthHeader: row.AuthHeader, AuthFormat: row.AuthFormat,
|
||||||
|
Headers: row.Headers, AllowedRoles: row.AllowedRoles,
|
||||||
|
Timeout: time.Duration(row.TimeoutSeconds) * time.Second,
|
||||||
|
}}
|
||||||
|
if row.RatePerSecond > 0 {
|
||||||
|
d.Upstream.RateLimit = &RateLimit{PerSecond: row.RatePerSecond, Burst: row.RateBurst}
|
||||||
|
}
|
||||||
|
out = append(out, d)
|
||||||
|
}
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,269 @@
|
|||||||
|
package aiproxy
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/DATA-DOG/go-sqlmock"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/metrics"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
|
"github.com/stretchr/testify/require"
|
||||||
|
)
|
||||||
|
|
||||||
|
type memStore struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
defs []Definition
|
||||||
|
err error
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *memStore) List(context.Context) ([]Definition, error) {
|
||||||
|
m.mu.Lock()
|
||||||
|
defer m.mu.Unlock()
|
||||||
|
return append([]Definition(nil), m.defs...), m.err
|
||||||
|
}
|
||||||
|
|
||||||
|
func def(name, url, key string) Definition {
|
||||||
|
return Definition{Kind: KindOpenAI, Upstream: Upstream{Name: name, BaseURL: url + "/v1", APIKey: key}}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReloadAddChangeRemove(t *testing.T) {
|
||||||
|
up := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
io.WriteString(w, r.Header.Get("Authorization"))
|
||||||
|
}))
|
||||||
|
defer up.Close()
|
||||||
|
p := New(Config{Prefix: "/ai", Audit: NopAuditSink{}})
|
||||||
|
require.NoError(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "code", BaseURL: up.URL + "/v1", APIKey: "c"}}))
|
||||||
|
front := httptest.NewServer(p.Handler(testAuth))
|
||||||
|
defer front.Close()
|
||||||
|
get := func(name string) (int, string) {
|
||||||
|
resp, body := do(t, "GET", front.URL+"/ai/"+name+"/v1/models", "", "bob:u", nil)
|
||||||
|
return resp.StatusCode, body
|
||||||
|
}
|
||||||
|
|
||||||
|
st := &memStore{defs: []Definition{def("a", up.URL, "k1"), def("code", up.URL, "stolen")}}
|
||||||
|
err := p.Reload(context.Background(), st)
|
||||||
|
assert.ErrorContains(t, err, "conflicts with a code-registered")
|
||||||
|
code, body := get("a")
|
||||||
|
assert.Equal(t, 200, code)
|
||||||
|
assert.Equal(t, "Bearer k1", body)
|
||||||
|
_, body = get("code")
|
||||||
|
assert.Equal(t, "Bearer c", body) // code-registered untouched
|
||||||
|
|
||||||
|
st.defs = []Definition{def("a", up.URL, "k2"), def("b", up.URL, "kb")}
|
||||||
|
require.NoError(t, p.Reload(context.Background(), st))
|
||||||
|
_, body = get("a")
|
||||||
|
assert.Equal(t, "Bearer k2", body)
|
||||||
|
_, body = get("b")
|
||||||
|
assert.Equal(t, "Bearer kb", body)
|
||||||
|
|
||||||
|
st.defs = []Definition{def("b", up.URL, "kb")}
|
||||||
|
require.NoError(t, p.Reload(context.Background(), st))
|
||||||
|
code, _ = get("a")
|
||||||
|
assert.Equal(t, 404, code)
|
||||||
|
code, _ = get("code")
|
||||||
|
assert.Equal(t, 200, code)
|
||||||
|
|
||||||
|
// store failure changes nothing; invalid def keeps last good one
|
||||||
|
st.err = errors.New("db down")
|
||||||
|
assert.Error(t, p.Reload(context.Background(), st))
|
||||||
|
code, _ = get("b")
|
||||||
|
assert.Equal(t, 200, code)
|
||||||
|
st.err = nil
|
||||||
|
st.defs = []Definition{{Kind: KindOpenAI, Upstream: Upstream{Name: "b", BaseURL: "ftp://x"}}}
|
||||||
|
assert.Error(t, p.Reload(context.Background(), st))
|
||||||
|
_, body = get("b")
|
||||||
|
assert.Equal(t, "Bearer kb", body)
|
||||||
|
|
||||||
|
assert.True(t, p.Unregister("code"))
|
||||||
|
assert.False(t, p.Unregister("code"))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestReloadKeepsRateLimitStateWhenUnchanged(t *testing.T) {
|
||||||
|
up := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {}))
|
||||||
|
defer up.Close()
|
||||||
|
p := New(Config{Prefix: "/ai", Audit: NopAuditSink{}})
|
||||||
|
d := def("a", up.URL, "")
|
||||||
|
d.Upstream.RateLimit = &RateLimit{PerSecond: 0.01, Burst: 1}
|
||||||
|
st := &memStore{defs: []Definition{d}}
|
||||||
|
require.NoError(t, p.Reload(context.Background(), st))
|
||||||
|
front := httptest.NewServer(p.Handler(testAuth))
|
||||||
|
defer front.Close()
|
||||||
|
|
||||||
|
resp, _ := do(t, "GET", front.URL+"/ai/a/v1/x", "", "bob:u", nil)
|
||||||
|
assert.Equal(t, 200, resp.StatusCode)
|
||||||
|
require.NoError(t, p.Reload(context.Background(), st))
|
||||||
|
resp, _ = do(t, "GET", front.URL+"/ai/a/v1/x", "", "bob:u", nil)
|
||||||
|
assert.Equal(t, 429, resp.StatusCode)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAutoReload(t *testing.T) {
|
||||||
|
p := New(Config{Audit: NopAuditSink{}})
|
||||||
|
st := &memStore{}
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
defer cancel()
|
||||||
|
require.NoError(t, p.AutoReload(ctx, st, 20*time.Millisecond))
|
||||||
|
assert.Empty(t, p.Names())
|
||||||
|
st.mu.Lock()
|
||||||
|
st.defs = []Definition{def("a", "http://x", "")}
|
||||||
|
st.mu.Unlock()
|
||||||
|
assert.Eventually(t, func() bool { return len(p.Names()) == 1 }, 2*time.Second, 10*time.Millisecond)
|
||||||
|
assert.Error(t, p.AutoReload(ctx, st, 0))
|
||||||
|
}
|
||||||
|
|
||||||
|
func procRows(data string) *sqlmock.Rows {
|
||||||
|
return sqlmock.NewRows([]string{"p_success", "p_error", "p_data"}).AddRow(true, nil, []byte(data))
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProcStore(t *testing.T) {
|
||||||
|
db, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(sqlmock.QueryMatcherRegexp))
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer db.Close()
|
||||||
|
|
||||||
|
_, err = NewProcStore(db, "x; drop table y")
|
||||||
|
assert.Error(t, err)
|
||||||
|
_, err = NewProcStore(nil, "")
|
||||||
|
assert.Error(t, err)
|
||||||
|
s, err := NewProcStore(db, "")
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
mock.ExpectQuery(`SELECT p_success, p_error, p_data FROM resolvespec_ai_proxies\(\)`).WillReturnRows(procRows(`[
|
||||||
|
{"name":"gpt","kind":"openai","base_url":"https://api.openai.com/v1","api_key":"sk","auth_header":"api-key",
|
||||||
|
"auth_format":"%s","headers":{"X-A":"1"},"allowed_roles":["ai"],"allowed":["m1","m2"],
|
||||||
|
"rate_per_second":2.5,"rate_burst":5,"timeout_seconds":30},
|
||||||
|
{"name":"tools","kind":"mcp","base_url":"http://localhost:3000/mcp","allowed":["t1"],"enabled":true},
|
||||||
|
{"name":"off","kind":"mcp","base_url":"http://localhost:3001/mcp","enabled":false},
|
||||||
|
{"name":5}
|
||||||
|
]`))
|
||||||
|
defs, err := s.List(context.Background())
|
||||||
|
require.NoError(t, err)
|
||||||
|
require.Len(t, defs, 2)
|
||||||
|
|
||||||
|
gpt := defs[0]
|
||||||
|
assert.Equal(t, KindOpenAI, gpt.Kind)
|
||||||
|
assert.Equal(t, "sk", gpt.Upstream.APIKey)
|
||||||
|
assert.Equal(t, "api-key", gpt.Upstream.AuthHeader)
|
||||||
|
assert.Equal(t, map[string]string{"X-A": "1"}, gpt.Upstream.Headers)
|
||||||
|
assert.Equal(t, []string{"ai"}, gpt.Upstream.AllowedRoles)
|
||||||
|
assert.Equal(t, []string{"m1", "m2"}, gpt.Allowed)
|
||||||
|
assert.Equal(t, &RateLimit{PerSecond: 2.5, Burst: 5}, gpt.Upstream.RateLimit)
|
||||||
|
assert.Equal(t, 30*time.Second, gpt.Upstream.Timeout)
|
||||||
|
assert.Equal(t, KindMCP, defs[1].Kind)
|
||||||
|
assert.Nil(t, defs[1].Upstream.RateLimit)
|
||||||
|
|
||||||
|
// feeds Reload
|
||||||
|
mock.ExpectQuery(`FROM resolvespec_ai_proxies`).WillReturnRows(procRows(`[{"name":"a","kind":"mcp","base_url":"http://x/mcp"},{"name":"bad","kind":"nope","base_url":"http://x"}]`))
|
||||||
|
p := New(Config{Audit: NopAuditSink{}})
|
||||||
|
assert.Error(t, p.Reload(context.Background(), s)) // "bad" reported
|
||||||
|
assert.Equal(t, []string{"a"}, p.Names())
|
||||||
|
assert.NoError(t, mock.ExpectationsWereMet())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestProcStoreFailures(t *testing.T) {
|
||||||
|
db, mock, err := sqlmock.New(sqlmock.QueryMatcherOption(sqlmock.QueryMatcherRegexp))
|
||||||
|
require.NoError(t, err)
|
||||||
|
defer db.Close()
|
||||||
|
s, err := NewProcStore(db, "myschema.my_proxies")
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
mock.ExpectQuery(`FROM myschema.my_proxies\(\)`).WillReturnRows(
|
||||||
|
sqlmock.NewRows([]string{"p_success", "p_error", "p_data"}).AddRow(false, "boom", []byte(`[]`)))
|
||||||
|
_, err = s.List(context.Background())
|
||||||
|
assert.EqualError(t, err, "boom")
|
||||||
|
|
||||||
|
mock.ExpectQuery(`FROM myschema.my_proxies`).WillReturnRows(procRows(`{not json`))
|
||||||
|
_, err = s.List(context.Background())
|
||||||
|
assert.ErrorContains(t, err, "parse AI proxies")
|
||||||
|
|
||||||
|
mock.ExpectQuery(`FROM myschema.my_proxies`).WillReturnError(errors.New("conn lost"))
|
||||||
|
_, err = s.List(context.Background())
|
||||||
|
assert.ErrorContains(t, err, "conn lost")
|
||||||
|
|
||||||
|
mock.ExpectQuery(`FROM myschema.my_proxies`).WillReturnRows(
|
||||||
|
sqlmock.NewRows([]string{"p_success", "p_error", "p_data"}).AddRow(true, nil, nil))
|
||||||
|
defs, err := s.List(context.Background())
|
||||||
|
assert.NoError(t, err)
|
||||||
|
assert.Empty(t, defs)
|
||||||
|
}
|
||||||
|
|
||||||
|
type auditCollector struct {
|
||||||
|
mu sync.Mutex
|
||||||
|
recs []AuditRecord
|
||||||
|
}
|
||||||
|
|
||||||
|
func (a *auditCollector) Record(r AuditRecord) {
|
||||||
|
a.mu.Lock()
|
||||||
|
a.recs = append(a.recs, r)
|
||||||
|
a.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
type metricsCollector struct {
|
||||||
|
metrics.Provider
|
||||||
|
mu sync.Mutex
|
||||||
|
calls []string
|
||||||
|
toks int64
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *metricsCollector) RecordAIProxy(up, kind, model, class, outcome string, d time.Duration, pt, ct int64) {
|
||||||
|
m.mu.Lock()
|
||||||
|
m.calls = append(m.calls, up+"|"+kind+"|"+model+"|"+class+"|"+outcome)
|
||||||
|
m.toks += pt + ct
|
||||||
|
m.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAuditAndMetrics(t *testing.T) {
|
||||||
|
old := metrics.GetProvider()
|
||||||
|
mc := &metricsCollector{Provider: old}
|
||||||
|
metrics.SetProvider(mc)
|
||||||
|
defer metrics.SetProvider(old)
|
||||||
|
|
||||||
|
ac := &auditCollector{}
|
||||||
|
up := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
io.WriteString(w, `{"usage":{"prompt_tokens":1,"completion_tokens":2}}`)
|
||||||
|
}))
|
||||||
|
defer up.Close()
|
||||||
|
p := New(Config{Prefix: "/ai", Audit: ac})
|
||||||
|
require.NoError(t, p.RegisterOpenAI(OpenAIUpstream{Upstream: Upstream{Name: "g", BaseURL: up.URL + "/v1", APIKey: "sk-secret", AllowedRoles: []string{"ai"}}, AllowedModels: []string{"m"}}))
|
||||||
|
front := httptest.NewServer(p.Handler(testAuth))
|
||||||
|
defer front.Close()
|
||||||
|
|
||||||
|
do(t, "POST", front.URL+"/ai/g/v1/chat/completions?k=v", `{"model":"m"}`, "bob:ai", nil)
|
||||||
|
do(t, "POST", front.URL+"/ai/g/v1/chat/completions", `{"model":"zzz"}`, "bob:ai", nil)
|
||||||
|
do(t, "GET", front.URL+"/ai/g/v1/models", "", "bob:user", nil)
|
||||||
|
|
||||||
|
require.Eventually(t, func() bool { ac.mu.Lock(); defer ac.mu.Unlock(); return len(ac.recs) == 3 }, 2*time.Second, 10*time.Millisecond)
|
||||||
|
ac.mu.Lock()
|
||||||
|
r := ac.recs
|
||||||
|
ac.mu.Unlock()
|
||||||
|
assert.Equal(t, OutcomeOK, r[0].Outcome)
|
||||||
|
assert.Equal(t, "bob", r[0].User)
|
||||||
|
assert.Equal(t, "/ai/g/v1/chat/completions", r[0].Path) // no query
|
||||||
|
assert.Equal(t, "m", r[0].Model)
|
||||||
|
assert.Equal(t, int64(3), r[0].Usage.TotalTokens)
|
||||||
|
assert.Equal(t, OutcomeDenied, r[1].Outcome)
|
||||||
|
assert.Equal(t, 403, r[1].Status)
|
||||||
|
assert.Equal(t, OutcomeDenied, r[2].Outcome)
|
||||||
|
for _, rec := range r {
|
||||||
|
assert.NotContains(t, rec.Reason+rec.Error+rec.Path, "sk-secret")
|
||||||
|
}
|
||||||
|
|
||||||
|
mc.mu.Lock()
|
||||||
|
defer mc.mu.Unlock()
|
||||||
|
assert.Equal(t, []string{"g|openai|m|2xx|ok", "g|openai||4xx|denied", "g|openai||4xx|denied"}, mc.calls)
|
||||||
|
assert.Equal(t, int64(3), mc.toks)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestModelLabelsAreBounded(t *testing.T) {
|
||||||
|
var m modelLabels
|
||||||
|
for i := 0; i < maxModelLabels+50; i++ {
|
||||||
|
m.label("u", "model-"+string(rune('a'+i%26))+string(rune('a'+(i/26)%26))+string(rune('a'+(i/676)%26)))
|
||||||
|
}
|
||||||
|
assert.Equal(t, "other", m.label("u", "brand-new"))
|
||||||
|
assert.Equal(t, "", m.label("u", ""))
|
||||||
|
}
|
||||||
@@ -0,0 +1,143 @@
|
|||||||
|
package aiproxy
|
||||||
|
|
||||||
|
import (
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
"net/url"
|
||||||
|
"regexp"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Kind is the type of an upstream.
|
||||||
|
type Kind string
|
||||||
|
|
||||||
|
const (
|
||||||
|
KindOpenAI Kind = "openai"
|
||||||
|
KindMCP Kind = "mcp"
|
||||||
|
)
|
||||||
|
|
||||||
|
const redacted = "[redacted]"
|
||||||
|
|
||||||
|
var nameRe = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9_.-]*$`)
|
||||||
|
|
||||||
|
// RateLimit limits requests per user per upstream (token bucket).
|
||||||
|
type RateLimit struct {
|
||||||
|
PerSecond float64 // sustained requests per second
|
||||||
|
Burst int // bucket size; defaults to ceil(PerSecond), min 1
|
||||||
|
}
|
||||||
|
|
||||||
|
// Upstream holds the settings shared by every upstream kind.
|
||||||
|
type Upstream struct {
|
||||||
|
// Name is the route segment clients use: {prefix}/{Name}/...
|
||||||
|
Name string
|
||||||
|
// BaseURL of the upstream. OpenAI: include the version path (https://api.openai.com/v1).
|
||||||
|
// MCP: the full streamable HTTP endpoint (http://localhost:3000/mcp).
|
||||||
|
BaseURL string
|
||||||
|
// APIKey is sent to the upstream. Never exposed to clients or logs.
|
||||||
|
APIKey string
|
||||||
|
// AuthHeader carries the key. Default "Authorization". Applies to OpenAI and MCP upstreams.
|
||||||
|
AuthHeader string
|
||||||
|
// AuthFormat renders the header value; %s is the key (e.g. "Bearer %s", "Token %s", "%s").
|
||||||
|
// Default: "Bearer %s" for the Authorization header, "%s" (raw key) for any other header.
|
||||||
|
AuthFormat string
|
||||||
|
// Headers are extra static headers sent to the upstream (applied before the key header).
|
||||||
|
// Per-request values can be set in a BeforeProxy hook.
|
||||||
|
Headers map[string]string
|
||||||
|
// AllowedRoles: caller needs at least one. Empty means any authenticated user.
|
||||||
|
AllowedRoles []string
|
||||||
|
// RateLimit per user for this upstream. Nil means unlimited.
|
||||||
|
RateLimit *RateLimit
|
||||||
|
// Timeout waiting for the upstream response headers. Default 120s. Streams are not cut off.
|
||||||
|
Timeout time.Duration
|
||||||
|
}
|
||||||
|
|
||||||
|
// OpenAIUpstream is an OpenAI-compatible API (OpenAI, Azure, Ollama, vLLM, ...).
|
||||||
|
type OpenAIUpstream struct {
|
||||||
|
Upstream
|
||||||
|
// AllowedModels restricts the "model" a request may use. Empty means any.
|
||||||
|
// When set, write requests must be JSON bodies carrying an allowed model.
|
||||||
|
AllowedModels []string
|
||||||
|
}
|
||||||
|
|
||||||
|
// MCPUpstream is an MCP server speaking streamable HTTP.
|
||||||
|
type MCPUpstream struct {
|
||||||
|
Upstream
|
||||||
|
// AllowedTools restricts tools/call to these tool names. Empty means any.
|
||||||
|
AllowedTools []string
|
||||||
|
}
|
||||||
|
|
||||||
|
// String redacts the key.
|
||||||
|
func (u Upstream) String() string {
|
||||||
|
return fmt.Sprintf("Upstream{Name:%s BaseURL:%s APIKey:%s}", u.Name, u.BaseURL, keyMask(u.APIKey))
|
||||||
|
}
|
||||||
|
|
||||||
|
// GoString redacts the key (%#v).
|
||||||
|
func (u Upstream) GoString() string { return u.String() }
|
||||||
|
|
||||||
|
// MarshalJSON redacts the key.
|
||||||
|
func (u Upstream) MarshalJSON() ([]byte, error) {
|
||||||
|
return json.Marshal(struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
BaseURL string `json:"base_url"`
|
||||||
|
APIKey string `json:"api_key,omitempty"`
|
||||||
|
AllowedRoles []string `json:"allowed_roles,omitempty"`
|
||||||
|
}{u.Name, u.BaseURL, keyMask(u.APIKey), u.AllowedRoles})
|
||||||
|
}
|
||||||
|
|
||||||
|
func keyMask(k string) string {
|
||||||
|
if k == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
return redacted
|
||||||
|
}
|
||||||
|
|
||||||
|
// validate checks the shared fields and returns the parsed base URL.
|
||||||
|
func (u Upstream) validate() (*url.URL, error) {
|
||||||
|
if !nameRe.MatchString(u.Name) {
|
||||||
|
return nil, fmt.Errorf("aiproxy: invalid upstream name %q", u.Name)
|
||||||
|
}
|
||||||
|
base, err := url.Parse(strings.TrimRight(u.BaseURL, "/"))
|
||||||
|
if err != nil || (base.Scheme != "http" && base.Scheme != "https") || base.Host == "" {
|
||||||
|
return nil, fmt.Errorf("aiproxy: upstream %q has an invalid base URL", u.Name)
|
||||||
|
}
|
||||||
|
if u.AuthFormat != "" && (strings.Count(u.AuthFormat, "%s") != 1 || strings.Count(u.AuthFormat, "%") != 1) {
|
||||||
|
return nil, fmt.Errorf("aiproxy: upstream %q auth format must contain exactly one %%s", u.Name)
|
||||||
|
}
|
||||||
|
if u.RateLimit != nil && u.RateLimit.PerSecond <= 0 {
|
||||||
|
return nil, fmt.Errorf("aiproxy: upstream %q rate limit must be > 0", u.Name)
|
||||||
|
}
|
||||||
|
return base, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// authValue returns the header name and value carrying the key, or "" when no key is set.
|
||||||
|
func (u Upstream) authValue() (name, value string) {
|
||||||
|
if u.APIKey == "" {
|
||||||
|
return "", ""
|
||||||
|
}
|
||||||
|
name = u.AuthHeader
|
||||||
|
if name == "" {
|
||||||
|
name = "Authorization"
|
||||||
|
}
|
||||||
|
format := u.AuthFormat
|
||||||
|
if format == "" {
|
||||||
|
if http.CanonicalHeaderKey(name) == "Authorization" {
|
||||||
|
format = "Bearer %s"
|
||||||
|
} else {
|
||||||
|
format = "%s"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return name, fmt.Sprintf(format, u.APIKey)
|
||||||
|
}
|
||||||
|
|
||||||
|
func toSet(items []string) map[string]struct{} {
|
||||||
|
if len(items) == 0 {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
m := make(map[string]struct{}, len(items))
|
||||||
|
for _, i := range items {
|
||||||
|
m[i] = struct{}{}
|
||||||
|
}
|
||||||
|
return m
|
||||||
|
}
|
||||||
@@ -0,0 +1,118 @@
|
|||||||
|
package aiproxy
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"io"
|
||||||
|
"sync"
|
||||||
|
|
||||||
|
"github.com/tidwall/gjson"
|
||||||
|
)
|
||||||
|
|
||||||
|
const maxCapture = 1 << 20
|
||||||
|
|
||||||
|
// watchBody passes the response through while it looks for OpenAI usage, and fires
|
||||||
|
// AfterProxy once the body ends or is closed.
|
||||||
|
type watchBody struct {
|
||||||
|
rc io.ReadCloser
|
||||||
|
capture bool // read usage
|
||||||
|
sse bool
|
||||||
|
status int
|
||||||
|
c *call
|
||||||
|
|
||||||
|
buf bytes.Buffer // JSON body (capped)
|
||||||
|
overflow bool
|
||||||
|
line []byte // partial SSE line
|
||||||
|
usage Usage
|
||||||
|
once sync.Once
|
||||||
|
}
|
||||||
|
|
||||||
|
func newWatchBody(rc io.ReadCloser, openai, sse bool, status int, c *call) io.ReadCloser {
|
||||||
|
return &watchBody{rc: rc, capture: openai, sse: sse, status: status, c: c}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *watchBody) Read(p []byte) (int, error) {
|
||||||
|
n, err := w.rc.Read(p)
|
||||||
|
if w.capture && n > 0 {
|
||||||
|
w.feed(p[:n])
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
w.end(err)
|
||||||
|
}
|
||||||
|
return n, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *watchBody) Close() error {
|
||||||
|
err := w.rc.Close()
|
||||||
|
w.end(nil)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *watchBody) end(err error) {
|
||||||
|
w.once.Do(func() {
|
||||||
|
if w.capture && !w.sse && !w.overflow {
|
||||||
|
if u, ok := parseUsage(w.buf.Bytes()); ok {
|
||||||
|
w.usage = u
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err == io.EOF {
|
||||||
|
err = nil
|
||||||
|
}
|
||||||
|
w.c.finish(w.status, w.usage, err)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *watchBody) feed(b []byte) {
|
||||||
|
if !w.sse {
|
||||||
|
if w.overflow || w.buf.Len()+len(b) > maxCapture {
|
||||||
|
w.overflow = true
|
||||||
|
return
|
||||||
|
}
|
||||||
|
w.buf.Write(b)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
w.line = append(w.line, b...)
|
||||||
|
for {
|
||||||
|
i := bytes.IndexByte(w.line, '\n')
|
||||||
|
if i < 0 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
line := bytes.TrimSpace(w.line[:i])
|
||||||
|
w.line = w.line[i+1:]
|
||||||
|
if rest, ok := bytes.CutPrefix(line, []byte("data:")); ok && bytes.Contains(rest, []byte("usage")) {
|
||||||
|
if u, ok := parseUsage(bytes.TrimSpace(rest)); ok {
|
||||||
|
w.usage = u
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(w.line) > maxCapture {
|
||||||
|
w.line = nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// parseUsage reads chat/embeddings ("usage") and responses ("response.usage") shapes.
|
||||||
|
func parseUsage(data []byte) (Usage, bool) {
|
||||||
|
u := gjson.GetBytes(data, "usage")
|
||||||
|
if !u.IsObject() {
|
||||||
|
u = gjson.GetBytes(data, "response.usage")
|
||||||
|
}
|
||||||
|
if !u.IsObject() {
|
||||||
|
return Usage{}, false
|
||||||
|
}
|
||||||
|
first := func(keys ...string) int64 {
|
||||||
|
for _, k := range keys {
|
||||||
|
if v := u.Get(k); v.Exists() {
|
||||||
|
return v.Int()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
out := Usage{
|
||||||
|
PromptTokens: first("prompt_tokens", "input_tokens"),
|
||||||
|
CompletionTokens: first("completion_tokens", "output_tokens"),
|
||||||
|
TotalTokens: first("total_tokens"),
|
||||||
|
}
|
||||||
|
if out.TotalTokens == 0 {
|
||||||
|
out.TotalTokens = out.PromptTokens + out.CompletionTokens
|
||||||
|
}
|
||||||
|
return out, true
|
||||||
|
}
|
||||||
@@ -47,6 +47,14 @@ type Provider interface {
|
|||||||
Handler() http.Handler
|
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.
|
// Resetter is optionally implemented by providers that can clear their recorded stats.
|
||||||
type Resetter interface {
|
type Resetter interface {
|
||||||
Reset()
|
Reset()
|
||||||
|
|||||||
@@ -32,6 +32,9 @@ type PrometheusProvider struct {
|
|||||||
eventDuration *prometheus.HistogramVec
|
eventDuration *prometheus.HistogramVec
|
||||||
eventQueueSize prometheus.Gauge
|
eventQueueSize prometheus.Gauge
|
||||||
panicsTotal *prometheus.CounterVec
|
panicsTotal *prometheus.CounterVec
|
||||||
|
aiRequests *prometheus.CounterVec
|
||||||
|
aiDuration *prometheus.HistogramVec
|
||||||
|
aiTokens *prometheus.CounterVec
|
||||||
|
|
||||||
pathLimiter *pathLimiter
|
pathLimiter *pathLimiter
|
||||||
pathNormalizer func(*http.Request) string
|
pathNormalizer func(*http.Request) string
|
||||||
@@ -162,6 +165,28 @@ func NewPrometheusProvider(cfg *Config) *PrometheusProvider {
|
|||||||
},
|
},
|
||||||
[]string{"method"},
|
[]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),
|
pathLimiter: newPathLimiter(cfg.HTTPMaxPaths),
|
||||||
pathNormalizer: cfg.HTTPPathNormalizer,
|
pathNormalizer: cfg.HTTPPathNormalizer,
|
||||||
@@ -310,6 +335,21 @@ func (p *PrometheusProvider) RecordPanic(methodName string) {
|
|||||||
p.panicsTotal.WithLabelValues(methodName).Inc()
|
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
|
// Handler implements Provider interface
|
||||||
// It responds 404 when metrics are disabled.
|
// It responds 404 when metrics are disabled.
|
||||||
func (p *PrometheusProvider) Handler() http.Handler {
|
func (p *PrometheusProvider) Handler() http.Handler {
|
||||||
@@ -437,6 +477,9 @@ func (p *PrometheusProvider) Reset() {
|
|||||||
p.eventProcessed.Reset()
|
p.eventProcessed.Reset()
|
||||||
p.eventDuration.Reset()
|
p.eventDuration.Reset()
|
||||||
p.panicsTotal.Reset()
|
p.panicsTotal.Reset()
|
||||||
|
p.aiRequests.Reset()
|
||||||
|
p.aiDuration.Reset()
|
||||||
|
p.aiTokens.Reset()
|
||||||
p.pathLimiter.reset()
|
p.pathLimiter.reset()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -389,6 +389,8 @@ func TestXFilesRecursivePreloadDepth(t *testing.T) {
|
|||||||
// TestXFilesResponseStructure validates the actual structure of the response
|
// TestXFilesResponseStructure validates the actual structure of the response
|
||||||
// This test can be expanded when we have a full database integration test environment
|
// This test can be expanded when we have a full database integration test environment
|
||||||
func TestXFilesResponseStructure(t *testing.T) {
|
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
|
// Load the expected correct response
|
||||||
correctResponsePath := filepath.Join("..", "..", "tests", "data", "xfiles.response.correct.json")
|
correctResponsePath := filepath.Join("..", "..", "tests", "data", "xfiles.response.correct.json")
|
||||||
correctData, err := os.ReadFile(correctResponsePath)
|
correctData, err := os.ReadFile(correctResponsePath)
|
||||||
|
|||||||
Reference in New Issue
Block a user