mirror of
https://github.com/bitechdev/ResolveSpec.git
synced 2026-10-08 14:26:28 +00:00
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.
316 lines
8.8 KiB
Go
316 lines
8.8 KiB
Go
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))
|
|
}
|