Files
LocalAI/core/services/failover/prober.go
T
localai-org-maint-bot e62854c340 fix(failover): pass credential lookup from CLI
Pass the API key environment lookup through ApplicationConfig to satisfy
core configuration lint. Keep credential resolution dynamic and exclude
the callback from serialization.

Handle the five close results reported by errcheck.

Assisted-by: Codex:gpt-6 golangci-lint
Signed-off-by: Ettore Di Giacinto <mudler@localai.io>
2026-09-27 07:42:20 +00:00

300 lines
9.5 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
EnvLookup func(string) string
}
func NewProber(loaded LoadedFunc, envLookup func(string) string) *DefaultProber {
return &DefaultProber{HTTP: &http.Client{
// A redirect is a failed probe, not something to follow: Go resends
// custom headers such as x-api-key to any host, and the target's
// API key must reach only the configured upstream.
CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse },
}, Loaded: loaded, EnvLookup: envLookup}
}
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(p.EnvLookup)
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 func() { _ = 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)
}