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