mirror of
https://github.com/mudler/LocalAI.git
synced 2026-09-29 17:44:30 -04:00
Go resends custom headers such as x-api-key when it follows a redirect, also to another host, so a redirecting upstream could receive the target's API key elsewhere. The probe client now treats a redirect as the response, which fails the probe as a non-2xx status. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto <mudler@localai.io>
299 lines
9.4 KiB
Go
299 lines
9.4 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{
|
|
// 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}
|
|
}
|
|
|
|
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)
|
|
}
|