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 | ||
|
|
e4c4315f4b | ||
|
|
3efd539e0f |
@@ -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
|
||||||
|
}
|
||||||
@@ -1,3 +1,4 @@
|
|||||||
|
//go:build integration
|
||||||
// +build integration
|
// +build integration
|
||||||
|
|
||||||
package database
|
package database
|
||||||
@@ -18,11 +19,11 @@ import (
|
|||||||
|
|
||||||
// Integration test models
|
// Integration test models
|
||||||
type IntegrationUser struct {
|
type IntegrationUser struct {
|
||||||
ID int `db:"id"`
|
ID int `db:"id"`
|
||||||
Name string `db:"name"`
|
Name string `db:"name"`
|
||||||
Email string `db:"email"`
|
Email string `db:"email"`
|
||||||
Age int `db:"age"`
|
Age int `db:"age"`
|
||||||
CreatedAt time.Time `db:"created_at"`
|
CreatedAt time.Time `db:"created_at"`
|
||||||
Posts []*IntegrationPost `bun:"rel:has-many,join:id=user_id"`
|
Posts []*IntegrationPost `bun:"rel:has-many,join:id=user_id"`
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -46,10 +47,10 @@ func (p IntegrationPost) TableName() string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type IntegrationComment struct {
|
type IntegrationComment struct {
|
||||||
ID int `db:"id"`
|
ID int `db:"id"`
|
||||||
Content string `db:"content"`
|
Content string `db:"content"`
|
||||||
PostID int `db:"post_id"`
|
PostID int `db:"post_id"`
|
||||||
CreatedAt time.Time `db:"created_at"`
|
CreatedAt time.Time `db:"created_at"`
|
||||||
Post *IntegrationPost `bun:"rel:belongs-to,join:post_id=id"`
|
Post *IntegrationPost `bun:"rel:belongs-to,join:post_id=id"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -26,11 +26,11 @@ func (u TestUser) TableName() string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type TestPost struct {
|
type TestPost struct {
|
||||||
ID int `db:"id"`
|
ID int `db:"id"`
|
||||||
Title string `db:"title"`
|
Title string `db:"title"`
|
||||||
Content string `db:"content"`
|
Content string `db:"content"`
|
||||||
UserID int `db:"user_id"`
|
UserID int `db:"user_id"`
|
||||||
User *TestUser `bun:"rel:belongs-to,join:user_id=id"`
|
User *TestUser `bun:"rel:belongs-to,join:user_id=id"`
|
||||||
Comments []TestComment `bun:"rel:has-many,join:id=post_id"`
|
Comments []TestComment `bun:"rel:has-many,join:id=post_id"`
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,77 @@
|
|||||||
|
package common
|
||||||
|
|
||||||
|
import (
|
||||||
|
"regexp"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/reflection"
|
||||||
|
)
|
||||||
|
|
||||||
|
var rePlainIdent = regexp.MustCompile(`^[A-Za-z_][A-Za-z0-9_]*$`)
|
||||||
|
|
||||||
|
// MainTableAlias returns the alias the main table gets in the SELECT:
|
||||||
|
// the model's TableAlias() when provided, otherwise the bare table name.
|
||||||
|
func MainTableAlias(model interface{}, tableName string) string {
|
||||||
|
if p, ok := model.(TableAliasProvider); ok {
|
||||||
|
if a := p.TableAlias(); a != "" {
|
||||||
|
return a
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return reflection.ExtractTableNameOnly(tableName)
|
||||||
|
}
|
||||||
|
|
||||||
|
func isModelSQLColumn(model interface{}, column string) bool {
|
||||||
|
for _, c := range reflection.GetSQLModelColumns(model) {
|
||||||
|
if strings.EqualFold(c, column) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
// QualifyModelColumn returns "alias"."column" when column is a plain identifier
|
||||||
|
// that exists on the model, so it stays unambiguous once joins are added.
|
||||||
|
// Anything else (expressions, JSON paths, other tables' columns) is returned unchanged.
|
||||||
|
func QualifyModelColumn(model interface{}, alias, column string) string {
|
||||||
|
if alias == "" || model == nil || !rePlainIdent.MatchString(column) || !isModelSQLColumn(model, column) {
|
||||||
|
return column
|
||||||
|
}
|
||||||
|
return QuoteIdent(alias) + "." + QuoteIdent(column)
|
||||||
|
}
|
||||||
|
|
||||||
|
// StripMainTablePrefix turns "<prefix>.<column>" into "<column>" when prefix is
|
||||||
|
// one of the main table's names/aliases and column exists on the model.
|
||||||
|
// Any other input is returned unchanged.
|
||||||
|
func StripMainTablePrefix(model interface{}, column string, prefixes ...string) string {
|
||||||
|
idx := strings.Index(column, ".")
|
||||||
|
if idx <= 0 || model == nil {
|
||||||
|
return column
|
||||||
|
}
|
||||||
|
prefix := strings.Trim(column[:idx], `"`)
|
||||||
|
col := strings.Trim(column[idx+1:], `"`)
|
||||||
|
if !rePlainIdent.MatchString(prefix) || !rePlainIdent.MatchString(col) {
|
||||||
|
return column
|
||||||
|
}
|
||||||
|
for _, p := range prefixes {
|
||||||
|
if p != "" && strings.EqualFold(p, prefix) && isModelSQLColumn(model, col) {
|
||||||
|
return col
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return column
|
||||||
|
}
|
||||||
|
|
||||||
|
// StripMainTablePrefixFromFilters applies StripMainTablePrefix to every filter column.
|
||||||
|
func StripMainTablePrefixFromFilters(model interface{}, filters []FilterOption, prefixes ...string) {
|
||||||
|
for i := range filters {
|
||||||
|
filters[i].Column = StripMainTablePrefix(model, filters[i].Column, prefixes...)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// NormalizeMainTableFilters rewrites "<main table or alias>.<column>" filter
|
||||||
|
// columns on opts to the bare model column (on a copy, never the caller's
|
||||||
|
// slice) so the column validator keeps them.
|
||||||
|
func NormalizeMainTableFilters(model interface{}, tableName string, opts *RequestOptions) {
|
||||||
|
opts.Filters = append([]FilterOption(nil), opts.Filters...)
|
||||||
|
StripMainTablePrefixFromFilters(model, opts.Filters,
|
||||||
|
MainTableAlias(model, tableName), reflection.ExtractTableNameOnly(tableName))
|
||||||
|
}
|
||||||
@@ -0,0 +1,55 @@
|
|||||||
|
package common
|
||||||
|
|
||||||
|
import "testing"
|
||||||
|
|
||||||
|
type qualifyModel struct {
|
||||||
|
ID int64 `bun:"id,pk"`
|
||||||
|
Name string `bun:"name"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestQualifyModelColumn(t *testing.T) {
|
||||||
|
m := qualifyModel{}
|
||||||
|
cases := []struct{ alias, col, want string }{
|
||||||
|
{"t", "name", `"t"."name"`},
|
||||||
|
{"t", "NAME", `"t"."NAME"`},
|
||||||
|
{"t", "missing", "missing"},
|
||||||
|
{"t", "data->>'x'", "data->>'x'"},
|
||||||
|
{"t", "rel.name", "rel.name"},
|
||||||
|
{"", "name", "name"},
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
if got := QualifyModelColumn(m, c.alias, c.col); got != c.want {
|
||||||
|
t.Errorf("Qualify(%q,%q) = %q, want %q", c.alias, c.col, got, c.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if got := QualifyModelColumn(nil, "t", "name"); got != "name" {
|
||||||
|
t.Errorf("nil model: %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStripMainTablePrefix(t *testing.T) {
|
||||||
|
m := qualifyModel{}
|
||||||
|
cases := []struct{ col, want string }{
|
||||||
|
{"province_state.name", "name"},
|
||||||
|
{`"province_state"."name"`, "name"},
|
||||||
|
{"PROVINCE_STATE.name", "name"},
|
||||||
|
{"rel_rid_country.name", "rel_rid_country.name"},
|
||||||
|
{"province_state.missing", "province_state.missing"},
|
||||||
|
{"name", "name"},
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
if got := StripMainTablePrefix(m, c.col, "province_state"); got != c.want {
|
||||||
|
t.Errorf("Strip(%q) = %q, want %q", c.col, got, c.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestStripMainTablePrefix_ValidatorKeepsFilter(t *testing.T) {
|
||||||
|
m := qualifyModel{}
|
||||||
|
filters := []FilterOption{{Column: "province_state.name", Operator: "eq", Value: "x"}}
|
||||||
|
StripMainTablePrefixFromFilters(m, filters, "province_state")
|
||||||
|
out := NewColumnValidator(m).FilterRequestOptions(RequestOptions{Filters: filters})
|
||||||
|
if len(out.Filters) != 1 || out.Filters[0].Column != "name" {
|
||||||
|
t.Fatalf("filters = %+v", out.Filters)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -25,10 +25,10 @@ func newMockDatabase() *mockDatabase {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (m *mockDatabase) NewSelect() SelectQuery { return &mockSelectQuery{} }
|
func (m *mockDatabase) NewSelect() SelectQuery { return &mockSelectQuery{} }
|
||||||
func (m *mockDatabase) NewInsert() InsertQuery { return &mockInsertQuery{db: m} }
|
func (m *mockDatabase) NewInsert() InsertQuery { return &mockInsertQuery{db: m} }
|
||||||
func (m *mockDatabase) NewUpdate() UpdateQuery { return &mockUpdateQuery{db: m} }
|
func (m *mockDatabase) NewUpdate() UpdateQuery { return &mockUpdateQuery{db: m} }
|
||||||
func (m *mockDatabase) NewDelete() DeleteQuery { return &mockDeleteQuery{db: m} }
|
func (m *mockDatabase) NewDelete() DeleteQuery { return &mockDeleteQuery{db: m} }
|
||||||
func (m *mockDatabase) RunInTransaction(ctx context.Context, fn func(Database) error) error {
|
func (m *mockDatabase) RunInTransaction(ctx context.Context, fn func(Database) error) error {
|
||||||
return fn(m)
|
return fn(m)
|
||||||
}
|
}
|
||||||
@@ -57,27 +57,31 @@ func (m *mockDatabase) DriverName() string {
|
|||||||
// Mock SelectQuery
|
// Mock SelectQuery
|
||||||
type mockSelectQuery struct{}
|
type mockSelectQuery struct{}
|
||||||
|
|
||||||
func (m *mockSelectQuery) Model(model interface{}) SelectQuery { return m }
|
func (m *mockSelectQuery) Model(model interface{}) SelectQuery { return m }
|
||||||
func (m *mockSelectQuery) Table(name string) SelectQuery { return m }
|
func (m *mockSelectQuery) Table(name string) SelectQuery { return m }
|
||||||
func (m *mockSelectQuery) Column(columns ...string) SelectQuery { return m }
|
func (m *mockSelectQuery) Column(columns ...string) SelectQuery { return m }
|
||||||
func (m *mockSelectQuery) ColumnExpr(query string, args ...interface{}) SelectQuery { return m }
|
func (m *mockSelectQuery) ColumnExpr(query string, args ...interface{}) SelectQuery { return m }
|
||||||
func (m *mockSelectQuery) Where(condition string, args ...interface{}) SelectQuery { return m }
|
func (m *mockSelectQuery) Where(condition string, args ...interface{}) SelectQuery { return m }
|
||||||
func (m *mockSelectQuery) WhereOr(query string, args ...interface{}) SelectQuery { return m }
|
func (m *mockSelectQuery) WhereOr(query string, args ...interface{}) SelectQuery { return m }
|
||||||
func (m *mockSelectQuery) Join(query string, args ...interface{}) SelectQuery { return m }
|
func (m *mockSelectQuery) Join(query string, args ...interface{}) SelectQuery { return m }
|
||||||
func (m *mockSelectQuery) LeftJoin(query string, args ...interface{}) SelectQuery { return m }
|
func (m *mockSelectQuery) LeftJoin(query string, args ...interface{}) SelectQuery { return m }
|
||||||
func (m *mockSelectQuery) Preload(relation string, conditions ...interface{}) SelectQuery { return m }
|
func (m *mockSelectQuery) Preload(relation string, conditions ...interface{}) SelectQuery { return m }
|
||||||
func (m *mockSelectQuery) PreloadRelation(relation string, apply ...func(SelectQuery) SelectQuery) SelectQuery { return m }
|
func (m *mockSelectQuery) PreloadRelation(relation string, apply ...func(SelectQuery) SelectQuery) SelectQuery {
|
||||||
func (m *mockSelectQuery) JoinRelation(relation string, apply ...func(SelectQuery) SelectQuery) SelectQuery { return m }
|
return m
|
||||||
func (m *mockSelectQuery) Order(order string) SelectQuery { return m }
|
}
|
||||||
func (m *mockSelectQuery) OrderExpr(order string, args ...interface{}) SelectQuery { return m }
|
func (m *mockSelectQuery) JoinRelation(relation string, apply ...func(SelectQuery) SelectQuery) SelectQuery {
|
||||||
func (m *mockSelectQuery) Limit(n int) SelectQuery { return m }
|
return m
|
||||||
func (m *mockSelectQuery) Offset(n int) SelectQuery { return m }
|
}
|
||||||
func (m *mockSelectQuery) Group(group string) SelectQuery { return m }
|
func (m *mockSelectQuery) Order(order string) SelectQuery { return m }
|
||||||
|
func (m *mockSelectQuery) OrderExpr(order string, args ...interface{}) SelectQuery { return m }
|
||||||
|
func (m *mockSelectQuery) Limit(n int) SelectQuery { return m }
|
||||||
|
func (m *mockSelectQuery) Offset(n int) SelectQuery { return m }
|
||||||
|
func (m *mockSelectQuery) Group(group string) SelectQuery { return m }
|
||||||
func (m *mockSelectQuery) Having(condition string, args ...interface{}) SelectQuery { return m }
|
func (m *mockSelectQuery) Having(condition string, args ...interface{}) SelectQuery { return m }
|
||||||
func (m *mockSelectQuery) Scan(ctx context.Context, dest interface{}) error { return nil }
|
func (m *mockSelectQuery) Scan(ctx context.Context, dest interface{}) error { return nil }
|
||||||
func (m *mockSelectQuery) ScanModel(ctx context.Context) error { return nil }
|
func (m *mockSelectQuery) ScanModel(ctx context.Context) error { return nil }
|
||||||
func (m *mockSelectQuery) Count(ctx context.Context) (int, error) { return 0, nil }
|
func (m *mockSelectQuery) Count(ctx context.Context) (int, error) { return 0, nil }
|
||||||
func (m *mockSelectQuery) Exists(ctx context.Context) (bool, error) { return false, nil }
|
func (m *mockSelectQuery) Exists(ctx context.Context) (bool, error) { return false, nil }
|
||||||
|
|
||||||
// Mock InsertQuery
|
// Mock InsertQuery
|
||||||
type mockInsertQuery struct {
|
type mockInsertQuery struct {
|
||||||
@@ -98,9 +102,9 @@ func (m *mockInsertQuery) Value(column string, value interface{}) InsertQuery {
|
|||||||
m.values[column] = value
|
m.values[column] = value
|
||||||
return m
|
return m
|
||||||
}
|
}
|
||||||
func (m *mockInsertQuery) OnConflict(action string) InsertQuery { return m }
|
func (m *mockInsertQuery) OnConflict(action string) InsertQuery { return m }
|
||||||
func (m *mockInsertQuery) ExcludeColumn(columns ...string) InsertQuery { return m }
|
func (m *mockInsertQuery) ExcludeColumn(columns ...string) InsertQuery { return m }
|
||||||
func (m *mockInsertQuery) Returning(columns ...string) InsertQuery { return m }
|
func (m *mockInsertQuery) Returning(columns ...string) InsertQuery { return m }
|
||||||
func (m *mockInsertQuery) Exec(ctx context.Context) (Result, error) {
|
func (m *mockInsertQuery) Exec(ctx context.Context) (Result, error) {
|
||||||
m.db.insertCalls = append(m.db.insertCalls, m.values)
|
m.db.insertCalls = append(m.db.insertCalls, m.values)
|
||||||
m.db.lastID++
|
m.db.lastID++
|
||||||
@@ -132,8 +136,8 @@ func (m *mockUpdateQuery) SetMap(values map[string]interface{}) UpdateQuery {
|
|||||||
return m
|
return m
|
||||||
}
|
}
|
||||||
func (m *mockUpdateQuery) Where(condition string, args ...interface{}) UpdateQuery { return m }
|
func (m *mockUpdateQuery) Where(condition string, args ...interface{}) UpdateQuery { return m }
|
||||||
func (m *mockUpdateQuery) ExcludeColumn(columns ...string) UpdateQuery { return m }
|
func (m *mockUpdateQuery) ExcludeColumn(columns ...string) UpdateQuery { return m }
|
||||||
func (m *mockUpdateQuery) Returning(columns ...string) UpdateQuery { return m }
|
func (m *mockUpdateQuery) Returning(columns ...string) UpdateQuery { return m }
|
||||||
func (m *mockUpdateQuery) Exec(ctx context.Context) (Result, error) {
|
func (m *mockUpdateQuery) Exec(ctx context.Context) (Result, error) {
|
||||||
// Record the update call
|
// Record the update call
|
||||||
m.db.updateCalls = append(m.db.updateCalls, m.setValues)
|
m.db.updateCalls = append(m.db.updateCalls, m.setValues)
|
||||||
@@ -171,9 +175,13 @@ func (m *mockResult) RowsAffected() int64 { return m.rowsAffected }
|
|||||||
type mockModelRegistry struct{}
|
type mockModelRegistry struct{}
|
||||||
|
|
||||||
func (m *mockModelRegistry) GetModel(name string) (interface{}, error) { return nil, nil }
|
func (m *mockModelRegistry) GetModel(name string) (interface{}, error) { return nil, nil }
|
||||||
func (m *mockModelRegistry) GetModelByEntity(schema, entity string) (interface{}, error) { return nil, nil }
|
func (m *mockModelRegistry) GetModelByEntity(schema, entity string) (interface{}, error) {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
func (m *mockModelRegistry) RegisterModel(name string, model interface{}) error { return nil }
|
func (m *mockModelRegistry) RegisterModel(name string, model interface{}) error { return nil }
|
||||||
func (m *mockModelRegistry) GetAllModels() map[string]interface{} { return make(map[string]interface{}) }
|
func (m *mockModelRegistry) GetAllModels() map[string]interface{} {
|
||||||
|
return make(map[string]interface{})
|
||||||
|
}
|
||||||
|
|
||||||
// Mock RelationshipInfoProvider
|
// Mock RelationshipInfoProvider
|
||||||
type mockRelationshipProvider struct {
|
type mockRelationshipProvider struct {
|
||||||
@@ -198,9 +206,9 @@ func (m *mockRelationshipProvider) RegisterRelation(modelTypeName, relationName
|
|||||||
|
|
||||||
// Test Models
|
// Test Models
|
||||||
type Department struct {
|
type Department struct {
|
||||||
ID int64 `json:"id" bun:"id,pk"`
|
ID int64 `json:"id" bun:"id,pk"`
|
||||||
Name string `json:"name"`
|
Name string `json:"name"`
|
||||||
Employees []*Employee `json:"employees,omitempty"`
|
Employees []*Employee `json:"employees,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func (d Department) TableName() string { return "departments" }
|
func (d Department) TableName() string { return "departments" }
|
||||||
@@ -227,9 +235,9 @@ func (t Task) TableName() string { return "tasks" }
|
|||||||
func (t Task) GetIDName() string { return "ID" }
|
func (t Task) GetIDName() string { return "ID" }
|
||||||
|
|
||||||
type Comment struct {
|
type Comment struct {
|
||||||
ID int64 `json:"id" bun:"id,pk"`
|
ID int64 `json:"id" bun:"id,pk"`
|
||||||
Text string `json:"text"`
|
Text string `json:"text"`
|
||||||
TaskID int64 `json:"task_id"`
|
TaskID int64 `json:"task_id"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c Comment) TableName() string { return "comments" }
|
func (c Comment) TableName() string { return "comments" }
|
||||||
|
|||||||
@@ -0,0 +1,92 @@
|
|||||||
|
package funcspec
|
||||||
|
|
||||||
|
import (
|
||||||
|
"database/sql"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/uptrace/bun/driver/sqliteshim"
|
||||||
|
)
|
||||||
|
|
||||||
|
const joinedBaseSQL = "SELECT p.id, p.name, c.name AS country_name FROM province_state p LEFT JOIN country c ON c.id = p.rid_country"
|
||||||
|
|
||||||
|
// funcspec filters are appended to author-written SQL with no model, so columns
|
||||||
|
// are used verbatim: clients disambiguate by sending "alias.column", and the
|
||||||
|
// dot must survive ValidSQL and every filter path.
|
||||||
|
func TestApplyFilters_QualifiedColumnsPreservedWithJoin(t *testing.T) {
|
||||||
|
h := NewHandler(&MockDatabase{})
|
||||||
|
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
params *RequestParameters
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"field filter", &RequestParameters{FieldFilters: map[string]string{"p.name": "abc"}}, "p.name = abc"},
|
||||||
|
{"search filter", &RequestParameters{SearchFilters: map[string]string{"p.name": "abc"}}, "CAST(p.name AS TEXT) ILIKE '%abc%'"},
|
||||||
|
{"search op eq", &RequestParameters{SearchOps: map[string]FilterOperator{"p.name": {Operator: "eq", Value: "abc", Logic: "AND"}}}, "p.name = 'abc'"},
|
||||||
|
{"search op contains", &RequestParameters{SearchOps: map[string]FilterOperator{"p.name": {Operator: "contains", Value: "abc", Logic: "AND"}}}, "CAST(p.name AS TEXT) ILIKE '%abc%'"},
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
t.Run(c.name, func(t *testing.T) {
|
||||||
|
got := h.ApplyFilters(joinedBaseSQL, c.params)
|
||||||
|
if !strings.Contains(got, c.want) {
|
||||||
|
t.Fatalf("SQL %q does not contain %q", got, c.want)
|
||||||
|
}
|
||||||
|
if !strings.Contains(got, "LEFT JOIN country") || !strings.Contains(got, " WHERE ") {
|
||||||
|
t.Fatalf("join/where lost: %q", got)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestApplyFilters_WithAndWithoutJoin_RealQuery(t *testing.T) {
|
||||||
|
sqldb, err := sql.Open(sqliteshim.ShimName, "file:funcspecjoin?mode=memory&cache=private")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer sqldb.Close()
|
||||||
|
for _, stmt := range []string{
|
||||||
|
"CREATE TABLE country (id INTEGER PRIMARY KEY, name TEXT)",
|
||||||
|
"CREATE TABLE province_state (id INTEGER PRIMARY KEY, name TEXT, rid_country INTEGER)",
|
||||||
|
"INSERT INTO country VALUES (1,'Abcland'),(2,'Other')",
|
||||||
|
"INSERT INTO province_state VALUES (1,'abc one',1),(2,'xyz two',1),(3,'nope',2)",
|
||||||
|
} {
|
||||||
|
if _, err := sqldb.Exec(stmt); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewHandler(&MockDatabase{})
|
||||||
|
count := func(base string, params *RequestParameters) (int, error) {
|
||||||
|
rows, err := sqldb.Query(h.ApplyFilters(base, params))
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
n := 0
|
||||||
|
for rows.Next() {
|
||||||
|
n++
|
||||||
|
}
|
||||||
|
return n, rows.Err()
|
||||||
|
}
|
||||||
|
|
||||||
|
const noJoin = "SELECT p.id, p.name FROM province_state p"
|
||||||
|
params := func(col string) *RequestParameters {
|
||||||
|
return &RequestParameters{SearchOps: map[string]FilterOperator{col: {Operator: "eq", Value: "nope", Logic: "AND"}}}
|
||||||
|
}
|
||||||
|
|
||||||
|
if n, err := count(noJoin, params("name")); err != nil || n != 1 {
|
||||||
|
t.Fatalf("no join, unqualified: n=%d err=%v", n, err)
|
||||||
|
}
|
||||||
|
if n, err := count(joinedBaseSQL, params("p.name")); err != nil || n != 1 {
|
||||||
|
t.Fatalf("join, qualified: n=%d err=%v", n, err)
|
||||||
|
}
|
||||||
|
if n, err := count(joinedBaseSQL, params("c.name")); err != nil || n != 0 {
|
||||||
|
t.Fatalf("join, joined-table column: n=%d err=%v", n, err)
|
||||||
|
}
|
||||||
|
// Documented limit: with a join in the author's SQL, an unqualified shared
|
||||||
|
// column is ambiguous and the client must send "alias.column".
|
||||||
|
if _, err := count(joinedBaseSQL, params("name")); err == nil || !strings.Contains(strings.ToLower(err.Error()), "ambiguous") {
|
||||||
|
t.Fatalf("expected ambiguous error for unqualified shared column, got %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,183 @@
|
|||||||
|
package resolvespec
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/uptrace/bun"
|
||||||
|
"github.com/uptrace/bun/dialect/sqlitedialect"
|
||||||
|
"github.com/uptrace/bun/driver/sqliteshim"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
|
||||||
|
)
|
||||||
|
|
||||||
|
type joinCountry struct {
|
||||||
|
bun.BaseModel `bun:"table:country,alias:country"`
|
||||||
|
ID int64 `bun:"id,pk"`
|
||||||
|
Name string `bun:"name"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type joinProvince struct {
|
||||||
|
bun.BaseModel `bun:"table:province_state,alias:province_state"`
|
||||||
|
ID int64 `bun:"id,pk"`
|
||||||
|
Name string `bun:"name"`
|
||||||
|
Abbreviation string `bun:"abbreviation"`
|
||||||
|
RidCountry int64 `bun:"rid_country"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// SQLite has no ILIKE; the ilike operator is covered by SQL-string tests, and
|
||||||
|
// LIKE here proves the CAST(... AS TEXT) wrapping executes with qualified columns.
|
||||||
|
func setupJoinDB(t *testing.T) *bun.DB {
|
||||||
|
t.Helper()
|
||||||
|
sqldb, err := sql.Open(sqliteshim.ShimName, "file:filterjoin?mode=memory&cache=private")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
db := bun.NewDB(sqldb, sqlitedialect.New())
|
||||||
|
t.Cleanup(func() { _ = db.Close() })
|
||||||
|
ctx := context.Background()
|
||||||
|
for _, m := range []interface{}{(*joinCountry)(nil), (*joinProvince)(nil)} {
|
||||||
|
if _, err := db.NewCreateTable().Model(m).IfNotExists().Exec(ctx); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if _, err := db.NewInsert().Model(&[]joinCountry{{ID: 1, Name: "Abcland"}, {ID: 2, Name: "Other"}}).Exec(ctx); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := db.NewInsert().Model(&[]joinProvince{
|
||||||
|
{ID: 1, Name: "abc one", Abbreviation: "A1", RidCountry: 1},
|
||||||
|
{ID: 2, Name: "xyz two", Abbreviation: "ABC", RidCountry: 1},
|
||||||
|
{ID: 3, Name: "nope", Abbreviation: "N3", RidCountry: 2},
|
||||||
|
}).Exec(ctx); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return db
|
||||||
|
}
|
||||||
|
|
||||||
|
const joinSQL = "LEFT JOIN country AS rel_rid_country ON rel_rid_country.id = province_state.rid_country"
|
||||||
|
|
||||||
|
func countWith(t *testing.T, db *bun.DB, withJoin bool, filters []common.FilterOption, alias string) (int, error) {
|
||||||
|
t.Helper()
|
||||||
|
h := &Handler{}
|
||||||
|
var q common.SelectQuery = database.NewBunAdapter(db).NewSelect().Model(&[]*joinProvince{})
|
||||||
|
if withJoin {
|
||||||
|
q = q.Join(joinSQL)
|
||||||
|
}
|
||||||
|
q = h.applyFilters(q, filters, &joinProvince{}, alias)
|
||||||
|
return q.Count(context.Background())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFilters_AmbiguousColumnWithJoin_RealQuery(t *testing.T) {
|
||||||
|
db := setupJoinDB(t)
|
||||||
|
|
||||||
|
filters := []common.FilterOption{
|
||||||
|
{Column: "name", Operator: "like", Value: "%abc%"},
|
||||||
|
{Column: "abbreviation", Operator: "like", Value: "%abc%", LogicOperator: "OR"},
|
||||||
|
}
|
||||||
|
|
||||||
|
t.Run("unqualified with join is ambiguous (the bug)", func(t *testing.T) {
|
||||||
|
_, err := countWith(t, db, true, filters, "")
|
||||||
|
if err == nil || !strings.Contains(strings.ToLower(err.Error()), "ambiguous") {
|
||||||
|
t.Fatalf("expected ambiguous column error, got %v", err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("qualified with join", func(t *testing.T) {
|
||||||
|
n, err := countWith(t, db, true, filters, "province_state")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if n != 2 {
|
||||||
|
t.Fatalf("count = %d, want 2", n)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("qualified without join", func(t *testing.T) {
|
||||||
|
n, err := countWith(t, db, false, filters, "province_state")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if n != 2 {
|
||||||
|
t.Fatalf("count = %d, want 2", n)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("unqualified without join still works", func(t *testing.T) {
|
||||||
|
n, err := countWith(t, db, false, filters, "")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if n != 2 {
|
||||||
|
t.Fatalf("count = %d, want 2", n)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestFilters_AllOperatorsWithJoin_RealQuery(t *testing.T) {
|
||||||
|
db := setupJoinDB(t)
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
filter common.FilterOption
|
||||||
|
want int
|
||||||
|
}{
|
||||||
|
{"eq", common.FilterOption{Column: "name", Operator: "eq", Value: "nope"}, 1},
|
||||||
|
{"neq", common.FilterOption{Column: "name", Operator: "neq", Value: "nope"}, 2},
|
||||||
|
{"like", common.FilterOption{Column: "name", Operator: "like", Value: "abc%"}, 1},
|
||||||
|
{"in", common.FilterOption{Column: "name", Operator: "in", Value: []string{"nope", "xyz two"}}, 2},
|
||||||
|
{"gt", common.FilterOption{Column: "id", Operator: "gt", Value: 1}, 2},
|
||||||
|
{"gte", common.FilterOption{Column: "id", Operator: "gte", Value: 2}, 2},
|
||||||
|
{"lt", common.FilterOption{Column: "id", Operator: "lt", Value: 3}, 2},
|
||||||
|
{"lte", common.FilterOption{Column: "id", Operator: "lte", Value: 1}, 1},
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
t.Run(c.name, func(t *testing.T) {
|
||||||
|
n, err := countWith(t, db, true, []common.FilterOption{c.filter}, "province_state")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if n != c.want {
|
||||||
|
t.Fatalf("count = %d, want %d", n, c.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// A filter on the joined table's own column, sent already qualified, must pass through untouched.
|
||||||
|
func TestFilters_JoinedTableColumnPassesThrough_RealQuery(t *testing.T) {
|
||||||
|
db := setupJoinDB(t)
|
||||||
|
n, err := countWith(t, db, true, []common.FilterOption{
|
||||||
|
{Column: "rel_rid_country.name", Operator: "eq", Value: "Abcland"},
|
||||||
|
}, "province_state")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if n != 2 {
|
||||||
|
t.Fatalf("count = %d, want 2", n)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Client sends "province_state.name": it must survive validation and be applied.
|
||||||
|
func TestFilters_ClientQualifiedColumn_NotDropped(t *testing.T) {
|
||||||
|
db := setupJoinDB(t)
|
||||||
|
model := &joinProvince{}
|
||||||
|
|
||||||
|
opts := common.RequestOptions{Filters: []common.FilterOption{
|
||||||
|
{Column: "province_state.name", Operator: "like", Value: "%abc%"},
|
||||||
|
}}
|
||||||
|
common.NormalizeMainTableFilters(model, "public.province_state", &opts)
|
||||||
|
opts = common.NewColumnValidator(model).FilterRequestOptions(opts)
|
||||||
|
if len(opts.Filters) != 1 {
|
||||||
|
t.Fatalf("filter was dropped: %+v", opts.Filters)
|
||||||
|
}
|
||||||
|
|
||||||
|
n, err := countWith(t, db, true, opts.Filters, "province_state")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if n != 1 {
|
||||||
|
t.Fatalf("count = %d, want 1 (filter must narrow the result)", n)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,65 @@
|
|||||||
|
package resolvespec
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestBuildFilterConditionAlias_QualifiesModelColumns(t *testing.T) {
|
||||||
|
h := &Handler{}
|
||||||
|
model := jsonColModel{}
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
filter common.FilterOption
|
||||||
|
want string
|
||||||
|
}{
|
||||||
|
{"eq", common.FilterOption{Column: "name", Operator: "eq", Value: "x"}, `"province_state"."name" = ?`},
|
||||||
|
{"ilike", common.FilterOption{Column: "name", Operator: "ilike", Value: "%a%"}, `CAST("province_state"."name" AS TEXT) ILIKE ?`},
|
||||||
|
{"like", common.FilterOption{Column: "name", Operator: "like", Value: "%a%"}, `CAST("province_state"."name" AS TEXT) LIKE ?`},
|
||||||
|
{"in", common.FilterOption{Column: "name", Operator: "in", Value: []string{"a", "b"}}, `"province_state"."name" IN (?,?)`},
|
||||||
|
{"non-model column untouched", common.FilterOption{Column: "other", Operator: "eq", Value: 1}, `other = ?`},
|
||||||
|
{"already qualified untouched", common.FilterOption{Column: "rel.name", Operator: "eq", Value: 1}, `rel.name = ?`},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
got, _ := h.buildFilterConditionAlias(tt.filter, model, "province_state")
|
||||||
|
if got != tt.want {
|
||||||
|
t.Fatalf("got %q, want %q", got, tt.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildFilterConditionAlias_NoAliasUnchanged(t *testing.T) {
|
||||||
|
h := &Handler{}
|
||||||
|
got, _ := h.buildFilterConditionAlias(common.FilterOption{Column: "name", Operator: "ilike", Value: "%a%"}, jsonColModel{}, "")
|
||||||
|
if got != "CAST(name AS TEXT) ILIKE ?" {
|
||||||
|
t.Fatalf("got %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestApplyFilter_QualifiesWithAlias(t *testing.T) {
|
||||||
|
h := &Handler{}
|
||||||
|
q := &jsonCapQuery{}
|
||||||
|
h.applyFilter(q, common.FilterOption{Column: "name", Operator: "ilike", Value: "%a%"}, jsonColModel{}, "province_state")
|
||||||
|
c := q.only(t)
|
||||||
|
if c.query != `CAST("province_state"."name" AS TEXT) ILIKE ?` {
|
||||||
|
t.Fatalf("query = %q", c.query)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestApplyFilters_OrGroupQualified(t *testing.T) {
|
||||||
|
h := &Handler{}
|
||||||
|
q := &jsonCapQuery{}
|
||||||
|
h.applyFilters(q, []common.FilterOption{
|
||||||
|
{Column: "name", Operator: "ilike", Value: "%a%"},
|
||||||
|
{Column: "id", Operator: "eq", Value: 1, LogicOperator: "OR"},
|
||||||
|
}, jsonColModel{}, "t")
|
||||||
|
c := q.only(t)
|
||||||
|
want := `(CAST("t"."name" AS TEXT) ILIKE ? OR "t"."id" = ?)`
|
||||||
|
if c.query != want {
|
||||||
|
t.Fatalf("query = %q, want %q", c.query, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
+30
-15
@@ -172,6 +172,9 @@ func (h *Handler) Handle(w common.ResponseWriter, r common.Request, params map[s
|
|||||||
// Add request-scoped data to context
|
// Add request-scoped data to context
|
||||||
ctx = WithRequestData(ctx, schema, entity, tableName, model, modelPtr)
|
ctx = WithRequestData(ctx, schema, entity, tableName, model, modelPtr)
|
||||||
|
|
||||||
|
// Accept "<main table or alias>.<column>" for model columns before validation drops them
|
||||||
|
common.NormalizeMainTableFilters(model, tableName, &req.Options)
|
||||||
|
|
||||||
// Validate and filter columns in options (log warnings for invalid columns)
|
// Validate and filter columns in options (log warnings for invalid columns)
|
||||||
validator := common.NewColumnValidator(model)
|
validator := common.NewColumnValidator(model)
|
||||||
req.Options = validator.FilterRequestOptions(req.Options)
|
req.Options = validator.FilterRequestOptions(req.Options)
|
||||||
@@ -411,7 +414,7 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Apply filters with proper grouping for OR logic
|
// Apply filters with proper grouping for OR logic
|
||||||
query = h.applyFilters(query, options.Filters, model)
|
query = h.applyFilters(query, options.Filters, model, common.MainTableAlias(model, tableName))
|
||||||
|
|
||||||
// Apply custom operators
|
// Apply custom operators
|
||||||
for _, customOp := range options.CustomOperators {
|
for _, customOp := range options.CustomOperators {
|
||||||
@@ -558,7 +561,7 @@ func (h *Handler) handleRead(ctx context.Context, w common.ResponseWriter, id st
|
|||||||
|
|
||||||
// Apply the same filters as the main query
|
// Apply the same filters as the main query
|
||||||
for _, filter := range options.Filters {
|
for _, filter := range options.Filters {
|
||||||
rowNumQuery = h.applyFilter(rowNumQuery, filter, model)
|
rowNumQuery = h.applyFilter(rowNumQuery, filter, model, common.MainTableAlias(model, tableName))
|
||||||
}
|
}
|
||||||
|
|
||||||
// Apply custom operators
|
// Apply custom operators
|
||||||
@@ -1932,7 +1935,7 @@ func (h *Handler) executeDelete(ctx context.Context, tx common.Database, hookCtx
|
|||||||
// applyFilters applies all filters with proper grouping for OR logic
|
// applyFilters applies all filters with proper grouping for OR logic
|
||||||
// Groups consecutive OR filters together to ensure proper query precedence
|
// Groups consecutive OR filters together to ensure proper query precedence
|
||||||
// Example: [A, B(OR), C(OR), D(AND)] => WHERE (A OR B OR C) AND D
|
// Example: [A, B(OR), C(OR), D(AND)] => WHERE (A OR B OR C) AND D
|
||||||
func (h *Handler) applyFilters(query common.SelectQuery, filters []common.FilterOption, model interface{}) common.SelectQuery {
|
func (h *Handler) applyFilters(query common.SelectQuery, filters []common.FilterOption, model interface{}, alias string) common.SelectQuery {
|
||||||
if len(filters) == 0 {
|
if len(filters) == 0 {
|
||||||
return query
|
return query
|
||||||
}
|
}
|
||||||
@@ -1952,11 +1955,11 @@ func (h *Handler) applyFilters(query common.SelectQuery, filters []common.Filter
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Apply the OR group as a single grouped WHERE clause
|
// Apply the OR group as a single grouped WHERE clause
|
||||||
query = h.applyFilterGroup(query, orGroup, model)
|
query = h.applyFilterGroup(query, orGroup, model, alias)
|
||||||
i = j
|
i = j
|
||||||
} else {
|
} else {
|
||||||
// Single filter with AND logic (or first filter)
|
// Single filter with AND logic (or first filter)
|
||||||
condition, args := h.buildFilterCondition(filters[i], model)
|
condition, args := h.buildFilterConditionAlias(filters[i], model, alias)
|
||||||
if condition != "" {
|
if condition != "" {
|
||||||
query = query.Where(condition, args...)
|
query = query.Where(condition, args...)
|
||||||
}
|
}
|
||||||
@@ -1969,7 +1972,7 @@ func (h *Handler) applyFilters(query common.SelectQuery, filters []common.Filter
|
|||||||
|
|
||||||
// applyFilterGroup applies a group of filters that should be OR'd together
|
// applyFilterGroup applies a group of filters that should be OR'd together
|
||||||
// Always wraps them in parentheses and applies as a single WHERE clause
|
// Always wraps them in parentheses and applies as a single WHERE clause
|
||||||
func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.FilterOption, model interface{}) common.SelectQuery {
|
func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.FilterOption, model interface{}, alias string) common.SelectQuery {
|
||||||
if len(filters) == 0 {
|
if len(filters) == 0 {
|
||||||
return query
|
return query
|
||||||
}
|
}
|
||||||
@@ -1979,7 +1982,7 @@ func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.Fi
|
|||||||
var args []interface{}
|
var args []interface{}
|
||||||
|
|
||||||
for _, filter := range filters {
|
for _, filter := range filters {
|
||||||
condition, filterArgs := h.buildFilterCondition(filter, model)
|
condition, filterArgs := h.buildFilterConditionAlias(filter, model, alias)
|
||||||
if condition != "" {
|
if condition != "" {
|
||||||
conditions = append(conditions, condition)
|
conditions = append(conditions, condition)
|
||||||
args = append(args, filterArgs...)
|
args = append(args, filterArgs...)
|
||||||
@@ -2005,6 +2008,12 @@ func (h *Handler) applyFilterGroup(query common.SelectQuery, filters []common.Fi
|
|||||||
// or the dotted data.x shorthand for a JSON column) resolve to a safe,
|
// or the dotted data.x shorthand for a JSON column) resolve to a safe,
|
||||||
// parameterised expression before the ordinary operator handling below.
|
// parameterised expression before the ordinary operator handling below.
|
||||||
func (h *Handler) buildFilterCondition(filter common.FilterOption, model interface{}) (conditionString string, conditionArgs []interface{}) {
|
func (h *Handler) buildFilterCondition(filter common.FilterOption, model interface{}) (conditionString string, conditionArgs []interface{}) {
|
||||||
|
return h.buildFilterConditionAlias(filter, model, "")
|
||||||
|
}
|
||||||
|
|
||||||
|
// buildFilterConditionAlias is buildFilterCondition with plain model columns
|
||||||
|
// qualified by the main table alias, so joins from preloads can't make them ambiguous.
|
||||||
|
func (h *Handler) buildFilterConditionAlias(filter common.FilterOption, model interface{}, alias string) (conditionString string, conditionArgs []interface{}) {
|
||||||
var condition string
|
var condition string
|
||||||
var args []interface{}
|
var args []interface{}
|
||||||
|
|
||||||
@@ -2012,6 +2021,9 @@ func (h *Handler) buildFilterCondition(filter common.FilterOption, model interfa
|
|||||||
return cond, jargs
|
return cond, jargs
|
||||||
}
|
}
|
||||||
|
|
||||||
|
rawColumn := filter.Column
|
||||||
|
filter.Column = common.QualifyModelColumn(model, alias, filter.Column)
|
||||||
|
|
||||||
switch filter.Operator {
|
switch filter.Operator {
|
||||||
case "eq", "=":
|
case "eq", "=":
|
||||||
condition = fmt.Sprintf("%s = ?", filter.Column)
|
condition = fmt.Sprintf("%s = ?", filter.Column)
|
||||||
@@ -2032,10 +2044,10 @@ func (h *Handler) buildFilterCondition(filter common.FilterOption, model interfa
|
|||||||
condition = fmt.Sprintf("%s <= ?", filter.Column)
|
condition = fmt.Sprintf("%s <= ?", filter.Column)
|
||||||
args = []interface{}{filter.Value}
|
args = []interface{}{filter.Value}
|
||||||
case "like":
|
case "like":
|
||||||
condition = fmt.Sprintf("%s LIKE ?", likeColumn(filter.Column, model))
|
condition = fmt.Sprintf("%s LIKE ?", likeColumn(filter.Column, rawColumn, model))
|
||||||
args = []interface{}{filter.Value}
|
args = []interface{}{filter.Value}
|
||||||
case "ilike":
|
case "ilike":
|
||||||
condition = fmt.Sprintf("%s ILIKE ?", likeColumn(filter.Column, model))
|
condition = fmt.Sprintf("%s ILIKE ?", likeColumn(filter.Column, rawColumn, model))
|
||||||
args = []interface{}{filter.Value}
|
args = []interface{}{filter.Value}
|
||||||
case "in":
|
case "in":
|
||||||
condition, args = common.BuildInCondition(filter.Column, filter.Value)
|
condition, args = common.BuildInCondition(filter.Column, filter.Value)
|
||||||
@@ -2073,14 +2085,14 @@ func (h *Handler) buildFilterCondition(filter common.FilterOption, model interfa
|
|||||||
// CAST(... AS TEXT) would switch to case-sensitive matching and defeat a
|
// CAST(... AS TEXT) would switch to case-sensitive matching and defeat a
|
||||||
// citext index. Every other column is cast to TEXT so LIKE/ILIKE also works
|
// citext index. Every other column is cast to TEXT so LIKE/ILIKE also works
|
||||||
// against date/time/timestamp and numeric columns.
|
// against date/time/timestamp and numeric columns.
|
||||||
func likeColumn(column string, model interface{}) string {
|
func likeColumn(column, rawColumn string, model interface{}) string {
|
||||||
if reflection.IsCitextColumn(model, column) {
|
if reflection.IsCitextColumn(model, rawColumn) {
|
||||||
return column
|
return column
|
||||||
}
|
}
|
||||||
return fmt.Sprintf("CAST(%s AS TEXT)", column)
|
return fmt.Sprintf("CAST(%s AS TEXT)", column)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOption, model interface{}) common.SelectQuery {
|
func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOption, model interface{}, alias string) common.SelectQuery {
|
||||||
// Determine which method to use based on LogicOperator
|
// Determine which method to use based on LogicOperator
|
||||||
useOrLogic := strings.EqualFold(filter.LogicOperator, "OR")
|
useOrLogic := strings.EqualFold(filter.LogicOperator, "OR")
|
||||||
|
|
||||||
@@ -2094,6 +2106,9 @@ func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOpti
|
|||||||
return query.Where(cond, jargs...)
|
return query.Where(cond, jargs...)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
rawColumn := filter.Column
|
||||||
|
filter.Column = common.QualifyModelColumn(model, alias, filter.Column)
|
||||||
|
|
||||||
switch filter.Operator {
|
switch filter.Operator {
|
||||||
case "eq", "=":
|
case "eq", "=":
|
||||||
condition = fmt.Sprintf("%s = ?", filter.Column)
|
condition = fmt.Sprintf("%s = ?", filter.Column)
|
||||||
@@ -2114,10 +2129,10 @@ func (h *Handler) applyFilter(query common.SelectQuery, filter common.FilterOpti
|
|||||||
condition = fmt.Sprintf("%s <= ?", filter.Column)
|
condition = fmt.Sprintf("%s <= ?", filter.Column)
|
||||||
args = []interface{}{filter.Value}
|
args = []interface{}{filter.Value}
|
||||||
case "like":
|
case "like":
|
||||||
condition = fmt.Sprintf("%s LIKE ?", likeColumn(filter.Column, model))
|
condition = fmt.Sprintf("%s LIKE ?", likeColumn(filter.Column, rawColumn, model))
|
||||||
args = []interface{}{filter.Value}
|
args = []interface{}{filter.Value}
|
||||||
case "ilike":
|
case "ilike":
|
||||||
condition = fmt.Sprintf("%s ILIKE ?", likeColumn(filter.Column, model))
|
condition = fmt.Sprintf("%s ILIKE ?", likeColumn(filter.Column, rawColumn, model))
|
||||||
args = []interface{}{filter.Value}
|
args = []interface{}{filter.Value}
|
||||||
case "in":
|
case "in":
|
||||||
condition, args = common.BuildInCondition(filter.Column, filter.Value)
|
condition, args = common.BuildInCondition(filter.Column, filter.Value)
|
||||||
@@ -2523,7 +2538,7 @@ func (h *Handler) applyPreloads(model interface{}, query common.SelectQuery, pre
|
|||||||
|
|
||||||
if len(preload.Filters) > 0 {
|
if len(preload.Filters) > 0 {
|
||||||
for _, filter := range preload.Filters {
|
for _, filter := range preload.Filters {
|
||||||
sq = h.applyFilter(sq, filter, nil)
|
sq = h.applyFilter(sq, filter, nil, "")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if len(preload.Sort) > 0 {
|
if len(preload.Sort) > 0 {
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
//go:build integration
|
||||||
// +build integration
|
// +build integration
|
||||||
|
|
||||||
package resolvespec
|
package resolvespec
|
||||||
@@ -22,12 +23,12 @@ import (
|
|||||||
|
|
||||||
// Test models
|
// Test models
|
||||||
type TestUser struct {
|
type TestUser struct {
|
||||||
ID uint `gorm:"primaryKey" json:"id"`
|
ID uint `gorm:"primaryKey" json:"id"`
|
||||||
Name string `gorm:"not null" json:"name"`
|
Name string `gorm:"not null" json:"name"`
|
||||||
Email string `gorm:"uniqueIndex;not null" json:"email"`
|
Email string `gorm:"uniqueIndex;not null" json:"email"`
|
||||||
Age int `json:"age"`
|
Age int `json:"age"`
|
||||||
Active bool `gorm:"default:true" json:"active"`
|
Active bool `gorm:"default:true" json:"active"`
|
||||||
CreatedAt time.Time `json:"created_at"`
|
CreatedAt time.Time `json:"created_at"`
|
||||||
Posts []TestPost `gorm:"foreignKey:UserID" json:"posts,omitempty"`
|
Posts []TestPost `gorm:"foreignKey:UserID" json:"posts,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -36,13 +37,13 @@ func (TestUser) TableName() string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type TestPost struct {
|
type TestPost struct {
|
||||||
ID uint `gorm:"primaryKey" json:"id"`
|
ID uint `gorm:"primaryKey" json:"id"`
|
||||||
UserID uint `gorm:"not null" json:"user_id"`
|
UserID uint `gorm:"not null" json:"user_id"`
|
||||||
Title string `gorm:"not null" json:"title"`
|
Title string `gorm:"not null" json:"title"`
|
||||||
Content string `json:"content"`
|
Content string `json:"content"`
|
||||||
Published bool `gorm:"default:false" json:"published"`
|
Published bool `gorm:"default:false" json:"published"`
|
||||||
CreatedAt time.Time `json:"created_at"`
|
CreatedAt time.Time `json:"created_at"`
|
||||||
User *TestUser `gorm:"foreignKey:UserID" json:"user,omitempty"`
|
User *TestUser `gorm:"foreignKey:UserID" json:"user,omitempty"`
|
||||||
Comments []TestComment `gorm:"foreignKey:PostID" json:"comments,omitempty"`
|
Comments []TestComment `gorm:"foreignKey:PostID" json:"comments,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -55,7 +56,7 @@ type TestComment struct {
|
|||||||
PostID uint `gorm:"not null" json:"post_id"`
|
PostID uint `gorm:"not null" json:"post_id"`
|
||||||
Content string `gorm:"not null" json:"content"`
|
Content string `gorm:"not null" json:"content"`
|
||||||
CreatedAt time.Time `json:"created_at"`
|
CreatedAt time.Time `json:"created_at"`
|
||||||
Post *TestPost `gorm:"foreignKey:PostID" json:"post,omitempty"`
|
Post *TestPost `gorm:"foreignKey:PostID" json:"post,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func (TestComment) TableName() string {
|
func (TestComment) TableName() string {
|
||||||
|
|||||||
@@ -142,7 +142,7 @@ func TestApplyFilter_JSONColumn(t *testing.T) {
|
|||||||
q := &jsonCapQuery{}
|
q := &jsonCapQuery{}
|
||||||
h.applyFilter(q, common.FilterOption{
|
h.applyFilter(q, common.FilterOption{
|
||||||
Column: "data->>'tier'", Operator: "in", Value: []string{"a", "b"}, LogicOperator: "OR",
|
Column: "data->>'tier'", Operator: "in", Value: []string{"a", "b"}, LogicOperator: "OR",
|
||||||
}, model)
|
}, model, "")
|
||||||
c := q.only(t)
|
c := q.only(t)
|
||||||
if c.method != "WhereOr" || c.query != `("data" #>> ?::text[]) IN (?,?)` {
|
if c.method != "WhereOr" || c.query != `("data" #>> ?::text[]) IN (?,?)` {
|
||||||
t.Fatalf("call = %+v", c)
|
t.Fatalf("call = %+v", c)
|
||||||
|
|||||||
@@ -0,0 +1,180 @@
|
|||||||
|
package restheadspec
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"database/sql"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/uptrace/bun"
|
||||||
|
"github.com/uptrace/bun/dialect/sqlitedialect"
|
||||||
|
"github.com/uptrace/bun/driver/sqliteshim"
|
||||||
|
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common"
|
||||||
|
"github.com/bitechdev/ResolveSpec/pkg/common/adapters/database"
|
||||||
|
)
|
||||||
|
|
||||||
|
type joinCountry struct {
|
||||||
|
bun.BaseModel `bun:"table:country,alias:country"`
|
||||||
|
ID int64 `bun:"id,pk"`
|
||||||
|
Name string `bun:"name"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type joinProvince struct {
|
||||||
|
bun.BaseModel `bun:"table:province_state,alias:province_state"`
|
||||||
|
ID int64 `bun:"id,pk"`
|
||||||
|
Name string `bun:"name"`
|
||||||
|
Abbreviation string `bun:"abbreviation"`
|
||||||
|
RidCountry int64 `bun:"rid_country"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func setupJoinDB(t *testing.T) *bun.DB {
|
||||||
|
t.Helper()
|
||||||
|
sqldb, err := sql.Open(sqliteshim.ShimName, "file:rhsfilterjoin?mode=memory&cache=private")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
db := bun.NewDB(sqldb, sqlitedialect.New())
|
||||||
|
t.Cleanup(func() { _ = db.Close() })
|
||||||
|
ctx := context.Background()
|
||||||
|
for _, m := range []interface{}{(*joinCountry)(nil), (*joinProvince)(nil)} {
|
||||||
|
if _, err := db.NewCreateTable().Model(m).IfNotExists().Exec(ctx); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if _, err := db.NewInsert().Model(&[]joinCountry{{ID: 1, Name: "Abcland"}, {ID: 2, Name: "Other"}}).Exec(ctx); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := db.NewInsert().Model(&[]joinProvince{
|
||||||
|
{ID: 1, Name: "abc one", Abbreviation: "A1", RidCountry: 1},
|
||||||
|
{ID: 2, Name: "xyz two", Abbreviation: "ABC", RidCountry: 1},
|
||||||
|
{ID: 3, Name: "nope", Abbreviation: "N3", RidCountry: 2},
|
||||||
|
}).Exec(ctx); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return db
|
||||||
|
}
|
||||||
|
|
||||||
|
const joinSQL = "LEFT JOIN country AS rel_rid_country ON rel_rid_country.id = province_state.rid_country"
|
||||||
|
|
||||||
|
// countFilters applies filters the way handleRead does (single AND filters via
|
||||||
|
// applyFilter, consecutive OR filters via applyOrFilterGroup) and counts.
|
||||||
|
func countFilters(t *testing.T, db *bun.DB, withJoin bool, filters []common.FilterOption) (int, error) {
|
||||||
|
t.Helper()
|
||||||
|
h := &Handler{}
|
||||||
|
model := &joinProvince{}
|
||||||
|
var q common.SelectQuery = database.NewBunAdapter(db).NewSelect().Model(&[]*joinProvince{})
|
||||||
|
if withJoin {
|
||||||
|
q = q.Join(joinSQL)
|
||||||
|
}
|
||||||
|
for i := 0; i < len(filters); {
|
||||||
|
f := filters[i]
|
||||||
|
castInfo := h.ValidateAndAdjustFilterForColumnType(&f, model)
|
||||||
|
if f.LogicOperator == "OR" {
|
||||||
|
group := []*common.FilterOption{&f}
|
||||||
|
info := []ColumnCastInfo{castInfo}
|
||||||
|
j := i + 1
|
||||||
|
for j < len(filters) && filters[j].LogicOperator == "OR" {
|
||||||
|
g := filters[j]
|
||||||
|
info = append(info, h.ValidateAndAdjustFilterForColumnType(&g, model))
|
||||||
|
group = append(group, &g)
|
||||||
|
j++
|
||||||
|
}
|
||||||
|
q = h.applyOrFilterGroup(q, group, info, "public.province_state", model)
|
||||||
|
i = j
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
q = h.applyFilter(q, f, "public.province_state", castInfo.NeedsCast, "AND", model)
|
||||||
|
i++
|
||||||
|
}
|
||||||
|
return q.Count(context.Background())
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRHSFilters_WithAndWithoutJoin_RealQuery(t *testing.T) {
|
||||||
|
db := setupJoinDB(t)
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
filters []common.FilterOption
|
||||||
|
want int
|
||||||
|
}{
|
||||||
|
{"eq", []common.FilterOption{{Column: "name", Operator: "eq", Value: "nope"}}, 1},
|
||||||
|
{"neq", []common.FilterOption{{Column: "name", Operator: "neq", Value: "nope"}}, 2},
|
||||||
|
{"like", []common.FilterOption{{Column: "name", Operator: "like", Value: "%abc%"}}, 1},
|
||||||
|
{"in", []common.FilterOption{{Column: "name", Operator: "in", Value: []string{"nope", "xyz two"}}}, 2},
|
||||||
|
{"gt", []common.FilterOption{{Column: "id", Operator: "gt", Value: 1}}, 2},
|
||||||
|
{"between", []common.FilterOption{{Column: "id", Operator: "between", Value: []interface{}{0, 3}}}, 2},
|
||||||
|
{"between_inclusive", []common.FilterOption{{Column: "id", Operator: "between_inclusive", Value: []interface{}{1, 3}}}, 3},
|
||||||
|
{"is_not_null", []common.FilterOption{{Column: "name", Operator: "is_not_null"}}, 3},
|
||||||
|
{"or group", []common.FilterOption{
|
||||||
|
{Column: "name", Operator: "like", Value: "%abc%", LogicOperator: "OR"},
|
||||||
|
{Column: "abbreviation", Operator: "like", Value: "%abc%", LogicOperator: "OR"},
|
||||||
|
}, 2},
|
||||||
|
{"or group then and", []common.FilterOption{
|
||||||
|
{Column: "name", Operator: "like", Value: "%abc%", LogicOperator: "OR"},
|
||||||
|
{Column: "abbreviation", Operator: "like", Value: "%abc%", LogicOperator: "OR"},
|
||||||
|
{Column: "id", Operator: "eq", Value: 2},
|
||||||
|
}, 1},
|
||||||
|
}
|
||||||
|
for _, c := range cases {
|
||||||
|
for _, withJoin := range []bool{false, true} {
|
||||||
|
name := c.name + "/no_join"
|
||||||
|
if withJoin {
|
||||||
|
name = c.name + "/join"
|
||||||
|
}
|
||||||
|
t.Run(name, func(t *testing.T) {
|
||||||
|
n, err := countFilters(t, db, withJoin, c.filters)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("query failed (ambiguous column?): %v", err)
|
||||||
|
}
|
||||||
|
if n != c.want {
|
||||||
|
t.Fatalf("count = %d, want %d", n, c.want)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRHSFilters_JoinedTableColumnPassesThrough(t *testing.T) {
|
||||||
|
db := setupJoinDB(t)
|
||||||
|
n, err := countFilters(t, db, true, []common.FilterOption{
|
||||||
|
{Column: "rel_rid_country.name", Operator: "eq", Value: "Abcland"},
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if n != 2 {
|
||||||
|
t.Fatalf("count = %d, want 2", n)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Client sends "province_state.name": it must survive validation and be applied.
|
||||||
|
func TestRHSFilters_ClientQualifiedColumn_NotDropped(t *testing.T) {
|
||||||
|
db := setupJoinDB(t)
|
||||||
|
h := &Handler{}
|
||||||
|
model := &joinProvince{}
|
||||||
|
|
||||||
|
opts := ExtendedRequestOptions{}
|
||||||
|
opts.Filters = []common.FilterOption{{Column: "province_state.name", Operator: "like", Value: "%abc%"}}
|
||||||
|
common.NormalizeMainTableFilters(model, "public.province_state", &opts.RequestOptions)
|
||||||
|
opts = h.filterExtendedOptions(common.NewColumnValidator(model), opts, model)
|
||||||
|
if len(opts.Filters) != 1 || opts.Filters[0].Column != "name" {
|
||||||
|
t.Fatalf("filter was dropped or not normalised: %+v", opts.Filters)
|
||||||
|
}
|
||||||
|
|
||||||
|
n, err := countFilters(t, db, true, opts.Filters)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if n != 1 {
|
||||||
|
t.Fatalf("count = %d, want 1", n)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRHSFilters_ILikeSQLQualified(t *testing.T) {
|
||||||
|
h := &Handler{}
|
||||||
|
q := &jsonCapQuery{}
|
||||||
|
f := common.FilterOption{Column: "name", Operator: "ilike", Value: "%abc%"}
|
||||||
|
h.applyFilter(q, f, "info.province_state", false, "AND", jsonColModel{})
|
||||||
|
if got := q.only(t).query; got != "CAST(province_state.name AS TEXT) ILIKE ?" {
|
||||||
|
t.Fatalf("query = %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -164,6 +164,9 @@ func (h *Handler) Handle(w common.ResponseWriter, r common.Request, params map[s
|
|||||||
// Parse options from headers - this now includes relation name resolution
|
// Parse options from headers - this now includes relation name resolution
|
||||||
options := h.parseOptionsFromHeaders(r, model)
|
options := h.parseOptionsFromHeaders(r, model)
|
||||||
|
|
||||||
|
// Accept "<main table or alias>.<column>" for model columns before validation drops them
|
||||||
|
common.NormalizeMainTableFilters(model, tableName, &options.RequestOptions)
|
||||||
|
|
||||||
// Validate and filter columns in options (log warnings for invalid columns)
|
// Validate and filter columns in options (log warnings for invalid columns)
|
||||||
validator := common.NewColumnValidator(model)
|
validator := common.NewColumnValidator(model)
|
||||||
options = h.filterExtendedOptions(validator, options, model)
|
options = h.filterExtendedOptions(validator, options, model)
|
||||||
|
|||||||
@@ -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