Files
LocalAI/core/http/middleware/failover.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

329 lines
9.3 KiB
Go

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 && 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 {
failover.PrepareTarget(cfg) // cfg is a copy
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 (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 {
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 := 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 {
w.release()
return err
}
att := st.attempt
status := w.held
if err == nil && status == 0 {
// A 4xx says nothing about the target's health.
if w.status < http.StatusBadRequest {
att.Succeed()
}
return nil
}
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
}
}
// 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)
}
}
}
// 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
}
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
// status is the code sent to the client, 0 until one is sent.
status int
}
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.status = code
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
stopped bool
}
func (r *replayBody) Read(p []byte) (int, error) {
n, err := r.src.Read(p)
if n > 0 && !r.overflow && !r.stopped {
if r.buf.Len()+n > r.limit {
r.overflow = true
// A new buffer, not Reset: Reset keeps the memory.
r.buf = bytes.Buffer{}
} 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 && !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.
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
}