mirror of
https://gitcode.com/JianFeeeee/ModelRouter.git
synced 2026-10-03 23:54:06 +00:00
feat: AUTO chain rewrite — silent failover+busy skip+pref round-robin+503 tier summary; chain edits reset slot cooldowns (P0/P1); stats by_status + audit jsonl rotation; UI priority-page health badges & status-code card; ctx-menu capture-phase close (outside-press guard); main.go ops warnings; local bundled-Lua verified tests (3 latent bugs fixed); plan.md
This commit is contained in:
@ -1,16 +1,29 @@
|
||||
// Package scheduler implements request scheduling across providers: per-source
|
||||
// concurrency caps (acquire with wait = queuing), AUTO model fallback chains,
|
||||
// and exponential backoff via provider health.
|
||||
// Package scheduler implements request scheduling across providers: direct
|
||||
// fallback scheduling over candidate lists, and the AUTO chain (tiers with
|
||||
// per-tier round-robin cursors, preference ordering, token-quota windows and
|
||||
// per-(source,model) cooldown awareness) per the target architecture in
|
||||
// plan.md.
|
||||
package scheduler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
"llmsproxy/internal/provider"
|
||||
"llmsproxy/internal/types"
|
||||
)
|
||||
|
||||
// busyWait is how long a fully-busy tier is polled for a free slot before the
|
||||
// request falls through to the next tier (bounded wait, plan 2.3).
|
||||
var busyWait = 2 * time.Second
|
||||
|
||||
// busyPoll is the polling interval while waiting for a busy tier.
|
||||
var busyPoll = 100 * time.Millisecond
|
||||
|
||||
// Scheduler drives one chat tool call across the candidate provider chain.
|
||||
type Scheduler struct {
|
||||
// MaxRetries how many fallback providers to try before failing.
|
||||
@ -27,29 +40,297 @@ func New(maxRetries int) *Scheduler {
|
||||
// Provider is the minimal interface the scheduler needs to schedule over.
|
||||
type Provider interface {
|
||||
Name() string
|
||||
Available() bool
|
||||
ModelFor(reqModel string) string
|
||||
ModelAvailable(model string) bool
|
||||
Pref(model string) int64
|
||||
Chat(ctx context.Context, req *types.ChatRequest) (*types.UnifiedResponse, error)
|
||||
ChatStream(ctx context.Context, req *types.ChatRequest) (<-chan types.UnifiedChunk, error)
|
||||
Image(ctx context.Context, req *types.ImageGenRequest) (*types.UnifiedResponse, error)
|
||||
}
|
||||
|
||||
// FromRegistry converts *provider.Provider slices to the scheduler interface.
|
||||
func FromRegistry(ps []*provider.Provider) []Provider {
|
||||
out := make([]Provider, len(ps))
|
||||
for i, p := range ps {
|
||||
out[i] = p
|
||||
}
|
||||
return out
|
||||
// ---- AUTO chain ----
|
||||
|
||||
// Rule is one persisted AUTO chain slot (mirror of config.ModelScope).
|
||||
type Rule struct {
|
||||
Model string
|
||||
Source string
|
||||
Tier int
|
||||
Quota int64
|
||||
Period string
|
||||
Hours int64
|
||||
}
|
||||
|
||||
// Slot is one schedulable chain position: a model pinned to its provider,
|
||||
// with an optional token-quota window. Slots are immutable after build.
|
||||
type Slot struct {
|
||||
Model string
|
||||
Source string
|
||||
Quota int64
|
||||
Period string
|
||||
Hours int64
|
||||
Prov Provider
|
||||
}
|
||||
|
||||
// TierNode is one priority tier. Slots keep their configured order (the
|
||||
// stable base for preference ordering). next is the round-robin cursor: it
|
||||
// holds the last used slot index (-1 = none yet), so the very first request
|
||||
// starts at the configured order and later ones rotate.
|
||||
type TierNode struct {
|
||||
Tier int
|
||||
Slots []*Slot
|
||||
next atomic.Int64
|
||||
}
|
||||
|
||||
// NextStart advances the tier cursor and returns the start index for the next
|
||||
// scheduling run (first run: index 0).
|
||||
func (tn *TierNode) NextStart() int64 {
|
||||
return tn.next.Add(1)
|
||||
}
|
||||
|
||||
// Chain is the immutable AUTO scheduling plan. A rebuilt chain is swapped in
|
||||
// atomically; per-tier cursors live inside the chain and are shared across
|
||||
// requests (rotation state resets when the chain is rebuilt, e.g. after
|
||||
// editing the rules — acceptable, the swap also resets cooldowns).
|
||||
type Chain struct {
|
||||
Tiers []*TierNode // descending tier order
|
||||
}
|
||||
|
||||
// TierErrors is the per-tier failure summary carried by ChainErr. Errors
|
||||
// (TierError or skipped-tier reasons) are collected in tier order.
|
||||
type TierError struct {
|
||||
Tier int
|
||||
Source string
|
||||
Model string
|
||||
Err error
|
||||
}
|
||||
|
||||
// ChainErr is returned by chain scheduling when every AUTO tier failed. Its
|
||||
// message summarizes each failed tier (which source/model and why) so a 503
|
||||
// names the culprits instead of the bare "no provider available".
|
||||
type ChainErr struct {
|
||||
Tiers []TierError
|
||||
Skipped []string // whole-tier reasons (cooling / quota / all busy)
|
||||
}
|
||||
|
||||
func (e *ChainErr) Error() string {
|
||||
var b strings.Builder
|
||||
b.WriteString("all auto tiers failed: ")
|
||||
first := true
|
||||
for _, t := range e.Tiers {
|
||||
if !first {
|
||||
b.WriteString("; ")
|
||||
}
|
||||
first = false
|
||||
fmt.Fprintf(&b, "tier %d %s/%s: %v", t.Tier, t.Source, t.Model, t.Err)
|
||||
}
|
||||
for _, s := range e.Skipped {
|
||||
if !first {
|
||||
b.WriteString("; ")
|
||||
}
|
||||
first = false
|
||||
b.WriteString(s)
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
|
||||
// BuildChain groups rules into descending tiers and resolves each slot's
|
||||
// provider via prov. Rules whose provider resolves to nil are dropped (the
|
||||
// source no longer serves the model). Slot order within a tier follows the
|
||||
// configured rule order.
|
||||
func BuildChain(rules []Rule, prov func(model, source string) Provider) *Chain {
|
||||
byTier := map[int][]*Slot{}
|
||||
var tiers []int
|
||||
for _, r := range rules {
|
||||
p := prov(r.Model, r.Source)
|
||||
if p == nil {
|
||||
continue
|
||||
}
|
||||
if _, ok := byTier[r.Tier]; !ok {
|
||||
tiers = append(tiers, r.Tier)
|
||||
}
|
||||
byTier[r.Tier] = append(byTier[r.Tier], &Slot{
|
||||
Model: r.Model,
|
||||
Source: r.Source,
|
||||
Quota: r.Quota,
|
||||
Period: r.Period,
|
||||
Hours: r.Hours,
|
||||
Prov: p,
|
||||
})
|
||||
}
|
||||
sort.Slice(tiers, func(i, j int) bool { return tiers[i] > tiers[j] })
|
||||
ch := &Chain{}
|
||||
for _, t := range tiers {
|
||||
tn := &TierNode{Tier: t, Slots: byTier[t]}
|
||||
tn.next.Store(-1)
|
||||
ch.Tiers = append(ch.Tiers, tn)
|
||||
}
|
||||
return ch
|
||||
}
|
||||
|
||||
// tierResult is the outcome of one scheduling run over one tier.
|
||||
type tierResult struct {
|
||||
resp *types.UnifiedResponse
|
||||
chunks <-chan types.UnifiedChunk
|
||||
src string
|
||||
model string
|
||||
hard []TierError // hard failures seen in this pass (nil = none)
|
||||
}
|
||||
|
||||
// runTier executes one tier pass starting at the round-robin base index.
|
||||
// Cooldown is the only hard skip (re-verified per slot); a busy slot is
|
||||
// skipped without any penalty; a hard failure is recorded and the pass moves
|
||||
// on to the next slot (plan 2.3: "单请求内不重试已失败槽" — the failed slot is
|
||||
// not retried, the others still are). hard == nil and no success means every
|
||||
// candidate was merely busy/cooling, so the caller may wait a bounded time.
|
||||
func runTier(ctx context.Context, tn *TierNode, cands []*Slot, base int64, req *types.ChatRequest, stream bool) tierResult {
|
||||
n := len(cands)
|
||||
var hard []TierError
|
||||
for i := 0; i < n; i++ {
|
||||
sl := cands[(int(base)+i)%n]
|
||||
if !sl.Prov.ModelAvailable(sl.Model) {
|
||||
continue
|
||||
}
|
||||
r := *req
|
||||
r.Model = sl.Model
|
||||
if stream {
|
||||
chunks, err := sl.Prov.ChatStream(ctx, &r)
|
||||
if err == nil {
|
||||
return tierResult{chunks: chunks, src: sl.Source, model: sl.Model}
|
||||
}
|
||||
if ctx.Err() != nil {
|
||||
return tierResult{}
|
||||
}
|
||||
if errors.Is(err, types.ErrBusy) {
|
||||
continue
|
||||
}
|
||||
hard = append(hard, TierError{Tier: tn.Tier, Source: sl.Source, Model: sl.Model, Err: err})
|
||||
continue
|
||||
}
|
||||
resp, err := sl.Prov.Chat(ctx, &r)
|
||||
if err == nil {
|
||||
return tierResult{resp: resp, src: sl.Source, model: sl.Model}
|
||||
}
|
||||
if ctx.Err() != nil {
|
||||
return tierResult{}
|
||||
}
|
||||
if errors.Is(err, types.ErrBusy) {
|
||||
continue
|
||||
}
|
||||
hard = append(hard, TierError{Tier: tn.Tier, Source: sl.Source, Model: sl.Model, Err: err})
|
||||
}
|
||||
return tierResult{hard: hard}
|
||||
}
|
||||
|
||||
// chainDrive runs a request down the chain (plan 2.3): tiers descending,
|
||||
// per-tier round-robin starting at the tier cursor, same-tier runs ordered by
|
||||
// preference (negative prefs sink but stay reachable). Quota-exhausted and
|
||||
// cooling slots are filtered up front; a fully busy tier is polled for a
|
||||
// bounded time before falling through. Failures are summarized in *ChainErr
|
||||
// for the caller to map to HTTP 503.
|
||||
func (s *Scheduler) chainDrive(ctx context.Context, chain *Chain, req *types.ChatRequest, exhausted func(*Slot) bool, stream bool) (*types.UnifiedResponse, <-chan types.UnifiedChunk, string, string, error) {
|
||||
if chain == nil || len(chain.Tiers) == 0 {
|
||||
return nil, nil, "", "", fmt.Errorf("no auto slot configured")
|
||||
}
|
||||
var ce ChainErr
|
||||
for _, tn := range chain.Tiers {
|
||||
// initial filter: quota-exhausted and cooling slots are dropped
|
||||
var cands []*Slot
|
||||
for _, sl := range tn.Slots {
|
||||
if exhausted != nil && exhausted(sl) {
|
||||
continue
|
||||
}
|
||||
if !sl.Prov.ModelAvailable(sl.Model) {
|
||||
continue
|
||||
}
|
||||
cands = append(cands, sl)
|
||||
}
|
||||
if len(cands) == 0 {
|
||||
ce.Skipped = append(ce.Skipped, fmt.Sprintf("tier %d: no schedulable slot (cooling or quota exhausted)", tn.Tier))
|
||||
continue
|
||||
}
|
||||
// preference orders a same-tier run; stable so equal prefs keep order
|
||||
sort.SliceStable(cands, func(i, j int) bool {
|
||||
return cands[i].Prov.Pref(cands[i].Model) > cands[j].Prov.Pref(cands[j].Model)
|
||||
})
|
||||
base := tn.NextStart()
|
||||
res := runTier(ctx, tn, cands, base, req, stream)
|
||||
if res.resp != nil || res.chunks != nil {
|
||||
return res.resp, res.chunks, res.src, res.model, nil
|
||||
}
|
||||
if ctx.Err() != nil {
|
||||
return nil, nil, "", "", ctx.Err()
|
||||
}
|
||||
if len(res.hard) > 0 {
|
||||
ce.Tiers = append(ce.Tiers, res.hard...)
|
||||
continue // hard failures: fall through to the next tier, no waiting
|
||||
}
|
||||
// every candidate was busy or cooling: bounded poll before downgrading
|
||||
deadline := time.Now().Add(busyWait)
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, nil, "", "", ctx.Err()
|
||||
case <-time.After(busyPoll):
|
||||
}
|
||||
done := time.Now().After(deadline)
|
||||
if done {
|
||||
ce.Skipped = append(ce.Skipped, fmt.Sprintf("tier %d: no free slot within %v", tn.Tier, busyWait))
|
||||
break
|
||||
}
|
||||
// refresh candidates: cooldowns may have expired meanwhile
|
||||
var again []*Slot
|
||||
for _, sl := range cands {
|
||||
if sl.Prov.ModelAvailable(sl.Model) {
|
||||
again = append(again, sl)
|
||||
}
|
||||
}
|
||||
if len(again) == 0 {
|
||||
ce.Skipped = append(ce.Skipped, fmt.Sprintf("tier %d: no free slot within %v", tn.Tier, busyWait))
|
||||
break
|
||||
}
|
||||
res = runTier(ctx, tn, again, base, req, stream)
|
||||
if res.resp != nil || res.chunks != nil {
|
||||
return res.resp, res.chunks, res.src, res.model, nil
|
||||
}
|
||||
if ctx.Err() != nil {
|
||||
return nil, nil, "", "", ctx.Err()
|
||||
}
|
||||
if len(res.hard) > 0 {
|
||||
ce.Tiers = append(ce.Tiers, res.hard...)
|
||||
break // hard failure while waiting: stop waiting, fall through
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(ce.Tiers) == 0 && len(ce.Skipped) == 0 {
|
||||
return nil, nil, "", "", fmt.Errorf("no auto slot configured")
|
||||
}
|
||||
return nil, nil, "", "", &ce
|
||||
}
|
||||
|
||||
// ChainChat runs a non-streaming AUTO request down the chain. exhausted, when
|
||||
// non-nil, decides slot token-quota exhaustion. Returns the response, the
|
||||
// serving source and the exact model id used; on total failure a *ChainErr
|
||||
// summarizing every tier.
|
||||
func (s *Scheduler) ChainChat(ctx context.Context, chain *Chain, req *types.ChatRequest, exhausted func(*Slot) bool) (*types.UnifiedResponse, string, string, error) {
|
||||
resp, _, src, model, err := s.chainDrive(ctx, chain, req, exhausted, false)
|
||||
return resp, src, model, err
|
||||
}
|
||||
|
||||
// ChainChatStream runs a streaming AUTO request down the chain. A slot is
|
||||
// abandoned only on connect failures / busy (before its first chunk); after a
|
||||
// stream starts it is pinned. Same return contract as ChainChat.
|
||||
func (s *Scheduler) ChainChatStream(ctx context.Context, chain *Chain, req *types.ChatRequest, exhausted func(*Slot) bool) (<-chan types.UnifiedChunk, string, string, error) {
|
||||
_, chunks, src, model, err := s.chainDrive(ctx, chain, req, exhausted, true)
|
||||
return chunks, src, model, err
|
||||
}
|
||||
|
||||
// ---- direct scheduling ----
|
||||
|
||||
// Chat runs a chat request across cands, falling back on failure. Each
|
||||
// candidate receives a request pinned to its own model (ModelFor), so an AUTO
|
||||
// chain fallback switches the model id per provider instead of reusing the
|
||||
// first candidate's model name.
|
||||
//
|
||||
// On success it returns the response together with the name of the provider
|
||||
// and the exact model id that actually served the request (used for stats).
|
||||
// candidate receives a request pinned to its own model (ModelFor), so a
|
||||
// fallback switches the model id per provider instead of reusing the first
|
||||
// candidate's model name. On success it returns the response together with
|
||||
// the name of the provider and the exact model id that served the request.
|
||||
func (s *Scheduler) Chat(ctx context.Context, cands []Provider, req *types.ChatRequest) (*types.UnifiedResponse, string, string, error) {
|
||||
attempts := s.MaxRetries + 1
|
||||
var lastErr error
|
||||
@ -74,9 +355,10 @@ func (s *Scheduler) Chat(ctx context.Context, cands []Provider, req *types.ChatR
|
||||
return nil, "", "", lastErr
|
||||
}
|
||||
|
||||
// ChatStream runs a streaming chat across cands, falling back early on connect
|
||||
// errors. The request model is pinned per candidate like Chat. On success it
|
||||
// returns the chunk channel plus the serving provider name and model id.
|
||||
// ChatStream runs a streaming chat across cands, falling back early on
|
||||
// connect errors. The request model is pinned per candidate like Chat. On
|
||||
// success it returns the chunk channel plus the serving provider name and
|
||||
// model id.
|
||||
func (s *Scheduler) ChatStream(ctx context.Context, cands []Provider, req *types.ChatRequest) (<-chan types.UnifiedChunk, string, string, error) {
|
||||
attempts := s.MaxRetries + 1
|
||||
var lastErr error
|
||||
@ -113,4 +395,4 @@ func (s *Scheduler) Image(ctx context.Context, cands []Provider, req *types.Imag
|
||||
lastErr = fmt.Errorf("no provider available")
|
||||
}
|
||||
return nil, "", lastErr
|
||||
}
|
||||
}
|
||||
|
||||
283
internal/scheduler/scheduler_test.go
Normal file
283
internal/scheduler/scheduler_test.go
Normal file
@ -0,0 +1,283 @@
|
||||
package scheduler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"llmsproxy/internal/types"
|
||||
)
|
||||
|
||||
// fakeProvider is an in-memory Provider used to exercise chain scheduling
|
||||
// deterministically without a Lua runtime.
|
||||
type fakeProvider struct {
|
||||
name string
|
||||
model string
|
||||
pref atomic.Int64
|
||||
available atomic.Bool
|
||||
busy atomic.Bool
|
||||
fail atomic.Bool
|
||||
chatHits atomic.Int64
|
||||
}
|
||||
|
||||
func fakeProv(name, model string) *fakeProvider {
|
||||
f := &fakeProvider{name: name, model: model}
|
||||
f.available.Store(true)
|
||||
return f
|
||||
}
|
||||
|
||||
func (f *fakeProvider) Name() string { return f.name }
|
||||
|
||||
func (f *fakeProvider) ModelFor(reqModel string) string { return f.model }
|
||||
|
||||
func (f *fakeProvider) ModelAvailable(model string) bool { return f.available.Load() }
|
||||
|
||||
func (f *fakeProvider) Pref(model string) int64 { return f.pref.Load() }
|
||||
|
||||
func (f *fakeProvider) Chat(ctx context.Context, req *types.ChatRequest) (*types.UnifiedResponse, error) {
|
||||
f.chatHits.Add(1)
|
||||
if f.busy.Load() {
|
||||
return nil, types.ErrBusy
|
||||
}
|
||||
if f.fail.Load() {
|
||||
return nil, fmt.Errorf("upstream error")
|
||||
}
|
||||
return &types.UnifiedResponse{Content: f.name, FinishReason: "stop", TokenUsage: types.TokenUsage{Prompt: 1, Completion: 1, Total: 2}}, nil
|
||||
}
|
||||
|
||||
func (f *fakeProvider) ChatStream(ctx context.Context, req *types.ChatRequest) (<-chan types.UnifiedChunk, error) {
|
||||
if f.busy.Load() {
|
||||
return nil, types.ErrBusy
|
||||
}
|
||||
if f.fail.Load() {
|
||||
return nil, fmt.Errorf("upstream error")
|
||||
}
|
||||
ch := make(chan types.UnifiedChunk, 2)
|
||||
ch <- types.UnifiedChunk{Content: f.name}
|
||||
ch <- types.UnifiedChunk{Done: true}
|
||||
close(ch)
|
||||
return ch, nil
|
||||
}
|
||||
|
||||
func (f *fakeProvider) Image(ctx context.Context, req *types.ImageGenRequest) (*types.UnifiedResponse, error) {
|
||||
return nil, errors.New("no image")
|
||||
}
|
||||
|
||||
// lookup resolves (model, source) -> provider for chain builders in tests.
|
||||
type lookup func(m, s string) Provider
|
||||
|
||||
func bySource(ps ...*fakeProvider) lookup {
|
||||
return func(m, s string) Provider {
|
||||
for _, p := range ps {
|
||||
if p.name == s {
|
||||
return p
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func chatReq() *types.ChatRequest {
|
||||
return &types.ChatRequest{Messages: []types.ChatMessage{{Role: "user", Content: types.StringContent("hi")}}}
|
||||
}
|
||||
|
||||
func TestBuildChainTierOrdering(t *testing.T) {
|
||||
a, b, c, d := fakeProv("s1", "a"), fakeProv("s2", "b"), fakeProv("s3", "c"), fakeProv("s4", "d")
|
||||
ch := BuildChain([]Rule{
|
||||
{Tier: 1, Model: "a", Source: "s1"},
|
||||
{Tier: 3, Model: "c", Source: "s3"},
|
||||
{Tier: 2, Model: "b", Source: "s2"},
|
||||
{Tier: 1, Model: "d", Source: "s4"},
|
||||
{Tier: 9, Model: "gone", Source: "missing"},
|
||||
}, bySource(a, b, c, d))
|
||||
if len(ch.Tiers) != 3 {
|
||||
t.Fatalf("tiers = %d, want 3", len(ch.Tiers))
|
||||
}
|
||||
if ch.Tiers[0].Tier != 3 || ch.Tiers[1].Tier != 2 || ch.Tiers[2].Tier != 1 {
|
||||
t.Fatalf("tier order = %d,%d,%d, want 3,2,1", ch.Tiers[0].Tier, ch.Tiers[1].Tier, ch.Tiers[2].Tier)
|
||||
}
|
||||
// same-tier slots keep configured order, unresolvable rule is dropped
|
||||
if len(ch.Tiers[2].Slots) != 2 || ch.Tiers[2].Slots[0].Model != "a" || ch.Tiers[2].Slots[1].Model != "d" {
|
||||
t.Fatalf("tier 1 slots = %+v", ch.Tiers[2].Slots)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChainRoundRobin(t *testing.T) {
|
||||
a, b := fakeProv("s1", "a"), fakeProv("s2", "b")
|
||||
ch := BuildChain([]Rule{
|
||||
{Tier: 0, Model: "a", Source: "s1"},
|
||||
{Tier: 0, Model: "b", Source: "s2"},
|
||||
}, bySource(a, b))
|
||||
s := New(0)
|
||||
var got []string
|
||||
for i := 0; i < 4; i++ {
|
||||
_, src, _, err := s.ChainChat(context.Background(), ch, chatReq(), nil)
|
||||
if err != nil {
|
||||
t.Fatalf("iter %d: %v", i, err)
|
||||
}
|
||||
got = append(got, src)
|
||||
}
|
||||
want := []string{"s1", "s2", "s1", "s2"}
|
||||
for i := range want {
|
||||
if got[i] != want[i] {
|
||||
t.Fatalf("rr order = %v, want %v", got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestChainPreferenceSinksButStaysReachable(t *testing.T) {
|
||||
neg := fakeProv("neg", "n")
|
||||
good := fakeProv("good", "g")
|
||||
neg.pref.Store(-5)
|
||||
good.fail.Store(true)
|
||||
ch := BuildChain([]Rule{
|
||||
{Tier: 0, Model: "n", Source: "neg"},
|
||||
{Tier: 0, Model: "g", Source: "good"},
|
||||
}, bySource(neg, good))
|
||||
s := New(0)
|
||||
resp, src, model, err := s.ChainChat(context.Background(), ch, chatReq(), nil)
|
||||
if err != nil {
|
||||
t.Fatalf("chain: %v", err)
|
||||
}
|
||||
// higher pref (good) is tried first and hard-fails; the pass moves on to
|
||||
// the negative-pref slot, which sinks but stays reachable: with the real
|
||||
// provider its success would RecordSuccess (+1 pref, self-heal)
|
||||
if src != "neg" || model != "n" || resp.Content != "neg" {
|
||||
t.Fatalf("served src=%q model=%q content=%q", src, model, resp.Content)
|
||||
}
|
||||
if good.chatHits.Load() == 0 {
|
||||
t.Fatal("higher-pref slot must be tried first")
|
||||
}
|
||||
}
|
||||
|
||||
func TestChainBusySkipsWithoutPenalty(t *testing.T) {
|
||||
a, b := fakeProv("s1", "a"), fakeProv("s2", "b")
|
||||
a.busy.Store(true)
|
||||
ch := BuildChain([]Rule{
|
||||
{Tier: 0, Model: "a", Source: "s1"},
|
||||
{Tier: 0, Model: "b", Source: "s2"},
|
||||
}, bySource(a, b))
|
||||
s := New(0)
|
||||
_, src, _, err := s.ChainChat(context.Background(), ch, chatReq(), nil)
|
||||
if err != nil {
|
||||
t.Fatalf("chain: %v", err)
|
||||
}
|
||||
if src != "s2" {
|
||||
t.Fatalf("src = %q, want s2", src)
|
||||
}
|
||||
if a.chatHits.Load() == 0 {
|
||||
t.Fatal("busy slot must have been attempted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestChainAllBusyBoundedWaitThenNextTier(t *testing.T) {
|
||||
oldWait, oldPoll := busyWait, busyPoll
|
||||
busyWait, busyPoll = 60*time.Millisecond, 10*time.Millisecond
|
||||
t.Cleanup(func() { busyWait, busyPoll = oldWait, oldPoll })
|
||||
a, b, c := fakeProv("s1", "a"), fakeProv("s2", "b"), fakeProv("s3", "c")
|
||||
a.busy.Store(true)
|
||||
b.busy.Store(true)
|
||||
ch := BuildChain([]Rule{
|
||||
{Tier: 5, Model: "a", Source: "s1"},
|
||||
{Tier: 5, Model: "b", Source: "s2"},
|
||||
{Tier: 4, Model: "c", Source: "s3"},
|
||||
}, bySource(a, b, c))
|
||||
s := New(0)
|
||||
t0 := time.Now()
|
||||
_, src, _, err := s.ChainChat(context.Background(), ch, chatReq(), nil)
|
||||
el := time.Since(t0)
|
||||
if err != nil {
|
||||
t.Fatalf("chain: %v", err)
|
||||
}
|
||||
if src != "s3" {
|
||||
t.Fatalf("src = %q, want s3 (downgrade after bounded wait)", src)
|
||||
}
|
||||
if el > time.Second {
|
||||
t.Fatalf("busy wait not bounded: %v", el)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChainQuotaExhausted(t *testing.T) {
|
||||
a, b := fakeProv("s1", "a"), fakeProv("s2", "b")
|
||||
ch := BuildChain([]Rule{
|
||||
{Tier: 0, Model: "a", Source: "s1", Quota: 100, Period: "hour"},
|
||||
{Tier: 0, Model: "b", Source: "s2"},
|
||||
}, bySource(a, b))
|
||||
s := New(0)
|
||||
exhausted := func(sl *Slot) bool { return sl.Source == "s1" && sl.Quota > 0 }
|
||||
_, src, _, err := s.ChainChat(context.Background(), ch, chatReq(), exhausted)
|
||||
if err != nil {
|
||||
t.Fatalf("chain: %v", err)
|
||||
}
|
||||
if src != "s2" {
|
||||
t.Fatalf("src = %q, want s2 (quota slot dropped)", src)
|
||||
}
|
||||
if a.chatHits.Load() != 0 {
|
||||
t.Fatal("quota-exhausted slot must not be called")
|
||||
}
|
||||
}
|
||||
|
||||
func TestChainErrSummary(t *testing.T) {
|
||||
a, b, c := fakeProv("s1", "a"), fakeProv("s2", "b"), fakeProv("s3", "c")
|
||||
a.fail.Store(true)
|
||||
b.fail.Store(true)
|
||||
c.available.Store(false) // whole tier 1 cooling
|
||||
ch := BuildChain([]Rule{
|
||||
{Tier: 0, Model: "a", Source: "s1"},
|
||||
{Tier: 0, Model: "b", Source: "s2"},
|
||||
{Tier: 1, Model: "c", Source: "s3"},
|
||||
}, bySource(a, b, c))
|
||||
s := New(0)
|
||||
_, _, _, err := s.ChainChat(context.Background(), ch, chatReq(), nil)
|
||||
var ce *ChainErr
|
||||
if !errors.As(err, &ce) {
|
||||
t.Fatalf("err = %v, want *ChainErr", err)
|
||||
}
|
||||
if len(ce.Tiers) != 2 || ce.Tiers[0].Source != "s1" || ce.Tiers[1].Source != "s2" {
|
||||
t.Fatalf("tiers = %+v", ce.Tiers)
|
||||
}
|
||||
if len(ce.Skipped) != 1 {
|
||||
t.Fatalf("skipped = %+v", ce.Skipped)
|
||||
}
|
||||
msg := ce.Error()
|
||||
if !strings.Contains(msg, "s1") || !strings.Contains(msg, "s2") || !strings.Contains(msg, "tier 1: no schedulable slot") {
|
||||
t.Fatalf("summary = %q", msg)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChainStreamFallsBackBeforeFirstChunk(t *testing.T) {
|
||||
a, b := fakeProv("s1", "a"), fakeProv("s2", "b")
|
||||
a.fail.Store(true)
|
||||
ch := BuildChain([]Rule{
|
||||
{Tier: 0, Model: "a", Source: "s1"},
|
||||
{Tier: 0, Model: "b", Source: "s2"},
|
||||
}, bySource(a, b))
|
||||
s := New(0)
|
||||
chunks, src, model, err := s.ChainChatStream(context.Background(), ch, chatReq(), nil)
|
||||
if err != nil {
|
||||
t.Fatalf("chain stream: %v", err)
|
||||
}
|
||||
if src != "s2" || model != "b" {
|
||||
t.Fatalf("src=%q model=%q", src, model)
|
||||
}
|
||||
var text string
|
||||
for ck := range chunks {
|
||||
text += ck.Content
|
||||
}
|
||||
if text != "s2" {
|
||||
t.Fatalf("text = %q", text)
|
||||
}
|
||||
}
|
||||
|
||||
func TestChainResetCooldownAfterSwap(t *testing.T) {
|
||||
a := fakeProv("s1", "a")
|
||||
ch := BuildChain([]Rule{{Tier: 0, Model: "a", Source: "s1"}}, bySource(a))
|
||||
// a freshly built chain must schedule from index 0 (cursor starts at -1)
|
||||
if base := ch.Tiers[0].NextStart(); base != 0 {
|
||||
t.Fatalf("first start = %d, want 0", base)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user