241 lines
6.5 KiB
Go
241 lines
6.5 KiB
Go
// Package proxy is the routing reverse proxy. It takes /{route}/v1/…, picks a host from the
|
|
// route's ordered list using the health table, forwards the request, streams the answer back as it
|
|
// arrives, and tells the health table when a host fails.
|
|
package proxy
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"encoding/json"
|
|
"errors"
|
|
"io"
|
|
"net/http"
|
|
"net/http/httputil"
|
|
"net/url"
|
|
"strings"
|
|
"time"
|
|
|
|
"log/slog"
|
|
|
|
"git.wntrmute.dev/kyle/crossbar/internal/config"
|
|
"git.wntrmute.dev/kyle/crossbar/internal/health"
|
|
)
|
|
|
|
// MaxBody is the largest request body we look at for a top-level "model" field.
|
|
const MaxBody = 16 << 20
|
|
|
|
// HostHeader is set on every proxied response: the name of the host that answered.
|
|
const HostHeader = "X-Crossbar-Host"
|
|
|
|
// errBodyTooLarge is returned when a request body exceeds MaxBody during the model peek.
|
|
var errBodyTooLarge = errors.New("body too large")
|
|
|
|
// Health is what the proxy needs from the health table.
|
|
type Health interface {
|
|
Get(name string) (health.Status, bool)
|
|
MarkDown(name, reason string)
|
|
}
|
|
|
|
// Handler forwards requests for a route to one of the route's healthy hosts.
|
|
type Handler struct {
|
|
cfg *config.Config
|
|
health Health
|
|
log *slog.Logger
|
|
}
|
|
|
|
// New builds a Handler. A nil logger becomes slog.Default().
|
|
func New(cfg *config.Config, h Health, log *slog.Logger) *Handler {
|
|
if log == nil {
|
|
log = slog.Default()
|
|
}
|
|
return &Handler{cfg: cfg, health: h, log: log}
|
|
}
|
|
|
|
// SplitRoute takes the first path segment as the route. "/a/v1/x" -> ("a", "/v1/x", true); "/a" and
|
|
// "/a/" -> ("a", "/", true); "/", "//x", "", "noslash/v1" -> ("", "", false). The query string, if
|
|
// present, is kept in rest.
|
|
func SplitRoute(path string) (route, rest string, ok bool) {
|
|
if path == "" || path[0] != '/' {
|
|
return "", "", false
|
|
}
|
|
after := path[1:]
|
|
slash := strings.IndexByte(after, '/')
|
|
if slash == -1 {
|
|
if after == "" {
|
|
return "", "", false
|
|
}
|
|
return after, "/", true
|
|
}
|
|
if after[:slash] == "" {
|
|
return "", "", false
|
|
}
|
|
return after[:slash], after[slash:], true
|
|
}
|
|
|
|
// Choose returns the first host in order that is healthy and lists model in Loaded; failing that,
|
|
// the first healthy host. ok is false if none. model may be "".
|
|
func Choose(hosts []string, model string, h Health) (string, bool) {
|
|
for _, name := range hosts {
|
|
s, ok := h.Get(name)
|
|
if !ok || !s.Healthy || !hasModel(s.Loaded, model) {
|
|
continue
|
|
}
|
|
return name, true
|
|
}
|
|
for _, name := range hosts {
|
|
s, ok := h.Get(name)
|
|
if ok && s.Healthy {
|
|
return name, true
|
|
}
|
|
}
|
|
return "", false
|
|
}
|
|
|
|
func hasModel(loaded []string, model string) bool {
|
|
if model == "" {
|
|
return false
|
|
}
|
|
for _, m := range loaded {
|
|
if m == model {
|
|
return true
|
|
}
|
|
}
|
|
return false
|
|
}
|
|
|
|
// allowedPath reports whether rest may be proxied: under /v1/, or the two admin paths.
|
|
func allowedPath(rest string) bool {
|
|
return strings.HasPrefix(rest, "/v1/") || rest == "/health" || rest == "/props"
|
|
}
|
|
|
|
// peekModel reads a non-GET/HEAD body up to MaxBody+1 bytes, restores it on the request, and
|
|
// returns the top-level "model". A non-JSON body or one without a model gives "". A body larger
|
|
// than MaxBody returns errBodyTooLarge.
|
|
func peekModel(r *http.Request) (string, error) {
|
|
if r.Method == http.MethodGet || r.Method == http.MethodHead {
|
|
return "", nil
|
|
}
|
|
if r.Body == nil || r.Body == http.NoBody {
|
|
return "", nil
|
|
}
|
|
body, err := io.ReadAll(io.LimitReader(r.Body, MaxBody+1))
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
if len(body) > MaxBody {
|
|
return "", errBodyTooLarge
|
|
}
|
|
r.Body = io.NopCloser(bytes.NewReader(body))
|
|
r.ContentLength = int64(len(body))
|
|
|
|
var req struct {
|
|
Model string `json:"model"`
|
|
}
|
|
_ = json.Unmarshal(body, &req)
|
|
return req.Model, nil
|
|
}
|
|
|
|
// statusRecorder records the status written and forwards Flush so the reverse proxy can stream.
|
|
type statusRecorder struct {
|
|
http.ResponseWriter
|
|
status int
|
|
}
|
|
|
|
func (r *statusRecorder) WriteHeader(code int) {
|
|
r.status = code
|
|
r.ResponseWriter.WriteHeader(code)
|
|
}
|
|
|
|
func (r *statusRecorder) Flush() {
|
|
r.ResponseWriter.(http.Flusher).Flush()
|
|
}
|
|
|
|
func (p *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
|
route, rest, ok := SplitRoute(r.URL.Path)
|
|
if !ok {
|
|
p.writeError(w, http.StatusBadRequest, "missing route")
|
|
return
|
|
}
|
|
routeCfg, ok := p.cfg.Routes[route]
|
|
if !ok {
|
|
p.writeError(w, http.StatusNotFound, "unknown route")
|
|
return
|
|
}
|
|
if !allowedPath(rest) {
|
|
p.writeError(w, http.StatusNotFound, "not found")
|
|
return
|
|
}
|
|
model, err := peekModel(r)
|
|
if err != nil {
|
|
p.writeError(w, http.StatusRequestEntityTooLarge, "body too large")
|
|
return
|
|
}
|
|
if model == "" {
|
|
model = routeCfg.DefaultModel
|
|
}
|
|
name, ok := Choose(routeCfg.Hosts, model, p.health)
|
|
if !ok {
|
|
p.writeError(w, http.StatusServiceUnavailable, "no healthy host")
|
|
return
|
|
}
|
|
|
|
host, ok := p.cfg.Hosts[name]
|
|
if !ok {
|
|
p.writeError(w, http.StatusBadGateway, "upstream failed")
|
|
return
|
|
}
|
|
target, err := url.Parse(host.BaseURL)
|
|
if err != nil {
|
|
p.writeError(w, http.StatusBadGateway, "upstream failed")
|
|
return
|
|
}
|
|
|
|
pr := newReverseProxy(p.health, name, target, rest)
|
|
rec := &statusRecorder{ResponseWriter: w, status: http.StatusOK}
|
|
start := time.Now()
|
|
pr.ServeHTTP(rec, r)
|
|
p.log.Info("request",
|
|
"route", route,
|
|
"host", name,
|
|
"method", r.Method,
|
|
"path", rest,
|
|
"status", rec.status,
|
|
"ms", time.Since(start).Milliseconds(),
|
|
)
|
|
}
|
|
|
|
func (p *Handler) writeError(w http.ResponseWriter, status int, msg string) {
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(status)
|
|
_ = json.NewEncoder(w).Encode(map[string]string{"error": msg})
|
|
}
|
|
|
|
// newReverseProxy forwards to a single host, rewriting the path to target.Path+rest and keeping the
|
|
// original query string. It flushes after every write so long server-sent-event streams are not
|
|
// buffered, and marks the host down on any transport error other than a client disconnect.
|
|
func newReverseProxy(h Health, name string, target *url.URL, rest string) *httputil.ReverseProxy {
|
|
return &httputil.ReverseProxy{
|
|
Rewrite: func(pr *httputil.ProxyRequest) {
|
|
pr.SetURL(target)
|
|
pr.Out.URL.Path = target.Path + rest
|
|
pr.Out.URL.RawPath = ""
|
|
pr.Out.Host = target.Host
|
|
pr.SetXForwarded()
|
|
},
|
|
FlushInterval: -1,
|
|
ModifyResponse: func(resp *http.Response) error {
|
|
resp.Header.Set(HostHeader, name)
|
|
return nil
|
|
},
|
|
ErrorHandler: func(w http.ResponseWriter, req *http.Request, err error) {
|
|
if errors.Is(err, context.Canceled) {
|
|
return
|
|
}
|
|
h.MarkDown(name, err.Error())
|
|
w.Header().Set("Content-Type", "application/json")
|
|
w.WriteHeader(http.StatusBadGateway)
|
|
_ = json.NewEncoder(w).Encode(map[string]string{"error": "upstream failed", "host": name})
|
|
},
|
|
}
|
|
}
|