From 1148d29b41bd31e9188852c9c9ce9e04ff1bac3d Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 17:13:09 +0000 Subject: [PATCH] feat(failover): resolve chains per request and retry on the next target The retry wraps SetModelAndConfig, so each attempt binds the request again from a replayed body. A 5xx of a chain request is held back until the handler returns, and a streamed response is never retried. Assisted-by: Claude:claude-opus-5-5 --- core/http/app.go | 1 + core/http/middleware/context_keys.go | 4 + core/http/middleware/failover.go | 282 ++++++++++++++++++++++++++ core/http/middleware/failover_test.go | 233 +++++++++++++++++++++ core/http/middleware/request.go | 22 +- 5 files changed, 540 insertions(+), 2 deletions(-) create mode 100644 core/http/middleware/failover.go create mode 100644 core/http/middleware/failover_test.go diff --git a/core/http/app.go b/core/http/app.go index f03522a45..db1317275 100644 --- a/core/http/app.go +++ b/core/http/app.go @@ -455,6 +455,7 @@ func API(application *application.Application) (*echo.Echo, error) { mcpJobsMw := auth.RequireFeature(application.AuthDB(), auth.FeatureMCPJobs) requestExtractor := httpMiddleware.NewRequestExtractor(application.ModelConfigLoader(), application.ModelLoader(), application.ApplicationConfig()) + requestExtractor.SetFailoverManager(application.FailoverManager()) // Register auth routes (login, callback, API keys, user management) routes.RegisterAuthRoutes(e, application) diff --git a/core/http/middleware/context_keys.go b/core/http/middleware/context_keys.go index d1983c882..f8b5582da 100644 --- a/core/http/middleware/context_keys.go +++ b/core/http/middleware/context_keys.go @@ -47,4 +47,8 @@ const ( // router nor the body-parse path has produced one. Distinct from // ContextKeyServedModel, which is the router's resolved choice. ContextKeyResponseModel = "routing.response_model" + + // ContextKeyFailoverAttempt holds the *failoverState of a request whose + // model is a failover chain. + ContextKeyFailoverAttempt = "failover.attempt" ) diff --git a/core/http/middleware/failover.go b/core/http/middleware/failover.go new file mode 100644 index 000000000..930b5e209 --- /dev/null +++ b/core/http/middleware/failover.go @@ -0,0 +1,282 @@ +package middleware + +import ( + "bufio" + "bytes" + "fmt" + "io" + "net" + "net/http" + "slices" + "strings" + + "github.com/labstack/echo/v4" + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/services/failover" +) + +const ( + HeaderServedModel = "X-LocalAI-Served-Model" + HeaderFailover = "X-LocalAI-Failover" +) + +// MaxFailoverReplayBody caps the request body kept for a retry. A larger body +// is still served, by one target only. +const MaxFailoverReplayBody = 32 << 20 + +type failoverState struct { + attempt *failover.Attempt +} + +// SetFailoverManager enables failover chains. Without it, a request for a +// chain fails with 503. +func (re *RequestExtractor) SetFailoverManager(m *failover.Manager) { re.failover = m } + +// resolveFailover returns the config of the target that should serve this +// attempt. The first attempt plans the chain; retries reuse the plan. +func (re *RequestExtractor) resolveFailover(c echo.Context, requested string, chain *config.ModelConfig) (*config.ModelConfig, error) { + st, _ := c.Get(ContextKeyFailoverAttempt).(*failoverState) + if st == nil || st.attempt.Chain() != chain.Name { + if re.failover == nil { + return nil, fmt.Errorf("model %q is a failover chain, but failover is not running", chain.Name) + } + att, err := re.failover.Plan(chain.Name) + if err != nil { + return nil, err + } + st = &failoverState{attempt: att} + c.Set(ContextKeyFailoverAttempt, st) + } + for { + cfg, err := re.loadFailoverTarget(st.attempt.Target()) + if err == nil { + c.Set(ContextKeyRequestedModel, requested) + c.Set(ContextKeyServedModel, cfg.Name) + setFailoverHeaders(c.Response().Header(), st.attempt) + return cfg, nil + } + if !st.attempt.Fail(err) { + // Clear the state so the retry wrapper sends this 503 as is. + c.Set(ContextKeyFailoverAttempt, nil) + return nil, err + } + } +} + +func (re *RequestExtractor) loadFailoverTarget(name string) (*config.ModelConfig, error) { + cfg, err := re.modelConfigLoader.LoadModelConfigFileByNameDefaultOptions(name, re.applicationConfig) + if err != nil { + return nil, err + } + resolved, _, err := re.modelConfigLoader.ResolveAlias(cfg) + return resolved, err +} + +func setFailoverHeaders(h http.Header, att *failover.Attempt) { + h.Set(HeaderServedModel, att.Target()) + switch { + case att.Degraded(): + h.Set(HeaderFailover, "degraded") + case att.Target() != att.Primary(): + h.Set(HeaderFailover, "fallback") + default: + h.Del(HeaderFailover) + } +} + +// failoverRetry runs h again on the next target while the response is not +// committed. h is SetModelAndConfig's body plus the rest of the chain, so +// every attempt binds the request again from the replayed body. +func failoverRetry(appConfig *config.ApplicationConfig, h echo.HandlerFunc) echo.HandlerFunc { + return func(c echo.Context) error { + req := c.Request() + src := req.Body + if src == nil { + src = http.NoBody + } + rec := &replayBody{src: src, limit: MaxFailoverReplayBody} + req.Body = rec + // A default-model middleware may have parsed a multipart form before + // this point, draining the body. That form stays valid for every + // attempt; a form parsed during an attempt is dropped and parsed again + // from the replayed body. + entryMultipart, entryForm, entryPostForm := req.MultipartForm, req.Form, req.PostForm + resp := c.Response() + orig := resp.Writer + baseHeader := resp.Header().Clone() + defer func() { resp.Writer = orig }() + active := func() bool { + st, _ := c.Get(ContextKeyFailoverAttempt).(*failoverState) + return st != nil + } + tracing := appConfig != nil && appConfig.EnableTracing + for { + w := &failoverWriter{ResponseWriter: orig, active: active} + resp.Writer = w + err := h(c) + st, _ := c.Get(ContextKeyFailoverAttempt).(*failoverState) + if st == nil { + w.release() + return err + } + att := st.attempt + status := w.held + if err == nil && status == 0 { + att.Succeed() + return nil + } + cause := attemptError(err, status, w.body.Bytes()) + retryable := req.Context().Err() == nil && failover.IsRetryable(err, status) + if !retryable || w.committed || !rec.replayable() { + if retryable { + att.Report(cause) + } + w.release() + return err + } + failover.RecordAttemptTrace(tracing, att.Chain(), att.Target(), cause) + if !att.Fail(cause) { + w.release() + return err + } + // Temp files of a form parsed during this attempt would otherwise + // outlive the request: the server only cleans up the last form. + if mf := c.Request().MultipartForm; mf != nil && mf != entryMultipart { + _ = mf.RemoveAll() + } + req.Body = rec.replay() + req.MultipartForm, req.Form, req.PostForm = entryMultipart, entryForm, entryPostForm + // Middleware after this one may have replaced the request; the + // next attempt starts again from the request as it arrived here. + c.SetRequest(req) + resetResponse(resp, baseHeader) + } + } +} + +func attemptError(err error, status int, body []byte) error { + if err != nil { + return err + } + msg := strings.TrimSpace(string(body)) + if len(msg) > 200 { + msg = msg[:200] + } + return fmt.Errorf("HTTP %d: %s", status, msg) +} + +func resetResponse(resp *echo.Response, base http.Header) { + h := resp.Header() + for k := range h { + delete(h, k) + } + for k, v := range base { + h[k] = slices.Clone(v) + } + resp.Committed = false + resp.Status = http.StatusOK + resp.Size = 0 +} + +// failoverWriter holds back an error response (status >= 500) of a chain +// request until the handler returns, so the retry can drop it. +type failoverWriter struct { + http.ResponseWriter + active func() bool + held int + body bytes.Buffer + committed bool +} + +func (w *failoverWriter) WriteHeader(code int) { + if w.held != 0 { + return + } + if !w.committed && code >= 500 && w.active() { + w.held = code + return + } + w.committed = true + w.ResponseWriter.WriteHeader(code) +} + +func (w *failoverWriter) Write(b []byte) (int, error) { + if w.held != 0 { + return w.body.Write(b) + } + w.committed = true + return w.ResponseWriter.Write(b) +} + +// FlushError and Hijack go through a ResponseController so the capabilities of +// wrapped writers further down stay reachable, as they were before this writer +// was inserted. +func (w *failoverWriter) FlushError() error { + w.release() + w.committed = true + return http.NewResponseController(w.ResponseWriter).Flush() +} + +func (w *failoverWriter) Flush() { _ = w.FlushError() } + +func (w *failoverWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) { + w.committed = true + return http.NewResponseController(w.ResponseWriter).Hijack() +} + +func (w *failoverWriter) Unwrap() http.ResponseWriter { return w.ResponseWriter } + +// release sends a held error response to the client. +func (w *failoverWriter) release() { + if w.held == 0 { + return + } + code := w.held + w.held = 0 + w.committed = true + w.ResponseWriter.WriteHeader(code) + _, _ = w.ResponseWriter.Write(w.body.Bytes()) + w.body.Reset() +} + +// replayBody records what the handler reads, up to limit, so the body can be +// sent again to the next target. Decoders often stop at the end of the value +// without reading to EOF, so a replay is the recorded bytes followed by +// whatever the previous attempt left unread. +type replayBody struct { + src io.ReadCloser + buf bytes.Buffer + limit int + overflow bool +} + +func (r *replayBody) Read(p []byte) (int, error) { + n, err := r.src.Read(p) + if n > 0 && !r.overflow { + if r.buf.Len()+n > r.limit { + r.overflow = true + r.buf.Reset() + } else { + r.buf.Write(p[:n]) + } + } + return n, err +} + +func (r *replayBody) Close() error { return r.src.Close() } + +// replayable reports whether everything read so far was kept. +func (r *replayBody) replayable() bool { return !r.overflow } + +// replay rewinds to the start of the body and keeps recording, so a third +// attempt can replay too. +func (r *replayBody) replay() io.ReadCloser { + data := bytes.Clone(r.buf.Bytes()) + rest := r.src + r.src = struct { + io.Reader + io.Closer + }{io.MultiReader(bytes.NewReader(data), rest), rest} + r.buf.Reset() + return r +} diff --git a/core/http/middleware/failover_test.go b/core/http/middleware/failover_test.go new file mode 100644 index 000000000..94168d23f --- /dev/null +++ b/core/http/middleware/failover_test.go @@ -0,0 +1,233 @@ +package middleware + +import ( + "bytes" + "context" + "errors" + "io" + "mime/multipart" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "sync" + + "github.com/labstack/echo/v4" + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/schema" + "github.com/mudler/LocalAI/core/services/failover" + "github.com/mudler/LocalAI/pkg/model" + "github.com/mudler/LocalAI/pkg/system" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +var _ = Describe("failover chains in the request pipeline", func() { + var ( + app *echo.Echo + fm *failover.Manager + mu sync.Mutex + calls []string + behavior map[string]func(c echo.Context) error + ) + + served := func(c echo.Context) error { + cfg := c.Get(CONTEXT_LOCALS_KEY_MODEL_CONFIG).(*config.ModelConfig) + return c.JSON(http.StatusOK, map[string]string{"served": cfg.Name}) + } + + handler := func(c echo.Context) error { + cfg := c.Get(CONTEXT_LOCALS_KEY_MODEL_CONFIG).(*config.ModelConfig) + mu.Lock() + calls = append(calls, cfg.Name) + b := behavior[cfg.Name] + mu.Unlock() + if b == nil { + return served(c) + } + return b(c) + } + + post := func(path, body string) *httptest.ResponseRecorder { + req := httptest.NewRequest(http.MethodPost, path, strings.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + app.ServeHTTP(rec, req) + return rec + } + chat := func(model string) *httptest.ResponseRecorder { + return post("/v1/chat/completions", `{"model":"`+model+`","messages":[{"role":"user","content":"hi"}]}`) + } + + BeforeEach(func() { + calls = nil + behavior = map[string]func(c echo.Context) error{} + dir := GinkgoT().TempDir() + write := func(name, body string) { + Expect(os.WriteFile(filepath.Join(dir, name+".yaml"), []byte(body), 0o600)).To(Succeed()) + } + write("a", "name: a\nbackend: fake-a\n") + write("b", "name: b\nbackend: fake-b\n") + write("plain", "name: plain\nbackend: fake-p\n") + write("chain", "name: chain\nfailover:\n targets:\n - model: a\n - model: b\n") + + ss := &system.SystemState{Model: system.Model{ModelsPath: dir}} + appConfig := config.NewApplicationConfig() + appConfig.SystemState = ss + mcl := config.NewModelConfigLoader(dir) + Expect(mcl.LoadModelConfigsFromPath(dir)).To(Succeed()) + re := NewRequestExtractor(mcl, model.NewModelLoader(ss), appConfig) + fm = failover.New(mcl) + re.SetFailoverManager(fm) + + app = echo.New() + // echo's default handler hides internal error messages; the specs + // below check which target's error reached the client. + app.HTTPErrorHandler = func(err error, c echo.Context) { + code := http.StatusInternalServerError + var he *echo.HTTPError + if errors.As(err, &he) { + code = he.Code + } + _ = c.JSON(code, map[string]string{"error": err.Error()}) + } + app.POST("/v1/chat/completions", handler, + re.SetModelAndConfig(func() schema.LocalAIRequest { return new(schema.OpenAIRequest) })) + app.POST("/v1/audio/transcriptions", handler, + re.SetModelAndConfig(func() schema.LocalAIRequest { return new(schema.OpenAIRequest) })) + // The real transcription route resolves a default model first, which + // parses the multipart form before SetModelAndConfig runs. + app.POST("/v1/audio/transcriptions-default", handler, + re.BuildFilteredFirstAvailableDefaultModel(config.NoFilterFn), + re.SetModelAndConfig(func() schema.LocalAIRequest { return new(schema.OpenAIRequest) })) + }) + + It("serves from the next target when the first fails before responding", func() { + behavior["a"] = func(echo.Context) error { return errors.New("dial tcp: connection refused") } + rec := chat("chain") + Expect(rec.Code).To(Equal(http.StatusOK)) + Expect(rec.Body.String()).To(ContainSubstring(`"served":"b"`)) + Expect(rec.Header().Get(HeaderServedModel)).To(Equal("b")) + Expect(rec.Header().Get(HeaderFailover)).To(Equal("fallback")) + Expect(calls).To(Equal([]string{"a", "b"})) + st, _ := fm.ChainStatus("chain") + Expect(st.Active).To(Equal("b")) + }) + + It("serves the primary without the failover header", func() { + rec := chat("chain") + Expect(rec.Header().Get(HeaderServedModel)).To(Equal("a")) + Expect(rec.Header().Get(HeaderFailover)).To(BeEmpty()) + }) + + It("drops a buffered 5xx response and retries", func() { + behavior["a"] = func(c echo.Context) error { + return c.JSON(http.StatusServiceUnavailable, map[string]string{"error": "no healthy nodes"}) + } + rec := chat("chain") + Expect(rec.Code).To(Equal(http.StatusOK)) + Expect(rec.Body.String()).ToNot(ContainSubstring("no healthy nodes")) + }) + + It("does not retry after streaming started, and trips the target", func() { + behavior["a"] = func(c echo.Context) error { + c.Response().Header().Set("Content-Type", "text/event-stream") + _, _ = c.Response().Write([]byte("data: x\n\n")) + c.Response().Flush() + return errors.New("connection reset by peer") + } + rec := chat("chain") + Expect(rec.Body.String()).To(HavePrefix("data: x")) + Expect(calls).To(Equal([]string{"a"})) + st, _ := fm.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(failover.StateDown)) + }) + + It("does not retry or trip on 4xx", func() { + behavior["a"] = func(echo.Context) error { return echo.NewHTTPError(http.StatusBadRequest, "bad") } + rec := chat("chain") + Expect(rec.Code).To(Equal(http.StatusBadRequest)) + Expect(calls).To(Equal([]string{"a"})) + st, _ := fm.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(failover.StateHealthy)) + }) + + It("does not retry when the client cancelled", func() { + ctx, cancel := context.WithCancel(context.Background()) + behavior["a"] = func(echo.Context) error { cancel(); return context.Canceled } + req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", + strings.NewReader(`{"model":"chain","messages":[{"role":"user","content":"hi"}]}`)).WithContext(ctx) + req.Header.Set("Content-Type", "application/json") + app.ServeHTTP(httptest.NewRecorder(), req) + Expect(calls).To(Equal([]string{"a"})) + st, _ := fm.ChainStatus("chain") + Expect(st.Targets[0].State).To(Equal(failover.StateHealthy)) + }) + + It("gives each attempt a fresh request", func() { + behavior["a"] = func(c echo.Context) error { + in := c.Get(CONTEXT_LOCALS_KEY_LOCALAI_REQUEST).(*schema.OpenAIRequest) + in.Messages = nil + return errors.New("dial tcp: connection refused") + } + behavior["b"] = func(c echo.Context) error { + in := c.Get(CONTEXT_LOCALS_KEY_LOCALAI_REQUEST).(*schema.OpenAIRequest) + Expect(in.Messages).To(HaveLen(1)) + return served(c) + } + Expect(chat("chain").Code).To(Equal(http.StatusOK)) + }) + + It("degraded: tries every target in priority order and returns the last error", func() { + behavior["a"] = func(echo.Context) error { return errors.New("dial tcp: a down") } + behavior["b"] = func(echo.Context) error { return errors.New("dial tcp: b down") } + chat("chain") // trips both + calls = nil + rec := chat("chain") + Expect(calls).To(Equal([]string{"a", "b"})) + Expect(rec.Code).To(Equal(http.StatusInternalServerError)) + Expect(rec.Body.String()).To(ContainSubstring("b down")) + Expect(rec.Header().Get(HeaderFailover)).To(Equal("degraded")) + }) + + DescribeTable("replays a multipart body for the next target", func(path string) { + var body bytes.Buffer + mw := multipart.NewWriter(&body) + _ = mw.WriteField("model", "chain") + fw, _ := mw.CreateFormFile("file", "a.wav") + _, _ = fw.Write(bytes.Repeat([]byte{7}, 4096)) + Expect(mw.Close()).To(Succeed()) + size := func(c echo.Context) int64 { + fh, err := c.FormFile("file") + Expect(err).ToNot(HaveOccurred()) + f, _ := fh.Open() + n, _ := io.Copy(io.Discard, f) + return n + } + behavior["a"] = func(c echo.Context) error { size(c); return errors.New("dial tcp: refused") } + behavior["b"] = func(c echo.Context) error { + Expect(size(c)).To(Equal(int64(4096))) + return served(c) + } + req := httptest.NewRequest(http.MethodPost, path, &body) + req.Header.Set("Content-Type", mw.FormDataContentType()) + rec := httptest.NewRecorder() + app.ServeHTTP(rec, req) + Expect(rec.Code).To(Equal(http.StatusOK), rec.Body.String()) + Expect(calls).To(Equal([]string{"a", "b"})) + }, + Entry("read by the handler", "/v1/audio/transcriptions"), + Entry("parsed before SetModelAndConfig", "/v1/audio/transcriptions-default"), + ) + + It("leaves plain models untouched", func() { + behavior["plain"] = func(c echo.Context) error { + return c.JSON(http.StatusInternalServerError, map[string]string{"error": "boom"}) + } + rec := chat("plain") + Expect(rec.Code).To(Equal(http.StatusInternalServerError)) + Expect(rec.Body.String()).To(ContainSubstring("boom")) + Expect(rec.Header().Get(HeaderServedModel)).To(BeEmpty()) + }) +}) diff --git a/core/http/middleware/request.go b/core/http/middleware/request.go index 080a0b73c..86ab9a66c 100644 --- a/core/http/middleware/request.go +++ b/core/http/middleware/request.go @@ -12,6 +12,7 @@ import ( "github.com/labstack/echo/v4" "github.com/mudler/LocalAI/core/config" "github.com/mudler/LocalAI/core/schema" + "github.com/mudler/LocalAI/core/services/failover" "github.com/mudler/LocalAI/core/services/galleryop" "github.com/mudler/LocalAI/core/templates" "github.com/mudler/LocalAI/pkg/distributedhdr" @@ -29,6 +30,7 @@ type RequestExtractor struct { modelConfigLoader *config.ModelConfigLoader modelLoader *model.ModelLoader applicationConfig *config.ApplicationConfig + failover *failover.Manager } func NewRequestExtractor(modelConfigLoader *config.ModelConfigLoader, modelLoader *model.ModelLoader, applicationConfig *config.ApplicationConfig) *RequestExtractor { @@ -122,7 +124,7 @@ func (re *RequestExtractor) BuildFilteredFirstAvailableDefaultModel(filterFn con // Otherwise, it's in its own method below for now func (re *RequestExtractor) SetModelAndConfig(initializer func() schema.LocalAIRequest) echo.MiddlewareFunc { return func(next echo.HandlerFunc) echo.HandlerFunc { - return func(c echo.Context) error { + return failoverRetry(re.applicationConfig, func(c echo.Context) error { input := initializer() if input == nil { return echo.NewHTTPError(http.StatusBadRequest, "unable to initialize body") @@ -194,6 +196,22 @@ func (re *RequestExtractor) SetModelAndConfig(initializer func() schema.LocalAIR cfg = resolved } + // A failover chain resolves to one of its targets, like an alias. + // failoverRetry re-runs this middleware for the next target. + if cfg != nil && cfg.IsFailover() { + resolved, fErr := re.resolveFailover(c, modelName, cfg) + if fErr != nil { + return c.JSON(http.StatusServiceUnavailable, schema.ErrorResponse{ + Error: &schema.APIError{ + Message: fErr.Error(), + Code: http.StatusServiceUnavailable, + Type: "failover_unavailable", + }, + }) + } + cfg = resolved + } + // Check if the model is disabled if cfg != nil && cfg.IsDisabled() { return c.JSON(http.StatusForbidden, schema.ErrorResponse{ @@ -209,7 +227,7 @@ func (re *RequestExtractor) SetModelAndConfig(initializer func() schema.LocalAIR c.Set(CONTEXT_LOCALS_KEY_MODEL_CONFIG, cfg) return next(c) - } + }) } }