Files
crossbar/docs/plans/v1/_files/internal/proxy/proxy_test.go
T

431 lines
15 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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).
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"
)
// 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) {
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)
}
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())
}
}
// 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
weight = 1.0
models = { "shared" = { parallel = 2 } }
[hosts.beta]
base_url = %q
weight = 10.0
models = { "shared" = { parallel = 2 } }
[routes.r]
hosts = ["alpha", "beta"]
default_model = "shared"
`
func TestDifferentConversationsSpreadByFreeSlots(t *testing.T) {
alpha, beta := newUpstream(t, "alpha"), newUpstream(t, "beta")
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("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())
}
}