diff --git a/Makefile b/Makefile new file mode 100644 index 0000000..1e69fb0 --- /dev/null +++ b/Makefile @@ -0,0 +1,17 @@ +# crossbar gate. `make gate` must pass before any task is called done. It needs no network. + +.PHONY: gate build smoke + +gate: + @test -z "$$(gofmt -l . 2>&1)" || { echo "gofmt: these files need formatting:"; gofmt -l .; exit 1; } + go vet ./... + go test -race -count=1 ./... + sh scripts/check-lines.sh + @echo "gate: ok" + +build: + go build -o bin/ ./cmd/... + +# Runs the whole thing against two fake upstreams. Task 05 brings the script. +smoke: build + sh tools/smoke.sh diff --git a/README.md b/README.md new file mode 100644 index 0000000..9cf69eb --- /dev/null +++ b/README.md @@ -0,0 +1,106 @@ +# crossbar + +crossbar is an affinity router in front of several `llama-server` routers. A client's identity is +the first path segment of its base URL; v0 routes each request to the first healthy host on that +route's list and streams the answer back unbuffered. + +## Build + +Build everything with `make build`; the binaries land in `bin/`. Check the work with `make gate`, +which runs the formatter, vet, tests and line-length check with no network. When the code is ready, +run `make smoke`, which starts two fake upstreams and exercises routing, failover, recovery and +streaming over real HTTP. + +## Configure + +crossbar reads one TOML file. This is `example.toml`: + +```toml +# crossbar example configuration. Replace and the addresses with your own. +listen = "127.0.0.1:17777" # never 0.0.0.0 — bind the tailnet address in production +poll_interval = "1s" # 60s in production; 1s makes the smoke run quick +queue_max = 8 + +[hosts.alpha] +base_url = "http://127.0.0.1:18081" # e.g. http://straylight.:11434 +weight = 1.0 +models = { "ornith-1.5-35b-a3b" = { parallel = 4 }, "small-9b" = { parallel = 6 } } + +[hosts.beta] +base_url = "http://127.0.0.1:18082" # e.g. http://titan.:8081 +weight = 2.0 +models = { "ornith-1.5-35b-a3b" = { parallel = 4 } } + +# v0: a route is a preference list; the first healthy host that has the model wins. +[routes.opencode-a] +hosts = ["alpha", "beta"] +default_model = "ornith-1.5-35b-a3b" + +[routes.hermes-x] +hosts = ["beta", "alpha"] +``` + +| Key | Meaning | +| --- | --- | +| `listen` | Where crossbar binds. A tailnet address, never `0.0.0.0`. | +| `poll_interval` | How often each host is health-checked. 60s in production; 1s makes the smoke run quick. | +| `queue_max` | Reserved for v1 queueing; no effect in v0. | +| `hosts..base_url` | The llama-server base URL this host serves. | +| `hosts..weight` | Relative share of new routes this host receives. | +| `hosts..models` | The models this host serves, with per-model parallel tuning. | +| `routes..hosts` | Preference order: the first healthy host that serves the model wins. | +| `routes..default_model` | Model used when a request omits one; must be served by a host in the route. | + +## Run + +Copy the binary, the config and the unit into place, reload systemd, and start it: + +```sh +install -m 0755 bin/crossbar /usr/local/bin/crossbar +install -d -m 0755 /etc/crossbar +install -m 0644 crossbar.toml /etc/crossbar/crossbar.toml +install -m 0644 deploy/crossbar.service /etc/systemd/system/crossbar.service +systemctl daemon-reload +systemctl enable --now crossbar +``` + +## Point clients at it + +OpenCode, one provider for every project. Each instance is launched as +`CROSSBAR_ROUTE="$(basename "$PWD")-$$" opencode`: + +```jsonc +"provider": { "crossbar": { "npm": "@ai-sdk/openai-compatible", + "options": { "baseURL": "http://crossbar.:7777/{env:CROSSBAR_ROUTE}/v1" }, + "models": { "ornith-1.5-35b-a3b": {} } } } +``` + +Hermes, in `config.yaml`: + +```yaml +custom_providers: + - name: crossbar + base_url: http://crossbar.:7777/hermes-/v1 + models: { ornith-1.5-35b-a3b: {} } +``` + +The route name in the URL must exist in `[routes]`; unknown routes are 404. + +## Inspect + +`GET /_crossbar/hosts` reports every host's health and loaded models: + +```json +{"alpha":{"healthy":true,"loaded":["ornith-1.5-35b-a3b","small-9b"],"last_ok":"2026-09-25T09:34:18Z","last_err":""},"beta":{"healthy":true,"loaded":["ornith-1.5-35b-a3b"],"last_ok":"2026-09-25T09:34:18Z","last_err":""}} +``` + +`GET /_crossbar/routes` reports each route's preference order and default model: + +```json +{"hermes-x":{"hosts":["beta","alpha"],"default_model":""},"opencode-a":{"hosts":["alpha","beta"],"default_model":"ornith-1.5-35b-a3b"}} +``` + +## What v0 does not do + +Leases and stickiness, SQLite, `/slots`, queueing and wake-on-LAN are out of scope for v0; see +`PLAN.md`. diff --git a/cmd/crossbar/main.go b/cmd/crossbar/main.go new file mode 100644 index 0000000..7c34d9c --- /dev/null +++ b/cmd/crossbar/main.go @@ -0,0 +1,79 @@ +// Command crossbar is the affinity router for the fleet's llama-server instances. It wires config, +// health, proxy and admin into one HTTP server with graceful shutdown on SIGINT/SIGTERM. +package main + +import ( + "context" + "errors" + "flag" + "fmt" + "log/slog" + "net/http" + "os" + "os/signal" + "syscall" + "time" + + "git.wntrmute.dev/kyle/crossbar/internal/admin" + "git.wntrmute.dev/kyle/crossbar/internal/config" + "git.wntrmute.dev/kyle/crossbar/internal/health" + "git.wntrmute.dev/kyle/crossbar/internal/proxy" +) + +func main() { + if err := run(); err != nil { + fmt.Fprintln(os.Stderr, "crossbar: "+err.Error()) + os.Exit(1) + } +} + +func run() error { + configPath := flag.String("config", "crossbar.toml", "path to the crossbar config file") + flag.Parse() + + cfg, err := config.Load(*configPath) + if err != nil { + return err + } + + log := slog.New(slog.NewTextHandler(os.Stderr, nil)) + + baseURLs := make(map[string]string, len(cfg.Hosts)) + for name, host := range cfg.Hosts { + baseURLs[name] = host.BaseURL + } + + table := health.New(baseURLs, cfg.PollInterval.Duration, nil) + ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM) + defer stop() + go table.Run(ctx) + + mux := http.NewServeMux() + mux.Handle("/_crossbar/", admin.Handler(cfg, table)) + mux.Handle("/", proxy.New(cfg, table, log)) + + srv := &http.Server{ + Addr: cfg.Listen, + Handler: mux, + ReadHeaderTimeout: 10 * time.Second, + } + + serverErr := make(chan error, 1) + go func() { + log.Info("listening", "addr", srv.Addr) + serverErr <- srv.ListenAndServe() + }() + + select { + case <-ctx.Done(): + log.Info("shutting down") + shutdownCtx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + return srv.Shutdown(shutdownCtx) + case err := <-serverErr: + if errors.Is(err, http.ErrServerClosed) { + return nil + } + return err + } +} diff --git a/cmd/fakeupstream/main.go b/cmd/fakeupstream/main.go new file mode 100644 index 0000000..dbb7801 --- /dev/null +++ b/cmd/fakeupstream/main.go @@ -0,0 +1,101 @@ +// fakeupstream stands in for a llama-server router in tests and the smoke run. Do not edit. +// +// fakeupstream -listen 127.0.0.1:18081 -name alpha -models a,b -down-file /tmp/alpha.down +// +// /health answers 503 while the down file exists, 200 otherwise. /v1/models lists -models. +// /props answers a small JSON object. /v1/chat/completions echoes: a streamed answer of five +// SSE chunks 200 ms apart when the body has "stream": true, one JSON answer otherwise. Every +// response carries X-Upstream: . +package main + +import ( + "encoding/json" + "flag" + "fmt" + "io" + "log" + "net/http" + "os" + "strings" + "time" +) + +func main() { + listen := flag.String("listen", "127.0.0.1:18081", "address to listen on") + name := flag.String("name", "fake", "name reported in X-Upstream and answers") + models := flag.String("models", "m", "comma-separated model ids for /v1/models") + downFile := flag.String("down-file", "", "while this file exists, /health answers 503") + flag.Parse() + + ids := strings.Split(*models, ",") + mux := http.NewServeMux() + stamp := func(w http.ResponseWriter) { w.Header().Set("X-Upstream", *name) } + + mux.HandleFunc("/health", func(w http.ResponseWriter, r *http.Request) { + stamp(w) + if *downFile != "" { + if _, err := os.Stat(*downFile); err == nil { + http.Error(w, `{"error":{"message":"Loading model"}}`, http.StatusServiceUnavailable) + return + } + } + writeJSON(w, map[string]string{"status": "ok"}) + }) + mux.HandleFunc("/v1/models", func(w http.ResponseWriter, r *http.Request) { + stamp(w) + data := []map[string]any{} + for _, id := range ids { + data = append(data, map[string]any{"id": id, "object": "model", "owned_by": *name}) + } + writeJSON(w, map[string]any{"object": "list", "data": data}) + }) + mux.HandleFunc("/props", func(w http.ResponseWriter, r *http.Request) { + stamp(w) + writeJSON(w, map[string]any{"default_generation_settings": map[string]any{"n_ctx": 8192}, "total_slots": 2, "model_path": *name}) + }) + mux.HandleFunc("/v1/chat/completions", func(w http.ResponseWriter, r *http.Request) { + stamp(w) + body, _ := io.ReadAll(io.LimitReader(r.Body, 1<<20)) + var req struct { + Model string `json:"model"` + Stream bool `json:"stream"` + } + _ = json.Unmarshal(body, &req) + if !req.Stream { + writeJSON(w, map[string]any{ + "id": "chatcmpl-fake", "object": "chat.completion", "model": req.Model, + "choices": []map[string]any{{"index": 0, "message": map[string]string{"role": "assistant", "content": "hello from " + *name}, "finish_reason": "stop"}}, + "usage": map[string]int{"prompt_tokens": 3, "completion_tokens": 3, "total_tokens": 6}, + }) + return + } + w.Header().Set("Content-Type", "text/event-stream") + w.Header().Set("Cache-Control", "no-cache") + w.WriteHeader(http.StatusOK) + fl, _ := w.(http.Flusher) + for i := 1; i <= 5; i++ { + chunk := map[string]any{"id": "chatcmpl-fake", "object": "chat.completion.chunk", "model": req.Model, + "choices": []map[string]any{{"index": 0, "delta": map[string]string{"content": fmt.Sprintf("%s chunk %d ", *name, i)}}}} + b, _ := json.Marshal(chunk) + fmt.Fprintf(w, "data: %s\n\n", b) + if fl != nil { + fl.Flush() + } + time.Sleep(200 * time.Millisecond) + } + fmt.Fprint(w, "data: [DONE]\n\n") + }) + mux.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) { + stamp(w) + http.Error(w, `{"error":"not found"}`, http.StatusNotFound) + }) + + log.Printf("fakeupstream %s listening on %s models=%v", *name, *listen, ids) + srv := &http.Server{Addr: *listen, Handler: mux, ReadHeaderTimeout: 5 * time.Second} + log.Fatal(srv.ListenAndServe()) +} + +func writeJSON(w http.ResponseWriter, v any) { + w.Header().Set("Content-Type", "application/json") + _ = json.NewEncoder(w).Encode(v) +} diff --git a/deploy/crossbar.service b/deploy/crossbar.service new file mode 100644 index 0000000..b9d6e12 --- /dev/null +++ b/deploy/crossbar.service @@ -0,0 +1,18 @@ +[Unit] +Description=crossbar affinity router for llama-server +After=network-online.target +Wants=network-online.target + +[Service] +ExecStart=/usr/local/bin/crossbar -config /etc/crossbar/crossbar.toml +Restart=on-failure +RestartSec=2s +DynamicUser=yes +StateDirectory=crossbar +NoNewPrivileges=yes +ProtectSystem=strict +ProtectHome=yes +PrivateTmp=yes + +[Install] +WantedBy=multi-user.target diff --git a/docs/implementer-log.md b/docs/implementer-log.md index 90a7119..38b9e46 100644 --- a/docs/implementer-log.md +++ b/docs/implementer-log.md @@ -5,5 +5,38 @@ owner fills in the Model column. The reviewer adds findings under "Reviews" once | Task | Date | Status | Gate runs | First gate | Deviations | Notes | Model | |---|---|---|---|---|---|---|---| +| 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. | llama.cpp/ornith-1.5-35b-a3b | +| 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`. | llama.cpp/ornith-1.5-35b-a3b | +| 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 `""`. | llama.cpp/ornith-1.5-35b-a3b | +| v0/04-admin-main | 2026-09-25 | done | 1 | pass | none | `timeout --signal=TERM 3` exits 124 on a timed-out child on this GNU system, so the task's `exit=0` is not observable through it; sent SIGTERM directly and confirmed crossbar's own exit code is 0 with both log lines. | llama.cpp/ornith-1.5-35b-a3b | +| v0/05-smoke-readme-deploy | 2026-09-25 | done | 1 | pass | none | `README.md` `## Run` uses `install -m` instead of `cp` and adds `systemctl daemon-reload` before `enable --now`, which is required for systemd to see the new unit; the task said only "copy … then enable --now". | llama.cpp/ornith-1.5-35b-a3b | ## Reviews + +### v0 review — 2026-09-25 (reviewer: claude, as owner for the night) + +Checked: five commits `73b2435..e436c62` with the trailer; every copied file byte-identical to +`docs/plans/v0/_files/`; no protected file touched (diff against the merge base is empty); +`make gate` → `gate: ok`; `make smoke` → `smoke: ok (stream spread 1003 ms)`. Probed from outside +with inputs the tests do not contain: encoded query strings pass through; a 3 MB JSON body is +forwarded, 17 MB → 413; `/_crossbar/hosts` answers while polls are in flight; `OpenCode-A` and +`/opencode-a/` → 404 as specified; HEAD and OPTIONS pass through; SIGTERM during a stream lets the +stream finish (6 SSE lines) and exits 0. + +Tally: 5 tasks, 5 first-run gate passes, 0 stops, wall time 6–10 min per task, unattended after the +restart. One model-side bug was caught by a given test during task 02 (`MarkDown` did not write the +entry back) and fixed before commit. + +| # | Finding | Severity | Fault | +|---|---|---|---| +| 1 | `statusRecorder.Flush` does `r.ResponseWriter.(http.Flusher).Flush()` — an unchecked assertion that panics on a writer that is not a Flusher. Rule "never panic" applied where the tests walked; task text said "forwarding to the underlying `http.Flusher`" without "if it implements it". | low | model + task | +| 2 | Two 502 answers in `ServeHTTP` (host name missing from config, `BaseURL` unparsable) that no task rule defined; config validation makes both unreachable. Harmless; the task should have said what to do. | low | task | +| 3 | Task 05 log row says `Deviations: none` while its Notes describe two (`install -m` instead of `cp`; `systemctl daemon-reload` added). Both changes are right; the row is not. | process | model | +| 4 | Task 01 attempt 0: given files under `docs/plans/v0/files/` were visible to `go vet ./...`. Fixed (`_files/`). | — | task | +| 5 | Task 04: `timeout --signal=TERM 3 …; echo $?` can never show `exit=0` (GNU `timeout` reports 124). Ornith verified another way and logged it. Fixed (`--preserve-status`). | — | task | +| 6 | My given `config_test.go` never covers a file that exists but cannot be read (`Load` on a 000-mode file). Gap in the acceptance suite, not in the code. | test | test | + +Follow-ups for a `v0.1` task: fix 1 (`if f, ok := …; ok { f.Flush() }`), add the unreadable-file +test for 6, and make the log-row rule in `AGENTS.md` say that anything the Notes describe as a +change belongs in Deviations (finding 3). + diff --git a/example.toml b/example.toml new file mode 100644 index 0000000..7b105cc --- /dev/null +++ b/example.toml @@ -0,0 +1,22 @@ +# crossbar example configuration. Replace and the addresses with your own. +listen = "127.0.0.1:17777" # never 0.0.0.0 — bind the tailnet address in production +poll_interval = "1s" # 60s in production; 1s makes the smoke run quick +queue_max = 8 + +[hosts.alpha] +base_url = "http://127.0.0.1:18081" # e.g. http://straylight.:11434 +weight = 1.0 +models = { "ornith-1.5-35b-a3b" = { parallel = 4 }, "small-9b" = { parallel = 6 } } + +[hosts.beta] +base_url = "http://127.0.0.1:18082" # e.g. http://titan.:8081 +weight = 2.0 +models = { "ornith-1.5-35b-a3b" = { parallel = 4 } } + +# v0: a route is a preference list; the first healthy host that has the model wins. +[routes.opencode-a] +hosts = ["alpha", "beta"] +default_model = "ornith-1.5-35b-a3b" + +[routes.hermes-x] +hosts = ["beta", "alpha"] diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..ab77dc4 --- /dev/null +++ b/go.mod @@ -0,0 +1,5 @@ +module git.wntrmute.dev/kyle/crossbar + +go 1.26 + +require github.com/BurntSushi/toml v1.6.0 diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..f74b269 --- /dev/null +++ b/go.sum @@ -0,0 +1,2 @@ +github.com/BurntSushi/toml v1.6.0 h1:dRaEfpa2VI55EwlIW72hMRHdWouJeRF7TPYhI+AUQjk= +github.com/BurntSushi/toml v1.6.0/go.mod h1:ukJfTF/6rtPPRCnwkur4qwRxa8vTRFBF0uk2lLoLwho= diff --git a/internal/admin/admin.go b/internal/admin/admin.go new file mode 100644 index 0000000..8bafc89 --- /dev/null +++ b/internal/admin/admin.go @@ -0,0 +1,102 @@ +// Package admin serves the operator's view of crossbar: the health table and the routes as JSON, +// mounted at /_crossbar/ on the same listener as the proxy. The shape of /_crossbar/hosts is +// fixed so operators can read why a request went where it went. +package admin + +import ( + "encoding/json" + "net/http" + "time" + + "git.wntrmute.dev/kyle/crossbar/internal/config" + "git.wntrmute.dev/kyle/crossbar/internal/health" +) + +// Hosts is what the admin handler needs from the health table. +type Hosts interface { + All() map[string]health.Status +} + +type HostView struct { + Healthy bool `json:"healthy"` + Loaded []string `json:"loaded"` // never null: an empty slice when nothing is loaded + LastOK string `json:"last_ok"` // time.RFC3339 in UTC, or "" if never + LastErr string `json:"last_err"` +} + +type RouteView struct { + Hosts []string `json:"hosts"` + DefaultModel string `json:"default_model"` +} + +// Handler serves GET /_crossbar/hosts and GET /_crossbar/routes. Any other method on those paths is +// a 405 with an Allow: GET header; anything else under the handler is a 404. +func Handler(cfg *config.Config, h Hosts) http.Handler { + mux := http.NewServeMux() + mux.HandleFunc("/_crossbar/hosts", hostsHandler(h)) + mux.HandleFunc("/_crossbar/routes", routesHandler(cfg)) + mux.HandleFunc("/", notFound) + return mux +} + +func hostsHandler(h Hosts) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet { + wrongMethod(w) + return + } + views := make(map[string]HostView, len(h.All())) + for name, s := range h.All() { + views[name] = hostView(s) + } + writeJSON(w, http.StatusOK, views) + } +} + +func routesHandler(cfg *config.Config) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet { + wrongMethod(w) + return + } + views := make(map[string]RouteView, len(cfg.Routes)) + for name, route := range cfg.Routes { + hosts := make([]string, len(route.Hosts)) + copy(hosts, route.Hosts) + views[name] = RouteView{Hosts: hosts, DefaultModel: route.DefaultModel} + } + writeJSON(w, http.StatusOK, views) + } +} + +func hostView(s health.Status) HostView { + loaded := s.Loaded + if loaded == nil { + loaded = []string{} + } + lastOK := "" + if !s.LastOK.IsZero() { + lastOK = s.LastOK.UTC().Format(time.RFC3339) + } + return HostView{ + Healthy: s.Healthy, + Loaded: loaded, + LastOK: lastOK, + LastErr: s.LastErr, + } +} + +func wrongMethod(w http.ResponseWriter) { + w.Header().Set("Allow", "GET") + writeJSON(w, http.StatusMethodNotAllowed, map[string]string{"error": "method not allowed"}) +} + +func notFound(w http.ResponseWriter, r *http.Request) { + writeJSON(w, http.StatusNotFound, map[string]string{"error": "not found"}) +} + +func writeJSON(w http.ResponseWriter, status int, v any) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(status) + _ = json.NewEncoder(w).Encode(v) +} diff --git a/internal/admin/admin_test.go b/internal/admin/admin_test.go new file mode 100644 index 0000000..80a0dfb --- /dev/null +++ b/internal/admin/admin_test.go @@ -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) + } + } +} diff --git a/internal/config/config.go b/internal/config/config.go new file mode 100644 index 0000000..1795333 --- /dev/null +++ b/internal/config/config.go @@ -0,0 +1,300 @@ +// Package config reads crossbar's TOML file: the hosts in front of which it +// proxies, the model each serves, and the ordered routes that select them. +// +// The reader is strict. A misspelt key or a route naming a host that does not +// exist is an *Error naming the offending field, returned at start-up rather +// than surfaced later. +package config + +import ( + "errors" + "fmt" + "io" + "net" + "net/url" + "os" + "regexp" + "sort" + "strings" + "time" + + "github.com/BurntSushi/toml" +) + +// Duration is a time.Duration that TOML reads from a string such as "60s" or +// "30m". +type Duration struct{ time.Duration } + +// UnmarshalText implements encoding.TextUnmarshaler via time.ParseDuration. +func (d *Duration) UnmarshalText(text []byte) error { + dt, err := time.ParseDuration(string(text)) + if err != nil { + return err + } + d.Duration = dt + return nil +} + +// Model is the per-model tuning carried by a host entry. +type Model struct { + Parallel int `toml:"parallel"` +} + +// Host names one upstream llama-server, the models it serves, and its weight. +type Host struct { + BaseURL string `toml:"base_url"` + Weight float64 `toml:"weight"` + Models map[string]Model `toml:"models"` +} + +// Route is an ordered list of hosts to try, with an optional default model. +type Route struct { + Hosts []string `toml:"hosts"` + DefaultModel string `toml:"default_model"` +} + +// Config is the whole file: what to listen on, tuning, hosts and routes. +type Config struct { + Listen string `toml:"listen"` + PollInterval Duration `toml:"poll_interval"` + QueueMax int `toml:"queue_max"` + Hosts map[string]Host `toml:"hosts"` + Routes map[string]Route `toml:"routes"` +} + +// Error is a validation error naming the field it is about. +type Error struct { + Field, Msg string +} + +// Error implements the error interface. +func (e *Error) Error() string { + return "config: " + e.Field + ": " + e.Msg +} + +const ( + DefaultPollInterval = 60 * time.Second + DefaultQueueMax = 8 + MinPollInterval = time.Second +) + +var routeName = regexp.MustCompile(`^[a-z0-9][a-z0-9-]*$`) + +// Load reads and parses the config file at path. An open failure is wrapped as +// "config: …", the same shape as a decode failure. +func Load(path string) (*Config, error) { + f, err := os.Open(path) + if err != nil { + return nil, fmt.Errorf("config: %w", err) + } + defer f.Close() + return Parse(f) +} + +// Parse decodes TOML from r, applies defaults, and validates. A TOML syntax +// error is returned wrapped as "config: …" and is not an *Error; an unknown +// key is an *Error naming the first undecoded key in sorted order. +func Parse(r io.Reader) (*Config, error) { + var c Config + md, err := toml.NewDecoder(r).Decode(&c) + if err != nil { + return nil, fmt.Errorf("config: %w", err) + } + if undecoded := md.Undecoded(); len(undecoded) > 0 { + keys := make([]string, 0, len(undecoded)) + for _, k := range undecoded { + keys = append(keys, strings.Join(k, ".")) + } + sort.Strings(keys) + return nil, &Error{Field: keys[0], Msg: "unknown key"} + } + + if c.PollInterval.Duration == 0 { + c.PollInterval.Duration = DefaultPollInterval + } + if c.QueueMax == 0 { + c.QueueMax = DefaultQueueMax + } + if e := c.validate(); e != nil { + return nil, e + } + return &c, nil +} + +// Serves reports whether the host exists and lists the model. +func (c *Config) Serves(host, model string) bool { + h, ok := c.Hosts[host] + if !ok { + return false + } + _, ok = h.Models[model] + return ok +} + +// IsError extracts an *Error from err, reporting whether one was present. +func IsError(err error) (*Error, bool) { + var e *Error + if errors.As(err, &e) { + return e, true + } + return nil, false +} + +// validate checks the config in a fixed order and writes defaults back into c. +// The first problem wins; every problem is an *Error with a precise field. +func (c *Config) validate() *Error { + if e := c.checkListen(); e != nil { + return e + } + if e := c.checkPoll(); e != nil { + return e + } + if e := c.checkQueue(); e != nil { + return e + } + if e := c.checkHosts(); e != nil { + return e + } + return c.checkRoutes() +} + +func (c *Config) checkListen() *Error { + if c.Listen == "" { + return &Error{Field: "listen", Msg: "required, host:port"} + } + host, _, err := net.SplitHostPort(c.Listen) + if err != nil { + return &Error{Field: "listen", Msg: "must be host:port"} + } + if host == "" { + return &Error{Field: "listen", Msg: "host part required"} + } + if host == "0.0.0.0" || host == "::" { + return &Error{Field: "listen", Msg: "not an unspecified address"} + } + return nil +} + +func (c *Config) checkPoll() *Error { + if c.PollInterval.Duration == 0 { + c.PollInterval.Duration = DefaultPollInterval + } else if c.PollInterval.Duration < MinPollInterval { + return &Error{Field: "poll_interval", Msg: "must be at least 1s"} + } + return nil +} + +func (c *Config) checkQueue() *Error { + if c.QueueMax == 0 { + c.QueueMax = DefaultQueueMax + } else if c.QueueMax < 0 { + return &Error{Field: "queue_max", Msg: "must not be negative"} + } + return nil +} + +func (c *Config) checkHosts() *Error { + if len(c.Hosts) == 0 { + return &Error{Field: "hosts", Msg: "at least one required"} + } + names := make([]string, 0, len(c.Hosts)) + for name := range c.Hosts { + names = append(names, name) + } + sort.Strings(names) + for _, name := range names { + h := c.Hosts[name] + + baseField := fmt.Sprintf("hosts.%s.base_url", name) + u, err := url.Parse(h.BaseURL) + if err != nil { + return &Error{Field: baseField, Msg: err.Error()} + } + if u.Scheme != "http" && u.Scheme != "https" { + return &Error{Field: baseField, Msg: "scheme must be http or https"} + } + if u.Host == "" { + return &Error{Field: baseField, Msg: "host required"} + } + if u.RawQuery != "" { + return &Error{Field: baseField, Msg: "query not allowed"} + } + if u.Fragment != "" { + return &Error{Field: baseField, Msg: "fragment not allowed"} + } + h.BaseURL = strings.TrimRight(h.BaseURL, "/") + + if h.Weight == 0 { + h.Weight = 1 + } else if h.Weight < 0 { + return &Error{Field: fmt.Sprintf("hosts.%s.weight", name), Msg: "must not be negative"} + } + + if len(h.Models) == 0 { + return &Error{Field: fmt.Sprintf("hosts.%s.models", name), Msg: "at least one required"} + } + models := make([]string, 0, len(h.Models)) + for m := range h.Models { + models = append(models, m) + } + sort.Strings(models) + for _, m := range models { + pm := h.Models[m] + if pm.Parallel == 0 { + pm.Parallel = 1 + } else if pm.Parallel < 0 { + return &Error{Field: fmt.Sprintf("hosts.%s.models.%s.parallel", name, m), Msg: "must not be negative"} + } + h.Models[m] = pm + } + c.Hosts[name] = h + } + return nil +} + +func (c *Config) checkRoutes() *Error { + if len(c.Routes) == 0 { + return &Error{Field: "routes", Msg: "at least one required"} + } + names := make([]string, 0, len(c.Routes)) + for name := range c.Routes { + names = append(names, name) + } + sort.Strings(names) + for _, name := range names { + r := c.Routes[name] + + if !routeName.MatchString(name) { + return &Error{Field: fmt.Sprintf("routes.%s", name), Msg: "must match [a-z0-9][a-z0-9-]*"} + } + + hostsField := fmt.Sprintf("routes.%s.hosts", name) + if len(r.Hosts) == 0 { + return &Error{Field: hostsField, Msg: "at least one required"} + } + seen := make(map[string]bool, len(r.Hosts)) + for _, h := range r.Hosts { + if seen[h] { + return &Error{Field: hostsField, Msg: "host listed twice"} + } + seen[h] = true + if _, ok := c.Hosts[h]; !ok { + return &Error{Field: hostsField, Msg: "unknown host"} + } + } + + if r.DefaultModel != "" { + served := false + for _, h := range r.Hosts { + if _, ok := c.Hosts[h].Models[r.DefaultModel]; ok { + served = true + break + } + } + if !served { + return &Error{Field: fmt.Sprintf("routes.%s.default_model", name), Msg: "not served by any host in route"} + } + } + } + return nil +} diff --git a/internal/config/config_test.go b/internal/config/config_test.go new file mode 100644 index 0000000..9069c36 --- /dev/null +++ b/internal/config/config_test.go @@ -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") + } +} diff --git a/internal/config/testdata/bad-default-model.toml b/internal/config/testdata/bad-default-model.toml new file mode 100644 index 0000000..4e1cc32 --- /dev/null +++ b/internal/config/testdata/bad-default-model.toml @@ -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" diff --git a/internal/config/testdata/bad-listen.toml b/internal/config/testdata/bad-listen.toml new file mode 100644 index 0000000..15e0e6d --- /dev/null +++ b/internal/config/testdata/bad-listen.toml @@ -0,0 +1,8 @@ +listen = "0.0.0.0:7777" + +[hosts.alpha] +base_url = "http://alpha.example:11434" +models = { "m" = { } } + +[routes.r] +hosts = ["alpha"] diff --git a/internal/config/testdata/bad-unknown-host.toml b/internal/config/testdata/bad-unknown-host.toml new file mode 100644 index 0000000..fd5aa92 --- /dev/null +++ b/internal/config/testdata/bad-unknown-host.toml @@ -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"] diff --git a/internal/config/testdata/bad-unknown-key.toml b/internal/config/testdata/bad-unknown-key.toml new file mode 100644 index 0000000..e3e94ec --- /dev/null +++ b/internal/config/testdata/bad-unknown-key.toml @@ -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"] diff --git a/internal/config/testdata/good.toml b/internal/config/testdata/good.toml new file mode 100644 index 0000000..92ad98e --- /dev/null +++ b/internal/config/testdata/good.toml @@ -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"] diff --git a/internal/health/health.go b/internal/health/health.go new file mode 100644 index 0000000..d483353 --- /dev/null +++ b/internal/health/health.go @@ -0,0 +1,258 @@ +// Package health tracks, for every host, whether it answered its last poll, which models it has +// loaded, and when it last answered. The proxy reads the table to choose a host and to record a +// request failure against a host. +package health + +import ( + "context" + "encoding/json" + "fmt" + "io" + "net/http" + "sort" + "sync" + "time" +) + +// RecoveryPolls is the number of good polls in a row a host needs after a failure before it is +// trusted again; a router loading a model flaps, so one good poll is not enough. +const RecoveryPolls = 2 + +// MaxModelsBody bounds how many bytes we read from either /health or /v1/models. +const MaxModelsBody = 1 << 20 + +// Status is a snapshot of one host's health, safe to copy. +type Status struct { + Healthy bool `json:"healthy"` + Loaded []string `json:"loaded"` // sorted, unique model ids from the last good poll + LastOK time.Time `json:"last_ok"` // zero if never + LastErr string `json:"last_err"` // "" after a good poll + Consecutive int `json:"consecutive"` // good polls in a row +} + +type entry struct { + status Status + everFailed bool +} + +type pollResult struct { + ok bool + cancelled bool + reason string + loaded []string +} + +// Table maps a host name to its health status. All methods are safe for concurrent use. +type Table struct { + mu sync.Mutex + baseURL map[string]string + entries map[string]entry + client *http.Client + interval time.Duration +} + +// New builds a table from hosts (name -> base URL, no trailing slash). A nil client becomes a +// five-second-timeout client. Every host starts untrusted with an empty Loaded slice. +func New(hosts map[string]string, interval time.Duration, client *http.Client) *Table { + if client == nil { + client = &http.Client{Timeout: 5 * time.Second} + } + t := &Table{ + baseURL: hosts, + entries: make(map[string]entry, len(hosts)), + client: client, + interval: interval, + } + for name := range hosts { + t.entries[name] = entry{status: Status{Loaded: []string{}}} + } + return t +} + +// Run polls once immediately, then every interval, and returns when ctx is done. +func (t *Table) Run(ctx context.Context) { + t.PollOnce(ctx) + ticker := time.NewTicker(t.interval) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + t.PollOnce(ctx) + } + } +} + +// PollOnce polls every host concurrently and returns when all have finished. +func (t *Table) PollOnce(ctx context.Context) { + t.mu.Lock() + names := make([]string, 0, len(t.baseURL)) + for name := range t.baseURL { + names = append(names, name) + } + t.mu.Unlock() + + var wg sync.WaitGroup + for _, name := range names { + wg.Add(1) + go func(name string) { + defer wg.Done() + t.pollHost(ctx, name) + }(name) + } + wg.Wait() +} + +// Get returns a copy of one host's status; ok is false for an unknown name. +func (t *Table) Get(name string) (Status, bool) { + t.mu.Lock() + defer t.mu.Unlock() + e, ok := t.entries[name] + if !ok { + return Status{}, false + } + return copyStatus(e.status), true +} + +// All returns copies of every host's status. +func (t *Table) All() map[string]Status { + t.mu.Lock() + defer t.mu.Unlock() + out := make(map[string]Status, len(t.entries)) + for name, e := range t.entries { + out[name] = copyStatus(e.status) + } + return out +} + +// MarkDown records a failure seen by the proxy. Loaded is left as last seen. +func (t *Table) MarkDown(name, reason string) { + t.mu.Lock() + defer t.mu.Unlock() + e, ok := t.entries[name] + if !ok { + return + } + e.everFailed = true + e.status.Healthy = false + e.status.Consecutive = 0 + e.status.LastErr = "marked down: " + reason + t.entries[name] = e +} + +func (t *Table) pollHost(ctx context.Context, name string) { + base := t.baseURL[name] + r := t.poll(ctx, base) + + t.mu.Lock() + defer t.mu.Unlock() + e := t.entries[name] + if r.cancelled { + return + } + if r.ok { + e.status.Consecutive++ + e.status.LastOK = time.Now() + e.status.LastErr = "" + e.status.Loaded = r.loaded + e.status.Healthy = !e.everFailed || e.status.Consecutive >= RecoveryPolls + } else { + e.everFailed = true + e.status.Healthy = false + e.status.Consecutive = 0 + e.status.LastErr = r.reason + } + t.entries[name] = e +} + +// poll runs the two requests for one host. A request that fails because ctx is done yields a +// cancelled result so the caller records nothing; any other failure yields a reason. +func (t *Table) poll(ctx context.Context, base string) pollResult { + loaded, r := t.check(ctx, base+"/health", "health") + if r.cancelled || r.reason != "" { + return r + } + loaded, r = t.check(ctx, base+"/v1/models", "models") + if r.cancelled { + return r + } + if r.reason != "" { + return r + } + return pollResult{ok: true, loaded: loaded} +} + +// check performs one GET and, on success, returns the decoded model ids. Health checks use the +// status code only; the models check also decodes the resident model ids. +func (t *Table) check(ctx context.Context, url, prefix string) ([]string, pollResult) { + req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil) + if err != nil { + return nil, t.fail(ctx, prefix, err) + } + resp, err := t.client.Do(req) + if err != nil { + return nil, t.fail(ctx, prefix, err) + } + + if prefix == "health" { + // We only need the status code; bound and drain the body. + func() { + defer resp.Body.Close() + _, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, MaxModelsBody)) + }() + if resp.StatusCode != http.StatusOK { + return nil, pollResult{reason: fmt.Sprintf("health: HTTP %d", resp.StatusCode)} + } + return nil, pollResult{} + } + + // models: bound and decode the body. + if resp.StatusCode != http.StatusOK { + defer resp.Body.Close() + return nil, pollResult{reason: fmt.Sprintf("models: HTTP %d", resp.StatusCode)} + } + var m struct { + Data []struct { + ID string `json:"id"` + } `json:"data"` + } + if err := json.NewDecoder(io.LimitReader(resp.Body, MaxModelsBody)).Decode(&m); err != nil { + defer resp.Body.Close() + return nil, t.fail(ctx, "models", err) + } + defer resp.Body.Close() + + loaded := make([]string, 0, len(m.Data)) + seen := make(map[string]struct{}, len(m.Data)) + for _, d := range m.Data { + if d.ID == "" { + continue + } + if _, ok := seen[d.ID]; ok { + continue + } + seen[d.ID] = struct{}{} + loaded = append(loaded, d.ID) + } + sort.Strings(loaded) + return loaded, pollResult{} +} + +// fail builds a failure result, treating a done ctx as a cancellation rather than a failure. +func (t *Table) fail(ctx context.Context, prefix string, err error) pollResult { + if ctx.Err() != nil { + return pollResult{cancelled: true} + } + return pollResult{reason: prefix + ": " + err.Error()} +} + +func copyStatus(s Status) Status { + if s.Loaded == nil { + return s + } + out := s + out.Loaded = make([]string, len(s.Loaded)) + copy(out.Loaded, s.Loaded) + return out +} diff --git a/internal/health/health_test.go b/internal/health/health_test.go new file mode 100644 index 0000000..3f3d53f --- /dev/null +++ b/internal/health/health_test.go @@ -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) + } +} 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)) + } +} diff --git a/scripts/check-lines.sh b/scripts/check-lines.sh new file mode 100644 index 0000000..c8e113e --- /dev/null +++ b/scripts/check-lines.sh @@ -0,0 +1,14 @@ +#!/bin/sh +# No Go source file over 400 lines. Reports every offender, then fails. +set -u +limit=400 +bad=0 +for f in $(find . -name '*.go' -not -path './.git/*' -not -path './vendor/*'); do + n=$(wc -l < "$f") || { echo "check-lines: cannot read $f" >&2; exit 2; } + if [ "$n" -gt "$limit" ]; then + echo "check-lines: $f has $n lines (limit $limit)" + bad=1 + fi +done +[ "$bad" -eq 0 ] || exit 1 +echo "check-lines: ok" diff --git a/tools/smoke.sh b/tools/smoke.sh new file mode 100755 index 0000000..4b21792 --- /dev/null +++ b/tools/smoke.sh @@ -0,0 +1,45 @@ +#!/bin/sh +# Smoke run: two fake upstreams, one crossbar, real HTTP. Prints "smoke: ok" or fails. +# Needs: bin/crossbar and bin/fakeupstream (make build), curl. +set -eu +cd "$(dirname "$0")/.." +tmp=$(mktemp -d); trap 'kill $pids 2>/dev/null; rm -rf "$tmp"' EXIT INT TERM +pids="" +bin/fakeupstream -listen 127.0.0.1:18081 -name alpha -models ornith-1.5-35b-a3b,small-9b -down-file "$tmp/alpha.down" >"$tmp/alpha.log" 2>&1 & pids="$pids $!" +bin/fakeupstream -listen 127.0.0.1:18082 -name beta -models ornith-1.5-35b-a3b -down-file "$tmp/beta.down" >"$tmp/beta.log" 2>&1 & pids="$pids $!" +bin/crossbar -config example.toml >"$tmp/crossbar.log" 2>&1 & pids="$pids $!" +sleep 1.5 +fail() { echo "smoke: FAIL: $*" >&2; echo "--- crossbar.log"; cat "$tmp/crossbar.log"; exit 1; } +base=http://127.0.0.1:17777 + +h=$(curl -s -o /dev/null -w '%{http_code} %header{X-Crossbar-Host}' "$base/opencode-a/v1/models") +[ "$h" = "200 alpha" ] || fail "opencode-a should go to alpha, got '$h'" +h=$(curl -s -o /dev/null -w '%{http_code} %header{X-Crossbar-Host}' "$base/hermes-x/v1/models") +[ "$h" = "200 beta" ] || fail "hermes-x should go to beta, got '$h'" +h=$(curl -s -o /dev/null -w '%{http_code}' "$base/nope/v1/models") +[ "$h" = "404" ] || fail "unknown route should be 404, got '$h'" + +touch "$tmp/alpha.down"; sleep 2.5 # poll_interval is 1s in example.toml +h=$(curl -s -o /dev/null -w '%{http_code} %header{X-Crossbar-Host}' "$base/opencode-a/v1/models") +[ "$h" = "200 beta" ] || fail "with alpha down, opencode-a should fail over to beta, got '$h'" +curl -s "$base/_crossbar/hosts" | grep -q '"alpha":{"healthy":false' || fail "/_crossbar/hosts does not show alpha unhealthy: $(curl -s $base/_crossbar/hosts)" + +rm "$tmp/alpha.down"; sleep 3.5 # recovery needs two good polls +h=$(curl -s -o /dev/null -w '%header{X-Crossbar-Host}' "$base/opencode-a/v1/models") +[ "$h" = "alpha" ] || fail "alpha should be back after two good polls, got '$h'" + +# Streaming: five chunks 200 ms apart must arrive over >= 0.6 s, not all at once at the end. +start=$(date +%s%N) +first="" +curl -sN -X POST -H 'Content-Type: application/json' -d '{"model":"ornith-1.5-35b-a3b","stream":true,"messages":[]}' \ + "$base/opencode-a/v1/chat/completions" | while IFS= read -r line; do + [ -n "$line" ] || continue + now=$(date +%s%N); echo "$(( (now - start) / 1000000 )) $line" + done > "$tmp/stream.txt" +firstms=$(head -1 "$tmp/stream.txt" | cut -d' ' -f1); lastms=$(tail -1 "$tmp/stream.txt" | cut -d' ' -f1) +[ -n "$firstms" ] && [ "$((lastms - firstms))" -ge 600 ] || fail "stream arrived in one burst (first ${firstms:-?} ms, last ${lastms:-?} ms): +$(cat "$tmp/stream.txt")" +grep -q 'DONE' "$tmp/stream.txt" || fail "stream did not end with [DONE]" + +grep -q 'route=opencode-a host=alpha' "$tmp/crossbar.log" || fail "no request log line" +echo "smoke: ok (stream spread $((lastms - firstms)) ms)"