v0 plan: AGENTS.md, gate, five task files with given tests, run-plan driver

Tests were run against a private reference implementation: gate ok after every
task in order, smoke ok (stream spread ~1000 ms). The reference is not in the
repository.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
This commit is contained in:
2026-09-25 01:51:53 -07:00
co-authored by Claude Fable 5.1
parent b52126ba0b
commit 89e47d83f8
25 changed files with 1847 additions and 0 deletions
@@ -0,0 +1,99 @@
package admin_test
import (
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"git.wntrmute.dev/kyle/crossbar/internal/admin"
"git.wntrmute.dev/kyle/crossbar/internal/config"
"git.wntrmute.dev/kyle/crossbar/internal/health"
)
type fakeHosts map[string]health.Status
func (f fakeHosts) All() map[string]health.Status { return f }
func testConfig(t *testing.T) *config.Config {
c, err := config.Parse(strings.NewReader(`
listen = "127.0.0.1:1"
[hosts.alpha]
base_url = "http://alpha:1"
models = { "m" = { } }
[hosts.beta]
base_url = "http://beta:1"
models = { "m" = { } }
[routes.r]
hosts = ["alpha", "beta"]
default_model = "m"
`))
if err != nil {
t.Fatal(err)
}
return c
}
func TestHosts(t *testing.T) {
when := time.Date(2026, 9, 25, 8, 0, 0, 0, time.UTC)
h := admin.Handler(testConfig(t), fakeHosts{
"alpha": {Healthy: true, Loaded: []string{"m"}, LastOK: when},
"beta": {Healthy: false, LastErr: "HTTP 503"},
})
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/_crossbar/hosts", nil))
if rec.Code != 200 || !strings.HasPrefix(rec.Header().Get("Content-Type"), "application/json") {
t.Fatalf("status %d, content-type %q", rec.Code, rec.Header().Get("Content-Type"))
}
var out map[string]admin.HostView
if err := json.Unmarshal(rec.Body.Bytes(), &out); err != nil {
t.Fatal(err)
}
if a := out["alpha"]; !a.Healthy || len(a.Loaded) != 1 || a.LastOK != "2026-09-25T08:00:00Z" || a.LastErr != "" {
t.Errorf("alpha = %+v", a)
}
if b := out["beta"]; b.Healthy || b.LastOK != "" || b.LastErr != "HTTP 503" || b.Loaded == nil {
t.Errorf("beta = %+v (loaded must be [] not null)", b)
}
if !strings.Contains(rec.Body.String(), `"loaded":[]`) {
t.Errorf("beta.loaded must encode as []: %s", rec.Body.String())
}
}
func TestRoutes(t *testing.T) {
h := admin.Handler(testConfig(t), fakeHosts{})
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/_crossbar/routes", nil))
var out map[string]admin.RouteView
if err := json.Unmarshal(rec.Body.Bytes(), &out); err != nil {
t.Fatalf("%v: %s", err, rec.Body.String())
}
r := out["r"]
if len(r.Hosts) != 2 || r.Hosts[0] != "alpha" || r.DefaultModel != "m" {
t.Errorf("routes = %+v", out)
}
}
func TestMethodsAndUnknown(t *testing.T) {
h := admin.Handler(testConfig(t), fakeHosts{})
for _, tc := range []struct {
method, path string
want int
}{
{http.MethodPost, "/_crossbar/hosts", 405},
{http.MethodDelete, "/_crossbar/routes", 405},
{http.MethodGet, "/_crossbar/nope", 404},
{http.MethodGet, "/_crossbar/", 404},
} {
rec := httptest.NewRecorder()
h.ServeHTTP(rec, httptest.NewRequest(tc.method, tc.path, nil))
if rec.Code != tc.want {
t.Errorf("%s %s = %d, want %d", tc.method, tc.path, rec.Code, tc.want)
}
if !strings.HasPrefix(rec.Header().Get("Content-Type"), "application/json") {
t.Errorf("%s %s: errors are JSON too", tc.method, tc.path)
}
}
}
@@ -0,0 +1,191 @@
package config_test
import (
"fmt"
"path/filepath"
"strings"
"testing"
"time"
"git.wntrmute.dev/kyle/crossbar/internal/config"
)
func TestGoodFile(t *testing.T) {
c, err := config.Load(filepath.Join("testdata", "good.toml"))
if err != nil {
t.Fatalf("Load: %v", err)
}
if c.Listen != "100.64.0.9:7777" {
t.Errorf("Listen = %q", c.Listen)
}
if c.PollInterval.Duration != 5*time.Second {
t.Errorf("PollInterval = %v", c.PollInterval.Duration)
}
if c.QueueMax != 4 {
t.Errorf("QueueMax = %d", c.QueueMax)
}
alpha := c.Hosts["alpha"]
if alpha.BaseURL != "http://alpha.example:11434" {
t.Errorf("trailing slash not stripped: %q", alpha.BaseURL)
}
if alpha.Weight != 2 {
t.Errorf("alpha.Weight = %v", alpha.Weight)
}
if alpha.Models["ornith-1.5-35b-a3b"].Parallel != 4 || alpha.Models["small-9b"].Parallel != 6 {
t.Errorf("alpha.Models = %+v", alpha.Models)
}
beta := c.Hosts["beta"]
if beta.Weight != 1 {
t.Errorf("beta.Weight default = %v, want 1", beta.Weight)
}
if beta.Models["ornith-1.5-35b-a3b"].Parallel != 1 {
t.Errorf("beta parallel default = %d, want 1", beta.Models["ornith-1.5-35b-a3b"].Parallel)
}
r := c.Routes["opencode-a"]
if len(r.Hosts) != 2 || r.Hosts[0] != "alpha" || r.Hosts[1] != "beta" {
t.Errorf("route hosts = %v", r.Hosts)
}
if r.DefaultModel != "ornith-1.5-35b-a3b" {
t.Errorf("DefaultModel = %q", r.DefaultModel)
}
if c.Routes["hermes-x"].DefaultModel != "" {
t.Errorf("hermes-x DefaultModel should be empty")
}
if !c.Serves("alpha", "small-9b") || c.Serves("beta", "small-9b") || c.Serves("nope", "m") {
t.Errorf("Serves is wrong")
}
}
func TestDefaults(t *testing.T) {
c, err := config.Parse(strings.NewReader(`
listen = "127.0.0.1:1"
[hosts.a]
base_url = "http://a:1"
models = { "m" = { } }
[routes.r]
hosts = ["a"]
`))
if err != nil {
t.Fatalf("Parse: %v", err)
}
if c.PollInterval.Duration != config.DefaultPollInterval {
t.Errorf("PollInterval default = %v", c.PollInterval.Duration)
}
if c.QueueMax != config.DefaultQueueMax {
t.Errorf("QueueMax default = %d", c.QueueMax)
}
}
func TestBadFiles(t *testing.T) {
cases := []struct{ file, field string }{
{"bad-listen.toml", "listen"},
{"bad-unknown-host.toml", "routes.r.hosts"},
{"bad-default-model.toml", "routes.r.default_model"},
{"bad-unknown-key.toml", "lease_idle"},
}
for _, tc := range cases {
t.Run(tc.file, func(t *testing.T) {
_, err := config.Load(filepath.Join("testdata", tc.file))
if err == nil {
t.Fatalf("want error")
}
e, ok := config.IsError(err)
if !ok {
t.Fatalf("want *config.Error, got %T: %v", err, err)
}
if e.Field != tc.field {
t.Errorf("Field = %q, want %q (%v)", e.Field, tc.field, err)
}
if !strings.HasPrefix(err.Error(), "config: "+tc.field+": ") {
t.Errorf("Error() = %q", err.Error())
}
})
}
}
func TestBadValues(t *testing.T) {
base := `
listen = %q
poll_interval = %q
[hosts.a]
base_url = %q
weight = %v
models = { "m" = { parallel = %d } }
[routes.%s]
hosts = ["a"]
`
cases := []struct {
name string
listen, poll, url, route string
weight float64
parallel int
field string
}{
{"empty listen", "", "5s", "http://a:1", "r", 1, 1, "listen"},
{"no port", "127.0.0.1", "5s", "http://a:1", "r", 1, 1, "listen"},
{"v6 any", "[::]:7", "5s", "http://a:1", "r", 1, 1, "listen"},
{"poll too short", "127.0.0.1:7", "500ms", "http://a:1", "r", 1, 1, "poll_interval"},
{"ftp url", "127.0.0.1:7", "5s", "ftp://a:1", "r", 1, 1, "hosts.a.base_url"},
{"no host", "127.0.0.1:7", "5s", "http://", "r", 1, 1, "hosts.a.base_url"},
{"query", "127.0.0.1:7", "5s", "http://a:1/v1?x=1", "r", 1, 1, "hosts.a.base_url"},
{"negative weight", "127.0.0.1:7", "5s", "http://a:1", "r", -1, 1, "hosts.a.weight"},
{"negative parallel", "127.0.0.1:7", "5s", "http://a:1", "r", 1, -2, "hosts.a.models.m.parallel"},
{"route name", "127.0.0.1:7", "5s", "http://a:1", "Bad_Name", 1, 1, "routes.Bad_Name"},
}
for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
text := fmt.Sprintf(base, tc.listen, tc.poll, tc.url, tc.weight, tc.parallel, tc.route)
_, err := config.Parse(strings.NewReader(text))
if err == nil {
t.Fatalf("want error for %s", tc.name)
}
e, ok := config.IsError(err)
if !ok {
t.Fatalf("want *config.Error, got %T: %v", err, err)
}
if e.Field != tc.field {
t.Errorf("Field = %q, want %q (%v)", e.Field, tc.field, err)
}
})
}
}
func TestMissingSections(t *testing.T) {
for _, tc := range []struct{ name, text, field string }{
{"no hosts", "listen = \"127.0.0.1:7\"\n[routes.r]\nhosts = [\"a\"]\n", "hosts"},
{"no routes", "listen = \"127.0.0.1:7\"\n[hosts.a]\nbase_url = \"http://a:1\"\nmodels = { \"m\" = { } }\n", "routes"},
{"host without models", "listen = \"127.0.0.1:7\"\n[hosts.a]\nbase_url = \"http://a:1\"\n[routes.r]\nhosts = [\"a\"]\n", "hosts.a.models"},
{"route without hosts", "listen = \"127.0.0.1:7\"\n[hosts.a]\nbase_url = \"http://a:1\"\nmodels = { \"m\" = { } }\n[routes.r]\n", "routes.r.hosts"},
{"host twice", "listen = \"127.0.0.1:7\"\n[hosts.a]\nbase_url = \"http://a:1\"\nmodels = { \"m\" = { } }\n[routes.r]\nhosts = [\"a\", \"a\"]\n", "routes.r.hosts"},
} {
t.Run(tc.name, func(t *testing.T) {
_, err := config.Parse(strings.NewReader(tc.text))
e, ok := config.IsError(err)
if !ok {
t.Fatalf("want *config.Error, got %v", err)
}
if e.Field != tc.field {
t.Errorf("Field = %q, want %q", e.Field, tc.field)
}
})
}
}
func TestNotTOML(t *testing.T) {
_, err := config.Parse(strings.NewReader("listen = [unterminated"))
if err == nil {
t.Fatal("want error")
}
if _, ok := config.IsError(err); ok {
t.Errorf("a syntax error is not a validation Error")
}
if !strings.HasPrefix(err.Error(), "config: ") {
t.Errorf("Error() = %q", err.Error())
}
}
func TestMissingFile(t *testing.T) {
if _, err := config.Load(filepath.Join("testdata", "does-not-exist.toml")); err == nil {
t.Fatal("want error")
}
}
@@ -0,0 +1,9 @@
listen = "127.0.0.1:7777"
[hosts.alpha]
base_url = "http://alpha.example:11434"
models = { "m" = { } }
[routes.r]
hosts = ["alpha"]
default_model = "not-served"
@@ -0,0 +1,8 @@
listen = "0.0.0.0:7777"
[hosts.alpha]
base_url = "http://alpha.example:11434"
models = { "m" = { } }
[routes.r]
hosts = ["alpha"]
@@ -0,0 +1,8 @@
listen = "127.0.0.1:7777"
[hosts.alpha]
base_url = "http://alpha.example:11434"
models = { "m" = { } }
[routes.r]
hosts = ["alpha", "gamma"]
@@ -0,0 +1,9 @@
listen = "127.0.0.1:7777"
lease_idle = "30m"
[hosts.alpha]
base_url = "http://alpha.example:11434"
models = { "m" = { } }
[routes.r]
hosts = ["alpha"]
+19
View File
@@ -0,0 +1,19 @@
listen = "100.64.0.9:7777"
poll_interval = "5s"
queue_max = 4
[hosts.alpha]
base_url = "http://alpha.example:11434/"
weight = 2.0
models = { "ornith-1.5-35b-a3b" = { parallel = 4 }, "small-9b" = { parallel = 6 } }
[hosts.beta]
base_url = "https://beta.example:8081"
models = { "ornith-1.5-35b-a3b" = { } }
[routes.opencode-a]
hosts = ["alpha", "beta"]
default_model = "ornith-1.5-35b-a3b"
[routes.hermes-x]
hosts = ["beta"]
@@ -0,0 +1,154 @@
package health_test
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"sync/atomic"
"testing"
"time"
"git.wntrmute.dev/kyle/crossbar/internal/health"
)
// fake is a llama-server stand-in whose /health can be flipped and whose model list is fixed.
type fake struct {
srv *httptest.Server
down atomic.Bool
models []string
hits atomic.Int32
}
func newFake(t *testing.T, models ...string) *fake {
f := &fake{models: models}
mux := http.NewServeMux()
mux.HandleFunc("/health", func(w http.ResponseWriter, r *http.Request) {
f.hits.Add(1)
if f.down.Load() {
http.Error(w, "loading", http.StatusServiceUnavailable)
return
}
_, _ = w.Write([]byte(`{"status":"ok"}`))
})
mux.HandleFunc("/v1/models", func(w http.ResponseWriter, r *http.Request) {
type m struct {
ID string `json:"id"`
}
var data []m
for _, id := range f.models {
data = append(data, m{ID: id})
}
_ = json.NewEncoder(w).Encode(map[string]any{"object": "list", "data": data})
})
f.srv = httptest.NewServer(mux)
t.Cleanup(f.srv.Close)
return f
}
func TestFirstPollMakesHealthy(t *testing.T) {
a := newFake(t, "zeta", "alpha", "alpha")
tbl := health.New(map[string]string{"a": a.srv.URL}, time.Hour, nil)
if s, ok := tbl.Get("a"); !ok || s.Healthy || len(s.Loaded) != 0 {
t.Fatalf("before any poll: %+v %v", s, ok)
}
tbl.PollOnce(context.Background())
s, _ := tbl.Get("a")
if !s.Healthy || s.Consecutive != 1 || s.LastErr != "" || s.LastOK.IsZero() {
t.Errorf("after one good poll: %+v", s)
}
if len(s.Loaded) != 2 || s.Loaded[0] != "alpha" || s.Loaded[1] != "zeta" {
t.Errorf("Loaded = %v, want sorted, unique [alpha zeta]", s.Loaded)
}
}
func TestFailureThenRecoveryNeedsTwoPolls(t *testing.T) {
a := newFake(t, "m")
tbl := health.New(map[string]string{"a": a.srv.URL}, time.Hour, nil)
ctx := context.Background()
tbl.PollOnce(ctx)
a.down.Store(true)
tbl.PollOnce(ctx)
s, _ := tbl.Get("a")
if s.Healthy || s.Consecutive != 0 || s.LastErr == "" {
t.Fatalf("after failure: %+v", s)
}
if len(s.Loaded) != 1 {
t.Errorf("Loaded is left as last seen; got %v", s.Loaded)
}
a.down.Store(false)
tbl.PollOnce(ctx)
if s, _ := tbl.Get("a"); s.Healthy || s.Consecutive != 1 {
t.Errorf("one good poll after a failure must not be healthy yet: %+v", s)
}
tbl.PollOnce(ctx)
if s, _ := tbl.Get("a"); !s.Healthy || s.Consecutive != 2 || s.LastErr != "" {
t.Errorf("two good polls: %+v", s)
}
}
func TestMarkDown(t *testing.T) {
a := newFake(t, "m")
tbl := health.New(map[string]string{"a": a.srv.URL}, time.Hour, nil)
tbl.PollOnce(context.Background())
tbl.MarkDown("a", "connection refused")
s, _ := tbl.Get("a")
if s.Healthy || s.Consecutive != 0 || s.LastErr != "marked down: connection refused" {
t.Errorf("after MarkDown: %+v", s)
}
tbl.MarkDown("nobody", "x") // unknown hosts are ignored, not a panic
tbl.PollOnce(context.Background())
if s, _ := tbl.Get("a"); s.Healthy {
t.Errorf("one poll after MarkDown must not be healthy: %+v", s)
}
}
func TestUnreachableAndUnknown(t *testing.T) {
tbl := health.New(map[string]string{"a": "http://127.0.0.1:1"}, time.Hour, &http.Client{Timeout: time.Second})
tbl.PollOnce(context.Background())
s, ok := tbl.Get("a")
if !ok || s.Healthy || s.LastErr == "" {
t.Errorf("unreachable host: %+v %v", s, ok)
}
if _, ok := tbl.Get("zzz"); ok {
t.Errorf("unknown host must report ok=false")
}
}
func TestAllIsACopy(t *testing.T) {
a := newFake(t, "m")
tbl := health.New(map[string]string{"a": a.srv.URL}, time.Hour, nil)
tbl.PollOnce(context.Background())
all := tbl.All()
all["a"].Loaded[0] = "changed"
if s, _ := tbl.Get("a"); s.Loaded[0] != "m" {
t.Errorf("All must return copies")
}
if len(all) != 1 {
t.Errorf("All = %v", all)
}
}
func TestRunPollsOnStart(t *testing.T) {
a := newFake(t, "m")
tbl := health.New(map[string]string{"a": a.srv.URL}, 20*time.Millisecond, nil)
ctx, cancel := context.WithCancel(context.Background())
done := make(chan struct{})
go func() { tbl.Run(ctx); close(done) }()
deadline := time.Now().Add(2 * time.Second)
for a.hits.Load() < 3 && time.Now().Before(deadline) {
time.Sleep(5 * time.Millisecond)
}
cancel()
select {
case <-done:
case <-time.After(time.Second):
t.Fatal("Run did not return after cancel")
}
if a.hits.Load() < 3 {
t.Errorf("Run polled %d times in 2s at 20ms interval", a.hits.Load())
}
if s, _ := tbl.Get("a"); !s.Healthy {
t.Errorf("not healthy after Run: %+v", s)
}
}
@@ -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))
}
}