mirror of
https://github.com/mudler/LocalAI.git
synced 2026-09-29 17:44:30 -04:00
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 Signed-off-by: Ettore Di Giacinto <mudler@localai.io>
329 lines
9.3 KiB
Go
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
|
|
}
|