Files
LocalAI/core/http/middleware/failover_test.go
T
Ettore Di Giacinto e183ea11d1 fix(failover): send remote targets their own upstream model
A chain request reached a cloud-proxy target with the client's model,
the chain name, whenever the target set no upstream_model: passthrough
forwards the body's model and translate falls back to it. The upstream
answered 404, which neither retries nor trips, while the liveness
probe, which checks the target's own name, kept passing.

PrepareTarget now sets the upstream model of a remote target to
proxy.upstream_model or the target name, the same name the probe uses.
The request pipeline and realtime chain stages both call it.

Assisted-by: Claude:claude-opus-5-5
2026-09-26 19:26:13 +00:00

344 lines
13 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")
write("remote", "name: remote\nbackend: cloud-proxy\nproxy:\n mode: passthrough\n upstream_url: http://127.0.0.1:1/v1/chat/completions\n")
write("remote-mapped", "name: remote-mapped\nbackend: cloud-proxy\nproxy:\n mode: translate\n provider: openai\n upstream_url: http://127.0.0.1:1/v1/chat/completions\n upstream_model: big-llm\n")
write("chain-remote", "name: chain-remote\nfailover:\n targets:\n - model: remote\n - model: remote-mapped\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)
// The scheduler's first tick syncs in the application; HasChains
// answers from the last sync.
fm.Sync()
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("sends a remote target its own upstream model, not the chain name", func() {
upstream := map[string]string{}
record := func(c echo.Context) error {
cfg := c.Get(CONTEXT_LOCALS_KEY_MODEL_CONFIG).(*config.ModelConfig)
mu.Lock()
upstream[cfg.Name] = cfg.Proxy.UpstreamModel
mu.Unlock()
return nil
}
behavior["remote"] = func(c echo.Context) error {
_ = record(c)
return errors.New("dial tcp: connection refused")
}
behavior["remote-mapped"] = func(c echo.Context) error { _ = record(c); return served(c) }
Expect(chat("chain-remote").Code).To(Equal(http.StatusOK))
// The probe checks the same name, so a request cannot fail on a model
// the liveness probe just found.
Expect(upstream).To(Equal(map[string]string{"remote": "remote", "remote-mapped": "big-llm"}))
for _, name := range []string{"remote", "remote-mapped"} {
cfg, ok := re.modelConfigLoader.GetModelConfig(name)
Expect(ok).To(BeTrue())
Expect(upstream[name]).To(Equal(failover.UpstreamModel(cfg)))
}
stored, _ := re.modelConfigLoader.GetModelConfig("remote")
Expect(stored.Proxy.UpstreamModel).To(BeEmpty(), "the shared config must not change")
})
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())
})
})