mirror of
https://github.com/mudler/LocalAI.git
synced 2026-09-29 01:25:03 -04:00
fix(failover): spill capacity and disabled targets without tripping them
An admission rejection or a disabled target moves the request to the next target through Attempt.Skip, which records no failure. A 4xx response no longer counts as a success. Requests skip body recording when no chain is configured, and stop it once the model is known not to be a chain. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto <mudler@localai.io>
This commit is contained in:
1 parent
471bd8bf96
commit
f28d07b09d
7 files changed
+208
-20
No files matched your search
@@ -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{
|
||||
|
||||
@@ -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"
|
||||
)
|
||||
@@ -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.
|
||||
|
||||
@@ -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"})
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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) }
|
||||
|
||||
@@ -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())
|
||||
|
||||
Reference in new issue
Block a user