mirror of
https://github.com/mudler/LocalAI.git
synced 2026-09-30 18:14:32 -04:00
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
This commit is contained in:
1 parent
ba42a195b4
commit
1148d29b41
5 files changed
+540
-2
No files matched your search
@@ -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)
|
||||
|
||||
@@ -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"
|
||||
)
|
||||
@@ -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
|
||||
}
|
||||
@@ -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())
|
||||
})
|
||||
})
|
||||
@@ -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)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in new issue
Block a user