From aaabf7ee451036efd7ee5299fd7eb38da29022eb Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 17:42:00 +0000 Subject: [PATCH] feat(failover): expose chain status, pins and events over REST and SSE Assisted-by: Claude:claude-opus-5-5 --- core/http/auth/helpers_test.go | 7 + core/http/auth/middleware_test.go | 35 ++++ .../endpoints/localai/api_instructions.go | 6 + .../localai/api_instructions_test.go | 3 +- core/http/endpoints/localai/failover.go | 169 ++++++++++++++++++ core/http/endpoints/localai/failover_test.go | 110 ++++++++++++ core/http/routes/localai.go | 10 ++ 7 files changed, 339 insertions(+), 1 deletion(-) create mode 100644 core/http/endpoints/localai/failover.go create mode 100644 core/http/endpoints/localai/failover_test.go diff --git a/core/http/auth/helpers_test.go b/core/http/auth/helpers_test.go index 1e31ac27f..1b52d4500 100644 --- a/core/http/auth/helpers_test.go +++ b/core/http/auth/helpers_test.go @@ -78,6 +78,9 @@ func newAuthTestApp(db *gorm.DB, appConfig *config.ApplicationConfig) *echo.Echo e.GET("/api/settings", ok) e.POST("/api/settings", ok) + // Failover chain reads and the event stream: standard auth, no admin gate. + e.GET("/api/failover", ok) + // Auth routes (exempt) e.GET("/api/auth/status", ok) e.GET("/api/auth/github/login", ok) @@ -139,6 +142,10 @@ func newAdminTestApp(db *gorm.DB, appConfig *config.ApplicationConfig) *echo.Ech e.POST("/backends/apply", ok, adminMw) e.GET("/api/agents", ok, adminMw) + // Failover chain pin/unpin (admin only) + e.POST("/api/failover/:chain/pin", ok, adminMw) + e.DELETE("/api/failover/:chain/pin", ok, adminMw) + // Trace/log endpoints (admin only) e.GET("/api/traces", ok, adminMw) e.POST("/api/traces/clear", ok, adminMw) diff --git a/core/http/auth/middleware_test.go b/core/http/auth/middleware_test.go index bdfadaafe..bcbba0096 100644 --- a/core/http/auth/middleware_test.go +++ b/core/http/auth/middleware_test.go @@ -223,6 +223,17 @@ var _ = Describe("Auth Middleware", func() { Expect(rec.Code).To(Equal(http.StatusOK)) }) + It("allows requests to the failover status endpoint with a valid session", func() { + sessionID := createTestSession(db, user.ID) + rec := doRequest(app, http.MethodGet, "/api/failover", withSessionCookie(sessionID)) + Expect(rec.Code).To(Equal(http.StatusOK)) + }) + + It("returns 401 for the failover status endpoint without credentials", func() { + rec := doRequest(app, http.MethodGet, "/api/failover") + Expect(rec.Code).To(Equal(http.StatusUnauthorized)) + }) + It("allows authenticated users to call moderation by default", func() { sessionID := createTestSession(db, user.ID) rec := doRequest(app, http.MethodPost, "/v1/moderations", withSessionCookie(sessionID)) @@ -526,6 +537,30 @@ var _ = Describe("Auth Middleware", func() { Expect(rec.Code).To(Equal(http.StatusForbidden)) }) + It("allows admin to pin and unpin failover chains", func() { + admin := createTestUser(db, "admin5@example.com", auth.RoleAdmin, auth.ProviderGitHub) + sessionID := createTestSession(db, admin.ID) + app := newAdminTestApp(db, appConfig) + + rec := doRequest(app, http.MethodPost, "/api/failover/chain/pin", withSessionCookie(sessionID)) + Expect(rec.Code).To(Equal(http.StatusOK)) + + rec = doRequest(app, http.MethodDelete, "/api/failover/chain/pin", withSessionCookie(sessionID)) + Expect(rec.Code).To(Equal(http.StatusOK)) + }) + + It("blocks non-admin from pinning or unpinning failover chains", func() { + user := createTestUser(db, "user5@example.com", auth.RoleUser, auth.ProviderGitHub) + sessionID := createTestSession(db, user.ID) + app := newAdminTestApp(db, appConfig) + + rec := doRequest(app, http.MethodPost, "/api/failover/chain/pin", withSessionCookie(sessionID)) + Expect(rec.Code).To(Equal(http.StatusForbidden)) + + rec = doRequest(app, http.MethodDelete, "/api/failover/chain/pin", withSessionCookie(sessionID)) + Expect(rec.Code).To(Equal(http.StatusForbidden)) + }) + It("allows user to access regular inference endpoints", func() { user := createTestUser(db, "user@example.com", auth.RoleUser, auth.ProviderGitHub) sessionID := createTestSession(db, user.ID) diff --git a/core/http/endpoints/localai/api_instructions.go b/core/http/endpoints/localai/api_instructions.go index f8af73de2..8d0ea6d2f 100644 --- a/core/http/endpoints/localai/api_instructions.go +++ b/core/http/endpoints/localai/api_instructions.go @@ -129,6 +129,12 @@ var instructionDefs = []instructionDef{ Tags: []string{"middleware", "pii", "router"}, Intro: "GET /api/middleware/status is the single round-trip the /app/middleware admin page reads to render the current state: every model's resolved PII enabled state and the NER detector models it references, recent event count, and the active routing models with their classifier configurations. Admin-only (the synthetic local user is admin in no-auth mode). PII detection policy is edited on each detector model's `pii_detection:` block via the model-config tools/UI — there is no global pattern set to mutate. GET /api/router/decisions returns the routing decision log filtered by correlation_id / user_id / router_model. The same surface is exposed as MCP tools (`get_middleware_status`, `get_pii_events`, `get_router_decisions`) for agent-driven inspection.", }, + { + Name: "failover", + Description: "Model failover chains: target health, pinning and switch events", + Tags: []string{"failover"}, + Intro: "A failover chain is a model config with a failover block. Requests for the chain name are served by its highest-priority healthy target; the X-LocalAI-Served-Model response header names it. Subscribe to GET /api/failover/events (SSE) to follow switches.", + }, { Name: "intelligent-routing", Description: "Per-model `router:` configuration that classifies requests and rewrites the served model", diff --git a/core/http/endpoints/localai/api_instructions_test.go b/core/http/endpoints/localai/api_instructions_test.go index 710d4d982..f42e1c92d 100644 --- a/core/http/endpoints/localai/api_instructions_test.go +++ b/core/http/endpoints/localai/api_instructions_test.go @@ -39,7 +39,7 @@ var _ = Describe("API Instructions Endpoints", func() { instructions, ok := resp["instructions"].([]any) Expect(ok).To(BeTrue()) - Expect(instructions).To(HaveLen(19)) + Expect(instructions).To(HaveLen(20)) // Verify each instruction has required fields and correct URL format for _, s := range instructions { @@ -81,6 +81,7 @@ var _ = Describe("API Instructions Endpoints", func() { "intelligent-routing", "voice-library", "3d", + "failover", )) }) }) diff --git a/core/http/endpoints/localai/failover.go b/core/http/endpoints/localai/failover.go new file mode 100644 index 000000000..608fe69f2 --- /dev/null +++ b/core/http/endpoints/localai/failover.go @@ -0,0 +1,169 @@ +package localai + +import ( + "encoding/json" + "errors" + "fmt" + "net/http" + "time" + + "github.com/labstack/echo/v4" + "github.com/mudler/LocalAI/core/schema" + "github.com/mudler/LocalAI/core/services/failover" +) + +type FailoverChainsResponse struct { + Chains []failover.ChainStatus `json:"chains"` +} + +type FailoverPinRequest struct { + Target string `json:"target"` +} + +func failoverError(c echo.Context, code int, msg string) error { + return c.JSON(code, schema.ErrorResponse{Error: &schema.APIError{Message: msg, Code: code, Type: "failover_error"}}) +} + +// ListFailoverChainsEndpoint lists failover chains and the health of their targets +// +// @Summary List failover chains and the health of their targets +// @Tags failover +// @Produce json +// @Success 200 {object} FailoverChainsResponse +// @Router /api/failover [get] +func ListFailoverChainsEndpoint(fm *failover.Manager) echo.HandlerFunc { + return func(c echo.Context) error { + return c.JSON(http.StatusOK, FailoverChainsResponse{Chains: fm.Status()}) + } +} + +// GetFailoverChainEndpoint returns one failover chain +// +// @Summary Get one failover chain +// @Tags failover +// @Produce json +// @Param chain path string true "Chain name" +// @Success 200 {object} failover.ChainStatus +// @Failure 404 {object} schema.ErrorResponse +// @Router /api/failover/{chain} [get] +func GetFailoverChainEndpoint(fm *failover.Manager) echo.HandlerFunc { + return func(c echo.Context) error { + st, ok := fm.ChainStatus(c.Param("chain")) + if !ok { + return failoverError(c, http.StatusNotFound, fmt.Sprintf("failover chain %q not found", c.Param("chain"))) + } + return c.JSON(http.StatusOK, st) + } +} + +// PinFailoverTargetEndpoint forces a chain to one target +// +// @Summary Pin a failover chain to one target +// @Tags failover +// @Accept json +// @Produce json +// @Param chain path string true "Chain name" +// @Param request body FailoverPinRequest true "Target to pin" +// @Success 200 {object} failover.ChainStatus +// @Failure 400 {object} schema.ErrorResponse +// @Failure 404 {object} schema.ErrorResponse +// @Router /api/failover/{chain}/pin [post] +func PinFailoverTargetEndpoint(fm *failover.Manager) echo.HandlerFunc { + return func(c echo.Context) error { + var req FailoverPinRequest + if err := c.Bind(&req); err != nil || req.Target == "" { + return failoverError(c, http.StatusBadRequest, "request body must set \"target\"") + } + chain := c.Param("chain") + if err := fm.Pin(chain, req.Target); err != nil { + return pinError(c, err) + } + st, _ := fm.ChainStatus(chain) + return c.JSON(http.StatusOK, st) + } +} + +// UnpinFailoverTargetEndpoint removes a pin +// +// @Summary Remove the pin from a failover chain +// @Tags failover +// @Produce json +// @Param chain path string true "Chain name" +// @Success 200 {object} failover.ChainStatus +// @Failure 404 {object} schema.ErrorResponse +// @Router /api/failover/{chain}/pin [delete] +func UnpinFailoverTargetEndpoint(fm *failover.Manager) echo.HandlerFunc { + return func(c echo.Context) error { + chain := c.Param("chain") + if err := fm.Unpin(chain); err != nil { + return pinError(c, err) + } + st, _ := fm.ChainStatus(chain) + return c.JSON(http.StatusOK, st) + } +} + +func pinError(c echo.Context, err error) error { + switch { + case errors.Is(err, failover.ErrChainNotFound): + return failoverError(c, http.StatusNotFound, err.Error()) + case errors.Is(err, failover.ErrTargetNotInChain): + return failoverError(c, http.StatusBadRequest, err.Error()) + } + return failoverError(c, http.StatusInternalServerError, err.Error()) +} + +// FailoverEventsEndpoint streams failover events +// +// @Summary Stream failover events (server-sent events) +// @Description The first event is "snapshot" with the full state, then "chain.switched" and "target.state" events. +// @Tags failover +// @Produce text/event-stream +// @Success 200 +// @Router /api/failover/events [get] +func FailoverEventsEndpoint(fm *failover.Manager) echo.HandlerFunc { + return func(c echo.Context) error { + // Subscribe before the snapshot so no event falls between the two. + events, cancel := fm.Subscribe(64) + defer cancel() + w := c.Response() + w.Header().Set("Content-Type", "text/event-stream") + w.Header().Set("Cache-Control", "no-cache") + w.Header().Set("Connection", "keep-alive") + w.WriteHeader(http.StatusOK) + send := func(name string, v any) error { + data, err := json.Marshal(v) + if err != nil { + return err + } + if _, err := fmt.Fprintf(w, "event: %s\ndata: %s\n\n", name, data); err != nil { + return err + } + w.Flush() + return nil + } + if err := send("snapshot", FailoverChainsResponse{Chains: fm.Status()}); err != nil { + return nil + } + keepalive := time.NewTicker(15 * time.Second) + defer keepalive.Stop() + for { + select { + case <-c.Request().Context().Done(): + return nil + case <-keepalive.C: + if _, err := fmt.Fprint(w, ": keepalive\n\n"); err != nil { + return nil + } + w.Flush() + case ev, ok := <-events: + if !ok { + return nil + } + if err := send(string(ev.Type), ev); err != nil { + return nil + } + } + } + } +} diff --git a/core/http/endpoints/localai/failover_test.go b/core/http/endpoints/localai/failover_test.go new file mode 100644 index 000000000..fd3b6a733 --- /dev/null +++ b/core/http/endpoints/localai/failover_test.go @@ -0,0 +1,110 @@ +package localai + +import ( + "bufio" + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "time" + + "github.com/labstack/echo/v4" + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/services/failover" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +type mapSource map[string]config.ModelConfig + +func (s mapSource) GetModelConfig(n string) (config.ModelConfig, bool) { c, ok := s[n]; return c, ok } +func (s mapSource) GetAllModelsConfigs() []config.ModelConfig { + var out []config.ModelConfig + for _, c := range s { + out = append(out, c) + } + return out +} + +var _ = Describe("failover endpoints", func() { + var ( + e *echo.Echo + fm *failover.Manager + ) + + BeforeEach(func() { + src := mapSource{ + "a": {Name: "a", Backend: "cloud-proxy"}, + "b": {Name: "b", Backend: "llama-cpp"}, + "chain": {Name: "chain", Failover: &config.FailoverConfig{Targets: []config.FailoverTarget{{Model: "a"}, {Model: "b"}}}}, + } + fm = failover.New(src) + e = echo.New() + e.GET("/api/failover", ListFailoverChainsEndpoint(fm)) + e.GET("/api/failover/events", FailoverEventsEndpoint(fm)) + e.GET("/api/failover/:chain", GetFailoverChainEndpoint(fm)) + e.POST("/api/failover/:chain/pin", PinFailoverTargetEndpoint(fm)) + e.DELETE("/api/failover/:chain/pin", UnpinFailoverTargetEndpoint(fm)) + }) + + do := func(method, path, body string) *httptest.ResponseRecorder { + req := httptest.NewRequest(method, path, strings.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + e.ServeHTTP(rec, req) + return rec + } + + It("lists chains", func() { + rec := do(http.MethodGet, "/api/failover", "") + Expect(rec.Code).To(Equal(http.StatusOK)) + var out FailoverChainsResponse + Expect(json.Unmarshal(rec.Body.Bytes(), &out)).To(Succeed()) + Expect(out.Chains).To(HaveLen(1)) + Expect(out.Chains[0].Active).To(Equal("a")) + }) + + It("gets one chain or 404", func() { + Expect(do(http.MethodGet, "/api/failover/chain", "").Code).To(Equal(http.StatusOK)) + Expect(do(http.MethodGet, "/api/failover/nope", "").Code).To(Equal(http.StatusNotFound)) + }) + + It("pins and unpins", func() { + rec := do(http.MethodPost, "/api/failover/chain/pin", `{"target":"b"}`) + Expect(rec.Code).To(Equal(http.StatusOK)) + st, _ := fm.ChainStatus("chain") + Expect(st.Active).To(Equal("b")) + Expect(do(http.MethodPost, "/api/failover/chain/pin", `{"target":"zzz"}`).Code).To(Equal(http.StatusBadRequest)) + Expect(do(http.MethodPost, "/api/failover/chain/pin", `{}`).Code).To(Equal(http.StatusBadRequest)) + Expect(do(http.MethodPost, "/api/failover/nope/pin", `{"target":"a"}`).Code).To(Equal(http.StatusNotFound)) + Expect(do(http.MethodDelete, "/api/failover/chain/pin", "").Code).To(Equal(http.StatusOK)) + st, _ = fm.ChainStatus("chain") + Expect(st.Pinned).To(BeNil()) + }) + + It("streams a snapshot, then switch events", func() { + srv := httptest.NewServer(e) + defer srv.Close() + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + req, _ := http.NewRequestWithContext(ctx, http.MethodGet, srv.URL+"/api/failover/events", nil) + resp, err := http.DefaultClient.Do(req) + Expect(err).ToNot(HaveOccurred()) + defer resp.Body.Close() + Expect(resp.Header.Get("Content-Type")).To(HavePrefix("text/event-stream")) + r := bufio.NewReader(resp.Body) + next := func() string { + for { + line, err := r.ReadString('\n') + Expect(err).ToNot(HaveOccurred()) + if strings.HasPrefix(line, "event: ") { + return strings.TrimSpace(strings.TrimPrefix(line, "event: ")) + } + } + } + Expect(next()).To(Equal("snapshot")) + Expect(fm.Pin("chain", "b")).To(Succeed()) + Expect(next()).To(Equal("chain.switched")) + }) +}) diff --git a/core/http/routes/localai.go b/core/http/routes/localai.go index 1d3c7adaf..7a6f901ec 100644 --- a/core/http/routes/localai.go +++ b/core/http/routes/localai.go @@ -171,6 +171,16 @@ func RegisterLocalAIRoutes(router *echo.Echo, return nil })) + // Failover chains: reads and the event stream use standard auth (any + // authenticated caller may watch chain health); pin/unpin are admin-only + // since they override the routing decision for every caller of the chain. + fm := app.FailoverManager() + router.GET("/api/failover", localai.ListFailoverChainsEndpoint(fm)) + router.GET("/api/failover/events", localai.FailoverEventsEndpoint(fm)) + router.GET("/api/failover/:chain", localai.GetFailoverChainEndpoint(fm)) + router.POST("/api/failover/:chain/pin", localai.PinFailoverTargetEndpoint(fm), adminMiddleware) + router.DELETE("/api/failover/:chain/pin", localai.UnpinFailoverTargetEndpoint(fm), adminMiddleware) + voiceProfiles := app.VoiceProfileStore() router.GET("/api/voice-profiles", localai.ListVoiceProfilesEndpoint(voiceProfiles)) router.GET("/api/voice-profiles/:id/audio", localai.ServeVoiceProfileAudioEndpoint(voiceProfiles))