291 lines
11 KiB
Go
291 lines
11 KiB
Go
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"
|
||
"net/http"
|
||
"strings"
|
||
"sync"
|
||
"testing"
|
||
"time"
|
||
|
||
"git.wntrmute.dev/kyle/crossbar/internal/proxy"
|
||
"git.wntrmute.dev/kyle/crossbar/internal/store"
|
||
)
|
||
|
||
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())
|
||
}
|
||
}
|
||
|
||
// waitUntil polls cond every 5 ms for up to two seconds and fails the test if it never holds.
|
||
func waitUntil(t *testing.T, cond func() bool) {
|
||
t.Helper()
|
||
deadline := time.Now().Add(2 * time.Second)
|
||
for time.Now().Before(deadline) {
|
||
if cond() {
|
||
return
|
||
}
|
||
time.Sleep(5 * time.Millisecond)
|
||
}
|
||
t.Fatal("condition not reached within two seconds")
|
||
}
|
||
|
||
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)
|
||
fire := func(i int) {
|
||
go func() {
|
||
resp := r.post("/r/v1/chat/completions", conversation(i, 1))
|
||
drain(resp)
|
||
codes <- resp.StatusCode
|
||
}()
|
||
}
|
||
// Arrival order is enforced by watching the limiter, not by sleeping: 1 runs, 2 queues,
|
||
// 3 finds the queue full.
|
||
fire(1)
|
||
waitUntil(t, func() bool { return r.lim.InFlight("alpha", "shared") == 1 })
|
||
fire(2)
|
||
waitUntil(t, func() bool { return r.lim.Queued("alpha", "shared") == 1 })
|
||
fire(3)
|
||
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())
|
||
}
|
||
}
|