diff --git a/core/application/application.go b/core/application/application.go index b15d848a0..56e22eeed 100644 --- a/core/application/application.go +++ b/core/application/application.go @@ -13,6 +13,7 @@ import ( "github.com/mudler/LocalAI/core/config" "github.com/mudler/LocalAI/core/http/auth" mcpTools "github.com/mudler/LocalAI/core/http/endpoints/mcp" + "github.com/mudler/LocalAI/core/services/advisorylock" "github.com/mudler/LocalAI/core/services/agentpool" "github.com/mudler/LocalAI/core/services/cloudproxy/mitm" "github.com/mudler/LocalAI/core/services/facerecognition" @@ -98,6 +99,9 @@ type Application struct { // failoverSync shares failover state between frontends; nil in // standalone mode or when it could not start. failoverSync *distsync.Sync + // failoverLock is the probe-leader lock this frontend holds or tries + // for; nil when failoverSync is nil. + failoverLock *advisorylock.HeldLock // Upgrade checker (background service for detecting backend upgrades) upgradeChecker *UpgradeChecker @@ -522,11 +526,7 @@ func (a *Application) Shutdown() error { a.shutdownOnce.Do(func() { // Before distributed shutdown: the sync's subscriptions live on the // NATS connection that closes there. - if a.failoverSync != nil { - if closeErr := a.failoverSync.Close(); closeErr != nil { - xlog.Warn("failover: closing state sync", "error", closeErr) - } - } + a.stopFailoverDistributed() a.distributed.Shutdown() if a.modelLoader != nil { err = a.modelLoader.StopAllGRPC() diff --git a/core/application/failover_distributed.go b/core/application/failover_distributed.go index 7c1b63b9f..37eb94edc 100644 --- a/core/application/failover_distributed.go +++ b/core/application/failover_distributed.go @@ -8,24 +8,31 @@ import ( "github.com/mudler/LocalAI/core/services/failover/distsync" "github.com/mudler/LocalAI/core/services/nodes" "github.com/mudler/xlog" - "gorm.io/gorm" ) // failoverLeaderGate elects the frontend that probes failover targets and // decides chains, so N frontends do not probe every target N times or -// disagree on which target is active. A lock error counts as "not leader": -// skipping one tick is safe, two leaders are not. -func failoverLeaderGate(db *gorm.DB) failover.LeaderGate { +// disagree on which target is active. +// +// Leadership is sticky: the leader keeps the lock across ticks and loses it +// only when its database session dies or it shuts down. A lock taken per tick +// would pass between frontends on almost every tick, and each change of +// leader redelivers the warm set and republishes all state. A lock error +// counts as "not leader": skipping one tick is safe, two leaders are not. +func failoverLeaderGate(l *advisorylock.HeldLock) failover.LeaderGate { return func(ctx context.Context, fn func()) bool { - ok, err := advisorylock.TryWithLockCtx(ctx, db, advisorylock.KeyFailoverProber, func() error { - fn() - return nil - }) - if err != nil { - xlog.Warn("failover: could not take the prober leader lock", "error", err) - return false + if !l.Held() || !l.Verify(ctx) { + ok, err := l.TryAcquire(ctx) + if err != nil { + xlog.Warn("failover: could not take the prober leader lock", "error", err) + return false + } + if !ok { + return false + } } - return ok + fn() + return true } } @@ -65,5 +72,21 @@ func (a *Application) startFailoverDistributed(ctx context.Context) { // Gate only once state is shared: a follower learns target health and // chain decisions solely from the leader's publishes, so gating without // the sync would leave every follower's chains frozen. - a.failoverManager.SetLeaderGate(failoverLeaderGate(db)) + a.failoverLock = advisorylock.NewHeldLock(db, advisorylock.KeyFailoverProber) + a.failoverManager.SetLeaderGate(failoverLeaderGate(a.failoverLock)) +} + +// stopFailoverDistributed detaches the manager from the sync before closing +// it, so nothing publishes into closed maps, and gives up leadership at once +// instead of making another frontend wait for this session to time out. +func (a *Application) stopFailoverDistributed() { + if a.failoverSync != nil { + a.failoverManager.SetStateSync(nil) + if err := a.failoverSync.Close(); err != nil { + xlog.Warn("failover: closing state sync", "error", err) + } + } + if a.failoverLock != nil { + a.failoverLock.Release() + } } diff --git a/core/application/failover_distributed_test.go b/core/application/failover_distributed_test.go index b4cf84bf0..1021d0105 100644 --- a/core/application/failover_distributed_test.go +++ b/core/application/failover_distributed_test.go @@ -5,6 +5,7 @@ import ( "time" "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/services/advisorylock" "github.com/mudler/LocalAI/core/services/failover" "github.com/mudler/LocalAI/pkg/model" . "github.com/onsi/ginkgo/v2" @@ -57,23 +58,30 @@ var _ = Describe("failoverPinnedResolver", func() { }) var _ = Describe("failoverLeaderGate", func() { - It("lets only one frontend lead at a time on the same database", func() { + It("keeps leadership with the first frontend until it releases", func() { // Not PostgreSQL, so advisorylock falls back to its in-process lock, // which has the same try-lock semantics as pg_try_advisory_lock. db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{}) Expect(err).ToNot(HaveOccurred()) - first, second := failoverLeaderGate(db), failoverLeaderGate(db) + firstLock := advisorylock.NewHeldLock(db, advisorylock.KeyFailoverProber) + secondLock := advisorylock.NewHeldLock(db, advisorylock.KeyFailoverProber) + DeferCleanup(firstLock.Release) + DeferCleanup(secondLock.Release) + first, second := failoverLeaderGate(firstLock), failoverLeaderGate(secondLock) ctx := context.Background() - var secondInside, secondRan bool - Expect(first(ctx, func() { - secondInside = second(ctx, func() { secondRan = true }) - })).To(BeTrue()) - Expect(secondInside).To(BeFalse(), "the second gate must not lead while the first holds the lock") - Expect(secondRan).To(BeFalse()) + var firstRuns, secondRuns int + for range 5 { + Expect(first(ctx, func() { firstRuns++ })).To(BeTrue()) + Expect(second(ctx, func() { secondRuns++ })).To(BeFalse(), "leadership must not flip between ticks") + } + Expect(firstRuns).To(Equal(5)) + Expect(secondRuns).To(BeZero()) - Expect(second(ctx, func() { secondRan = true })).To(BeTrue()) - Expect(secondRan).To(BeTrue()) + firstLock.Release() + Expect(second(ctx, func() { secondRuns++ })).To(BeTrue(), "a released leadership passes to the next frontend") + Expect(secondRuns).To(Equal(1)) + Expect(first(ctx, func() { firstRuns++ })).To(BeFalse()) }) }) @@ -99,4 +107,28 @@ var _ = Describe("applyFailoverWarmTargets in distributed mode", func() { app.applyFailoverWarmTargets([]string{"b"}) Consistently(preloaded, 200*time.Millisecond).ShouldNot(Receive(), "only the probe leader preloads warm targets") }) + + It("preloads warm targets on the probe leader", func() { + preloaded := make(chan string, 4) + orig := preloadModelByName + preloadModelByName = func(_ context.Context, _ *config.ModelConfigLoader, _ *model.ModelLoader, _ *config.ApplicationConfig, name string) ([]string, error) { + preloaded <- name + return nil, nil + } + DeferCleanup(func() { preloadModelByName = orig }) + + app := &Application{ + applicationConfig: &config.ApplicationConfig{Context: context.Background()}, + distributed: &DistributedServices{}, + } + // A gate that always grants: this frontend is the leader. The warm + // set is delivered on the first tick that wins it. + app.failoverManager = failover.New(warmChainSource(), + failover.WithLeaderGate(func(_ context.Context, fn func()) bool { fn(); return true }), + failover.WithOnWarmChanged(app.applyFailoverWarmTargets)) + + app.failoverManager.Tick(context.Background()) + Eventually(preloaded, time.Second).Should(Receive(Equal("b"))) + Eventually(preloaded, time.Second).Should(Receive(Equal("c"))) + }) }) diff --git a/core/services/advisorylock/held_lock.go b/core/services/advisorylock/held_lock.go new file mode 100644 index 000000000..92e8acbe8 --- /dev/null +++ b/core/services/advisorylock/held_lock.go @@ -0,0 +1,147 @@ +package advisorylock + +import ( + "context" + "database/sql" + "database/sql/driver" + "fmt" + "sync" + "time" + + "github.com/mudler/xlog" + "gorm.io/gorm" +) + +// heldLockCheckTimeout bounds the liveness check and the unlock, so a hung +// database cannot stall the caller's loop. +const heldLockCheckTimeout = 5 * time.Second + +// HeldLock is an advisory lock that stays taken across calls, for leader +// election where leadership must be sticky. TryWithLockCtx releases the lock +// when fn returns; with a short fn, every contender wins it in turn and +// leadership flips on each tick. +// +// On PostgreSQL the lock belongs to one session, so HeldLock keeps a +// dedicated connection out of the pool while it holds the lock. If that +// session dies, the server drops the lock and another instance can take it. +// On other dialects it holds the package's in-process lock for the key. +type HeldLock struct { + db *gorm.DB + key int64 + + mu sync.Mutex + conn *sql.Conn // PostgreSQL: the session that holds the lock + local bool // other dialects: this lock holds the in-process slot +} + +// NewHeldLock returns a lock for key on db. It takes nothing until TryAcquire. +func NewHeldLock(db *gorm.DB, key int64) *HeldLock { + return &HeldLock{db: db, key: key} +} + +// TryAcquire takes the lock without blocking. It returns true if this +// HeldLock holds the lock afterwards, including when it already held it. +func (l *HeldLock) TryAcquire(ctx context.Context) (bool, error) { + l.mu.Lock() + defer l.mu.Unlock() + if l.conn != nil || l.local { + return true, nil + } + + if !isPostgres(l.db) { + select { + case localLockChan(l.key) <- struct{}{}: + l.local = true + return true, nil + default: + return false, nil + } + } + + sqlDB, err := l.db.DB() + if err != nil { + return false, fmt.Errorf("get sql.DB: %w", err) + } + conn, err := sqlDB.Conn(ctx) + if err != nil { + return false, fmt.Errorf("advisory lock conn: %w", err) + } + var acquired bool + if err := conn.QueryRowContext(ctx, "SELECT pg_try_advisory_lock($1)", l.key).Scan(&acquired); err != nil { + // The lock may have been granted before the error (a cancelled + // context, say); discarding the session is the only way to be sure + // it is not left held. + discardConn(conn) + return false, fmt.Errorf("pg_try_advisory_lock: %w", err) + } + if !acquired { + _ = conn.Close() // holds nothing, so it may go back to the pool + return false, nil + } + l.conn = conn + return true, nil +} + +// Held reports whether this HeldLock believes it holds the lock. It does not +// touch the database; use Verify for that. +func (l *HeldLock) Held() bool { + l.mu.Lock() + defer l.mu.Unlock() + return l.conn != nil || l.local +} + +// Verify reports whether the lock is still held. On PostgreSQL it checks +// that the holding session is alive; if it is not, the lock is gone +// server-side, so Verify drops it here too and returns false. +func (l *HeldLock) Verify(ctx context.Context) bool { + l.mu.Lock() + defer l.mu.Unlock() + if l.local { + return true + } + if l.conn == nil { + return false + } + pctx, cancel := context.WithTimeout(ctx, heldLockCheckTimeout) + defer cancel() + if err := l.conn.PingContext(pctx); err != nil { + xlog.Warn("advisory lock session lost, giving up the lock", "key", l.key, "error", err) + // A ping that only timed out may leave the session (and the lock) + // alive; discarding closes it, so the server releases the lock. + discardConn(l.conn) + l.conn = nil + return false + } + return true +} + +// Release gives up the lock. It is safe to call when the lock is not held. +func (l *HeldLock) Release() { + l.mu.Lock() + defer l.mu.Unlock() + if l.local { + <-localLockChan(l.key) + l.local = false + return + } + if l.conn == nil { + return + } + ctx, cancel := context.WithTimeout(context.Background(), heldLockCheckTimeout) + defer cancel() + if _, err := l.conn.ExecContext(ctx, "SELECT pg_advisory_unlock($1)", l.key); err != nil { + xlog.Warn("advisory lock unlock failed, closing its session instead", "key", l.key, "error", err) + } + // Never hand the session back to the pool: if the unlock failed it still + // holds the lock, and a later pool user would silently inherit it. + discardConn(l.conn) + l.conn = nil +} + +// discardConn closes conn's underlying session instead of returning it to +// the pool. database/sql drops a connection when Raw's callback returns +// driver.ErrBadConn. +func discardConn(conn *sql.Conn) { + _ = conn.Raw(func(any) error { return driver.ErrBadConn }) + _ = conn.Close() +} diff --git a/core/services/advisorylock/held_lock_test.go b/core/services/advisorylock/held_lock_test.go new file mode 100644 index 000000000..fa8322d10 --- /dev/null +++ b/core/services/advisorylock/held_lock_test.go @@ -0,0 +1,105 @@ +package advisorylock + +import ( + "context" + + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + + "github.com/mudler/LocalAI/core/services/testutil" + "gorm.io/driver/sqlite" + "gorm.io/gorm" +) + +// expectStickyHandover checks the contract both backends share: the first +// holder keeps the lock across repeated checks while a rival is denied, and +// the rival gets it once the holder releases. +func expectStickyHandover(db *gorm.DB, key int64) { + ctx := context.Background() + first, second := NewHeldLock(db, key), NewHeldLock(db, key) + DeferCleanup(first.Release) + DeferCleanup(second.Release) + + ok, err := first.TryAcquire(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(ok).To(BeTrue()) + + for range 5 { + Expect(first.Held()).To(BeTrue()) + Expect(first.Verify(ctx)).To(BeTrue()) + ok, err = second.TryAcquire(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(ok).To(BeFalse(), "a rival must not take a held lock") + Expect(second.Held()).To(BeFalse()) + } + + first.Release() + Expect(first.Held()).To(BeFalse()) + Expect(first.Verify(ctx)).To(BeFalse()) + + ok, err = second.TryAcquire(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(ok).To(BeTrue(), "the lock is free once the holder releases it") + Expect(second.Verify(ctx)).To(BeTrue()) +} + +var _ = Describe("HeldLock (SQLite fallback)", Label("sqlite"), func() { + It("stays with its holder until Release", func() { + db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{}) + Expect(err).ToNot(HaveOccurred()) + expectStickyHandover(db, 12101) + }) + + It("is idempotent: acquiring twice and releasing twice is safe", func() { + db, err := gorm.Open(sqlite.Open("file::memory:?cache=shared"), &gorm.Config{}) + Expect(err).ToNot(HaveOccurred()) + l := NewHeldLock(db, 12102) + ok, err := l.TryAcquire(context.Background()) + Expect(err).ToNot(HaveOccurred()) + Expect(ok).To(BeTrue()) + ok, err = l.TryAcquire(context.Background()) + Expect(err).ToNot(HaveOccurred()) + Expect(ok).To(BeTrue(), "a holder asking again keeps the lock") + l.Release() + l.Release() + ok, err = l.TryAcquire(context.Background()) + Expect(err).ToNot(HaveOccurred()) + Expect(ok).To(BeTrue()) + l.Release() + }) +}) + +var _ = Describe("HeldLock (PostgreSQL)", func() { + var db *gorm.DB + + BeforeEach(func() { + db = testutil.SetupTestDB() + }) + + It("stays with its holder until Release, then hands over", func() { + expectStickyHandover(db, 12201) + }) + + It("drops leadership when the holding session dies, freeing the lock", func() { + ctx := context.Background() + first, second := NewHeldLock(db, 12202), NewHeldLock(db, 12202) + DeferCleanup(first.Release) + DeferCleanup(second.Release) + + ok, err := first.TryAcquire(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(ok).To(BeTrue()) + + // Kill the holder's backend, as a network cut or DB restart would. + var pid int + Expect(first.conn.QueryRowContext(ctx, "SELECT pg_backend_pid()").Scan(&pid)).To(Succeed()) + Expect(db.Exec("SELECT pg_terminate_backend(?)", pid).Error).ToNot(HaveOccurred()) + + Expect(first.Verify(ctx)).To(BeFalse(), "a dead session no longer holds the lock") + Expect(first.Held()).To(BeFalse()) + + ok, err = second.TryAcquire(ctx) + Expect(err).ToNot(HaveOccurred()) + Expect(ok).To(BeTrue(), "the server released the dead session's lock") + }) +}) diff --git a/docs/content/features/model-failover.md b/docs/content/features/model-failover.md index 24b9ae8dd..f4f42027b 100644 --- a/docs/content/features/model-failover.md +++ b/docs/content/features/model-failover.md @@ -183,15 +183,21 @@ share one failover state: - Pins apply to the whole cluster. LocalAI stores them in the database, so they persist across restarts. A pin set on one frontend applies on all. -- Target health and the active target of each chain are shared over NATS. - All frontends send a chain's requests to the same target. +- Target health and the active target of each chain are shared over NATS, so + all frontends converge on the same target for a chain. - One frontend, the probe leader, runs the health checks, decides fail-over - and fail-back, and loads warm targets. A PostgreSQL advisory lock selects - the leader. If the leader stops, another frontend takes over. + and fail-back, and loads warm targets. The leader holds a PostgreSQL + advisory lock and keeps it until it shuts down or its database connection + fails. Then another frontend takes the lock and becomes the leader. - Warm targets stay loaded on the workers. The router and the replica reconciler treat them like pinned models and do not evict them. - A frontend that starts late gets the current state within 10 seconds, because the leader sends its full state again every 10 seconds. +- If PostgreSQL is not available, no frontend holds the lock. Health checks + and fail-back stop until the database is back. Requests still fail over to + the next target when a target fails during the request. +- If a frontend cannot start the shared state, it logs an error and manages + failover alone, as a single LocalAI instance does. ## Limits