Add the routing reverse proxy

Implemented-By: OpenCode session (model recorded in docs/implementer-log.md)
This commit is contained in:
2026-09-25 02:26:19 -07:00
parent 68ad94d693
commit 5b347d9dd9
3 changed files with 559 additions and 0 deletions
+240
View File
@@ -0,0 +1,240 @@
// 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})
},
}
}
+318
View File
@@ -0,0 +1,318 @@
package proxy_test
import (
"encoding/json"
"fmt"
"io"
"net/http"
"net/http/httptest"
"strings"
"sync"
"testing"
"time"
"git.wntrmute.dev/kyle/crossbar/internal/config"
"git.wntrmute.dev/kyle/crossbar/internal/health"
"git.wntrmute.dev/kyle/crossbar/internal/proxy"
)
// fakeHealth is a hand-set health table that also records MarkDown calls.
type fakeHealth struct {
mu sync.Mutex
st map[string]health.Status
marked []string
}
func (f *fakeHealth) Get(name string) (health.Status, bool) {
f.mu.Lock()
defer f.mu.Unlock()
s, ok := f.st[name]
return s, ok
}
func (f *fakeHealth) MarkDown(name, reason string) {
f.mu.Lock()
defer f.mu.Unlock()
f.marked = append(f.marked, name)
s := f.st[name]
s.Healthy = false
s.LastErr = reason
f.st[name] = s
}
func (f *fakeHealth) markedHosts() []string {
f.mu.Lock()
defer f.mu.Unlock()
return append([]string{}, f.marked...)
}
// upstream records what it received and answers with its name.
type upstream struct {
name string
srv *httptest.Server
mu sync.Mutex
reqs []recorded
}
type recorded struct {
method, path, host, xff string
body string
}
func newUpstream(t *testing.T, name string) *upstream {
u := &upstream{name: name}
u.srv = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
b, _ := io.ReadAll(r.Body)
u.mu.Lock()
u.reqs = append(u.reqs, recorded{r.Method, r.URL.RequestURI(), r.Host, r.Header.Get("X-Forwarded-For"), string(b)})
u.mu.Unlock()
w.Header().Set("Content-Type", "application/json")
fmt.Fprintf(w, `{"from":%q}`, name)
}))
t.Cleanup(u.srv.Close)
return u
}
func (u *upstream) last(t *testing.T) recorded {
u.mu.Lock()
defer u.mu.Unlock()
if len(u.reqs) == 0 {
t.Fatalf("%s: no request received", u.name)
}
return u.reqs[len(u.reqs)-1]
}
func cfgFor(t *testing.T, alpha, beta string) *config.Config {
c, err := config.Parse(strings.NewReader(fmt.Sprintf(`
listen = "127.0.0.1:1"
[hosts.alpha]
base_url = %q
models = { "shared" = { }, "alpha-only" = { } }
[hosts.beta]
base_url = %q
models = { "shared" = { }, "beta-only" = { } }
[routes.r]
hosts = ["alpha", "beta"]
default_model = "shared"
[routes.beta-first]
hosts = ["beta", "alpha"]
`, alpha, beta)))
if err != nil {
t.Fatal(err)
}
return c
}
func healthy(loaded ...string) health.Status {
return health.Status{Healthy: true, Loaded: loaded, Consecutive: 1}
}
func TestSplitRoute(t *testing.T) {
for _, tc := range []struct {
path, route, rest string
ok bool
}{
{"/a/v1/x", "a", "/v1/x", true},
{"/a/v1/x?q=1", "a", "/v1/x?q=1", true},
{"/a", "a", "/", true},
{"/a/", "a", "/", true},
{"/opencode-a/v1/chat/completions", "opencode-a", "/v1/chat/completions", true},
{"/", "", "", false},
{"//x", "", "", false},
{"", "", "", false},
{"noslash/v1", "", "", false},
} {
route, rest, ok := proxy.SplitRoute(tc.path)
if route != tc.route || rest != tc.rest || ok != tc.ok {
t.Errorf("SplitRoute(%q) = %q %q %v, want %q %q %v", tc.path, route, rest, ok, tc.route, tc.rest, tc.ok)
}
}
}
func TestChoose(t *testing.T) {
h := &fakeHealth{st: map[string]health.Status{
"down": {Healthy: false, Loaded: []string{"m"}},
"alpha": healthy("shared", "alpha-only"),
"beta": healthy("shared", "beta-only"),
}}
hosts := []string{"down", "alpha", "beta"}
if got, ok := proxy.Choose(hosts, "", h); !ok || got != "alpha" {
t.Errorf("no model: %q %v, want alpha (first healthy)", got, ok)
}
if got, ok := proxy.Choose(hosts, "beta-only", h); !ok || got != "beta" {
t.Errorf("beta-only: %q %v, want beta (has the model loaded)", got, ok)
}
if got, ok := proxy.Choose(hosts, "nobody-has-it", h); !ok || got != "alpha" {
t.Errorf("unknown model falls back to the first healthy host: %q %v", got, ok)
}
if got, ok := proxy.Choose([]string{"down", "missing"}, "m", h); ok {
t.Errorf("no healthy host must give ok=false, got %q", got)
}
if got, ok := proxy.Choose(nil, "m", h); ok {
t.Errorf("empty hosts: %q %v", got, ok)
}
}
func TestRoutesToFirstHealthyAndRewrites(t *testing.T) {
alpha, beta := newUpstream(t, "alpha"), newUpstream(t, "beta")
h := &fakeHealth{st: map[string]health.Status{"alpha": healthy("shared"), "beta": healthy("shared")}}
p := proxy.New(cfgFor(t, alpha.srv.URL, beta.srv.URL), h, nil)
rec := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodGet, "http://crossbar.local:7777/r/v1/models?x=1", nil)
req.RemoteAddr = "10.9.8.7:5555"
p.ServeHTTP(rec, req)
if rec.Code != 200 || rec.Header().Get(proxy.HostHeader) != "alpha" {
t.Fatalf("status %d host %q body %s", rec.Code, rec.Header().Get(proxy.HostHeader), rec.Body.String())
}
got := alpha.last(t)
if got.path != "/v1/models?x=1" {
t.Errorf("upstream path = %q, want route stripped and query kept", got.path)
}
if got.host != strings.TrimPrefix(alpha.srv.URL, "http://") {
t.Errorf("Host header = %q, want the upstream's %q", got.host, strings.TrimPrefix(alpha.srv.URL, "http://"))
}
if got.xff != "10.9.8.7" {
t.Errorf("X-Forwarded-For = %q, want the client address", got.xff)
}
if !strings.Contains(rec.Body.String(), `"from":"alpha"`) {
t.Errorf("body = %s", rec.Body.String())
}
}
func TestModelPreferenceAndBodyPassThrough(t *testing.T) {
alpha, beta := newUpstream(t, "alpha"), newUpstream(t, "beta")
h := &fakeHealth{st: map[string]health.Status{"alpha": healthy("shared", "alpha-only"), "beta": healthy("shared", "beta-only")}}
p := proxy.New(cfgFor(t, alpha.srv.URL, beta.srv.URL), h, nil)
body := `{"model":"beta-only","messages":[{"role":"user","content":"hi"}],"stream":false}`
rec := httptest.NewRecorder()
p.ServeHTTP(rec, httptest.NewRequest(http.MethodPost, "/r/v1/chat/completions", strings.NewReader(body)))
if rec.Code != 200 || rec.Header().Get(proxy.HostHeader) != "beta" {
t.Fatalf("status %d host %q", rec.Code, rec.Header().Get(proxy.HostHeader))
}
if got := beta.last(t); got.body != body || got.method != http.MethodPost {
t.Errorf("upstream got %+v; the body must arrive unchanged after the model peek", got)
}
// Not JSON: no model, the route default ("shared") applies, first healthy wins.
rec = httptest.NewRecorder()
p.ServeHTTP(rec, httptest.NewRequest(http.MethodPost, "/r/v1/embeddings", strings.NewReader("plain text")))
if rec.Header().Get(proxy.HostHeader) != "alpha" {
t.Errorf("non-JSON body: host %q, want alpha", rec.Header().Get(proxy.HostHeader))
}
if got := alpha.last(t); got.body != "plain text" {
t.Errorf("non-JSON body must pass through unchanged, got %q", got.body)
}
}
func TestFailoverOnUpstreamError(t *testing.T) {
alpha, beta := newUpstream(t, "alpha"), newUpstream(t, "beta")
h := &fakeHealth{st: map[string]health.Status{"alpha": healthy("shared"), "beta": healthy("shared")}}
p := proxy.New(cfgFor(t, alpha.srv.URL, beta.srv.URL), h, nil)
alpha.srv.Close() // health still believes alpha is up
rec := httptest.NewRecorder()
p.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/r/v1/models", nil))
if rec.Code != http.StatusBadGateway {
t.Fatalf("first request after alpha died: %d, want 502", rec.Code)
}
var e map[string]string
if err := json.Unmarshal(rec.Body.Bytes(), &e); err != nil || e["error"] != "upstream failed" || e["host"] != "alpha" {
t.Errorf("502 body = %s", rec.Body.String())
}
if m := h.markedHosts(); len(m) != 1 || m[0] != "alpha" {
t.Errorf("MarkDown calls = %v, want [alpha]", m)
}
rec = httptest.NewRecorder()
p.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/r/v1/models", nil))
if rec.Code != 200 || rec.Header().Get(proxy.HostHeader) != "beta" {
t.Errorf("second request: %d %q, want 200 from beta", rec.Code, rec.Header().Get(proxy.HostHeader))
}
}
func TestErrors(t *testing.T) {
alpha, beta := newUpstream(t, "alpha"), newUpstream(t, "beta")
h := &fakeHealth{st: map[string]health.Status{"alpha": {Healthy: false}, "beta": {Healthy: false}}}
p := proxy.New(cfgFor(t, alpha.srv.URL, beta.srv.URL), h, nil)
for _, tc := range []struct {
name, method, path string
body io.Reader
want int
msg string
}{
{"bare slash", http.MethodGet, "/", nil, 400, "missing route"},
{"double slash", http.MethodGet, "//v1/models", nil, 400, "missing route"},
{"unknown route", http.MethodGet, "/nope/v1/models", nil, 404, "unknown route"},
{"disallowed path", http.MethodGet, "/r/slots", nil, 404, "not found"},
{"admin through proxy", http.MethodGet, "/r/_crossbar/hosts", nil, 404, "not found"},
{"no healthy host", http.MethodGet, "/r/v1/models", nil, 503, "no healthy host"},
{"body too large", http.MethodPost, "/r/v1/chat/completions", strings.NewReader(strings.Repeat("x", proxy.MaxBody+1)), 413, "body too large"},
} {
rec := httptest.NewRecorder()
p.ServeHTTP(rec, httptest.NewRequest(tc.method, tc.path, tc.body))
if rec.Code != tc.want {
t.Errorf("%s: status %d, want %d", tc.name, rec.Code, tc.want)
}
var e map[string]string
if err := json.Unmarshal(rec.Body.Bytes(), &e); err != nil || e["error"] != tc.msg {
t.Errorf("%s: body %s, want error %q", tc.name, rec.Body.String(), tc.msg)
}
if !strings.HasPrefix(rec.Header().Get("Content-Type"), "application/json") {
t.Errorf("%s: errors are JSON", tc.name)
}
}
if len(h.markedHosts()) != 0 {
t.Errorf("errors before choosing a host must not mark anything down: %v", h.markedHosts())
}
}
// TestStreamingIsNotBuffered: the upstream writes one chunk, flushes, and then waits until the
// test has *read* that chunk. If the proxy buffered, the read would never complete.
func TestStreamingIsNotBuffered(t *testing.T) {
release := make(chan struct{})
up := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "text/event-stream")
w.WriteHeader(200)
fmt.Fprint(w, "data: first\n\n")
w.(http.Flusher).Flush()
select {
case <-release:
case <-time.After(5 * time.Second):
}
fmt.Fprint(w, "data: second\n\n")
}))
t.Cleanup(up.Close)
beta := newUpstream(t, "beta")
h := &fakeHealth{st: map[string]health.Status{"alpha": healthy("shared"), "beta": healthy("shared")}}
front := httptest.NewServer(proxy.New(cfgFor(t, up.URL, beta.srv.URL), h, nil))
t.Cleanup(front.Close)
resp, err := http.Post(front.URL+"/r/v1/chat/completions", "application/json", strings.NewReader(`{"model":"shared","stream":true}`))
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
buf := make([]byte, 64)
done := make(chan string, 1)
go func() {
n, err := resp.Body.Read(buf)
if err != nil {
done <- "read error: " + err.Error()
return
}
done <- string(buf[:n])
}()
select {
case got := <-done:
if !strings.HasPrefix(got, "data: first") {
t.Fatalf("first read = %q", got)
}
case <-time.After(2 * time.Second):
t.Fatal("the first chunk did not arrive before the upstream finished: the proxy buffers")
}
close(release)
rest, _ := io.ReadAll(resp.Body)
if !strings.Contains(string(rest), "data: second") {
t.Errorf("rest = %q", rest)
}
if resp.Header.Get(proxy.HostHeader) != "alpha" {
t.Errorf("host header %q", resp.Header.Get(proxy.HostHeader))
}
}