diff --git a/core/http/middleware/admission.go b/core/http/middleware/admission.go index c79066925..13c97c711 100644 --- a/core/http/middleware/admission.go +++ b/core/http/middleware/admission.go @@ -38,6 +38,7 @@ func AdmissionControl(limiter *admission.Limiter, events pii.EventStore) echo.Mi if !ok { retryAfter := admission.RetryAfter(cfg.Limits.RetryAfterSeconds) recordAdmissionRejection(events, cfg.Name, retryAfter) + c.Set(ContextKeyAdmissionRejected, true) c.Response().Header().Set("Retry-After", strconv.Itoa(int(retryAfter.Seconds()))) return c.JSON(http.StatusServiceUnavailable, map[string]any{ "error": map[string]any{ diff --git a/core/http/middleware/context_keys.go b/core/http/middleware/context_keys.go index f8b5582da..1e1f3886c 100644 --- a/core/http/middleware/context_keys.go +++ b/core/http/middleware/context_keys.go @@ -51,4 +51,9 @@ const ( // ContextKeyFailoverAttempt holds the *failoverState of a request whose // model is a failover chain. ContextKeyFailoverAttempt = "failover.attempt" + + // ContextKeyAdmissionRejected is set to true by AdmissionControl when it + // turns a request away because the model is at capacity. The failover + // retry sends such a request to the next target without tripping this one. + ContextKeyAdmissionRejected = "admission.rejected" ) diff --git a/core/http/middleware/failover.go b/core/http/middleware/failover.go index 930b5e209..f3744fdd9 100644 --- a/core/http/middleware/failover.go +++ b/core/http/middleware/failover.go @@ -49,6 +49,14 @@ func (re *RequestExtractor) resolveFailover(c echo.Context, requested string, ch } for { cfg, err := re.loadFailoverTarget(st.attempt.Target()) + if err == nil && cfg.IsDisabled() { + // Disabled on purpose, not broken: move on without a trip. + if st.attempt.Skip() { + continue + } + c.Set(ContextKeyFailoverAttempt, nil) + return nil, fmt.Errorf("failover chain %q: target %q is disabled", chain.Name, cfg.Name) + } if err == nil { c.Set(ContextKeyRequestedModel, requested) c.Set(ContextKeyServedModel, cfg.Name) @@ -87,8 +95,13 @@ func setFailoverHeaders(h http.Header, att *failover.Attempt) { // 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 { +func (re *RequestExtractor) failoverRetry(h echo.HandlerFunc) echo.HandlerFunc { return func(c echo.Context) error { + // Without chains there is nothing to retry, so plain installations + // pay neither for the body copy nor for the writer. + if !re.failover.HasChains() { + return h(c) + } req := c.Request() src := req.Body if src == nil { @@ -109,10 +122,11 @@ func failoverRetry(appConfig *config.ApplicationConfig, h echo.HandlerFunc) echo st, _ := c.Get(ContextKeyFailoverAttempt).(*failoverState) return st != nil } - tracing := appConfig != nil && appConfig.EnableTracing + tracing := re.applicationConfig != nil && re.applicationConfig.EnableTracing for { w := &failoverWriter{ResponseWriter: orig, active: active} resp.Writer = w + c.Set(ContextKeyAdmissionRejected, nil) err := h(c) st, _ := c.Get(ContextKeyFailoverAttempt).(*failoverState) if st == nil { @@ -122,22 +136,34 @@ func failoverRetry(appConfig *config.ApplicationConfig, h echo.HandlerFunc) echo att := st.attempt status := w.held if err == nil && status == 0 { - att.Succeed() + // A 4xx says nothing about the target's health. + if w.status < http.StatusBadRequest { + 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) + if rejected, _ := c.Get(ContextKeyAdmissionRejected).(bool); rejected && !w.committed { + // The target is at capacity, not broken: spill this request to + // the next target without counting a failure. + if !rec.replayable() || !att.Skip() { + w.release() + return err + } + } else { + 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 } - 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. @@ -154,6 +180,14 @@ func failoverRetry(appConfig *config.ApplicationConfig, h echo.HandlerFunc) echo } } +// stopFailoverRecording releases the recorded body of a request whose model +// turned out not to be a chain: it will never be replayed. +func stopFailoverRecording(c echo.Context) { + if rb, ok := c.Request().Body.(*replayBody); ok { + rb.stop() + } +} + func attemptError(err error, status int, body []byte) error { if err != nil { return err @@ -186,6 +220,8 @@ type failoverWriter struct { held int body bytes.Buffer committed bool + // status is the code sent to the client, 0 until one is sent. + status int } func (w *failoverWriter) WriteHeader(code int) { @@ -197,6 +233,7 @@ func (w *failoverWriter) WriteHeader(code int) { return } w.committed = true + w.status = code w.ResponseWriter.WriteHeader(code) } @@ -248,14 +285,16 @@ type replayBody struct { buf bytes.Buffer limit int overflow bool + stopped bool } func (r *replayBody) Read(p []byte) (int, error) { n, err := r.src.Read(p) - if n > 0 && !r.overflow { + if n > 0 && !r.overflow && !r.stopped { if r.buf.Len()+n > r.limit { r.overflow = true - r.buf.Reset() + // A new buffer, not Reset: Reset keeps the memory. + r.buf = bytes.Buffer{} } else { r.buf.Write(p[:n]) } @@ -266,7 +305,13 @@ func (r *replayBody) Read(p []byte) (int, error) { 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 } +func (r *replayBody) replayable() bool { return !r.overflow && !r.stopped } + +// stop ends recording and frees what was kept. +func (r *replayBody) stop() { + r.stopped = true + r.buf = bytes.Buffer{} +} // replay rewinds to the start of the body and keeps recording, so a third // attempt can replay too. diff --git a/core/http/middleware/failover_test.go b/core/http/middleware/failover_test.go index 94168d23f..a16337545 100644 --- a/core/http/middleware/failover_test.go +++ b/core/http/middleware/failover_test.go @@ -12,11 +12,13 @@ import ( "path/filepath" "strings" "sync" + "testing/iotest" "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/routing/admission" "github.com/mudler/LocalAI/pkg/model" "github.com/mudler/LocalAI/pkg/system" . "github.com/onsi/ginkgo/v2" @@ -27,6 +29,8 @@ var _ = Describe("failover chains in the request pipeline", func() { var ( app *echo.Echo fm *failover.Manager + re *RequestExtractor + limiter *admission.Limiter mu sync.Mutex calls []string behavior map[string]func(c echo.Context) error @@ -71,13 +75,17 @@ var _ = Describe("failover chains in the request pipeline", func() { 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") + write("capped", "name: capped\nbackend: fake-c\nlimits:\n max_concurrent: 1\n") + write("off", "name: off\nbackend: fake-o\ndisabled: true\n") + write("chain-capped", "name: chain-capped\nfailover:\n targets:\n - model: capped\n - model: b\n") + write("chain-off", "name: chain-off\nfailover:\n targets:\n - model: off\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) + re = NewRequestExtractor(mcl, model.NewModelLoader(ss), appConfig) fm = failover.New(mcl) re.SetFailoverManager(fm) @@ -96,6 +104,10 @@ var _ = Describe("failover chains in the request pipeline", func() { re.SetModelAndConfig(func() schema.LocalAIRequest { return new(schema.OpenAIRequest) })) app.POST("/v1/audio/transcriptions", handler, re.SetModelAndConfig(func() schema.LocalAIRequest { return new(schema.OpenAIRequest) })) + limiter = admission.New() + app.POST("/v1/chat/admitted", handler, + re.SetModelAndConfig(func() schema.LocalAIRequest { return new(schema.OpenAIRequest) }), + AdmissionControl(limiter, nil)) // The real transcription route resolves a default model first, which // parses the multipart form before SetModelAndConfig runs. app.POST("/v1/audio/transcriptions-default", handler, @@ -221,6 +233,71 @@ var _ = Describe("failover chains in the request pipeline", func() { Entry("parsed before SetModelAndConfig", "/v1/audio/transcriptions-default"), ) + It("spills an admission rejection to the next target without tripping it", func() { + release, ok := limiter.Acquire("capped", 1) + Expect(ok).To(BeTrue()) + defer release() + rec := post("/v1/chat/admitted", `{"model":"chain-capped","messages":[{"role":"user","content":"hi"}]}`) + Expect(rec.Code).To(Equal(http.StatusOK), rec.Body.String()) + Expect(rec.Body.String()).To(ContainSubstring(`"served":"b"`)) + Expect(rec.Header().Get("Retry-After")).To(BeEmpty()) + st, _ := fm.ChainStatus("chain-capped") + Expect(st.Targets[0].State).To(Equal(failover.StateHealthy)) + Expect(st.Active).To(Equal("capped")) + }) + + It("skips a disabled target without tripping it", func() { + rec := chat("chain-off") + Expect(rec.Code).To(Equal(http.StatusOK), rec.Body.String()) + Expect(rec.Body.String()).To(ContainSubstring(`"served":"b"`)) + Expect(calls).To(Equal([]string{"b"})) + st, _ := fm.ChainStatus("chain-off") + Expect(st.Targets[0].State).To(Equal(failover.StateHealthy)) + }) + + It("counts a 4xx response as neither success nor failure", 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 + behavior["a"] = func(c echo.Context) error { + return c.JSON(http.StatusBadRequest, map[string]string{"error": "bad"}) + } + rec := chat("chain") + Expect(rec.Code).To(Equal(http.StatusBadRequest)) + st, _ := fm.ChainStatus("chain") + // a is a cold local target: a success would have recovered it. + Expect(st.Targets[0].State).To(Equal(failover.StateDown)) + }) + + It("stops recording the body once the model is known not to be a chain", func() { + behavior["plain"] = func(c echo.Context) error { + rb, ok := c.Request().Body.(*replayBody) + Expect(ok).To(BeTrue()) + Expect(rb.replayable()).To(BeFalse()) + Expect(rb.buf.Cap()).To(BeZero()) + return served(c) + } + Expect(chat("plain").Code).To(Equal(http.StatusOK)) + }) + + It("does not record the body without a failover manager", func() { + re.SetFailoverManager(nil) + behavior["plain"] = func(c echo.Context) error { + _, ok := c.Request().Body.(*replayBody) + Expect(ok).To(BeFalse()) + return served(c) + } + Expect(chat("plain").Code).To(Equal(http.StatusOK)) + }) + + It("releases the recorded bytes on overflow", func() { + // One byte per read, so some bytes are recorded before the limit is hit. + rb := &replayBody{src: io.NopCloser(iotest.OneByteReader(bytes.NewReader(make([]byte, 64)))), limit: 16} + _, _ = io.ReadAll(rb) + Expect(rb.replayable()).To(BeFalse()) + Expect(rb.buf.Cap()).To(BeZero()) + }) + It("leaves plain models untouched", func() { behavior["plain"] = func(c echo.Context) error { return c.JSON(http.StatusInternalServerError, map[string]string{"error": "boom"}) diff --git a/core/http/middleware/request.go b/core/http/middleware/request.go index 86ab9a66c..7bc702e20 100644 --- a/core/http/middleware/request.go +++ b/core/http/middleware/request.go @@ -124,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 failoverRetry(re.applicationConfig, func(c echo.Context) error { + return re.failoverRetry(func(c echo.Context) error { input := initializer() if input == nil { return echo.NewHTTPError(http.StatusBadRequest, "unable to initialize body") @@ -210,6 +210,8 @@ func (re *RequestExtractor) SetModelAndConfig(initializer func() schema.LocalAIR }) } cfg = resolved + } else { + stopFailoverRecording(c) } // Check if the model is disabled diff --git a/core/services/failover/manager.go b/core/services/failover/manager.go index 408c44621..ab130d1ee 100644 --- a/core/services/failover/manager.go +++ b/core/services/failover/manager.go @@ -201,6 +201,28 @@ func (m *Manager) lookupTarget(name string) (config.ModelConfig, bool) { return c, ok } +// HasChains reports whether any failover chain is configured. The request +// path uses it to skip chain bookkeeping on installations without chains. +// Chains are synced lazily, so with none known yet the config source is +// consulted, which catches a chain added since the last sync. +func (m *Manager) HasChains() bool { + if m == nil { + return false + } + m.mu.Lock() + known := len(m.chains) > 0 + m.mu.Unlock() + if known { + return true + } + for _, c := range m.src.GetAllModelsConfigs() { + if c.IsFailover() { + return true + } + } + return false +} + func (m *Manager) chainLocked(name string) *chainState { if ch := m.chains[name]; ch != nil { return ch @@ -405,6 +427,17 @@ func (a *Attempt) Fail(err error) bool { return true } +// Skip moves to the next target without recording a failure, for a target +// that could not take this request although nothing is wrong with it (at +// capacity, disabled). It returns false when no target is left. +func (a *Attempt) Skip() bool { + if a.i+1 >= len(a.targets) { + return false + } + a.i++ + return true +} + // Report records err against the current target without moving on: the // response was already committed, so nothing is left to retry. func (a *Attempt) Report(err error) { a.m.ReportFailure(a.Target(), err) } diff --git a/core/services/failover/manager_test.go b/core/services/failover/manager_test.go index 3bef7945e..d0736cf19 100644 --- a/core/services/failover/manager_test.go +++ b/core/services/failover/manager_test.go @@ -49,6 +49,31 @@ var _ = Describe("Manager", func() { Expect(st.Targets[1].Kind).To(Equal(KindLocal)) }) + It("skips to the next target without recording a failure", func() { + att, _ := m.Plan("chain") + Expect(att.Skip()).To(BeTrue()) + Expect(att.Target()).To(Equal("b")) + Expect(att.Skip()).To(BeFalse()) + Expect(att.Target()).To(Equal("b")) + st, _ := m.ChainStatus("chain") + Expect(st.Active).To(Equal("a")) + Expect(st.Targets[0].State).To(Equal(StateHealthy)) + }) + + It("reports whether any chain is configured, including one added since the last sync", func() { + empty := New(newFakeSource(remote("a")), WithClock(clock)) + Expect(empty.HasChains()).To(BeFalse()) + lateSrc := newFakeSource(remote("a"), local("b")) + late := New(lateSrc, WithClock(clock)) + late.Sync() + Expect(late.HasChains()).To(BeFalse()) + lateSrc.Put(chainCfg("chain", nil, t("a"), t("b"))) + Expect(late.HasChains()).To(BeTrue()) + Expect(m.HasChains()).To(BeTrue()) + var nilManager *Manager + Expect(nilManager.HasChains()).To(BeFalse()) + }) + It("returns ErrChainNotFound for an unknown chain", func() { _, err := m.Plan("nope") Expect(errors.Is(err, ErrChainNotFound)).To(BeTrue())