mirror of
https://github.com/mudler/LocalAI.git
synced 2026-09-29 17:44:30 -04:00
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>
311 lines
12 KiB
Go
311 lines
12 KiB
Go
package middleware
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"errors"
|
|
"io"
|
|
"mime/multipart"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"os"
|
|
"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"
|
|
. "github.com/onsi/gomega"
|
|
)
|
|
|
|
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
|
|
)
|
|
|
|
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")
|
|
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)
|
|
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) }))
|
|
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,
|
|
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("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"})
|
|
}
|
|
rec := chat("plain")
|
|
Expect(rec.Code).To(Equal(http.StatusInternalServerError))
|
|
Expect(rec.Body.String()).To(ContainSubstring("boom"))
|
|
Expect(rec.Header().Get(HeaderServedModel)).To(BeEmpty())
|
|
})
|
|
})
|