package scheduler import ( "context" "errors" "strings" "testing" ) // The chain trace is the only way a caller learns that a request was DEGRADED // — served by a lower tier than the one that should have taken it. chainDrive's // return value carries only the winner, so without these events "tier 1 was // cooling and we dropped to tier 2" is indistinguishable from "tier 1 served // it", which is the exact question a priority chain exists to answer. // recorder collects trace events for assertions. type recorder struct{ events []TraceEvent } func (r *recorder) sink(ev TraceEvent) { r.events = append(r.events, ev) } func (r *recorder) kinds() []TraceKind { out := make([]TraceKind, 0, len(r.events)) for _, e := range r.events { out = append(out, e.Kind) } return out } func (r *recorder) find(k TraceKind) *TraceEvent { for i := range r.events { if r.events[i].Kind == k { return &r.events[i] } } return nil } // TestTraceSelectedOnlyOnHappyPath: a clean walk emits exactly one event. func TestTraceSelectedOnlyOnHappyPath(t *testing.T) { p1 := fakeProv("p1", "m1") ch := BuildChain([]Rule{{Model: "m1", Source: "p1", Tier: 1}}, func(model, source string) Provider { return p1 }) var rec recorder _, _, _, err := New(3).ChainChat(context.Background(), ch, chatReq(), nil, rec.sink) if err != nil { t.Fatalf("chain: %v", err) } if got := rec.kinds(); len(got) != 1 || got[0] != TraceSelected { t.Errorf("events = %v, want a single selected", got) } e := rec.find(TraceSelected) if e.Tier != 1 || e.Source != "p1" || e.Model != "m1" { t.Errorf("selected event = %+v, want tier 1 / p1 / m1", e) } if e.Attempt != 1 { t.Errorf("Attempt = %d, want 1", e.Attempt) } } // TestTraceRecordsTierSkipAndDegradation is the core case: tier 1 is // unschedulable, tier 2 answers. The trace must show the skip AND the eventual // selection, so a consumer can see the request was served one tier down. func TestTraceRecordsTierSkipAndDegradation(t *testing.T) { // p1 is unavailable (not probeable), so tier 1 yields no candidates. p1 := fakeProv("p1", "m1") p1.available.Store(false) p2 := fakeProv("p2", "m2") ch := BuildChain([]Rule{ {Model: "m1", Source: "p1", Tier: 1}, {Model: "m2", Source: "p2", Tier: 2}, }, bySource(p1, p2)) var rec recorder _, src, model, err := New(3).ChainChat(context.Background(), ch, chatReq(), nil, rec.sink) if err != nil { t.Fatalf("chain: %v", err) } if src != "p2" || model != "m2" { t.Fatalf("served by %s/%s, want p2/m2", src, model) } skip := rec.find(TraceTierSkip) if skip == nil { t.Fatalf("no tier_skip event; events = %v", rec.kinds()) } if skip.Tier != 1 { t.Errorf("skip tier = %d, want 1", skip.Tier) } if !strings.Contains(skip.Reason, "cooling") { t.Errorf("skip reason = %q, want it to mention cooling", skip.Reason) } sel := rec.find(TraceSelected) if sel == nil || sel.Tier != 2 { t.Errorf("selected = %+v, want tier 2", sel) } // The order matters: the skip must be observable BEFORE the selection. if rec.events[0].Kind != TraceTierSkip || rec.events[len(rec.events)-1].Kind != TraceSelected { t.Errorf("event order = %v, want skip first and selected last", rec.kinds()) } } // TestTraceRecordsHardSlotFailures: a slot that returns an upstream error is a // different event from a skip — the request tried it and it failed. Losing that // distinction makes a flaky upstream look like an idle one. func TestTraceRecordsHardSlotFailures(t *testing.T) { p1 := fakeProv("p1", "m1") p1.fail.Store(true) // Chat returns "upstream error" p2 := fakeProv("p2", "m2") ch := BuildChain([]Rule{ {Model: "m1", Source: "p1", Tier: 1}, {Model: "m2", Source: "p2", Tier: 2}, }, bySource(p1, p2)) var rec recorder _, src, _, err := New(3).ChainChat(context.Background(), ch, chatReq(), nil, rec.sink) if err != nil { t.Fatalf("chain: %v", err) } if src != "p2" { t.Fatalf("served by %s, want p2", src) } fail := rec.find(TraceSlotFail) if fail == nil { t.Fatalf("no slot_fail event; events = %v", rec.kinds()) } if fail.Tier != 1 || fail.Source != "p1" || fail.Model != "m1" { t.Errorf("slot_fail = %+v, want tier 1 / p1 / m1", fail) } if !strings.Contains(fail.Err, "upstream error") { t.Errorf("slot_fail error = %q, want the upstream text", fail.Err) } // A hard failure must NOT be reported as a skip. if rec.find(TraceTierSkip) != nil { t.Error("a hard failure was also reported as a tier_skip") } } // TestTraceNilSinkIsSafe: the gateway passes nil when no plugin is loaded, so // every emit path must tolerate it. This is the "plugins are optional" property // on the scheduler side. func TestTraceNilSinkIsSafe(t *testing.T) { p1 := fakeProv("p1", "m1") p1.fail.Store(true) p2 := fakeProv("p2", "m2") ch := BuildChain([]Rule{ {Model: "m1", Source: "p1", Tier: 1}, {Model: "m2", Source: "p2", Tier: 2}, }, bySource(p1, p2)) if _, _, _, err := New(3).ChainChat(context.Background(), ch, chatReq(), nil, nil); err != nil { t.Fatalf("a nil trace sink broke the walk: %v", err) } } // TestTraceOnTotalFailure: when every tier fails, the walk still emits its // per-step events AND returns the ChainErr. The trace is additive — it must not // replace or disturb the error contract callers depend on for the 503. func TestTraceOnTotalFailure(t *testing.T) { p1 := fakeProv("p1", "m1") p1.fail.Store(true) p2 := fakeProv("p2", "m2") p2.fail.Store(true) ch := BuildChain([]Rule{ {Model: "m1", Source: "p1", Tier: 1}, {Model: "m2", Source: "p2", Tier: 2}, }, bySource(p1, p2)) var rec recorder _, _, _, err := New(3).ChainChat(context.Background(), ch, chatReq(), nil, rec.sink) var ce *ChainErr if !errors.As(err, &ce) { t.Fatalf("err = %v, want a *ChainErr so the gateway can answer 503", err) } if len(ce.Tiers) != 2 { t.Errorf("ChainErr.Tiers = %d, want 2 (the error contract must be unchanged)", len(ce.Tiers)) } if n := len(rec.kinds()); n != 2 { t.Errorf("events = %v, want two slot_fail and no selection", rec.kinds()) } if rec.find(TraceSelected) != nil { t.Error("a selected event was emitted for a walk that served nothing") } } // TestTraceSkipsEmptyChain: no chain configured must not emit anything; the // gateway answers 503 before scheduling in that case anyway. func TestTraceSkipsEmptyChain(t *testing.T) { var rec recorder _, _, _, err := New(3).ChainChat(context.Background(), &Chain{}, chatReq(), nil, rec.sink) if err == nil { t.Fatal("expected an error for an empty chain") } if len(rec.events) != 0 { t.Errorf("events = %v, want none", rec.kinds()) } }