First Commit
ci / go (push) Waiting to run
ci / go-db (agent) (push) Waiting to run
ci / go-db (config) (push) Waiting to run
ci / go-db (db) (push) Waiting to run
ci / go-db (evidence) (push) Waiting to run
ci / go-db (llmrec) (push) Waiting to run
ci / go-db (server) (push) Waiting to run
web / web (push) Waiting to run
docs / links (push) Canceled after 0s
detections / detections (push) Canceled after 0s

This commit is contained in:
dela
2026-10-09 08:38:16 +08:00
commit 0335d572de
756 changed files with 201663 additions and 0 deletions
+184
View File
@@ -0,0 +1,184 @@
// Package llmpool implements LLM failover ("轮询"): a Provider decorator that
// walks an ordered chain of LLM profiles and moves to the next one when the
// current one can't serve the request — out of credit, revoked key, rate-limited
// past the SDK's own retries, or down.
//
// The chain is built by the server from llm_profiles (active profile first, then
// by priority), so the whole process shares ONE chain and ONE circuit-breaker
// registry: when a task discovers a profile is out of credit, every other task
// skips it immediately.
package llmpool
import (
"sync"
"time"
)
// backoff is the cooling-off ladder, indexed by trip count: the 1st trip cools
// for a minute, the 2nd for five, everything after that for half an hour. A
// profile that is merely rate-limited recovers quickly; one that keeps failing
// stops being probed every round.
var backoff = []time.Duration{time.Minute, 5 * time.Minute, 30 * time.Minute}
// State is one profile's circuit-breaker state.
type State struct {
Fails int // consecutive failures; reset by a success
Trips int // total trips, indexes the backoff ladder
OpenUntil time.Time // zero / past = closed (usable)
LastError string
LastAt time.Time
}
// Open reports whether the breaker is currently open (profile should be skipped).
func (s State) Open() bool { return time.Now().Before(s.OpenUntil) }
// softTripAfter is the DEFAULT number of consecutive TRANSIENT failures (429 /
// 5xx / network) that trip the breaker. Deterministic failures (no credit, bad
// key) trip on the first regardless. Overridable via Registry.SetPolicy.
const softTripAfter = 3
// Registry holds the circuit-breaker state of every profile, keyed by profile id.
// It is process-wide and outlives any single Pool instance, so rebuilding the pool
// (saving an unrelated profile, toggling a setting) never clears what we've learned
// about which backends are broken.
type Registry struct {
mu sync.Mutex
m map[int64]*State
// persist / forget mirror state to the DB so a cooling-off window survives a
// restart. Both may be nil (no DB); both are called off the hot path.
persist func(id int64, st State)
forget func(id int64)
// softTrip / cooldown are the operator-set overrides (0 = use the built-in
// default / ladder). They live here rather than being read per failure because
// Trip runs on the failure path of every request.
softTrip int
cooldown time.Duration
}
// SetPolicy overrides the breaker's two knobs. softTrip: how many consecutive
// transient failures trip it (0 = default softTripAfter; negative = transient
// failures never trip it, leaving only the deterministic ones). cooldown: a fixed
// cooling-off window (0 = the 1min/5min/30min ladder). Safe to call at any time.
func (r *Registry) SetPolicy(softTrip int, cooldown time.Duration) {
r.mu.Lock()
defer r.mu.Unlock()
r.softTrip, r.cooldown = softTrip, cooldown
}
// tripAfter is the effective consecutive-transient-failure threshold. Callers
// hold r.mu. A negative override yields 0, which Trip reads as "never soft-trip".
func (r *Registry) tripAfter() int {
if r.softTrip == 0 {
return softTripAfter
}
return max(r.softTrip, 0)
}
// coolFor is the cooling-off window for the trips-th trip (1-based): the fixed
// override when set, else the ladder (last rung repeats). Callers hold r.mu.
func (r *Registry) coolFor(trips int) time.Duration {
if r.cooldown > 0 {
return r.cooldown
}
if trips-1 < len(backoff) {
return backoff[trips-1]
}
return backoff[len(backoff)-1]
}
// NewRegistry builds an empty registry. persist/forget may be nil.
func NewRegistry(persist func(id int64, st State), forget func(id int64)) *Registry {
return &Registry{m: map[int64]*State{}, persist: persist, forget: forget}
}
// Restore seeds state loaded from the DB at startup (bypasses persistence).
func (r *Registry) Restore(id int64, st State) {
r.mu.Lock()
defer r.mu.Unlock()
cp := st
r.m[id] = &cp
}
// Get returns a copy of one profile's state.
func (r *Registry) Get(id int64) State {
r.mu.Lock()
defer r.mu.Unlock()
if st := r.m[id]; st != nil {
return *st
}
return State{}
}
// IsOpen reports whether a profile is in its cooling-off window.
func (r *Registry) IsOpen(id int64) bool { return r.Get(id).Open() }
// Pass records a successful call: clears the failure counters so a profile that
// recovers is fully trusted again (and drops the persisted row).
func (r *Registry) Pass(id int64) {
r.mu.Lock()
st := r.m[id]
if st == nil || (st.Fails == 0 && st.Trips == 0 && st.OpenUntil.IsZero()) {
r.mu.Unlock()
return // already clean — nothing to write
}
delete(r.m, id)
forget := r.forget
r.mu.Unlock()
if forget != nil {
forget(id)
}
}
// Trip records a failed call. hard=true marks a deterministic failure (no credit,
// invalid key, missing model) which opens the breaker immediately; hard=false is a
// transient one (429 / 5xx / network) that needs softTripAfter in a row. Returns
// true when this call is what opened the breaker, so the caller can log it once.
func (r *Registry) Trip(id int64, errMsg string, hard bool) (tripped bool) {
r.mu.Lock()
st := r.m[id]
if st == nil {
st = &State{}
r.m[id] = st
}
st.Fails++
st.LastError = errMsg
st.LastAt = time.Now()
softTrip := r.tripAfter()
if hard || (softTrip > 0 && st.Fails >= softTrip) {
st.Trips++
st.OpenUntil = time.Now().Add(r.coolFor(st.Trips))
st.Fails = 0 // counted into this trip; start fresh for the half-open probe
tripped = true
}
snap := *st
persist := r.persist
r.mu.Unlock()
if persist != nil {
persist(id, snap)
}
return tripped
}
// Reset clears one profile's state — the UI's "立即恢复" action.
func (r *Registry) Reset(id int64) {
r.mu.Lock()
delete(r.m, id)
forget := r.forget
r.mu.Unlock()
if forget != nil {
forget(id)
}
}
// Snapshot returns a copy of all tracked state, for the status API.
func (r *Registry) Snapshot() map[int64]State {
r.mu.Lock()
defer r.mu.Unlock()
out := make(map[int64]State, len(r.m))
for id, st := range r.m {
out[id] = *st
}
return out
}
+67
View File
@@ -0,0 +1,67 @@
package llmpool
import (
"testing"
"time"
)
// A configured soft-trip threshold replaces the default: two transient failures
// are enough, and the fixed cooldown replaces the 1/5/30min ladder.
func TestSetPolicyOverridesThresholdAndCooldown(t *testing.T) {
reg := NewRegistry(nil, nil)
reg.SetPolicy(2, 90*time.Second)
if reg.Trip(1, "429", false) {
t.Fatal("tripped on the first transient failure, want the second")
}
if !reg.Trip(1, "429", false) {
t.Fatal("should trip on transient failure #2")
}
st := reg.Get(1)
if d := time.Until(st.OpenUntil); d < 80*time.Second || d > 90*time.Second {
t.Fatalf("cooldown=%v, want ~90s", d)
}
// Every later trip keeps the same fixed window instead of climbing the ladder.
reg.Trip(1, "429", true)
st = reg.Get(1)
if d := time.Until(st.OpenUntil); d < 80*time.Second || d > 90*time.Second {
t.Fatalf("second cooldown=%v, want ~90s (fixed)", d)
}
}
// A negative threshold turns off soft tripping entirely: transient failures never
// open the breaker, deterministic ones still do immediately.
func TestSetPolicyDisablesSoftTrip(t *testing.T) {
reg := NewRegistry(nil, nil)
reg.SetPolicy(-1, 0)
for i := range 10 {
if reg.Trip(1, "429", false) {
t.Fatalf("transient failure #%d tripped the breaker, want never", i+1)
}
}
if reg.IsOpen(1) {
t.Fatal("breaker should stay closed for transient failures")
}
if !reg.Trip(1, "no credit", true) {
t.Fatal("a hard failure must still trip immediately")
}
if d := time.Until(reg.Get(1).OpenUntil); d < 50*time.Second || d > 60*time.Second {
t.Fatalf("cooldown=%v, want the default first rung (~1min)", d)
}
}
// The zero policy is the historical behaviour: 3 transient failures, ladder cooldown.
func TestZeroPolicyKeepsDefaults(t *testing.T) {
reg := NewRegistry(nil, nil)
reg.SetPolicy(0, 0)
for i := 1; i < softTripAfter; i++ {
if reg.Trip(1, "429", false) {
t.Fatalf("tripped after %d transient failures, want %d", i, softTripAfter)
}
}
if !reg.Trip(1, "429", false) {
t.Fatalf("should trip on failure #%d", softTripAfter)
}
if d := time.Until(reg.Get(1).OpenUntil); d < 50*time.Second || d > 60*time.Second {
t.Fatalf("cooldown=%v, want the first ladder rung (~1min)", d)
}
}
+299
View File
@@ -0,0 +1,299 @@
package llmpool
import (
"context"
"errors"
"fmt"
"iter"
"log"
"regexp"
"strconv"
"strings"
"sync/atomic"
"github.com/Autumn-27/norma/llm"
)
// Member is one LLM profile in the chain, already built into a provider (wrapped
// with the recorder by the caller, so a failed attempt is still recorded under
// its own profile name).
type Member struct {
ID int64 // llm_profiles.id
Name string // profile name, for logs / UI
Model string
Format string // "anthropic" | "openai"
Priority int // the profile's configured priority (display only)
Active bool // is_default (display only)
// Rank is the ordering key the caller assigned: higher goes first, and members
// sharing a Rank take turns leading (load-spreading across duplicate keys).
// The caller encodes "active profile heads the chain" as a Rank above every
// user-settable priority, so this type needs no policy of its own.
Rank int
// WindowTokens is the profile's context window in tokens. A member whose
// window can't hold the request is skipped rather than made to fail on it.
WindowTokens int
Prov llm.Provider
}
// RankActive is the Rank the caller gives the chain head (the active profile, or
// an explicitly bound one) so it always outranks any configured priority.
const RankActive = int(^uint(0)>>1) - 1
// ErrExhausted is returned when every member of the chain failed.
var ErrExhausted = errors.New("LLM 폴백 체인: 사용 가능한 설정이 없습니다")
// Pool is an llm.Provider that fails over across an ordered chain of members.
// It is safe for concurrent use: members are immutable after construction and
// all mutable state lives in the shared Registry.
type Pool struct {
members []*Member // in chain order (active first, then priority DESC)
health *Registry
rr atomic.Uint64 // rotates the starting point within an equal-priority group
}
// New builds a Pool over members (already in chain order). Returns nil when the
// chain is empty. A single-member chain is still a valid Pool — it just behaves
// exactly like the bare provider.
func New(members []*Member, health *Registry) *Pool {
if len(members) == 0 {
return nil
}
if health == nil {
health = NewRegistry(nil, nil)
}
return &Pool{members: members, health: health}
}
// Members returns the chain in order (read-only).
func (p *Pool) Members() []*Member { return p.members }
// Head returns the first member of the chain.
func (p *Pool) Head() *Member { return p.members[0] }
// Stream implements llm.Provider with failover.
//
// The one hard rule: a member may only be abandoned BEFORE it has yielded any
// event. Once text or a tool_use has reached the caller, re-sending the same
// request to another model would duplicate output and corrupt the conversation
// history — so a mid-stream failure is surfaced as-is and left to the agent
// harness's resume logic. Fortunately the failures this exists for (402 no
// credit, 401 bad key, 429, 5xx) all surface during request establishment,
// before the body is read, so they always land in the safe window.
func (p *Pool) Stream(ctx context.Context, req llm.CompletionRequest) iter.Seq2[llm.StreamEvent, error] {
order := p.order(req)
return func(yield func(llm.StreamEvent, error) bool) {
var lastErr error
for i, m := range order {
emitted := false
var failed error
for ev, err := range m.Prov.Stream(ctx, req) {
if err != nil && !emitted && shouldFailover(ctx, err) {
failed = err
break // safe window: nothing reached the caller yet
}
emitted = true
if !yield(ev, err) {
return // caller stopped consuming (cancel / early exit)
}
if err != nil {
return // terminal error already handed to the caller
}
}
if failed == nil {
p.health.Pass(m.ID) // completed (or failed in a non-failover way)
return
}
lastErr = failed
hard := isHardFailure(failed)
if p.health.Trip(m.ID, trimErr(failed), hard) {
log.Printf("[llmpool] 配置 %q(%s) 已熔断:%s", m.Name, m.Model, trimErr(failed))
}
if i+1 < len(order) {
n := order[i+1]
log.Printf("[llmpool] LLM 故障转移:%q(%s) → %q(%s),原因:%s",
m.Name, m.Model, n.Name, n.Model, trimErr(failed))
}
}
if lastErr == nil {
lastErr = ErrExhausted
}
log.Printf("[llmpool] 轮询链已耗尽(%d 个配置全部失败),最后错误:%s", len(order), trimErr(lastErr))
yield(llm.StreamEvent{}, fmt.Errorf("%w: %v", ErrExhausted, lastErr))
}
}
// Complete implements llm.Provider with failover for non-streaming calls. A
// non-streaming request is atomic — it never delivers partial output — so every
// failover-eligible failure lands in the safe window and the next member can be
// tried without risk of duplicated output. Mirrors Stream's health-tripping and
// chain-exhaustion behavior.
func (p *Pool) Complete(ctx context.Context, req llm.CompletionRequest) (llm.Message, string, llm.Usage, error) {
order := p.order(req)
var lastErr error
for i, m := range order {
msg, sr, usage, err := m.Prov.Complete(ctx, req)
if err == nil {
p.health.Pass(m.ID)
return msg, sr, usage, nil
}
if !shouldFailover(ctx, err) {
// Non-failover error (e.g. ctx cancel, deterministic 4xx): surface as-is
// without tripping health, matching Stream's non-failover path.
p.health.Pass(m.ID)
return llm.Message{}, "", llm.Usage{}, err
}
lastErr = err
hard := isHardFailure(err)
if p.health.Trip(m.ID, trimErr(err), hard) {
log.Printf("[llmpool] 配置 %q(%s) 已熔断:%s", m.Name, m.Model, trimErr(err))
}
if i+1 < len(order) {
n := order[i+1]
log.Printf("[llmpool] LLM 故障转移:%q(%s) → %q(%s),原因:%s",
m.Name, m.Model, n.Name, n.Model, trimErr(err))
}
}
if lastErr == nil {
lastErr = ErrExhausted
}
log.Printf("[llmpool] 轮询链已耗尽(%d 个配置全部失败),最后错误:%s", len(order), trimErr(lastErr))
return llm.Message{}, "", llm.Usage{}, fmt.Errorf("%w: %v", ErrExhausted, lastErr)
}
// order picks the members to try, in order: skip those in a cooling-off window
// and those whose context window can't hold this request, then rotate within each
// equal-priority group so same-priority profiles share the load. Never returns an
// empty slice — if everything is filtered out, the head of the chain is tried
// anyway, since stalling the engine is worse than one more failed request.
func (p *Pool) order(req llm.CompletionRequest) []*Member {
est := estimateTokens(req)
var open []*Member
for _, m := range p.members {
if p.health.IsOpen(m.ID) {
continue
}
if m.WindowTokens > 0 && est > m.WindowTokens {
continue // would 400 on length — not a useful failover target
}
open = append(open, m)
}
if len(open) == 0 {
return p.members[:1] // last resort: probe the head rather than stall
}
return rotateGroups(open, p.rr.Add(1)-1)
}
// rotateGroups rotates each run of equal-ranked members by n, so profiles sharing
// a priority take turns going first (free load-spreading across duplicate keys).
// The chain head has its own Rank and always stays at the front.
func rotateGroups(in []*Member, n uint64) []*Member {
out := make([]*Member, 0, len(in))
for i := 0; i < len(in); {
j := i + 1
for j < len(in) && in[j].Rank == in[i].Rank {
j++
}
g := in[i:j]
if len(g) > 1 {
off := int(n % uint64(len(g)))
for k := range g {
out = append(out, g[(k+off)%len(g)])
}
} else {
out = append(out, g...)
}
i = j
}
return out
}
// statusRe pulls the HTTP status out of the SDK's error text, which is formatted
// as "<prefix>: status <code>: <body>" (norma/llm/retry.go). The SDK exposes no
// typed error, so the string is what we have.
var statusRe = regexp.MustCompile(`status (\d{3})`)
// statusOf returns the HTTP status carried by err, or 0 if it isn't one.
func statusOf(err error) int {
m := statusRe.FindStringSubmatch(err.Error())
if m == nil {
return 0
}
code, _ := strconv.Atoi(m[1])
return code
}
// shouldFailover reports whether err justifies trying the next profile.
//
// Never fails over on context cancellation — that's the user stopping a task or a
// task-level timeout, and burning a backup key on it would both waste credit and
// pollute the run's termination diagnosis. Never on 400 either: a malformed or
// over-long request fails identically everywhere.
func shouldFailover(ctx context.Context, err error) bool {
if err == nil {
return false
}
if ctx.Err() != nil || errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) {
return false
}
switch code := statusOf(err); {
case code == 0:
return true // no status → transport-level failure (reset / DNS / timeout)
case code == 400:
return false // bad or over-long request: identical everywhere
case code == 401, code == 402, code == 403, code == 404, code == 408, code == 429:
return true
case code >= 500:
return true
}
return false
}
// isHardFailure reports whether the failure is deterministic (the profile will
// keep failing until a human fixes it) rather than transient. Hard failures open
// the breaker on the first occurrence.
func isHardFailure(err error) bool {
switch statusOf(err) {
case 401, 402, 403, 404:
return true
}
return false
}
// trimErr shortens an error for logs/UI — provider bodies can be long.
func trimErr(err error) string {
s := strings.TrimSpace(strings.ReplaceAll(err.Error(), "\n", " "))
if len(s) > 300 {
s = s[:300] + "…"
}
return s
}
// estimateTokens roughly sizes a request so members whose context window clearly
// can't hold it are skipped. Deliberately crude (~3.5 chars/token) and biased to
// over-estimate slightly; it only needs to tell "fits" from "nowhere near".
func estimateTokens(req llm.CompletionRequest) int {
n := 0
for _, s := range req.System {
n += len(s)
}
for _, m := range req.Messages {
n += blocksLen(m.Content)
}
for _, t := range req.Tools {
// InputSchema is a decoded map, so its serialized size isn't available
// cheaply — charge a flat ~120 chars per top-level property instead.
n += len(t.Name) + len(t.Description) + len(t.InputSchema)*120
}
return n * 2 / 7 // ≈ len/3.5
}
func blocksLen(bs []llm.ContentBlock) int {
n := 0
for _, b := range bs {
n += len(b.Text) + len(b.Thinking) + len(b.Input) + len(b.Name)
if len(b.Content) > 0 {
n += blocksLen(b.Content)
}
}
return n
}
+90
View File
@@ -0,0 +1,90 @@
package llmpool
import (
"context"
"errors"
"strings"
"testing"
"unicode"
"github.com/Autumn-27/norma/llm"
)
// hasHangul reports whether s contains any Hangul character.
func hasHangul(s string) bool {
for _, r := range s {
if unicode.Is(unicode.Hangul, r) {
return true
}
}
return false
}
// hasHanzi reports whether s contains any CJK (Han) character.
func hasHanzi(s string) bool {
for _, r := range s {
if unicode.Is(unicode.Han, r) {
return true
}
}
return false
}
// assertSurfacedExhaustion checks the error an exhausted chain hands back: it must
// still wrap the ErrExhausted sentinel (callers rely on errors.Is), read as a
// Korean message (no Han characters, no full-width colon), and lead with the
// localized sentinel text. failProv's body is ASCII ("anthropic: status 402:
// nope"), so the whole surfaced string must be Han-free.
func assertSurfacedExhaustion(t *testing.T, label string, err error) {
t.Helper()
if err == nil {
t.Fatalf("%s: exhausted chain returned nil error", label)
}
if !errors.Is(err, ErrExhausted) {
t.Fatalf("%s: surfaced error no longer wraps ErrExhausted: %v", label, err)
}
msg := err.Error()
if hasHanzi(msg) {
t.Errorf("%s: surfaced error still carries Chinese characters: %q", label, msg)
}
if strings.ContainsRune(msg, ':') {
t.Errorf("%s: surfaced error still uses a full-width colon: %q", label, msg)
}
if !strings.HasPrefix(msg, ErrExhausted.Error()) {
t.Errorf("%s: surfaced error does not lead with the localized sentinel: %q", label, msg)
}
}
// The chain-exhaustion error is user-facing: a worker whose whole LLM chain fails
// records it through agent/capture.go as a "result" activity shown in the run
// transcript. So its text must be Korean, while its identity (errors.Is) must be
// preserved for callers that branch on the sentinel.
func TestExhaustedErrorLocalized(t *testing.T) {
assertKoreanErrText(t, "ErrExhausted", ErrExhausted.Error())
// Drive both surfaced paths so a future edit to either wrap site is caught.
streamChain := New([]*Member{
member(1, "a", 10, failProv("a", 402)),
member(2, "b", 5, failProv("b", 402)),
}, NewRegistry(nil, nil))
_, serr := drain(streamChain.Stream(context.Background(), llm.CompletionRequest{}))
assertSurfacedExhaustion(t, "Stream", serr)
completeChain := New([]*Member{
member(1, "a", 10, failProv("a", 402)),
member(2, "b", 5, failProv("b", 402)),
}, NewRegistry(nil, nil))
_, _, _, cerr := completeChain.Complete(context.Background(), llm.CompletionRequest{})
assertSurfacedExhaustion(t, "Complete", cerr)
}
// assertKoreanErrText fails unless s contains Hangul and no Han character.
func assertKoreanErrText(t *testing.T, label, s string) {
t.Helper()
if !hasHangul(s) {
t.Errorf("%s: 한글이 없습니다: %q", label, s)
}
if hasHanzi(s) {
t.Errorf("%s: 중국어 한자가 남아 있습니다: %q", label, s)
}
}
+370
View File
@@ -0,0 +1,370 @@
package llmpool
import (
"context"
"errors"
"fmt"
"iter"
"strings"
"testing"
"time"
"github.com/Autumn-27/norma/llm"
)
// fakeProv is a scripted provider: script[i] is what the i-th call yields —
// some events, then optionally an error.
type fakeProv struct {
name string
calls int
events [][]llm.StreamEvent // events emitted before the error, per call
errs []error // error to end each call with (nil = clean finish)
}
// at returns script entry n, repeating the last one once the script runs out, so
// a provider defined as "always succeeds" / "always 402" keeps behaving that way
// across repeated calls.
func at[T any](s []T, n int) (T, bool) {
var zero T
if len(s) == 0 {
return zero, false
}
if n >= len(s) {
n = len(s) - 1
}
return s[n], true
}
func (f *fakeProv) Stream(ctx context.Context, req llm.CompletionRequest) iter.Seq2[llm.StreamEvent, error] {
n := f.calls
f.calls++
return func(yield func(llm.StreamEvent, error) bool) {
if evs, ok := at(f.events, n); ok {
for _, e := range evs {
if !yield(e, nil) {
return
}
}
}
err, _ := at(f.errs, n)
if err != nil {
yield(llm.StreamEvent{}, err)
return
}
yield(llm.StreamEvent{Type: llm.SEMessageStop}, nil)
}
}
func (f *fakeProv) Complete(ctx context.Context, req llm.CompletionRequest) (llm.Message, string, llm.Usage, error) {
acc := llm.NewAccumulator()
for ev, err := range f.Stream(ctx, req) {
if err != nil {
return llm.Message{}, "", llm.Usage{}, err
}
acc.Add(ev)
}
return acc.Message(), acc.StopReason, acc.Usage, nil
}
// ok builds a provider that always succeeds with one text delta.
func okProv(name, text string) *fakeProv {
return &fakeProv{name: name, events: [][]llm.StreamEvent{{{Type: llm.SETextDelta, Text: text}}}}
}
// failProv builds a provider that fails immediately (before any event) with status.
func failProv(name string, status int) *fakeProv {
return &fakeProv{name: name, errs: []error{fmt.Errorf("anthropic: status %d: nope", status)}}
}
func member(id int64, name string, rank int, p llm.Provider) *Member {
return &Member{ID: id, Name: name, Model: "m" + name, Rank: rank, Prov: p}
}
// drain consumes a stream, returning the concatenated text and terminal error.
func drain(seq iter.Seq2[llm.StreamEvent, error]) (string, error) {
var sb strings.Builder
for ev, err := range seq {
if err != nil {
return sb.String(), err
}
if ev.Type == llm.SETextDelta {
sb.WriteString(ev.Text)
}
}
return sb.String(), nil
}
func TestFailoverOnNoCredit(t *testing.T) {
a, b := failProv("a", 402), okProv("b", "hello")
p := New([]*Member{member(1, "a", 10, a), member(2, "b", 5, b)}, NewRegistry(nil, nil))
got, err := drain(p.Stream(context.Background(), llm.CompletionRequest{}))
if err != nil {
t.Fatalf("expected failover to succeed, got %v", err)
}
if got != "hello" {
t.Fatalf("text = %q, want %q", got, "hello")
}
if a.calls != 1 || b.calls != 1 {
t.Fatalf("calls: a=%d b=%d, want 1/1", a.calls, b.calls)
}
}
// A 402 is deterministic: one failure must open the breaker, so the NEXT request
// skips that profile entirely instead of paying for another round-trip.
func TestHardFailureTripsBreakerImmediately(t *testing.T) {
a, b := failProv("a", 402), okProv("b", "x")
reg := NewRegistry(nil, nil)
p := New([]*Member{member(1, "a", 10, a), member(2, "b", 5, b)}, reg)
_, _ = drain(p.Stream(context.Background(), llm.CompletionRequest{}))
if !reg.IsOpen(1) {
t.Fatal("402 should have tripped the breaker on the first failure")
}
_, _ = drain(p.Stream(context.Background(), llm.CompletionRequest{}))
if a.calls != 1 {
t.Fatalf("tripped profile was called again: calls=%d, want 1", a.calls)
}
if b.calls != 2 {
t.Fatalf("fallback calls=%d, want 2", b.calls)
}
}
// 429 is transient: the SDK already retried, but we shouldn't write a profile off
// until it fails repeatedly.
func TestSoftFailureNeedsRepeats(t *testing.T) {
reg := NewRegistry(nil, nil)
for i := 1; i < softTripAfter; i++ {
if reg.Trip(1, "429", false) {
t.Fatalf("tripped after %d transient failures, want %d", i, softTripAfter)
}
}
if !reg.Trip(1, "429", false) {
t.Fatalf("should trip on failure #%d", softTripAfter)
}
if !reg.IsOpen(1) {
t.Fatal("breaker should be open")
}
}
// A success must fully clear the counters, so an intermittent profile never
// accumulates its way to a trip.
func TestPassResetsCounters(t *testing.T) {
reg := NewRegistry(nil, nil)
reg.Trip(1, "429", false)
reg.Trip(1, "429", false)
reg.Pass(1)
if got := reg.Get(1).Fails; got != 0 {
t.Fatalf("fails=%d after Pass, want 0", got)
}
if reg.Trip(1, "429", false) {
t.Fatal("tripped immediately after a success — counters were not reset")
}
}
func TestBackoffLadderGrows(t *testing.T) {
reg := NewRegistry(nil, nil)
var prev time.Duration
for i := range 4 {
reg.Reset(1)
st := State{Trips: i}
reg.Restore(1, st)
reg.Trip(1, "402", true)
d := time.Until(reg.Get(1).OpenUntil)
// Non-decreasing, with a second of slack: the ladder plateaus at its last
// rung, and each Trip stamps its own time.Now().
if i > 0 && d < prev-time.Second {
t.Fatalf("trip #%d cools for %v, shorter than the previous %v", i+1, d, prev)
}
prev = d
}
if prev < 25*time.Minute {
t.Fatalf("ladder tops out at %v, want ~30m", prev)
}
}
// The safety rule: once output has reached the caller, a mid-stream failure must
// NOT be retried on another model — that would duplicate the assistant turn.
func TestNoFailoverAfterEmit(t *testing.T) {
a := &fakeProv{
name: "a",
events: [][]llm.StreamEvent{{{Type: llm.SETextDelta, Text: "partial"}}},
errs: []error{fmt.Errorf("anthropic: status 500: mid-stream drop")},
}
b := okProv("b", "full")
p := New([]*Member{member(1, "a", 10, a), member(2, "b", 5, b)}, NewRegistry(nil, nil))
got, err := drain(p.Stream(context.Background(), llm.CompletionRequest{}))
if err == nil {
t.Fatal("mid-stream error should surface, not be swallowed by a failover")
}
if got != "partial" {
t.Fatalf("text = %q, want the partial output %q", got, "partial")
}
if b.calls != 0 {
t.Fatalf("fell over to the backup after emitting output (calls=%d) — would duplicate the turn", b.calls)
}
}
// Cancelling a task must not burn a backup key, and must not be diagnosed as an
// LLM fault.
func TestNoFailoverOnCancel(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
cancel()
a := &fakeProv{name: "a", errs: []error{context.Canceled}}
b := okProv("b", "x")
reg := NewRegistry(nil, nil)
p := New([]*Member{member(1, "a", 10, a), member(2, "b", 5, b)}, reg)
_, err := drain(p.Stream(ctx, llm.CompletionRequest{}))
if !errors.Is(err, context.Canceled) {
t.Fatalf("err = %v, want context.Canceled", err)
}
if b.calls != 0 {
t.Fatalf("failed over on cancellation (backup calls=%d)", b.calls)
}
if reg.Get(1).Fails != 0 {
t.Fatal("cancellation counted as a profile failure")
}
}
func TestShouldFailoverByStatus(t *testing.T) {
cases := []struct {
status int
want bool
}{
{400, false}, // bad/over-long request fails identically everywhere
{401, true}, {402, true}, {403, true}, {404, true},
{408, true}, {429, true}, {500, true}, {503, true},
{200, false},
}
for _, c := range cases {
err := fmt.Errorf("openai: status %d: body", c.status)
if got := shouldFailover(context.Background(), err); got != c.want {
t.Errorf("status %d: shouldFailover=%v, want %v", c.status, got, c.want)
}
}
// A transport error carries no status and must fail over.
if !shouldFailover(context.Background(), errors.New("dial tcp: connection reset by peer")) {
t.Error("network error should fail over")
}
}
func TestHardVsSoftClassification(t *testing.T) {
for _, s := range []int{401, 402, 403, 404} {
if !isHardFailure(fmt.Errorf("status %d: x", s)) {
t.Errorf("status %d should be a hard failure", s)
}
}
for _, s := range []int{408, 429, 500, 502, 503} {
if isHardFailure(fmt.Errorf("status %d: x", s)) {
t.Errorf("status %d should be transient, not hard", s)
}
}
}
// A member whose context window can't hold the request is a guaranteed 400 —
// skip it rather than spend a round-trip proving it.
func TestSkipsMembersTooSmallForRequest(t *testing.T) {
small, big := okProv("small", "s"), okProv("big", "b")
ms := member(1, "small", 10, small)
ms.WindowTokens = 1000
mb := member(2, "big", 5, big)
mb.WindowTokens = 1_000_000
p := New([]*Member{ms, mb}, NewRegistry(nil, nil))
req := llm.CompletionRequest{Messages: []llm.Message{llm.UserText(strings.Repeat("x", 100_000))}}
got, err := drain(p.Stream(context.Background(), req))
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if got != "b" {
t.Fatalf("text=%q, want the large-window member to serve it", got)
}
if small.calls != 0 {
t.Fatalf("sent a request to a member that can't hold it (calls=%d)", small.calls)
}
}
// Everything tripped: probing the head beats stalling the engine outright.
func TestAllTrippedStillProbesHead(t *testing.T) {
a, b := okProv("a", "a"), okProv("b", "b")
reg := NewRegistry(nil, nil)
reg.Trip(1, "402", true)
reg.Trip(2, "402", true)
p := New([]*Member{member(1, "a", 10, a), member(2, "b", 5, b)}, reg)
got, err := drain(p.Stream(context.Background(), llm.CompletionRequest{}))
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if got != "a" || a.calls != 1 {
t.Fatalf("text=%q a.calls=%d — want the head probed as a last resort", got, a.calls)
}
}
func TestExhaustedChainReportsClearly(t *testing.T) {
a, b := failProv("a", 402), failProv("b", 402)
p := New([]*Member{member(1, "a", 10, a), member(2, "b", 5, b)}, NewRegistry(nil, nil))
_, err := drain(p.Stream(context.Background(), llm.CompletionRequest{}))
if !errors.Is(err, ErrExhausted) {
t.Fatalf("err = %v, want it to wrap ErrExhausted", err)
}
}
// Equal-rank members take turns leading, so duplicate keys share the load.
func TestEqualRankRotates(t *testing.T) {
a, b := okProv("a", "a"), okProv("b", "b")
head := member(1, "head", RankActive, failProv("head", 402))
p := New([]*Member{head, member(2, "a", 5, a), member(3, "b", 5, b)}, NewRegistry(nil, nil))
first := map[string]int{}
for range 4 {
txt, _ := drain(p.Stream(context.Background(), llm.CompletionRequest{}))
first[txt]++
}
if first["a"] == 0 || first["b"] == 0 {
t.Fatalf("same-rank members did not rotate: %v", first)
}
}
// The head keeps its own rank and never joins a rotation group.
func TestHeadAlwaysFirst(t *testing.T) {
head := okProv("head", "H")
p := New([]*Member{
member(1, "head", RankActive, head),
member(2, "a", 5, okProv("a", "a")),
member(3, "b", 5, okProv("b", "b")),
}, NewRegistry(nil, nil))
for range 5 {
if txt, _ := drain(p.Stream(context.Background(), llm.CompletionRequest{})); txt != "H" {
t.Fatalf("head was not tried first, got %q", txt)
}
}
}
// A single-member chain must behave exactly like the bare provider.
func TestSingleMemberPassthrough(t *testing.T) {
a := okProv("a", "solo")
p := New([]*Member{member(1, "a", RankActive, a)}, NewRegistry(nil, nil))
got, err := drain(p.Stream(context.Background(), llm.CompletionRequest{}))
if err != nil || got != "solo" {
t.Fatalf("got %q / %v, want solo / nil", got, err)
}
}
func TestNewEmptyChainIsNil(t *testing.T) {
if New(nil, nil) != nil {
t.Fatal("empty chain should yield a nil pool so callers use the bare provider")
}
}
// An expired cooling-off window must not keep a profile out of the chain.
func TestExpiredWindowIsClosed(t *testing.T) {
reg := NewRegistry(nil, nil)
reg.Restore(1, State{Trips: 1, OpenUntil: time.Now().Add(-time.Second)})
if reg.IsOpen(1) {
t.Fatal("expired window should read as closed")
}
}