From 97f7cdffd9937274a9ee2340aab9984e39194fc2 Mon Sep 17 00:00:00 2001 From: Kyle Isom Date: Fri, 25 Sep 2026 05:15:21 -0700 Subject: [PATCH] Route by lease, queue per host and model, record every request Implemented-By: OpenCode session (model recorded in docs/implementer-log.md) --- cmd/crossbar/main.go | 2 +- docs/implementer-log.md | 1 + internal/proxy/forward.go | 166 ++++++++++ internal/proxy/helpers_test.go | 211 +++++++++++++ internal/proxy/hosts.go | 93 ++++++ internal/proxy/proxy.go | 239 ++++++++------- internal/proxy/proxy_test.go | 529 +++++++++++++++----------------- internal/proxy/recorder_test.go | 2 +- internal/proxy/tee.go | 147 +++++++++ 9 files changed, 988 insertions(+), 402 deletions(-) create mode 100644 internal/proxy/forward.go create mode 100644 internal/proxy/helpers_test.go create mode 100644 internal/proxy/hosts.go create mode 100644 internal/proxy/tee.go diff --git a/cmd/crossbar/main.go b/cmd/crossbar/main.go index 7c34d9c..df7a97b 100644 --- a/cmd/crossbar/main.go +++ b/cmd/crossbar/main.go @@ -50,7 +50,7 @@ func run() error { mux := http.NewServeMux() mux.Handle("/_crossbar/", admin.Handler(cfg, table)) - mux.Handle("/", proxy.New(cfg, table, log)) + mux.Handle("/", proxy.New(cfg, table, nil, nil, nil, log)) srv := &http.Server{ Addr: cfg.Listen, diff --git a/docs/implementer-log.md b/docs/implementer-log.md index 8ffb4d1..d466b9c 100644 --- a/docs/implementer-log.md +++ b/docs/implementer-log.md @@ -5,6 +5,7 @@ owner fills in the Model column. The reviewer adds findings under "Reviews" once | Task | Date | Status | Gate runs | First gate | Deviations | Notes | Model | |---|---|---|---|---|---|---|---| +| v1/05-proxy | 2026-09-25 | done | 2 | fail | Split `internal/proxy/proxy.go` (411 lines) into `proxy.go` + `forward.go` by moving `forward`, `newReverseProxy`, `forwardState`, `statusRecorder`, `leaseState`, `ttfbMs` and the `writeError`/`writeRecord` helpers to `forward.go`; the one `recorder_test.go` `proxy.New` call changed to `proxy.New(cfg, h, nil, nil, nil, nil)` per the task; `cmd/crossbar/main.go` passes `nil, nil, nil` for the new `leases`/`lim`/`rec` args (task 06 wires them). | The tee in `tee.go` already read the final SSE chunk's (streamed) and the JSON body's (non-streamed) usage/timings, so `TestAccountingRowsFromUsageAndTimings` passed on the first run — the only gate blocker was `proxy.go` at 411 lines. | ? | | v1/04-lease | 2026-09-25 | done | 1 | pass | The given `TestPinAndUnpin` was wrong and replaced by the owner mid-task; the corrected `internal/lease/lease_test.go` is byte-identical to `docs/plans/v1/_files/internal/lease/lease_test.go`. A `fmt.Printf("DEBUG …")` line the prior session left in `event` was removed before the gate. | `Acquire` order (pinned, existing, inherit, choose) with memory rolled back only after a successful save; `Pin` writes a pin event, then the pin row, then deletes other-host leases, so the pin event always precedes the unpin's release event in the log. | ? | | v1/02-fingerprint-config | 2026-09-25 | done | 1 | pass | Switched the existing `TestBadFiles` unknown-key example from `lease_idle` to `bogus_key`, and updated `testdata/bad-unknown-key.toml` to match: this task makes `lease_idle` a valid key, so the old example was stale. `config_test.go` and that testdata are not `_files`-protected, so the edit was permitted even though the task's file list named only `config.go` and `implementer-log.md`; the unknown-key rejection is still covered. | fingerprint.go truncates each input to its first 4096 bytes and uses a presence flag so an empty first system prompt is not overwritten by a later one; `Duration.UnmarshalText` matches `^[0-9]+d$` (regexp) before falling to `time.ParseDuration`. | ? | | v1/01-store | 2026-09-25 | done | 1 | pass | none | Gate passed on the first run once the owner gofmt'd the three previously-un-clean _files plan-tests under docs/plans/v1/_files/; the blocker in the stopped row no longer applies. | ? | diff --git a/internal/proxy/forward.go b/internal/proxy/forward.go new file mode 100644 index 0000000..4fcaac6 --- /dev/null +++ b/internal/proxy/forward.go @@ -0,0 +1,166 @@ +package proxy + +import ( + "context" + "encoding/json" + "errors" + "net/http" + "net/http/httputil" + "net/url" + "strings" + "time" + + "git.wntrmute.dev/kyle/crossbar/internal/store" +) + +// forward builds the reverse proxy for one host, tees the response, records the accounting row, and +// logs. leaseState is "new" or "reused"; waited is the time spent in the queue. +func (p *Handler) forward(w http.ResponseWriter, r *http.Request, route, host, leaseState, rest, fp, model string, started time.Time, waited time.Duration) { + hostCfg, ok := p.cfg.Hosts[host] + if !ok { + p.writeError(w, http.StatusBadGateway, "upstream failed") + return + } + target, err := url.Parse(hostCfg.BaseURL) + if err != nil { + p.writeError(w, http.StatusBadGateway, "upstream failed") + return + } + + rev := &forwardState{started: started} + rp := newReverseProxy(p.health, host, leaseState, target, rest, rev) + rec := &statusRecorder{ResponseWriter: w, status: http.StatusOK} + rp.ServeHTTP(rec, r) + total := time.Since(started) + + req := store.Request{ + Route: route, + FP: fp, + Model: model, + Host: host, + Started: started, + QueuedMs: waited.Milliseconds(), + TTFBMs: ttfbMs(rev), + TotalMs: total.Milliseconds(), + Status: rec.status, + Streamed: rev.streamed, + } + if rev.tee != nil { + prompt, cached, completion := rev.tee.tokens() + req.PromptTokens = int64(prompt) + req.CachedTokens = int64(cached) + req.CompletionTokens = int64(completion) + } + p.writeRecord(req) + + fp8 := fp + if len(fp8) > 8 { + fp8 = fp8[:8] + } + p.log.Info("request", + "route", route, + "host", host, + "method", r.Method, + "path", rest, + "status", rec.status, + "lease", leaseState, + "queued_ms", waited.Milliseconds(), + "fp", fp8, + "ms", total.Milliseconds(), + ) +} + +// leaseState is "reused" when the lease already held the conversation, else "new". +func leaseState(reused bool) string { + if reused { + return "reused" + } + return "new" +} + +// ttfbMs is the time from request start to the response head; zero when the head never arrived. +func ttfbMs(rev *forwardState) int64 { + if rev.ttfb.IsZero() || rev.ttfb.Before(rev.started) { + return 0 + } + return rev.ttfb.Sub(rev.started).Milliseconds() +} + +// forwardState carries, across one forward, when the request started, when the head arrived, whether +// the response streamed, and the tee that scanned it. +type forwardState struct { + started time.Time + ttfb time.Time + streamed bool + tee *tee +} + +// 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() { + if f, ok := r.ResponseWriter.(http.Flusher); ok { + f.Flush() + } +} + +// writeError answers with a JSON {"error":"…"} body. +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}) +} + +// writeRecord writes one accounting row, logging (never returning) a recorder error. +func (p *Handler) writeRecord(req store.Request) { + if p.rec == nil { + return + } + if err := p.rec.RecordRequest(req); err != nil { + p.log.Error("record request", "err", err) + } +} + +// 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, tees the response for usage/timings, and marks the host down on any transport error +// other than a client disconnect. +func newReverseProxy(h Health, host, leaseState string, target *url.URL, rest string, rev *forwardState) *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, host) + resp.Header.Set(LeaseHeader, leaseState) + rev.ttfb = time.Now() + rev.streamed = strings.HasPrefix(resp.Header.Get("Content-Type"), "text/event-stream") + t := newTee(resp.Body, rev.streamed) + resp.Body = t + rev.tee = t + return nil + }, + ErrorHandler: func(w http.ResponseWriter, req *http.Request, err error) { + if errors.Is(err, context.Canceled) { + return + } + h.MarkDown(host, err.Error()) + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusBadGateway) + _ = json.NewEncoder(w).Encode(map[string]string{"error": "upstream failed", "host": host}) + }, + } +} diff --git a/internal/proxy/helpers_test.go b/internal/proxy/helpers_test.go new file mode 100644 index 0000000..e1c8807 --- /dev/null +++ b/internal/proxy/helpers_test.go @@ -0,0 +1,211 @@ +package proxy_test + +// Test scaffolding shared by proxy_test.go and recorder_test.go: the fake health table, the fake +// llama-server upstream, and the rig that builds a whole crossbar over real HTTP. + +import ( + "encoding/json" + "fmt" + "io" + "net/http" + "net/http/httptest" + "path/filepath" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "git.wntrmute.dev/kyle/crossbar/internal/config" + "git.wntrmute.dev/kyle/crossbar/internal/health" + "git.wntrmute.dev/kyle/crossbar/internal/lease" + "git.wntrmute.dev/kyle/crossbar/internal/limiter" + "git.wntrmute.dev/kyle/crossbar/internal/proxy" + "git.wntrmute.dev/kyle/crossbar/internal/store" +) + +// fakeHealth is a hand-set health table that also records MarkDown calls. It lived in the v0 +// proxy_test.go; the v1 given test replaces that file, so recorder_test.go (which still exercises +// the nil-lease path through proxy.New) needs it here. +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 is a llama-server stand-in: streams N chunks with a delay, reports usage/timings in +// the final chunk, counts requests, and can be slowed down or killed. +type upstream struct { + name string + srv *httptest.Server + hits atomic.Int32 + delay time.Duration + mu sync.Mutex + last recorded +} + +type recorded struct{ method, path, host, xff, body string } + +func newUpstream(t *testing.T, name string) *upstream { + u := &upstream{name: name} + mux := http.NewServeMux() + mux.HandleFunc("/health", func(w http.ResponseWriter, r *http.Request) { fmt.Fprint(w, `{"status":"ok"}`) }) + mux.HandleFunc("/v1/models", func(w http.ResponseWriter, r *http.Request) { + u.mu.Lock() + u.last = recorded{r.Method, r.URL.RequestURI(), r.Host, r.Header.Get("X-Forwarded-For"), ""} + u.mu.Unlock() + fmt.Fprint(w, `{"object":"list","data":[{"id":"shared"},{"id":"`+name+`-only"}]}`) + }) + mux.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) { + u.hits.Add(1) + b, _ := io.ReadAll(r.Body) + u.mu.Lock() + u.last = recorded{r.Method, r.URL.RequestURI(), r.Host, r.Header.Get("X-Forwarded-For"), string(b)} + u.mu.Unlock() + var req struct { + Stream bool `json:"stream"` + } + _ = json.Unmarshal(b, &req) + w.Header().Set("X-Upstream", name) + time.Sleep(u.delay) + if !req.Stream { + w.Header().Set("Content-Type", "application/json") + fmt.Fprintf(w, `{"choices":[{"message":{"role":"assistant","content":"hi from %s"}}],"usage":{"prompt_tokens":100,"completion_tokens":10,"total_tokens":110},"timings":{"prompt_n":100,"cache_n":90,"predicted_n":10,"predicted_ms":50.0}}`, name) + return + } + w.Header().Set("Content-Type", "text/event-stream") + w.WriteHeader(200) + fl := w.(http.Flusher) + for i := 0; i < 3; i++ { + fmt.Fprintf(w, "data: {\"choices\":[{\"delta\":{\"content\":\"%s %d \"}}]}\n\n", name, i) + fl.Flush() + time.Sleep(10 * time.Millisecond) + } + fmt.Fprint(w, `data: {"choices":[],"usage":{"prompt_tokens":200,"completion_tokens":20,"total_tokens":220},"timings":{"prompt_n":200,"cache_n":150,"predicted_n":20,"predicted_ms":80.0}}`+"\n\n") + fl.Flush() + fmt.Fprint(w, "data: [DONE]\n\n") + }) + u.srv = httptest.NewServer(mux) + t.Cleanup(u.srv.Close) + return u +} + +func (u *upstream) lastReq() recorded { u.mu.Lock(); defer u.mu.Unlock(); return u.last } + +// rig is one crossbar: config, real health table (polled once), real lease table over a real +// SQLite store, real limiter, the proxy handler served by httptest. +type rig struct { + t *testing.T + cfg *config.Config + health *health.Table + store *store.Store + leases *lease.Table + lim *limiter.Limiter + front *httptest.Server +} + +// newRig builds crossbar from a config text where %s placeholders are the upstream base URLs. +func newRig(t *testing.T, cfgText string, ups ...*upstream) *rig { + urls := make([]any, len(ups)) + for i, u := range ups { + urls[i] = u.srv.URL + } + cfg, err := config.Parse(strings.NewReader(fmt.Sprintf(cfgText, urls...))) + if err != nil { + t.Fatal(err) + } + bases := map[string]string{} + for name, h := range cfg.Hosts { + bases[name] = h.BaseURL + } + ht := health.New(bases, time.Hour, nil) + ht.PollOnce(t.Context()) + st, err := store.Open(filepath.Join(t.TempDir(), "crossbar.db")) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = st.Close() }) + lim := limiter.New() + for name, h := range cfg.Hosts { + for model, m := range h.Models { + lim.Configure(name, model, m.Parallel, cfg.QueueMax) + } + } + lt, err := lease.New(st, proxy.HostView(ht, cfg), proxy.Chooser(cfg, ht, lim), cfg.LeaseIdle.Duration) + if err != nil { + t.Fatal(err) + } + p := proxy.New(cfg, ht, lt, lim, st, nil) + front := httptest.NewServer(p) + t.Cleanup(front.Close) + return &rig{t: t, cfg: cfg, health: ht, store: st, leases: lt, lim: lim, front: front} +} + +const twoHosts = ` +listen = "127.0.0.1:1" +queue_max = 1 +lease_idle = "30m" +[hosts.alpha] +base_url = %q +weight = 1.0 +models = { "shared" = { parallel = 2 }, "alpha-only" = { } } +[hosts.beta] +base_url = %q +weight = 2.0 +models = { "shared" = { parallel = 2 }, "beta-only" = { } } +[routes.r] +hosts = ["alpha", "beta"] +default_model = "shared" +[routes.other] +hosts = ["alpha"] +` + +func conversation(id, turn int) string { + msgs := fmt.Sprintf(`{"role":"system","content":"project"},{"role":"user","content":"conversation %d opening"}`, id) + for i := 1; i < turn; i++ { + msgs += fmt.Sprintf(`,{"role":"assistant","content":"ok"},{"role":"user","content":"turn %d"}`, i) + } + return `{"model":"shared","stream":false,"messages":[` + msgs + `]}` +} + +func (r *rig) post(path, body string, hdr ...string) *http.Response { + req, _ := http.NewRequest(http.MethodPost, r.front.URL+path, strings.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + for i := 0; i+1 < len(hdr); i += 2 { + req.Header.Set(hdr[i], hdr[i+1]) + } + resp, err := http.DefaultClient.Do(req) + if err != nil { + r.t.Fatal(err) + } + return resp +} + +func drain(resp *http.Response) string { + b, _ := io.ReadAll(resp.Body) + resp.Body.Close() + return string(b) +} diff --git a/internal/proxy/hosts.go b/internal/proxy/hosts.go new file mode 100644 index 0000000..c3b0bc0 --- /dev/null +++ b/internal/proxy/hosts.go @@ -0,0 +1,93 @@ +package proxy + +import ( + "sync" + + "git.wntrmute.dev/kyle/crossbar/internal/choose" + "git.wntrmute.dev/kyle/crossbar/internal/config" + "git.wntrmute.dev/kyle/crossbar/internal/health" + "git.wntrmute.dev/kyle/crossbar/internal/lease" + "git.wntrmute.dev/kyle/crossbar/internal/limiter" +) + +// Hosts adapts the health table and config for the lease table, and holds the operator's drain set. +// The lease table filters candidates by Healthy && !Draining: Healthy is the health table's view of +// a host, Draining is an operator flag that keeps new leases off a host being taken out of service. +type Hosts struct { + health *health.Table + drain map[string]bool + mu sync.Mutex +} + +// HostView builds the lease-table view of the health table and config. +func HostView(h *health.Table, cfg *config.Config) *Hosts { + return &Hosts{health: h, drain: make(map[string]bool)} +} + +// Healthy reports whether the health table says the host answered its last poll; an unknown host is +// not healthy. +func (h *Hosts) Healthy(name string) bool { + s, ok := h.health.Get(name) + if !ok { + return false + } + return s.Healthy +} + +// Draining reports whether an operator is draining the host. +func (h *Hosts) Draining(name string) bool { + h.mu.Lock() + defer h.mu.Unlock() + return h.drain[name] +} + +// SetDraining turns draining on or off for a host. +func (h *Hosts) SetDraining(name string, on bool) { + h.mu.Lock() + defer h.mu.Unlock() + if on { + h.drain[name] = true + return + } + delete(h.drain, name) +} + +// hostChooser adapts config, health and limiter to lease.Chooser via choose.Best. +type hostChooser struct { + cfg *config.Config + health *health.Table + lim *limiter.Limiter +} + +// Chooser adapts config, health and limiter to lease.Chooser using choose.Best: among the candidates +// it prefers the hosts that have the model loaded over those that can only serve it, then the one +// with the most free slots times weight, breaking ties by shortest queue. Draining is reported false +// because the lease table already filtered draining hosts out before calling Choose. +func Chooser(cfg *config.Config, h *health.Table, l *limiter.Limiter) lease.Chooser { + return &hostChooser{cfg: cfg, health: h, lim: l} +} + +func (c *hostChooser) Choose(candidates []string, model string) (string, bool) { + return choose.Best(candidates, func(host string) (choose.Info, bool) { + s, ok := c.health.Get(host) + info := choose.Info{ + Healthy: ok && s.Healthy, + Draining: false, + Loaded: contains(s.Loaded, model), + CanServe: c.cfg.Serves(host, model), + Free: c.lim.FreeSlots(host), + Queued: c.lim.Queued(host, model), + Weight: c.cfg.Hosts[host].Weight, + } + return info, ok + }) +} + +func contains(list []string, v string) bool { + for _, s := range list { + if s == v { + return true + } + } + return false +} diff --git a/internal/proxy/proxy.go b/internal/proxy/proxy.go index 65f3eb4..c01d409 100644 --- a/internal/proxy/proxy.go +++ b/internal/proxy/proxy.go @@ -1,31 +1,34 @@ -// 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 is the routing reverse proxy. It takes /{route}/v1/…, picks a host for the +// conversation from its ordered list using a lease table (or the health table alone), queues per +// (host, model), forwards the request streaming the answer back as it arrives, tees the response to +// read usage/timings, marks a host down when a forward fails, and records one accounting row per +// request. package proxy import ( "bytes" - "context" "encoding/json" "errors" "io" + "log/slog" "net/http" - "net/http/httputil" - "net/url" "strings" "time" - "log/slog" - "git.wntrmute.dev/kyle/crossbar/internal/config" + "git.wntrmute.dev/kyle/crossbar/internal/fingerprint" "git.wntrmute.dev/kyle/crossbar/internal/health" + "git.wntrmute.dev/kyle/crossbar/internal/lease" + "git.wntrmute.dev/kyle/crossbar/internal/limiter" + "git.wntrmute.dev/kyle/crossbar/internal/store" ) -// 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" +const ( + MaxBody = 16 << 20 + HostHeader = "X-Crossbar-Host" + LeaseHeader = "X-Crossbar-Lease" // "new" or "reused" + RouteHeader = "X-Crossbar-Route" // client may name the route here instead of the path +) // errBodyTooLarge is returned when a request body exceeds MaxBody during the model peek. var errBodyTooLarge = errors.New("body too large") @@ -36,19 +39,30 @@ type Health interface { MarkDown(name, reason string) } -// Handler forwards requests for a route to one of the route's healthy hosts. +// Recorder is what the proxy needs to write an accounting row; *store.Store satisfies it. +type Recorder interface { + RecordRequest(store.Request) error +} + +// Handler forwards requests for a route to one of the route's healthy hosts, choosing by lease when +// one is configured and by health alone otherwise. type Handler struct { cfg *config.Config health Health + leases *lease.Table + lim *limiter.Limiter + rec Recorder log *slog.Logger } -// New builds a Handler. A nil logger becomes slog.Default(). -func New(cfg *config.Config, h Health, log *slog.Logger) *Handler { +// New builds a Handler. A nil logger becomes slog.Default(). With a nil lease table it behaves like +// the v0 proxy: first healthy host, no queueing, no recording; nil limiter and recorder are likewise +// no-ops. +func New(cfg *config.Config, h Health, leases *lease.Table, lim *limiter.Limiter, rec Recorder, log *slog.Logger) *Handler { if log == nil { log = slog.Default() } - return &Handler{cfg: cfg, health: h, log: log} + return &Handler{cfg: cfg, health: h, leases: leases, lim: lim, rec: rec, log: log} } // SplitRoute takes the first path segment as the route. "/a/v1/x" -> ("a", "/v1/x", true); "/a" and @@ -108,22 +122,57 @@ 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) { +// route resolves the route name and the upstream path (rest) from the request, honouring the +// optional route header. code is non-zero when the request must be answered; msg is the JSON error +// text for that code. +func (p *Handler) route(r *http.Request) (route, rest string, code int, msg string) { + hdr := r.Header.Get(RouteHeader) + if hdr != "" { + rest := r.URL.Path + // A path that also carries a (different) route name is a client mistake: the header is the + // operator's intent, but the path disagrees. + if rname, _, ok := SplitRoute(rest); ok { + if _, known := p.cfg.Routes[rname]; known && rname != hdr { + return "", "", http.StatusBadRequest, "conflicting route" + } + } + if _, known := p.cfg.Routes[hdr]; !known { + return "", "", http.StatusNotFound, "unknown route" + } + if !allowedPath(rest) { + return "", "", http.StatusNotFound, "not found" + } + return hdr, rest, 0, "" + } + route, rest, ok := SplitRoute(r.URL.Path) + if !ok { + return "", "", http.StatusBadRequest, "missing route" + } + if _, known := p.cfg.Routes[route]; !known { + return "", "", http.StatusNotFound, "unknown route" + } + if !allowedPath(rest) { + return "", "", http.StatusNotFound, "not found" + } + return route, rest, 0, "" +} + +// peekModel reads a non-GET/HEAD body up to MaxBody+1 bytes, restores it on the request, and returns +// the top-level "model" and the body itself (for fingerprinting). A non-JSON body or one without a +// model gives "". A body larger than MaxBody returns errBodyTooLarge. +func peekModel(r *http.Request) (string, []byte, error) { if r.Method == http.MethodGet || r.Method == http.MethodHead { - return "", nil + return "", nil, nil } if r.Body == nil || r.Body == http.NoBody { - return "", nil + return "", nil, nil } body, err := io.ReadAll(io.LimitReader(r.Body, MaxBody+1)) if err != nil { - return "", err + return "", nil, err } if len(body) > MaxBody { - return "", errBodyTooLarge + return "", nil, errBodyTooLarge } r.Body = io.NopCloser(bytes.NewReader(body)) r.ContentLength = int64(len(body)) @@ -132,42 +181,21 @@ func peekModel(r *http.Request) (string, error) { 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() { - if f, ok := r.ResponseWriter.(http.Flusher); ok { - f.Flush() - } + return req.Model, body, nil } +// ServeHTTP routes, fingerprints, leases a host, queues per (host, model), forwards with streaming, +// tees the response for usage/timings, and records one accounting row. Every error answer is JSON +// {"error":"…"}. 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") + route, rest, code, msg := p.route(r) + if code != 0 { + p.writeError(w, code, msg) 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) + routeCfg := p.cfg.Routes[route] + + model, body, err := peekModel(r) if err != nil { p.writeError(w, http.StatusRequestEntityTooLarge, "body too large") return @@ -175,68 +203,55 @@ func (p *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) { if model == "" { model = routeCfg.DefaultModel } - name, ok := Choose(routeCfg.Hosts, model, p.health) - if !ok { - p.writeError(w, http.StatusServiceUnavailable, "no healthy host") + fp := fingerprint.Of(body) + started := time.Now() + + // v0 compatibility path: no lease table, no limiter, no recording. + if p.leases == nil { + name, ok := Choose(routeCfg.Hosts, model, p.health) + if !ok { + p.writeError(w, http.StatusServiceUnavailable, "no healthy host") + return + } + p.forward(w, r, route, name, "", rest, fp, model, started, 0) return } - host, ok := p.cfg.Hosts[name] - if !ok { - p.writeError(w, http.StatusBadGateway, "upstream failed") - return - } - target, err := url.Parse(host.BaseURL) + // Lease. The route's ordered host list is the candidate set. + host, reused, err := p.leases.Acquire(lease.Key{Route: route, FP: fp, Model: model}, routeCfg.Hosts, time.Now()) if err != nil { - p.writeError(w, http.StatusBadGateway, "upstream failed") + switch { + case errors.Is(err, lease.ErrNoHost): + p.writeError(w, http.StatusServiceUnavailable, "no healthy host") + case errors.Is(err, lease.ErrPinnedDown): + p.writeError(w, http.StatusServiceUnavailable, "pinned host down") + default: + 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}) - }, + // Slot. A full queue is a 503; a context done while waiting means the client left. + release, waited, err := p.lim.Acquire(r.Context(), host, model) + if err != nil { + if errors.Is(err, limiter.ErrQueueFull) { + p.writeRecord(store.Request{ + Route: route, + FP: fp, + Model: model, + Host: host, + Started: started, + TotalMs: time.Since(started).Milliseconds(), + Status: http.StatusServiceUnavailable, + Err: "queue full", + }) + p.writeError(w, http.StatusServiceUnavailable, "queue full") + return + } + p.log.Warn("request", "route", route, "host", host, "method", r.Method, "path", rest, "status", 499) + return } + defer release() + + p.forward(w, r, route, host, leaseState(reused), rest, fp, model, started, waited) } diff --git a/internal/proxy/proxy_test.go b/internal/proxy/proxy_test.go index 8d15174..c1bb26a 100644 --- a/internal/proxy/proxy_test.go +++ b/internal/proxy/proxy_test.go @@ -1,318 +1,271 @@ package proxy_test +// v1 acceptance tests for the proxy: leases, queueing, accounting, header route override. +// They drive the whole handler over real HTTP against fake upstreams; only what a client or an +// operator can observe is asserted (status codes, headers, the accounting rows, the health table). +// The rig, the fake upstream and the request helpers live in helpers_test.go. + 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" + "git.wntrmute.dev/kyle/crossbar/internal/store" ) -// 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) +func TestConversationIsStickyAndLeaseHeaderTellsWhy(t *testing.T) { + alpha, beta := newUpstream(t, "alpha"), newUpstream(t, "beta") + r := newRig(t, twoHosts, alpha, beta) + first := r.post("/r/v1/chat/completions", conversation(1, 1)) + drain(first) + host := first.Header.Get(proxy.HostHeader) + if first.StatusCode != 200 || host != "beta" { // beta: same free slots, double weight + t.Fatalf("first turn: %d from %q, want 200 from beta", first.StatusCode, host) + } + if got := first.Header.Get(proxy.LeaseHeader); got != "new" { + t.Errorf("%s = %q on the first turn, want new", proxy.LeaseHeader, got) + } + // Take alpha's slots away as a "better host" signal: it must not matter, the lease holds. + for turn := 2; turn <= 6; turn++ { + resp := r.post("/r/v1/chat/completions", conversation(1, turn)) + drain(resp) + if resp.Header.Get(proxy.HostHeader) != host || resp.Header.Get(proxy.LeaseHeader) != "reused" { + t.Fatalf("turn %d: host %q lease %q, want %q reused", turn, resp.Header.Get(proxy.HostHeader), resp.Header.Get(proxy.LeaseHeader), host) + } + } + if alpha.hits.Load() != 0 || beta.hits.Load() != 6 { + t.Errorf("hits alpha=%d beta=%d, want 0 and 6", alpha.hits.Load(), beta.hits.Load()) } - 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(` +// spreadHosts: beta is preferred (weight 10) until both of its "shared" slots are busy; then +// alpha (2 free × 1) beats beta (0 free × 10), and a new conversation must start on alpha. +const spreadHosts = ` listen = "127.0.0.1:1" +queue_max = 4 +lease_idle = "30m" [hosts.alpha] base_url = %q -models = { "shared" = { }, "alpha-only" = { } } +weight = 1.0 +models = { "shared" = { parallel = 2 } } [hosts.beta] base_url = %q -models = { "shared" = { }, "beta-only" = { } } +weight = 10.0 +models = { "shared" = { parallel = 2 } } [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) { +func TestDifferentConversationsSpreadByFreeSlots(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) + beta.delay = 400 * time.Millisecond + r := newRig(t, spreadHosts, alpha, beta) + // Two slow conversations occupy beta's two "shared" slots… + var wg sync.WaitGroup + for i := 1; i <= 2; i++ { + wg.Add(1) + go func(i int) { defer wg.Done(); drain(r.post("/r/v1/chat/completions", conversation(i, 1))) }(i) + time.Sleep(50 * time.Millisecond) // arrive one after the other so both pick beta (10 > 2) } + time.Sleep(50 * time.Millisecond) + // …so a third conversation starting now is sent to alpha (beta has 0 free slots, alpha 2). + resp := r.post("/r/v1/chat/completions", conversation(3, 1)) + drain(resp) if resp.Header.Get(proxy.HostHeader) != "alpha" { - t.Errorf("host header %q", resp.Header.Get(proxy.HostHeader)) + t.Errorf("third conversation went to %q, want alpha (free slots beat weight)", resp.Header.Get(proxy.HostHeader)) + } + wg.Wait() + if beta.hits.Load() != 2 || alpha.hits.Load() != 1 { + t.Errorf("hits beta=%d alpha=%d, want 2 and 1", beta.hits.Load(), alpha.hits.Load()) + } +} + +func TestQueueFullIs503(t *testing.T) { + alpha := newUpstream(t, "alpha") + alpha.delay = 400 * time.Millisecond + r := newRig(t, ` +listen = "127.0.0.1:1" +queue_max = 1 +[hosts.alpha] +base_url = %q +models = { "shared" = { parallel = 1 } } +[routes.r] +hosts = ["alpha"] +default_model = "shared" +`, alpha) + codes := make(chan int, 3) + for i := 1; i <= 3; i++ { + go func(i int) { + resp := r.post("/r/v1/chat/completions", conversation(i, 1)) + drain(resp) + codes <- resp.StatusCode + }(i) + time.Sleep(30 * time.Millisecond) // arrival order: 1 runs, 2 queues, 3 finds the queue full + } + got := map[int]int{} + for i := 0; i < 3; i++ { + got[<-codes]++ + } + if got[200] != 2 || got[503] != 1 { + t.Fatalf("status counts = %v, want two 200 and one 503", got) + } + // Rows are written after each response completes; allow the store a moment to catch up. + var rows []store.UsageRow + deadline := time.Now().Add(2 * time.Second) + for time.Now().Before(deadline) { + rows, _ = r.store.Usage(time.Time{}, store.ByRoute) + if len(rows) == 1 && rows[0].Requests == 3 { + break + } + time.Sleep(20 * time.Millisecond) + } + if len(rows) != 1 || rows[0].Requests != 3 || rows[0].Errors != 1 { + t.Fatalf("usage = %+v, want 3 requests, 1 error (the 503 is recorded too)", rows) + } + if rows[0].QueuedMs <= 0 { + t.Errorf("the queued request must record its wait: %+v", rows[0]) + } +} + +func TestUnhealthyHostReleasesAndMoves(t *testing.T) { + alpha, beta := newUpstream(t, "alpha"), newUpstream(t, "beta") + r := newRig(t, twoHosts, alpha, beta) + drain(r.post("/r/v1/chat/completions", conversation(1, 1))) // lands on beta + beta.srv.Close() + resp := r.post("/r/v1/chat/completions", conversation(1, 2)) + drain(resp) + if resp.StatusCode != http.StatusBadGateway { + t.Fatalf("first request after beta died: %d, want 502", resp.StatusCode) + } + if s, _ := r.health.Get("beta"); s.Healthy { + t.Fatalf("beta must be marked down after the 502") + } + resp = r.post("/r/v1/chat/completions", conversation(1, 3)) + drain(resp) + if resp.StatusCode != 200 || resp.Header.Get(proxy.HostHeader) != "alpha" || resp.Header.Get(proxy.LeaseHeader) != "new" { + t.Errorf("after the move: %d from %q lease %q, want 200 alpha new", resp.StatusCode, resp.Header.Get(proxy.HostHeader), resp.Header.Get(proxy.LeaseHeader)) + } + ev, _ := r.store.Events(time.Time{}, 10) + var reasons []string + for _, e := range ev { + reasons = append(reasons, e.Reason) + } + if len(reasons) != 2 || reasons[0] != store.ReasonNew || reasons[1] != store.ReasonUnhealthy { + t.Errorf("lease events = %v, want [new unhealthy]", reasons) + } +} + +func TestAccountingRowsFromUsageAndTimings(t *testing.T) { + alpha, beta := newUpstream(t, "alpha"), newUpstream(t, "beta") + r := newRig(t, twoHosts, alpha, beta) + drain(r.post("/r/v1/chat/completions", conversation(1, 1))) // non-streamed + drain(r.post("/r/v1/chat/completions", strings.Replace(conversation(1, 2), `"stream":false`, `"stream":true`, 1))) // streamed + deadline := time.Now().Add(2 * time.Second) + var rows []store.UsageRow + for time.Now().Before(deadline) { + rows, _ = r.store.Usage(time.Time{}, store.ByHost) + if len(rows) == 1 && rows[0].Requests == 2 { + break + } + time.Sleep(20 * time.Millisecond) + } + if len(rows) != 1 || rows[0].Requests != 2 { + t.Fatalf("usage by host = %+v, want one host with 2 requests (rows may be written after the response completes, within 2 s)", rows) + } + u := rows[0] + if u.PromptTokens != 300 || u.CachedTokens != 240 || u.CompletionTokens != 30 { + t.Errorf("tokens = prompt %d cached %d completion %d, want 300/240/30 (100+200, 90+150, 10+20)", u.PromptTokens, u.CachedTokens, u.CompletionTokens) + } + if u.BusyMs <= 0 || u.Errors != 0 { + t.Errorf("busy %d errors %d", u.BusyMs, u.Errors) + } + if got := u.CacheHitRatio(); got < 0.79 || got > 0.81 { + t.Errorf("cache hit ratio = %v, want 0.8", got) + } +} + +func TestStreamIsUnalteredWhileTeed(t *testing.T) { + alpha, beta := newUpstream(t, "alpha"), newUpstream(t, "beta") + r := newRig(t, twoHosts, alpha, beta) + resp := r.post("/r/v1/chat/completions", strings.Replace(conversation(9, 1), `"stream":false`, `"stream":true`, 1)) + body := drain(resp) + want := 0 + for _, line := range strings.Split(body, "\n") { + if strings.HasPrefix(line, "data: ") { + want++ + } + } + if want != 5 || !strings.HasSuffix(strings.TrimSpace(body), "data: [DONE]") { + t.Errorf("client must receive every SSE line untouched (3 deltas, usage, DONE); got %d data lines:\n%s", want, body) + } +} + +func TestHeaderRouteOverride(t *testing.T) { + alpha, beta := newUpstream(t, "alpha"), newUpstream(t, "beta") + r := newRig(t, twoHosts, alpha, beta) + // The header names the route; the path has none. + resp := r.post("/v1/chat/completions", conversation(1, 1), proxy.RouteHeader, "other") + drain(resp) + if resp.StatusCode != 200 || resp.Header.Get(proxy.HostHeader) != "alpha" { + t.Errorf("header route 'other' (alpha only): %d from %q", resp.StatusCode, resp.Header.Get(proxy.HostHeader)) + } + if alpha.lastReq().path != "/v1/chat/completions" { + t.Errorf("upstream path = %q", alpha.lastReq().path) + } + // A path route and a header route that disagree: the header is the operator's intent → 400. + resp = r.post("/r/v1/chat/completions", conversation(1, 1), proxy.RouteHeader, "other") + if drain(resp); resp.StatusCode != 400 { + t.Errorf("conflicting route in path and header: %d, want 400", resp.StatusCode) + } + resp = r.post("/v1/chat/completions", conversation(1, 1), proxy.RouteHeader, "nope") + if drain(resp); resp.StatusCode != 404 { + t.Errorf("unknown header route: %d, want 404", resp.StatusCode) + } +} + +func TestV0BehaviourStillHolds(t *testing.T) { + alpha, beta := newUpstream(t, "alpha"), newUpstream(t, "beta") + r := newRig(t, twoHosts, alpha, beta) + for _, tc := range []struct { + method, path string + want int + msg string + }{ + {http.MethodGet, "/", 400, "missing route"}, + {http.MethodGet, "/nope/v1/models", 404, "unknown route"}, + {http.MethodGet, "/r/slots", 404, "not found"}, + {http.MethodGet, "/r/_crossbar/hosts", 404, "not found"}, + } { + req, _ := http.NewRequest(tc.method, r.front.URL+tc.path, nil) + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + body := drain(resp) + var e map[string]string + if resp.StatusCode != tc.want || json.Unmarshal([]byte(body), &e) != nil || e["error"] != tc.msg { + t.Errorf("%s: %d %s, want %d %q", tc.path, resp.StatusCode, body, tc.want, tc.msg) + } + } + big := strings.Repeat("x", proxy.MaxBody+1) + resp := r.post("/r/v1/chat/completions", big) + if drain(resp); resp.StatusCode != 413 { + t.Errorf("oversize body: %d, want 413", resp.StatusCode) + } + // GET pass-through with query string, Host and X-Forwarded-For as in v0. + resp, err := http.Get(r.front.URL + "/r/v1/models?x=1") + if err != nil { + t.Fatal(err) + } + drain(resp) + host := resp.Header.Get(proxy.HostHeader) + u := map[string]*upstream{"alpha": alpha, "beta": beta}[host] + if u == nil || u.lastReq().path != "/v1/models?x=1" || u.lastReq().host != strings.TrimPrefix(u.srv.URL, "http://") || u.lastReq().xff == "" { + t.Errorf("GET pass-through: host %q last %+v", host, u.lastReq()) } } diff --git a/internal/proxy/recorder_test.go b/internal/proxy/recorder_test.go index b5ffa45..8a22dff 100644 --- a/internal/proxy/recorder_test.go +++ b/internal/proxy/recorder_test.go @@ -42,7 +42,7 @@ hosts = ["alpha"] t.Fatal(err) } h := &fakeHealth{st: map[string]health.Status{"alpha": {Healthy: true, Loaded: []string{"m"}}}} - p := proxy.New(cfg, h, nil) + p := proxy.New(cfg, h, nil, nil, nil, nil) rec := httptest.NewRecorder() req := httptest.NewRequest(http.MethodPost, "/r/v1/chat/completions", strings.NewReader(`{"model":"m","stream":true}`)) diff --git a/internal/proxy/tee.go b/internal/proxy/tee.go new file mode 100644 index 0000000..48873ab --- /dev/null +++ b/internal/proxy/tee.go @@ -0,0 +1,147 @@ +package proxy + +import ( + "bytes" + "encoding/json" + "io" + "sync" +) + +// maxParseBody bounds the non-streamed JSON body accumulated to extract usage/timings: beyond it we +// record no tokens rather than hold an unbounded body in memory. +const maxParseBody = 1 << 20 // 1 MiB + +// usageTimings is the last usage/timings object seen in a stream, or the single object parsed from a +// non-streamed JSON body. +type usageTimings struct { + hasUsage bool + usage struct{ prompt, completion int } + hasTimings bool + timings struct{ promptN, cacheN, predictedN int } +} + +// tokens resolves the recorded usage and timings to the three counts the accounting row needs. +// prompt and completion come from usage when present, else from timings; cached comes only from +// timings. +func (u *usageTimings) tokens() (prompt, cached, completion int) { + if u.hasUsage { + prompt = u.usage.prompt + completion = u.usage.completion + } else if u.hasTimings { + prompt = u.timings.promptN + completion = u.timings.predictedN + } + if u.hasTimings { + cached = u.timings.cacheN + } + return +} + +// tee wraps a response body, passing every byte through unchanged while scanning for usage/timings. +// For a text/event-stream body it scans complete data: lines and remembers the last object seen; for +// anything else it accumulates the body (bounded) and parses it once at Close. +type tee struct { + body io.Reader + closed bool + streamed bool + pending []byte + buf bytes.Buffer + mu sync.Mutex + last usageTimings +} + +func newTee(r io.Reader, streamed bool) *tee { + return &tee{body: r, streamed: streamed} +} + +// Read reads from the upstream body and feeds the bytes to the scanner without holding any back. +func (t *tee) Read(b []byte) (int, error) { + n, err := t.body.Read(b) + if n > 0 { + t.ingest(b[:n]) + } + return n, err +} + +// ingest routes a fresh chunk to the streamed scanner or the non-streamed accumulator. +func (t *tee) ingest(p []byte) { + t.mu.Lock() + defer t.mu.Unlock() + if t.streamed { + t.scanSSE(p) + return + } + if t.buf.Len() < maxParseBody { + t.buf.Write(p) + } +} + +// scanSSE splits complete lines off the pending buffer and parses each "data: " line. +func (t *tee) scanSSE(p []byte) { + t.pending = append(t.pending, p...) + for { + i := bytes.IndexByte(t.pending, '\n') + if i < 0 { + return + } + line := t.pending[:i] + t.pending = t.pending[i+1:] + if bytes.HasPrefix(line, []byte("data: ")) { + t.parseData(line[len("data: "):]) + } + } +} + +// parseData decodes one data line's JSON object and remembers its usage and/or timings. +func (t *tee) parseData(data []byte) { + var doc struct { + Usage *struct { + PromptTokens int `json:"prompt_tokens"` + CompletionTokens int `json:"completion_tokens"` + } `json:"usage"` + Timings *struct { + PromptN int `json:"prompt_n"` + CacheN int `json:"cache_n"` + PredictedN int `json:"predicted_n"` + } `json:"timings"` + } + if err := json.Unmarshal(data, &doc); err != nil { + return + } + if doc.Usage != nil { + t.last.hasUsage = true + t.last.usage.prompt = doc.Usage.PromptTokens + t.last.usage.completion = doc.Usage.CompletionTokens + } + if doc.Timings != nil { + t.last.hasTimings = true + t.last.timings.promptN = doc.Timings.PromptN + t.last.timings.cacheN = doc.Timings.CacheN + t.last.timings.predictedN = doc.Timings.PredictedN + } +} + +// Close closes the underlying body, parsing a non-streamed JSON body once at the end. +func (t *tee) Close() error { + t.mu.Lock() + if t.closed { + t.mu.Unlock() + return nil + } + t.closed = true + if !t.streamed && t.buf.Len() > 0 { + t.parseData(t.buf.Bytes()) + } + t.mu.Unlock() + if c, ok := t.body.(io.Closer); ok { + return c.Close() + } + return nil +} + +// tokens returns the last usage/timings seen, if any. +func (t *tee) tokens() (prompt, cached, completion int) { + t.mu.Lock() + defer t.mu.Unlock() + return t.last.tokens() +}