Files
LocalAI/core/services/failover/prober.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

294 lines
9.1 KiB
Go

package failover
import (
"bytes"
"context"
"encoding/binary"
"encoding/json"
"errors"
"fmt"
"io"
"mime/multipart"
"net/http"
"net/url"
"strings"
"github.com/mudler/LocalAI/core/config"
"github.com/mudler/LocalAI/pkg/grpc"
pb "github.com/mudler/LocalAI/pkg/grpc/proto"
)
// ErrNotLoaded is what Inference returns for a local target whose backend is
// not running. It neither confirms nor fails recovery: the manager judges the
// target like a cold one, by real requests once min_dwell has passed.
var ErrNotLoaded = errors.New("failover: target is not loaded")
// LoadedFunc returns the running backend of a local target, or nil when it is
// not loaded. It must never load the model: a probe that loads blocks until
// the load ends (while the warm preload loads the same model) and then judges
// the target on an expired context.
type LoadedFunc func(cfg config.ModelConfig) grpc.Backend
// DefaultProber probes remote targets over the upstream's OpenAI-compatible
// API and local targets through their gRPC backend.
type DefaultProber struct {
HTTP *http.Client
Loaded LoadedFunc
}
func NewProber(loaded LoadedFunc) *DefaultProber {
return &DefaultProber{HTTP: &http.Client{}, Loaded: loaded}
}
func (p *DefaultProber) Liveness(ctx context.Context, cfg config.ModelConfig, kind Kind, warm bool) error {
switch {
case kind == KindRemote:
return p.remoteLiveness(ctx, cfg)
case warm:
return p.localHealth(ctx, cfg)
}
// A cold target is judged only by real requests: loading it only to probe
// it could evict other models, and no cheaper check is reliable. A missing
// model file is not one: models download on first use, some backends need
// no file, and a dotted name like "Phi-3.5-mini" looks like a file path. A
// false "down" would take away the very fallback the chain exists for,
// while a real load failure still trips the target and the request moves on.
return nil
}
func (p *DefaultProber) Inference(ctx context.Context, cfg config.ModelConfig, kind Kind, warm bool) error {
if kind == KindRemote {
return p.remoteInference(ctx, cfg)
}
return p.localInference(ctx, cfg)
}
// UpstreamBase strips the endpoint path from a cloud-proxy upstream_url:
// everything from "/v1" on, so a path prefix before it survives.
func UpstreamBase(raw string) (string, error) {
u, err := url.Parse(raw)
if err != nil || u.Scheme == "" || u.Host == "" {
return "", fmt.Errorf("invalid upstream_url %q", raw)
}
path := u.Path
if i := strings.Index(path, "/v1"); i >= 0 {
path = path[:i]
}
return u.Scheme + "://" + u.Host + strings.TrimSuffix(path, "/"), nil
}
// UpstreamModel is the model name the upstream knows the target by.
func UpstreamModel(cfg config.ModelConfig) string {
if cfg.Proxy.UpstreamModel != "" {
return cfg.Proxy.UpstreamModel
}
return cfg.Name
}
// PrepareTarget readies a copy of a target's config to serve a chain request.
// A remote target gets its upstream model set explicitly: left empty,
// passthrough forwards the client's "model" (the chain name) and translate
// falls back to it, so the upstream would answer 404 for a model the liveness
// probe (which checks UpstreamModel) just found.
func PrepareTarget(cfg *config.ModelConfig) {
if KindOf(*cfg) == KindRemote {
cfg.Proxy.UpstreamModel = UpstreamModel(*cfg)
}
}
func (p *DefaultProber) authorize(req *http.Request, cfg config.ModelConfig) error {
key, err := cfg.Proxy.ResolveAPIKey()
if err != nil || key == "" {
return err
}
if cfg.Proxy.Provider == config.ProxyProviderAnthropic {
req.Header.Set("x-api-key", key)
req.Header.Set("anthropic-version", "2023-06-01")
return nil
}
req.Header.Set("Authorization", "Bearer "+key)
return nil
}
func (p *DefaultProber) do(req *http.Request, cfg config.ModelConfig) (*http.Response, error) {
if err := p.authorize(req, cfg); err != nil {
return nil, err
}
resp, err := p.HTTP.Do(req)
if err != nil {
return nil, err
}
if resp.StatusCode/100 != 2 {
resp.Body.Close()
return nil, fmt.Errorf("upstream %s: HTTP %d", req.URL.Path, resp.StatusCode)
}
return resp, nil
}
func (p *DefaultProber) remoteLiveness(ctx context.Context, cfg config.ModelConfig) error {
base, err := UpstreamBase(cfg.Proxy.UpstreamURL)
if err != nil {
return err
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, base+"/v1/models", nil)
if err != nil {
return err
}
resp, err := p.do(req, cfg)
if err != nil {
return err
}
defer resp.Body.Close()
var list struct {
Data []struct {
ID string `json:"id"`
} `json:"data"`
}
if err := json.NewDecoder(io.LimitReader(resp.Body, 4<<20)).Decode(&list); err != nil {
return fmt.Errorf("upstream /v1/models: %w", err)
}
want := UpstreamModel(cfg)
for _, d := range list.Data {
if d.ID == want {
return nil
}
}
return fmt.Errorf("upstream does not list model %q", want)
}
func (p *DefaultProber) remoteInference(ctx context.Context, cfg config.ModelConfig) error {
base, err := UpstreamBase(cfg.Proxy.UpstreamURL)
if err != nil {
return err
}
model := UpstreamModel(cfg)
ping := []map[string]string{{"role": "user", "content": "ping"}}
switch {
case cfg.HasUsecases(config.FLAG_CHAT) || cfg.HasUsecases(config.FLAG_COMPLETION):
if cfg.Proxy.Provider == config.ProxyProviderAnthropic {
return p.postJSON(ctx, cfg, base+"/v1/messages", map[string]any{"model": model, "max_tokens": 1, "messages": ping})
}
return p.postJSON(ctx, cfg, base+"/v1/chat/completions", map[string]any{"model": model, "max_tokens": 1, "messages": ping})
case cfg.HasUsecases(config.FLAG_EMBEDDINGS):
return p.postJSON(ctx, cfg, base+"/v1/embeddings", map[string]any{"model": model, "input": "ping"})
case cfg.HasUsecases(config.FLAG_TRANSCRIPT):
return p.postTranscription(ctx, cfg, base, model)
case cfg.HasUsecases(config.FLAG_TTS):
return p.postJSON(ctx, cfg, base+"/v1/audio/speech", map[string]any{"model": model, "input": "ok"})
}
// Image, video and other costly usecases: liveness is the confirmation.
return p.remoteLiveness(ctx, cfg)
}
func (p *DefaultProber) postJSON(ctx context.Context, cfg config.ModelConfig, endpoint string, body any) error {
b, err := json.Marshal(body)
if err != nil {
return err
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(b))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
resp, err := p.do(req, cfg)
if err != nil {
return err
}
_, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, 1<<20))
return resp.Body.Close()
}
func (p *DefaultProber) postTranscription(ctx context.Context, cfg config.ModelConfig, base, model string) error {
var buf bytes.Buffer
mw := multipart.NewWriter(&buf)
_ = mw.WriteField("model", model)
fw, err := mw.CreateFormFile("file", "probe.wav")
if err != nil {
return err
}
_, _ = fw.Write(silenceWAV())
if err := mw.Close(); err != nil {
return err
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, base+"/v1/audio/transcriptions", &buf)
if err != nil {
return err
}
req.Header.Set("Content-Type", mw.FormDataContentType())
resp, err := p.do(req, cfg)
if err != nil {
return err
}
_, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, 1<<20))
return resp.Body.Close()
}
// silenceWAV is 200 ms of 16 kHz mono 16-bit silence.
func silenceWAV() []byte {
const rate, samples = 16000, 3200
data := samples * 2
b := make([]byte, 44+data)
copy(b[0:], "RIFF")
binary.LittleEndian.PutUint32(b[4:], uint32(36+data))
copy(b[8:], "WAVE")
copy(b[12:], "fmt ")
binary.LittleEndian.PutUint32(b[16:], 16)
binary.LittleEndian.PutUint16(b[20:], 1) // PCM
binary.LittleEndian.PutUint16(b[22:], 1) // mono
binary.LittleEndian.PutUint32(b[24:], rate)
binary.LittleEndian.PutUint32(b[28:], rate*2)
binary.LittleEndian.PutUint16(b[32:], 2)
binary.LittleEndian.PutUint16(b[34:], 16)
copy(b[36:], "data")
binary.LittleEndian.PutUint32(b[40:], uint32(data))
return b
}
func (p *DefaultProber) loaded(cfg config.ModelConfig) grpc.Backend {
if p.Loaded == nil {
return nil
}
return p.Loaded(cfg)
}
// localHealth checks a warm target's running backend. A target that is not
// loaded passes: the warm preload is loading it, or a crash removed it and
// the next real request loads it again and judges it.
func (p *DefaultProber) localHealth(ctx context.Context, cfg config.ModelConfig) error {
b := p.loaded(cfg)
if b == nil {
return nil
}
return healthCheck(ctx, b)
}
func healthCheck(ctx context.Context, b grpc.Backend) error {
ok, err := b.HealthCheck(ctx)
if err != nil {
return err
}
if !ok {
return errors.New("backend health check failed")
}
return nil
}
func (p *DefaultProber) localInference(ctx context.Context, cfg config.ModelConfig) error {
b := p.loaded(cfg)
if b == nil {
return ErrNotLoaded
}
var err error
switch {
case cfg.HasUsecases(config.FLAG_CHAT) || cfg.HasUsecases(config.FLAG_COMPLETION):
_, err = b.Predict(ctx, &pb.PredictOptions{Prompt: "ping", Tokens: 1})
return err
case cfg.HasUsecases(config.FLAG_EMBEDDINGS):
_, err = b.Embeddings(ctx, &pb.PredictOptions{Embeddings: "ping"})
return err
}
// A backend process that answers HealthCheck rarely fails only for TTS or
// transcription, so a real request adds little here.
return healthCheck(ctx, b)
}