mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-08 22:36:28 +00:00
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.
This commit is contained in:
@@ -0,0 +1,315 @@
|
||||
package aiproxy
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"mime"
|
||||
"net/http"
|
||||
"net/http/httputil"
|
||||
"path"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/bitechdev/ResolveSpec/pkg/logger"
|
||||
"github.com/bitechdev/ResolveSpec/pkg/security"
|
||||
"github.com/tidwall/gjson"
|
||||
)
|
||||
|
||||
type ctxKey struct{}
|
||||
|
||||
// call is the per-request state shared between the handler and the proxy callbacks.
|
||||
type call struct {
|
||||
p *Proxy
|
||||
t *target
|
||||
hc *HookContext
|
||||
start time.Time
|
||||
once sync.Once
|
||||
}
|
||||
|
||||
func (c *call) finish(status int, usage Usage, err error) {
|
||||
c.once.Do(func() {
|
||||
c.hc.StatusCode = status
|
||||
c.hc.Usage = usage
|
||||
c.hc.Error = err
|
||||
c.hc.Duration = time.Since(c.start)
|
||||
c.p.hooks.executeAfter(c.hc)
|
||||
outcome := OutcomeOK
|
||||
if err != nil || status >= 400 {
|
||||
outcome = OutcomeUpstreamError
|
||||
}
|
||||
c.p.record(c.t, c.hc, status, outcome, "", err)
|
||||
})
|
||||
}
|
||||
|
||||
var defaultStrip = []string{"Authorization", "Cookie", "Proxy-Authorization", "X-Api-Key"}
|
||||
|
||||
// Handler returns the proxy handler. auth is the authentication middleware, normally
|
||||
// security.NewAuthMiddleware(securityList). It is required: without it every request is refused.
|
||||
func (p *Proxy) Handler(auth func(http.Handler) http.Handler) http.Handler {
|
||||
if auth == nil {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
writeError(w, http.StatusInternalServerError, "server_error", "authentication is not configured")
|
||||
})
|
||||
}
|
||||
return auth(http.HandlerFunc(p.serve))
|
||||
}
|
||||
|
||||
func (p *Proxy) serve(w http.ResponseWriter, r *http.Request) {
|
||||
// route: {prefix}/{name}/{rest}
|
||||
rel := strings.TrimPrefix(r.URL.Path, p.cfg.Prefix)
|
||||
if rel == r.URL.Path && p.cfg.Prefix != "" {
|
||||
writeError(w, http.StatusNotFound, "not_found", "not found")
|
||||
return
|
||||
}
|
||||
rel = strings.TrimPrefix(rel, "/")
|
||||
name, rest, _ := strings.Cut(rel, "/")
|
||||
t := p.get(name)
|
||||
if t == nil {
|
||||
writeError(w, http.StatusNotFound, "not_found", "unknown upstream")
|
||||
return
|
||||
}
|
||||
rest = "/" + rest
|
||||
if rest != "/" {
|
||||
rest = path.Clean(rest)
|
||||
}
|
||||
if t.kind == KindOpenAI {
|
||||
switch {
|
||||
case rest == "/v1":
|
||||
rest = ""
|
||||
case strings.HasPrefix(rest, "/v1/"):
|
||||
rest = strings.TrimPrefix(rest, "/v1")
|
||||
default:
|
||||
writeError(w, http.StatusNotFound, "not_found", "unknown path")
|
||||
return
|
||||
}
|
||||
} else if rest == "/" {
|
||||
rest = ""
|
||||
}
|
||||
|
||||
start := time.Now()
|
||||
user, _ := security.GetUserContext(r.Context())
|
||||
hc := &HookContext{
|
||||
Context: r.Context(),
|
||||
Request: r,
|
||||
UserContext: user,
|
||||
Upstream: t.up.Name,
|
||||
Kind: t.kind,
|
||||
}
|
||||
deny := func(code int, outcome Outcome, typ, msg string) {
|
||||
hc.Duration = time.Since(start)
|
||||
p.record(t, hc, code, outcome, msg, nil)
|
||||
writeError(w, code, typ, msg)
|
||||
}
|
||||
|
||||
// authenticated user, no guests
|
||||
if user == nil || hasRole(user.Roles, "guest") {
|
||||
deny(http.StatusUnauthorized, OutcomeDenied, "invalid_api_key", "authentication required")
|
||||
return
|
||||
}
|
||||
if len(t.roles) > 0 && !anyRole(user.Roles, t.roles) {
|
||||
deny(http.StatusForbidden, OutcomeDenied, "forbidden", "not allowed to use this upstream")
|
||||
return
|
||||
}
|
||||
if rl := t.up.RateLimit; rl != nil {
|
||||
key := t.up.Name + "|" + strconv.Itoa(user.UserID) + "|" + user.UserName
|
||||
if allowed, wait := p.limits.allow(key, rl); !allowed {
|
||||
w.Header().Set("Retry-After", strconv.Itoa(int(wait.Seconds())+1))
|
||||
deny(http.StatusTooManyRequests, OutcomeRateLimited, "rate_limit_exceeded", "rate limit exceeded")
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if code, typ, msg := p.inspect(t, r, hc); code != 0 {
|
||||
deny(code, OutcomeDenied, typ, msg)
|
||||
return
|
||||
}
|
||||
|
||||
if err := p.hooks.Execute(BeforeProxy, hc); err != nil {
|
||||
code := hc.AbortCode
|
||||
if code == 0 {
|
||||
code = http.StatusForbidden
|
||||
}
|
||||
msg := hc.AbortMessage
|
||||
if msg == "" {
|
||||
msg = "request rejected"
|
||||
}
|
||||
if !hc.Abort {
|
||||
logger.Error("aiproxy: %v", err)
|
||||
code, msg = http.StatusInternalServerError, "request failed"
|
||||
}
|
||||
deny(code, OutcomeDenied, "rejected", msg)
|
||||
return
|
||||
}
|
||||
|
||||
c := &call{p: p, t: t, hc: hc, start: start}
|
||||
ctx := context.WithValue(r.Context(), ctxKey{}, &routed{call: c, rest: rest})
|
||||
t.proxy.ServeHTTP(w, hc.Request.WithContext(ctx))
|
||||
}
|
||||
|
||||
// routed is stored on the request context for the ReverseProxy callbacks.
|
||||
type routed struct {
|
||||
call *call
|
||||
rest string
|
||||
}
|
||||
|
||||
func (p *Proxy) buildProxy(t *target) *httputil.ReverseProxy {
|
||||
strip := append(append([]string{}, defaultStrip...), p.cfg.StripHeaders...)
|
||||
if t.authName != "" {
|
||||
strip = append(strip, t.authName)
|
||||
}
|
||||
return &httputil.ReverseProxy{
|
||||
Transport: newTransport(t.up.Timeout),
|
||||
FlushInterval: -1,
|
||||
Rewrite: func(pr *httputil.ProxyRequest) {
|
||||
rt, _ := pr.In.Context().Value(ctxKey{}).(*routed)
|
||||
out := pr.Out
|
||||
u := *t.base
|
||||
u.Path = t.base.Path
|
||||
if rt != nil {
|
||||
u.Path += rt.rest
|
||||
}
|
||||
u.RawPath = ""
|
||||
u.RawQuery = pr.In.URL.RawQuery
|
||||
out.URL = &u
|
||||
out.Host = u.Host
|
||||
for _, h := range strip {
|
||||
out.Header.Del(h)
|
||||
}
|
||||
out.Header.Del("Accept-Encoding") // let the transport decompress so usage can be read
|
||||
for k, v := range t.up.Headers {
|
||||
out.Header.Set(k, v)
|
||||
}
|
||||
if t.authName != "" {
|
||||
out.Header.Set(t.authName, t.authVal)
|
||||
}
|
||||
},
|
||||
ModifyResponse: func(resp *http.Response) error {
|
||||
rt, _ := resp.Request.Context().Value(ctxKey{}).(*routed)
|
||||
h := resp.Header
|
||||
h.Del("Set-Cookie")
|
||||
h.Del("Www-Authenticate")
|
||||
if resp.StatusCode >= 300 && resp.StatusCode < 400 {
|
||||
h.Del("Location")
|
||||
}
|
||||
if resp.StatusCode == http.StatusUnauthorized {
|
||||
logger.Warn("aiproxy: upstream %q rejected its credentials", t.up.Name)
|
||||
resp.Body.Close()
|
||||
body, _ := json.Marshal(errorBody("bad_gateway", "upstream rejected the proxy credentials"))
|
||||
resp.StatusCode, resp.Status = http.StatusBadGateway, "502 Bad Gateway"
|
||||
resp.Body = io.NopCloser(bytes.NewReader(body))
|
||||
resp.ContentLength = int64(len(body))
|
||||
h.Set("Content-Type", "application/json")
|
||||
h.Set("Content-Length", strconv.Itoa(len(body)))
|
||||
}
|
||||
if rt != nil {
|
||||
sse := strings.HasPrefix(h.Get("Content-Type"), "text/event-stream")
|
||||
resp.Body = newWatchBody(resp.Body, t.kind == KindOpenAI, sse, resp.StatusCode, rt.call)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
ErrorHandler: func(w http.ResponseWriter, r *http.Request, err error) {
|
||||
if rt, _ := r.Context().Value(ctxKey{}).(*routed); rt != nil {
|
||||
rt.call.finish(http.StatusBadGateway, Usage{}, err)
|
||||
}
|
||||
if !errors.Is(err, context.Canceled) {
|
||||
logger.Warn("aiproxy: upstream %q error: %v", t.up.Name, err)
|
||||
}
|
||||
writeError(w, http.StatusBadGateway, "bad_gateway", "upstream request failed")
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// inspect reads JSON request bodies to extract/enforce model and tools.
|
||||
// It returns a non-zero status when the request must be rejected.
|
||||
func (p *Proxy) inspect(t *target, r *http.Request, hc *HookContext) (int, string, string) {
|
||||
switch r.Method {
|
||||
case http.MethodPost, http.MethodPut, http.MethodPatch:
|
||||
default:
|
||||
return 0, "", ""
|
||||
}
|
||||
if r.Body == nil || r.Body == http.NoBody {
|
||||
return 0, "", ""
|
||||
}
|
||||
mt, _, _ := mime.ParseMediaType(r.Header.Get("Content-Type"))
|
||||
isJSON := mt == "application/json" || strings.HasSuffix(mt, "+json")
|
||||
|
||||
if !isJSON {
|
||||
if t.kind == KindOpenAI && len(t.models) > 0 {
|
||||
return http.StatusBadRequest, "invalid_request_error", "model cannot be verified for this content type"
|
||||
}
|
||||
return 0, "", ""
|
||||
}
|
||||
|
||||
body, err := io.ReadAll(io.LimitReader(r.Body, p.cfg.MaxBodyBytes+1))
|
||||
r.Body.Close()
|
||||
if err != nil {
|
||||
return http.StatusBadRequest, "invalid_request_error", "could not read request body"
|
||||
}
|
||||
if int64(len(body)) > p.cfg.MaxBodyBytes {
|
||||
return http.StatusRequestEntityTooLarge, "invalid_request_error", "request body too large"
|
||||
}
|
||||
r.Body = io.NopCloser(bytes.NewReader(body))
|
||||
r.ContentLength = int64(len(body))
|
||||
r.Header.Set("Content-Length", strconv.Itoa(len(body)))
|
||||
|
||||
switch t.kind {
|
||||
case KindOpenAI:
|
||||
hc.Model = gjson.GetBytes(body, "model").String()
|
||||
if len(t.models) > 0 {
|
||||
if _, ok := t.models[hc.Model]; !ok {
|
||||
return http.StatusForbidden, "model_not_allowed", "model not allowed"
|
||||
}
|
||||
}
|
||||
case KindMCP:
|
||||
msgs := []gjson.Result{gjson.ParseBytes(body)}
|
||||
if msgs[0].IsArray() {
|
||||
msgs = msgs[0].Array()
|
||||
}
|
||||
for _, m := range msgs {
|
||||
if m.Get("method").String() != "tools/call" {
|
||||
continue
|
||||
}
|
||||
tool := m.Get("params.name").String()
|
||||
if len(t.tools) > 0 {
|
||||
if _, ok := t.tools[tool]; !ok {
|
||||
return http.StatusForbidden, "tool_not_allowed", "tool not allowed"
|
||||
}
|
||||
}
|
||||
hc.Tools = append(hc.Tools, tool)
|
||||
}
|
||||
}
|
||||
return 0, "", ""
|
||||
}
|
||||
|
||||
func hasRole(roles []string, want string) bool {
|
||||
for _, r := range roles {
|
||||
if r == want {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func anyRole(roles []string, allowed map[string]struct{}) bool {
|
||||
for _, r := range roles {
|
||||
if _, ok := allowed[r]; ok {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func errorBody(typ, msg string) map[string]any {
|
||||
return map[string]any{"error": map[string]any{"message": msg, "type": typ, "code": typ}}
|
||||
}
|
||||
|
||||
func writeError(w http.ResponseWriter, status int, typ, msg string) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
w.WriteHeader(status)
|
||||
_ = json.NewEncoder(w).Encode(errorBody(typ, msg))
|
||||
}
|
||||
Reference in New Issue
Block a user