Files
ResolveSpec/pkg/aiproxy/handler.go
T

317 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/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))
}