Files
ResolveSpec/pkg/aiproxy/store_test.go
T
warkanum 63cb0d22eb feat(aiproxy): add authenticated proxy for OpenAI-compatible APIs and MCP servers
Hides upstream URLs and keys behind the security layer. Per-upstream roles,
model/tool allowlists, per-user rate limits, Before/AfterProxy hooks, audit
sink, Prometheus metrics (metrics.AIProxyRecorder) and upstreams loaded from
the resolvespec_ai_proxies stored procedure.
2026-10-07 23:39:40 +02:00

270 lines
9.4 KiB
Go

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", ""))
}