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