mirror of
https://github.com/mudler/LocalAI.git
synced 2026-09-29 01:25:03 -04:00
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:
1 parent
f28d07b09d
commit
2a1213f1e2
7 files changed
+339
-1
No files matched your search
@@ -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)
|
||||
|
||||
@@ -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",
|
||||
))
|
||||
})
|
||||
})
|
||||
|
||||
@@ -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"))
|
||||
})
|
||||
})
|
||||
@@ -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))
|
||||
|
||||
Reference in new issue
Block a user