feat(failover): expose chain status, pins and events over REST and SSE

Assisted-by: Claude:claude-opus-5-5
Signed-off-by: Ettore Di Giacinto <mudler@localai.io>
This commit is contained in:
Ettore Di Giacinto committed 2026-09-27 07:42:20 +00:00
1 parent f28d07b09d
commit 2a1213f1e2
7 files changed
+339 -1

No files matched your search

+7
View File
@@ -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)
+35
View File
@@ -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)
@@ -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",
@@ -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",
))
})
})
+169
View File
@@ -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
}
}
}
}
}
@@ -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"))
})
})
+10
View File
@@ -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))