From 8cffc87c8b92ac4029b69ab106f9a9391d36a3ce Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 15:40:35 +0000 Subject: [PATCH] feat(failover): add chain manager with trip, fail-back and pins Health is tracked per target and the active target per chain. Fail-back waits for recovery probes and a minimum time on the fallback. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/services/failover/fakes_test.go | 95 ++++ core/services/failover/manager.go | 619 +++++++++++++++++++++++++ core/services/failover/manager_test.go | 232 +++++++++ 3 files changed, 946 insertions(+) create mode 100644 core/services/failover/fakes_test.go create mode 100644 core/services/failover/manager.go create mode 100644 core/services/failover/manager_test.go diff --git a/core/services/failover/fakes_test.go b/core/services/failover/fakes_test.go new file mode 100644 index 000000000..cad6b3552 --- /dev/null +++ b/core/services/failover/fakes_test.go @@ -0,0 +1,95 @@ +package failover + +import ( + "sort" + "sync" + "time" + + "github.com/mudler/LocalAI/core/config" +) + +type fakeClock struct { + mu sync.Mutex + now time.Time +} + +func newFakeClock() *fakeClock { return &fakeClock{now: time.Date(2026, 9, 26, 10, 0, 0, 0, time.UTC)} } +func (c *fakeClock) Now() time.Time { + c.mu.Lock() + defer c.mu.Unlock() + return c.now +} +func (c *fakeClock) Advance(d time.Duration) { + c.mu.Lock() + c.now = c.now.Add(d) + c.mu.Unlock() +} + +type fakeSource struct { + mu sync.Mutex + cfgs map[string]config.ModelConfig +} + +func newFakeSource(cfgs ...config.ModelConfig) *fakeSource { + s := &fakeSource{cfgs: map[string]config.ModelConfig{}} + for _, c := range cfgs { + s.cfgs[c.Name] = c + } + return s +} +func (s *fakeSource) Put(c config.ModelConfig) { s.mu.Lock(); s.cfgs[c.Name] = c; s.mu.Unlock() } +func (s *fakeSource) Delete(name string) { s.mu.Lock(); delete(s.cfgs, name); s.mu.Unlock() } +func (s *fakeSource) GetModelConfig(n string) (config.ModelConfig, bool) { + s.mu.Lock() + defer s.mu.Unlock() + c, ok := s.cfgs[n] + return c, ok +} +func (s *fakeSource) GetAllModelsConfigs() []config.ModelConfig { + s.mu.Lock() + defer s.mu.Unlock() + out := make([]config.ModelConfig, 0, len(s.cfgs)) + for _, c := range s.cfgs { + out = append(out, c) + } + sort.Slice(out, func(i, j int) bool { return out[i].Name < out[j].Name }) + return out +} + +func local(name string) config.ModelConfig { + return config.ModelConfig{Name: name, Backend: "llama-cpp"} +} +func remote(name string) config.ModelConfig { + return config.ModelConfig{Name: name, Backend: "cloud-proxy"} +} + +// chainCfg builds a chain; fc may be nil for defaults. +func chainCfg(name string, fc *config.FailoverConfig, targets ...config.FailoverTarget) config.ModelConfig { + f := config.FailoverConfig{} + if fc != nil { + f = *fc + } + f.Targets = targets + return config.ModelConfig{Name: name, Failover: &f} +} + +func t(model string) config.FailoverTarget { return config.FailoverTarget{Model: model} } +func warmT(model string) config.FailoverTarget { + return config.FailoverTarget{Model: model, Warm: true} +} + +// drain returns the events buffered so far without blocking. +func drain(ch <-chan Event) []Event { + var out []Event + for { + select { + case ev, ok := <-ch: + if !ok { + return out + } + out = append(out, ev) + default: + return out + } + } +} diff --git a/core/services/failover/manager.go b/core/services/failover/manager.go new file mode 100644 index 000000000..b38ec9a49 --- /dev/null +++ b/core/services/failover/manager.go @@ -0,0 +1,619 @@ +package failover + +import ( + "context" + "errors" + "fmt" + "slices" + "sort" + "sync" + "sync/atomic" + "time" + + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/xlog" +) + +var ( + ErrChainNotFound = errors.New("failover chain not found") + ErrTargetNotInChain = errors.New("target is not in this failover chain") + ErrNoTarget = errors.New("failover chain has no usable target") +) + +// ConfigSource is the part of ModelConfigLoader the manager reads. +type ConfigSource interface { + GetModelConfig(name string) (config.ModelConfig, bool) + GetAllModelsConfigs() []config.ModelConfig +} + +type Clock interface{ Now() time.Time } + +type realClock struct{} + +func (realClock) Now() time.Time { return time.Now() } + +type Option func(*Manager) + +func WithClock(c Clock) Option { return func(m *Manager) { m.clock = c } } +func WithProber(p Prober) Option { return func(m *Manager) { m.prober = p } } + +// WithOnWarmChanged is called outside the manager lock when the set of warm +// local targets changes. The application pins and preloads them. +func WithOnWarmChanged(fn func(warm []string)) Option { return func(m *Manager) { m.onWarm = fn } } + +// Manager tracks health per target and the active target per chain. +type Manager struct { + mu sync.Mutex + src ConfigSource + clock Clock + prober Prober + onWarm func([]string) + targets map[string]*targetState + chains map[string]*chainState + subs map[int]chan Event + nextSub int + warm []string + warmPending bool + closed bool +} + +type targetState struct { + name string + kind Kind + warm bool + state TargetState + failures []time.Time + consecutiveOK int + downSince time.Time + lastProbe time.Time + lastActivity time.Time + lastError string + // params come from the first chain, in name order, that lists the target. + params config.FailoverConfig +} + +func (ts *targetState) cold() bool { return ts.kind == KindLocal && !ts.warm } + +type chainState struct { + name string + cfg config.FailoverConfig + targets []string + active int + activeSince time.Time + pinned string + state ChainState +} + +func New(src ConfigSource, opts ...Option) *Manager { + m := &Manager{ + src: src, + clock: realClock{}, + targets: map[string]*targetState{}, + chains: map[string]*chainState{}, + subs: map[int]chan Event{}, + } + for _, o := range opts { + o(m) + } + return m +} + +// Sync reconciles chains with the config source. There is no config-change +// hook in the loader, so this runs on every tick and on a lookup miss. +func (m *Manager) Sync() { + m.mu.Lock() + m.syncLocked() + warm, deliver := m.takeWarmLocked() + m.mu.Unlock() + if deliver && m.onWarm != nil { + m.onWarm(warm) + } +} + +func (m *Manager) syncLocked() { + now := m.clock.Now() + seenChains := map[string]bool{} + claimed := map[string]bool{} + for _, c := range m.src.GetAllModelsConfigs() { + if !c.IsFailover() { + continue + } + seenChains[c.Name] = true + names := make([]string, 0, len(c.Failover.Targets)) + for _, t := range c.Failover.Targets { + names = append(names, t.Model) + } + ch := m.chains[c.Name] + if ch == nil || !slices.Equal(ch.targets, names) { + pinned := "" + if ch != nil && slices.Contains(names, ch.pinned) { + pinned = ch.pinned + } + ch = &chainState{name: c.Name, targets: names, activeSince: now, state: ChainPrimary, pinned: pinned} + m.chains[c.Name] = ch + } + ch.cfg = *c.Failover + for _, t := range c.Failover.Targets { + ts := m.targets[t.Model] + if ts == nil { + ts = &targetState{name: t.Model, state: StateHealthy} + m.targets[t.Model] = ts + } + if !claimed[t.Model] { + claimed[t.Model] = true + ts.params = *c.Failover + ts.warm = false + } + tc, ok := m.lookupTarget(t.Model) + if !ok { + m.setTargetLocked(ts, StateMissing, ReasonMissing, "target config not found") + continue + } + ts.kind = KindOf(tc) + if t.Warm && ts.kind == KindLocal { + ts.warm = true + } + if ts.state == StateMissing { + m.setTargetLocked(ts, StateHealthy, ReasonRecovery, "") + } + } + } + for name := range m.chains { + if !seenChains[name] { + delete(m.chains, name) + } + } + for name := range m.targets { + if !claimed[name] { + delete(m.targets, name) + } + } + for _, ch := range m.chains { + m.recomputeLocked(ch, "") + } + var warm []string + for name, ts := range m.targets { + if ts.warm { + warm = append(warm, name) + } + } + sort.Strings(warm) + if !slices.Equal(warm, m.warm) { + m.warm = warm + m.warmPending = true + } +} + +func (m *Manager) takeWarmLocked() ([]string, bool) { + if !m.warmPending { + return nil, false + } + m.warmPending = false + return slices.Clone(m.warm), true +} + +// lookupTarget returns the config that serves a target, one alias hop deep. +func (m *Manager) lookupTarget(name string) (config.ModelConfig, bool) { + c, ok := m.src.GetModelConfig(name) + if ok && c.IsAlias() { + return m.src.GetModelConfig(c.Alias) + } + return c, ok +} + +func (m *Manager) chainLocked(name string) *chainState { + if ch := m.chains[name]; ch != nil { + return ch + } + m.syncLocked() + return m.chains[name] +} + +// targetLocked looks up a target's health state, syncing lazily on a miss so +// ReportFailure/ReportSuccess work before the first Plan or Sync call. +func (m *Manager) targetLocked(name string) *targetState { + if ts := m.targets[name]; ts != nil { + return ts + } + m.syncLocked() + return m.targets[name] +} + +// WarmTargets returns the warm local targets, sorted. +func (m *Manager) WarmTargets() []string { + m.mu.Lock() + defer m.mu.Unlock() + return slices.Clone(m.warm) +} + +// Reevaluate recomputes every chain. Dwell-based fail-back needs no event, so +// the scheduler calls this on every tick. +func (m *Manager) Reevaluate() { + m.mu.Lock() + defer m.mu.Unlock() + for _, ch := range m.chains { + m.recomputeLocked(ch, "") + } +} + +func (m *Manager) setTargetLocked(ts *targetState, to TargetState, reason Reason, errMsg string) { + if ts.state == to { + return + } + from := ts.state + now := m.clock.Now() + ts.state = to + switch to { + case StateDown: + ts.downSince = now + ts.consecutiveOK = 0 + ts.failures = nil + case StateRecovering, StateHealthy: + ts.consecutiveOK = 0 + ts.failures = nil + } + m.emitLocked(Event{Type: EventTargetState, Target: ts.name, From: string(from), To: string(to), Reason: reason, Error: errMsg, At: now}) +} + +// recomputeLocked picks the active target. override replaces the reason of a +// resulting switch (pin and unpin are always "manual"). +func (m *Manager) recomputeLocked(ch *chainState, override Reason) { + now := m.clock.Now() + prev := ch.active + next := prev + reason := ReasonTrip + best := -1 + for i, name := range ch.targets { + if ts := m.targets[name]; ts != nil && ts.state == StateHealthy { + best = i + break + } + } + switch { + case ch.pinned != "": + next = slices.Index(ch.targets, ch.pinned) + reason = ReasonManual + case best == -1: + // Nothing is healthy: keep the active target, Plan tries all of them. + case best > prev: + next = best // the active target is not healthy + case best < prev: + cur := m.targets[ch.targets[prev]] + curHealthy := cur != nil && cur.state == StateHealthy + if !curHealthy { + next = best + } else if now.Sub(ch.activeSince) >= ch.cfg.MinDwell() { + next = best + reason = ReasonRecovery + } + } + if override != "" { + reason = override + } + var state ChainState + switch { + case ch.pinned == "" && best == -1: + state = ChainDegraded + case next == 0: + state = ChainPrimary + default: + state = ChainFallback + } + switch { + case next != prev: + ch.active = next + ch.activeSince = now + m.emitLocked(Event{Type: EventChainSwitched, Chain: ch.name, From: ch.targets[prev], To: ch.targets[next], State: string(state), Reason: reason, At: now}) + case state == ChainDegraded && ch.state != ChainDegraded: + m.emitLocked(Event{Type: EventChainSwitched, Chain: ch.name, From: ch.targets[prev], To: ch.targets[next], State: string(state), Reason: ReasonDegraded, At: now}) + } + ch.state = state +} + +func (m *Manager) recomputeForLocked(target string) { + for _, ch := range m.chains { + if slices.Contains(ch.targets, target) { + m.recomputeLocked(ch, "") + } + } +} + +// Attempt walks the targets of one request in order. +type Attempt struct { + m *Manager + chain string + primary string + degraded bool + targets []string + i int +} + +// Plan returns the attempt order for one request: the active target, then the +// other healthy targets. A degraded chain tries every target in priority +// order; a pinned chain only the pinned target. +func (m *Manager) Plan(chain string) (*Attempt, error) { + m.mu.Lock() + defer m.mu.Unlock() + ch := m.chainLocked(chain) + if ch == nil { + return nil, fmt.Errorf("%w: %q", ErrChainNotFound, chain) + } + att := &Attempt{m: m, chain: ch.name, primary: ch.targets[0], degraded: ch.state == ChainDegraded} + usable := func(name string) bool { + ts := m.targets[name] + if ts == nil || ts.state == StateMissing { + return false + } + return att.degraded || ts.state == StateHealthy + } + switch { + case ch.pinned != "": + att.targets = []string{ch.pinned} + case att.degraded: + for _, name := range ch.targets { + if usable(name) { + att.targets = append(att.targets, name) + } + } + default: + active := ch.targets[ch.active] + if usable(active) { + att.targets = append(att.targets, active) + } + for _, name := range ch.targets { + if name != active && usable(name) { + att.targets = append(att.targets, name) + } + } + } + if len(att.targets) == 0 { + return nil, fmt.Errorf("%w: %q", ErrNoTarget, chain) + } + return att, nil +} + +func (a *Attempt) Chain() string { return a.chain } +func (a *Attempt) Target() string { return a.targets[a.i] } +func (a *Attempt) Primary() string { return a.primary } +func (a *Attempt) Degraded() bool { return a.degraded } + +// Fail records err against the current target and moves to the next one. It +// returns false when no target is left. +func (a *Attempt) Fail(err error) bool { + a.m.ReportFailure(a.Target(), err) + if a.i+1 >= len(a.targets) { + return false + } + a.i++ + return true +} + +// Report records err against the current target without moving on: the +// response was already committed, so nothing is left to retry. +func (a *Attempt) Report(err error) { a.m.ReportFailure(a.Target(), err) } + +func (a *Attempt) Succeed() { a.m.ReportSuccess(a.Target()) } + +func (m *Manager) ReportFailure(target string, err error) { + m.mu.Lock() + defer m.mu.Unlock() + ts := m.targetLocked(target) + if ts == nil { + return + } + msg := "" + if err != nil { + msg = err.Error() + } + m.recordFailureLocked(ts, msg) + m.recomputeForLocked(target) +} + +func (m *Manager) recordFailureLocked(ts *targetState, msg string) { + now := m.clock.Now() + ts.lastError = msg + switch ts.state { + case StateRecovering: + m.setTargetLocked(ts, StateDown, ReasonTrip, msg) + case StateHealthy: + cut := now.Add(-ts.params.TripWindow()) + kept := ts.failures[:0] + for _, f := range ts.failures { + if f.After(cut) { + kept = append(kept, f) + } + } + ts.failures = append(kept, now) + if len(ts.failures) >= ts.params.TripErrors() { + m.setTargetLocked(ts, StateDown, ReasonTrip, msg) + } + } +} + +func (m *Manager) ReportSuccess(target string) { + m.mu.Lock() + defer m.mu.Unlock() + ts := m.targetLocked(target) + if ts == nil { + return + } + ts.lastActivity = m.clock.Now() + m.recordPassLocked(ts) + m.recomputeForLocked(target) +} + +// recordPassLocked counts a served request or a passed inference probe. +func (m *Manager) recordPassLocked(ts *targetState) { + switch ts.state { + case StateHealthy: + ts.failures = nil + return + case StateMissing: + return + case StateDown: + if ts.cold() { + // Cold targets are never probed; a served request is proof enough. + m.setTargetLocked(ts, StateHealthy, ReasonRecovery, "") + return + } + m.setTargetLocked(ts, StateRecovering, ReasonRecovery, "") + } + ts.consecutiveOK++ + if ts.consecutiveOK >= ts.params.RecoveryProbes() { + m.setTargetLocked(ts, StateHealthy, ReasonRecovery, "") + } +} + +func (m *Manager) Pin(chain, target string) error { + m.mu.Lock() + defer m.mu.Unlock() + ch := m.chainLocked(chain) + if ch == nil { + return fmt.Errorf("%w: %q", ErrChainNotFound, chain) + } + if !slices.Contains(ch.targets, target) { + return fmt.Errorf("%w: %q", ErrTargetNotInChain, target) + } + ch.pinned = target + m.recomputeLocked(ch, ReasonManual) + return nil +} + +func (m *Manager) Unpin(chain string) error { + m.mu.Lock() + defer m.mu.Unlock() + ch := m.chainLocked(chain) + if ch == nil { + return fmt.Errorf("%w: %q", ErrChainNotFound, chain) + } + ch.pinned = "" + m.recomputeLocked(ch, ReasonManual) + return nil +} + +// Status returns every chain, sorted by name. +func (m *Manager) Status() []ChainStatus { + m.mu.Lock() + defer m.mu.Unlock() + m.syncLocked() + names := make([]string, 0, len(m.chains)) + for name := range m.chains { + names = append(names, name) + } + sort.Strings(names) + out := make([]ChainStatus, 0, len(names)) + for _, name := range names { + out = append(out, m.statusLocked(m.chains[name])) + } + return out +} + +func (m *Manager) ChainStatus(name string) (ChainStatus, bool) { + m.mu.Lock() + defer m.mu.Unlock() + ch := m.chainLocked(name) + if ch == nil { + return ChainStatus{}, false + } + return m.statusLocked(ch), true +} + +func (m *Manager) statusLocked(ch *chainState) ChainStatus { + cs := ChainStatus{Name: ch.name, State: ch.state, Active: ch.targets[ch.active], ActiveSince: ch.activeSince} + if ch.pinned != "" { + p := ch.pinned + cs.Pinned = &p + } + for _, name := range ch.targets { + st := TargetStatus{Model: name} + if ts := m.targets[name]; ts != nil { + st.Kind, st.Warm, st.State = ts.kind, ts.warm, ts.state + st.ConsecutiveOK, st.LastError = ts.consecutiveOK, ts.lastError + if !ts.lastProbe.IsZero() { + lp := ts.lastProbe + st.LastProbe = &lp + } + } + cs.Targets = append(cs.Targets, st) + } + return cs +} + +// Subscribe returns a buffered event channel and a cancel func. A subscriber +// that does not keep up loses events rather than blocking the manager. +func (m *Manager) Subscribe(buffer int) (<-chan Event, func()) { + m.mu.Lock() + defer m.mu.Unlock() + ch := make(chan Event, buffer) + if m.closed { + close(ch) + return ch, func() {} + } + id := m.nextSub + m.nextSub++ + m.subs[id] = ch + var once sync.Once + return ch, func() { + once.Do(func() { + m.mu.Lock() + defer m.mu.Unlock() + if c, ok := m.subs[id]; ok { + delete(m.subs, id) + close(c) + } + }) + } +} + +func (m *Manager) emitLocked(ev Event) { + for _, c := range m.subs { + select { + case c <- ev: + default: + xlog.Warn("failover: dropping event for a slow subscriber", "type", ev.Type, "chain", ev.Chain, "target", ev.Target) + } + } +} + +func (m *Manager) close() { + m.mu.Lock() + defer m.mu.Unlock() + m.closed = true + for id, c := range m.subs { + close(c) + delete(m.subs, id) + } +} + +// Do runs fn against the chain's targets in plan order. fn calls commit once +// output has reached the client; after that a failure is not retried. +func (m *Manager) Do(ctx context.Context, chain string, fn func(ctx context.Context, target string, commit func()) error) error { + att, err := m.Plan(chain) + if err != nil { + return err + } + for { + var committed atomic.Bool + err := fn(ctx, att.Target(), func() { committed.Store(true) }) + switch { + case err == nil: + att.Succeed() + return nil + case ctx.Err() != nil || !IsRetryable(err, 0): + return err + case committed.Load(): + att.Report(err) + return err + case !att.Fail(err): + return err + } + } +} + +// Prober checks targets. Implemented by DefaultProber (prober.go). +type Prober interface { + // Liveness is the cheap steady-state check. + Liveness(ctx context.Context, target config.ModelConfig, kind Kind, warm bool) error + // Inference sends one minimal real request to confirm recovery. + Inference(ctx context.Context, target config.ModelConfig, kind Kind, warm bool) error +} diff --git a/core/services/failover/manager_test.go b/core/services/failover/manager_test.go new file mode 100644 index 000000000..0ea98c2a0 --- /dev/null +++ b/core/services/failover/manager_test.go @@ -0,0 +1,232 @@ +package failover + +import ( + "context" + "errors" + "time" + + "github.com/mudler/LocalAI/core/config" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + . "github.com/onsi/gomega/gstruct" +) + +var errBoom = errors.New("dial tcp: connection refused") + +var _ = Describe("Manager", func() { + var ( + clock *fakeClock + src *fakeSource + m *Manager + ) + + BeforeEach(func() { + clock = newFakeClock() + src = newFakeSource(remote("a"), local("b"), chainCfg("chain", nil, t("a"), t("b"))) + m = New(src, WithClock(clock)) + }) + + switched := func(evs []Event) []Event { + var out []Event + for _, e := range evs { + if e.Type == EventChainSwitched { + out = append(out, e) + } + } + return out + } + + It("plans the primary first on a fresh chain", func() { + att, err := m.Plan("chain") + Expect(err).ToNot(HaveOccurred()) + Expect(att.Target()).To(Equal("a")) + Expect(att.Primary()).To(Equal("a")) + Expect(att.Degraded()).To(BeFalse()) + st, ok := m.ChainStatus("chain") + Expect(ok).To(BeTrue()) + Expect(st.State).To(Equal(ChainPrimary)) + Expect(st.Targets[0].Kind).To(Equal(KindRemote)) + Expect(st.Targets[1].Kind).To(Equal(KindLocal)) + }) + + It("returns ErrChainNotFound for an unknown chain", func() { + _, err := m.Plan("nope") + Expect(errors.Is(err, ErrChainNotFound)).To(BeTrue()) + }) + + It("trips on the first failure by default and switches with an event", func() { + events, cancel := m.Subscribe(16) + defer cancel() + att, _ := m.Plan("chain") + Expect(att.Fail(errBoom)).To(BeTrue()) + Expect(att.Target()).To(Equal("b")) + st, _ := m.ChainStatus("chain") + Expect(st.Active).To(Equal("b")) + Expect(st.State).To(Equal(ChainFallback)) + Expect(st.Targets[0].State).To(Equal(StateDown)) + Expect(st.Targets[0].LastError).To(ContainSubstring("connection refused")) + sw := switched(drain(events)) + Expect(sw).To(HaveLen(1)) + Expect(sw[0]).To(MatchFields(IgnoreExtras, Fields{ + "Chain": Equal("chain"), "From": Equal("a"), "To": Equal("b"), + "State": Equal("fallback"), "Reason": Equal(ReasonTrip), + })) + }) + + It("counts failures inside the trip window only", func() { + src.Put(chainCfg("chain", &config.FailoverConfig{Trip: config.FailoverTrip{Errors: 2, Window: "30s"}}, t("a"), t("b"))) + m.Sync() + m.ReportFailure("a", errBoom) + clock.Advance(31 * time.Second) + m.ReportFailure("a", errBoom) + st, _ := m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateHealthy)) + clock.Advance(time.Second) + m.ReportFailure("a", errBoom) + st, _ = m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateDown)) + }) + + It("fails back only after recovery probes and min_dwell", func() { + m.Plan("chain") + m.ReportFailure("a", errBoom) + for i := 0; i < 3; i++ { + m.ReportSuccess("a") // a real success counts like a passed inference probe + } + st, _ := m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateHealthy)) + Expect(st.Active).To(Equal("b"), "min_dwell has not passed") + clock.Advance(61 * time.Second) + events, cancel := m.Subscribe(16) + defer cancel() + m.Reevaluate() + st, _ = m.ChainStatus("chain") + Expect(st.Active).To(Equal("a")) + Expect(switched(drain(events))[0].Reason).To(Equal(ReasonRecovery)) + }) + + It("moves up at once when the active target itself goes down", func() { + src.Put(chainCfg("chain", nil, t("a"), t("b"), t("c"))) + src.Put(local("c")) + m.Sync() + m.ReportFailure("a", errBoom) // active: b + for i := 0; i < 3; i++ { + m.ReportSuccess("a") // a healthy again, but dwell not passed + } + m.ReportFailure("b", errBoom) // b down: go to a now, not c + st, _ := m.ChainStatus("chain") + Expect(st.Active).To(Equal("a")) + }) + + It("goes degraded when all targets are down and plans all of them in priority order", func() { + m.ReportFailure("a", errBoom) + m.ReportFailure("b", errBoom) + st, _ := m.ChainStatus("chain") + Expect(st.State).To(Equal(ChainDegraded)) + att, err := m.Plan("chain") + Expect(err).ToNot(HaveOccurred()) + Expect(att.Degraded()).To(BeTrue()) + Expect(att.Target()).To(Equal("a")) + Expect(att.Fail(errBoom)).To(BeTrue()) + Expect(att.Target()).To(Equal("b")) + Expect(att.Fail(errBoom)).To(BeFalse()) + }) + + It("pins a target regardless of health", func() { + Expect(m.Pin("chain", "b")).To(Succeed()) + att, _ := m.Plan("chain") + Expect(att.Target()).To(Equal("b")) + Expect(att.Fail(errBoom)).To(BeFalse(), "a pin allows only the pinned target") + st, _ := m.ChainStatus("chain") + Expect(*st.Pinned).To(Equal("b")) + Expect(st.Active).To(Equal("b")) + Expect(m.Unpin("chain")).To(Succeed()) + st, _ = m.ChainStatus("chain") + Expect(st.Pinned).To(BeNil()) + Expect(errors.Is(m.Pin("chain", "zzz"), ErrTargetNotInChain)).To(BeTrue()) + Expect(errors.Is(m.Pin("nope", "a"), ErrChainNotFound)).To(BeTrue()) + }) + + It("shares target health across chains", func() { + src.Put(chainCfg("chain2", nil, t("a"), t("b"))) + m.Sync() + m.ReportFailure("a", errBoom) + s1, _ := m.ChainStatus("chain") + s2, _ := m.ChainStatus("chain2") + Expect(s1.Active).To(Equal("b")) + Expect(s2.Active).To(Equal("b")) + }) + + It("marks a removed target missing and leaves it out of plans", func() { + m.Plan("chain") + src.Delete("a") + m.Sync() + st, _ := m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateMissing)) + att, _ := m.Plan("chain") + Expect(att.Target()).To(Equal("b")) + Expect(att.Fail(errBoom)).To(BeFalse()) + }) + + It("resets a chain whose target list changed", func() { + m.ReportFailure("a", errBoom) + src.Put(local("c")) + src.Put(chainCfg("chain", nil, t("c"), t("b"))) + m.Sync() + st, _ := m.ChainStatus("chain") + Expect(st.Active).To(Equal("c")) + }) + + It("reports warm local targets and ignores warm on remote ones", func() { + var got []string + m = New(src, WithClock(clock), WithOnWarmChanged(func(w []string) { got = w })) + src.Put(chainCfg("chain", nil, warmT("a"), warmT("b"))) + m.Sync() + Expect(got).To(Equal([]string{"b"})) + Expect(m.WarmTargets()).To(Equal([]string{"b"})) + }) + + It("closes a subscription on cancel", func() { + events, cancel := m.Subscribe(1) + cancel() + _, ok := <-events + Expect(ok).To(BeFalse()) + cancel() // idempotent + }) + + Describe("Do", func() { + It("retries on the next target until commit", func() { + var tried []string + err := m.Do(context.Background(), "chain", func(_ context.Context, target string, commit func()) error { + tried = append(tried, target) + if target == "a" { + return errBoom + } + return nil + }) + Expect(err).ToNot(HaveOccurred()) + Expect(tried).To(Equal([]string{"a", "b"})) + }) + + It("does not retry after commit but still trips the target", func() { + var tried []string + err := m.Do(context.Background(), "chain", func(_ context.Context, target string, commit func()) error { + tried = append(tried, target) + commit() + return errBoom + }) + Expect(err).To(MatchError(errBoom)) + Expect(tried).To(Equal([]string{"a"})) + st, _ := m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateDown)) + }) + + It("does not retry or trip on a non-retryable error", func() { + bad := errors.New("the request exceeds the available context size") + err := m.Do(context.Background(), "chain", func(_ context.Context, _ string, _ func()) error { return bad }) + Expect(err).To(MatchError(bad)) + st, _ := m.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(StateHealthy)) + }) + }) +})