Add the routing reverse proxy
Implemented-By: OpenCode session (model recorded in docs/implementer-log.md)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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})
|
||||
},
|
||||
}
|
||||
}
|
||||
@@ -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))
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user