diff --git a/docs/implementer-log.md b/docs/implementer-log.md index 9939db8..202db05 100644 --- a/docs/implementer-log.md +++ b/docs/implementer-log.md @@ -7,5 +7,6 @@ owner fills in the Model column. The reviewer adds findings under "Reviews" once |---|---|---|---|---|---|---|---| | v0/01-module-gate-config | 2026-09-25 | done | 1 | pass | none | `go mod download` fetched the module (network available); gate passed on the first run. | ? | | v0/02-health | 2026-09-25 | done | 1 | pass | none | First gate run passed. `MarkDown` initially forgot to write the entry back; caught by `TestMarkDown`. | ? | +| v0/03-proxy | 2026-09-25 | done | 1 | pass | none | `SplitRoute` must reject an empty first segment (`/`, `//x`) as `ok=false`; the model peek restores the body and leaves non-JSON/empty as `""`. | ? | ## Reviews diff --git a/internal/proxy/proxy.go b/internal/proxy/proxy.go new file mode 100644 index 0000000..bb59b26 --- /dev/null +++ b/internal/proxy/proxy.go @@ -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}) + }, + } +} diff --git a/internal/proxy/proxy_test.go b/internal/proxy/proxy_test.go new file mode 100644 index 0000000..8d15174 --- /dev/null +++ b/internal/proxy/proxy_test.go @@ -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)) + } +}