diff --git a/core/application/application.go b/core/application/application.go index b49851d47..ac66e75b6 100644 --- a/core/application/application.go +++ b/core/application/application.go @@ -16,6 +16,7 @@ import ( "github.com/mudler/LocalAI/core/services/agentpool" "github.com/mudler/LocalAI/core/services/cloudproxy/mitm" "github.com/mudler/LocalAI/core/services/facerecognition" + "github.com/mudler/LocalAI/core/services/failover" "github.com/mudler/LocalAI/core/services/galleryop" "github.com/mudler/LocalAI/core/services/monitoring" "github.com/mudler/LocalAI/core/services/nodes" @@ -82,6 +83,7 @@ type Application struct { routerRegistry *router.Registry routerCorpus *corpus.Manager admissionLimiter *admission.Limiter + failoverManager *failover.Manager watchdogMutex sync.Mutex watchdogStop chan bool p2pMutex sync.Mutex @@ -478,6 +480,9 @@ func (a *Application) AdmissionLimiter() *admission.Limiter { return a.admissionLimiter } +// FailoverManager serves failover chains. Never nil after New. +func (a *Application) FailoverManager() *failover.Manager { return a.failoverManager } + // StartupConfig returns the original startup configuration (from env vars, before file loading) func (a *Application) StartupConfig() *config.ApplicationConfig { return a.startupConfig diff --git a/core/application/failover.go b/core/application/failover.go new file mode 100644 index 000000000..3194e294b --- /dev/null +++ b/core/application/failover.go @@ -0,0 +1,17 @@ +package application + +import ( + "github.com/mudler/LocalAI/core/backend" + "github.com/mudler/xlog" +) + +// applyFailoverWarmTargets pins warm failover targets in the watchdog and +// loads them, so a switch does not wait for a cold load. +func (a *Application) applyFailoverWarmTargets(warm []string) { + a.SyncPinnedModelsToWatchdog() + for _, name := range warm { + if _, err := backend.PreloadModelByName(a.ApplicationConfig().Context, a.ModelConfigLoader(), a.ModelLoader(), a.ApplicationConfig(), name); err != nil { + xlog.Warn("failover: could not preload warm target", "model", name, "error", err) + } + } +} diff --git a/core/application/startup.go b/core/application/startup.go index abc2f4a17..1382835f6 100644 --- a/core/application/startup.go +++ b/core/application/startup.go @@ -1,6 +1,7 @@ package application import ( + "context" "crypto/rand" "encoding/hex" "fmt" @@ -12,6 +13,7 @@ import ( "github.com/mudler/LocalAI/core/config" "github.com/mudler/LocalAI/core/gallery" "github.com/mudler/LocalAI/core/http/auth" + "github.com/mudler/LocalAI/core/services/failover" "github.com/mudler/LocalAI/core/services/galleryop" "github.com/mudler/LocalAI/core/services/jobs" "github.com/mudler/LocalAI/core/services/messaging" @@ -27,6 +29,7 @@ import ( "github.com/mudler/LocalAI/core/trace" "github.com/mudler/LocalAI/internal" "github.com/mudler/LocalAI/pkg/downloader" + "github.com/mudler/LocalAI/pkg/grpc" "github.com/mudler/LocalAI/pkg/modelartifacts" "github.com/mudler/LocalAI/pkg/signals" "github.com/mudler/LocalAI/pkg/vram" @@ -250,6 +253,16 @@ func New(opts ...config.AppOption) (*Application, error) { // the embedding-cache stats endpoint sees a single source of truth. application.routerRegistry = router.NewRegistry() + // Failover chains: probe targets and track which one is active per + // chain. WithOnWarmChanged pins and preloads warm local targets so a + // switch to them does not wait for a cold load. + application.failoverManager = failover.New(application.ModelConfigLoader(), + failover.WithProber(failover.NewProber(func(ctx context.Context, cfg config.ModelConfig) (grpc.Backend, error) { + return application.ModelLoader().Load(backend.ModelOptions(cfg, options)...) + }, application.ModelLoader().ModelPath)), + failover.WithOnWarmChanged(application.applyFailoverWarmTargets), + ) + // Subsystem 5: admission control. Limiter is always wired so a // model that gains a limits: block via gallery install or YAML // edit takes effect on the next restart without conditional plumbing. @@ -548,6 +561,12 @@ func New(opts ...config.AppOption) (*Application, error) { } } + // Start the failover scheduler: it syncs chains from config, runs + // liveness/recovery probes and dwell-based fail-back. Run is the only + // caller of Sync in production so onWarm callbacks stay ordered. + failover.RegisterMetrics(application.failoverManager) + go application.failoverManager.Run(options.Context) + // Watch the configuration directory startWatcher(options) diff --git a/core/application/watchdog.go b/core/application/watchdog.go index 330c95353..4674cb470 100644 --- a/core/application/watchdog.go +++ b/core/application/watchdog.go @@ -2,6 +2,7 @@ package application import ( "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/services/failover" "github.com/mudler/LocalAI/pkg/model" "github.com/mudler/xlog" ) @@ -23,6 +24,9 @@ func (a *Application) SyncPinnedModelsToWatchdog() { pinned = append(pinned, cfg.Name) } } + if a.failoverManager != nil { + pinned = failover.MergePinned(pinned, a.failoverManager.WarmTargets()) + } wd.SetPinnedModels(pinned) xlog.Debug("Synced pinned models to watchdog", "count", len(pinned)) } diff --git a/core/http/react-ui/src/pages/Traces.jsx b/core/http/react-ui/src/pages/Traces.jsx index 6e8f3e800..3a44b9c21 100644 --- a/core/http/react-ui/src/pages/Traces.jsx +++ b/core/http/react-ui/src/pages/Traces.jsx @@ -105,6 +105,7 @@ const TYPE_COLORS = { vector_store: { bg: 'var(--color-accent-light)', color: 'var(--color-data-7)' }, token_classify: { bg: 'var(--color-info-light)', color: 'var(--color-data-3)' }, pattern_pii: { bg: 'var(--color-error-light)', color: 'var(--color-data-2)' }, + failover: { bg: 'var(--color-warning-light)', color: 'var(--color-data-2)' }, } function typeBadgeStyle(type) { diff --git a/core/services/failover/manager.go b/core/services/failover/manager.go index cfb2d8b66..408c44621 100644 --- a/core/services/failover/manager.go +++ b/core/services/failover/manager.go @@ -226,6 +226,17 @@ func (m *Manager) WarmTargets() []string { return slices.Clone(m.warm) } +// targetStates snapshots each target's health state for the metrics gauge. +func (m *Manager) targetStates() map[string]TargetState { + m.mu.Lock() + defer m.mu.Unlock() + out := make(map[string]TargetState, len(m.targets)) + for name, ts := range m.targets { + out[name] = ts.state + } + return out +} + // Reevaluate recomputes every chain. Dwell-based fail-back needs no event, so // the scheduler calls this on every tick. func (m *Manager) Reevaluate() { @@ -572,6 +583,9 @@ func (m *Manager) Subscribe(buffer int) (<-chan Event, func()) { } func (m *Manager) emitLocked(ev Event) { + if ev.Type == EventChainSwitched { + recordSwitch(ev) + } for _, c := range m.subs { select { case c <- ev: diff --git a/core/services/failover/metrics.go b/core/services/failover/metrics.go new file mode 100644 index 000000000..99d1a8212 --- /dev/null +++ b/core/services/failover/metrics.go @@ -0,0 +1,54 @@ +package failover + +import ( + "context" + "sync" + + "go.opentelemetry.io/otel" + "go.opentelemetry.io/otel/attribute" + "go.opentelemetry.io/otel/metric" +) + +var ( + metricsOnce sync.Once + switches metric.Int64Counter +) + +func initMetrics() { + metricsOnce.Do(func() { + meter := otel.Meter("github.com/mudler/LocalAI") + switches, _ = meter.Int64Counter("localai_failover_switches_total", + metric.WithDescription("Failover chain switches between targets")) + }) +} + +func recordSwitch(ev Event) { + initMetrics() + if switches == nil { + return + } + switches.Add(context.Background(), 1, metric.WithAttributes( + attribute.String("chain", ev.Chain), + attribute.String("from", ev.From), + attribute.String("to", ev.To), + attribute.String("reason", string(ev.Reason)), + )) +} + +// RegisterMetrics exports target health as a gauge. The application calls it +// once for its manager; tests create many managers and skip it. +func RegisterMetrics(m *Manager) { + meter := otel.Meter("github.com/mudler/LocalAI") + _, _ = meter.Int64ObservableGauge("localai_failover_target_up", + metric.WithDescription("1 when a failover target is healthy, 0 otherwise"), + metric.WithInt64Callback(func(_ context.Context, o metric.Int64Observer) error { + for name, state := range m.targetStates() { + v := int64(0) + if state == StateHealthy { + v = 1 + } + o.Observe(v, metric.WithAttributes(attribute.String("target", name))) + } + return nil + })) +} diff --git a/core/services/failover/metrics_test.go b/core/services/failover/metrics_test.go new file mode 100644 index 000000000..42dcd24b8 --- /dev/null +++ b/core/services/failover/metrics_test.go @@ -0,0 +1,19 @@ +package failover + +import ( + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("metrics", func() { + It("registers and records without a meter provider", func() { + src := newFakeSource(remote("a"), local("b"), chainCfg("chain", nil, t("a"), t("b"))) + m := New(src, WithClock(newFakeClock())) + Expect(func() { RegisterMetrics(m) }).ToNot(Panic()) + Expect(func() { m.ReportFailure("a", errBoom) }).ToNot(Panic()) + }) + It("records an attempt trace only when enabled", func() { + Expect(func() { RecordAttemptTrace(false, "chain", "a", errBoom) }).ToNot(Panic()) + Expect(func() { RecordAttemptTrace(true, "chain", "a", errBoom) }).ToNot(Panic()) + }) +}) diff --git a/core/services/failover/trace.go b/core/services/failover/trace.go new file mode 100644 index 000000000..8eb6618f1 --- /dev/null +++ b/core/services/failover/trace.go @@ -0,0 +1,23 @@ +package failover + +import ( + "fmt" + "time" + + "github.com/mudler/LocalAI/core/trace" +) + +// RecordAttemptTrace shows in the Traces UI why a target was skipped. +func RecordAttemptTrace(enabled bool, chain, target string, err error) { + if !enabled || err == nil { + return + } + trace.RecordBackendTrace(trace.BackendTrace{ + Timestamp: time.Now(), + Type: trace.BackendTraceFailover, + ModelName: target, + Summary: fmt.Sprintf("failover chain %s: %s failed, trying the next target", chain, target), + Error: err.Error(), + Data: map[string]any{"chain": chain}, + }) +} diff --git a/core/trace/backend_trace.go b/core/trace/backend_trace.go index 072e41ebf..e14370c45 100644 --- a/core/trace/backend_trace.go +++ b/core/trace/backend_trace.go @@ -48,6 +48,7 @@ const ( BackendTraceTokenClassify BackendTraceType = "token_classify" BackendTracePatternPII BackendTraceType = "pattern_pii" BackendTraceVectorStore BackendTraceType = "vector_store" + BackendTraceFailover BackendTraceType = "failover" ) const (