v0 plan: move given files to _files/ so go vet ./... ignores them; note the first-attempt finding
Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
This commit is contained in:
@@ -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"]
|
||||
@@ -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))
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user