From a7dc8bf42c7eef0604bb07714779f64d01abd6e1 Mon Sep 17 00:00:00 2001 From: Ettore Di Giacinto Date: Sat, 26 Sep 2026 15:55:43 +0000 Subject: [PATCH] feat(failover): probe remote targets over HTTP and local ones over gRPC Remote liveness uses /v1/models, which every OpenAI-compatible upstream serves. Recovery sends one minimal request for the target's usecase. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto --- core/config/model_config.go | 20 ++ core/config/model_config_failover_test.go | 21 ++ core/services/failover/prober.go | 283 ++++++++++++++++++ core/services/failover/prober_test.go | 169 +++++++++++ ...2026-09-26-model-failover-chains-design.md | 2 +- 5 files changed, 494 insertions(+), 1 deletion(-) create mode 100644 core/services/failover/prober.go create mode 100644 core/services/failover/prober_test.go diff --git a/core/config/model_config.go b/core/config/model_config.go index 33fe12d21..533da24ec 100644 --- a/core/config/model_config.go +++ b/core/config/model_config.go @@ -302,6 +302,26 @@ const ( ProxyProviderAnthropic = "anthropic" ) +// ResolveAPIKey returns the upstream key from api_key_env or api_key_file, or +// "" when neither is set. The cloud-proxy backend applies the same rules. +func (p ProxyConfig) ResolveAPIKey() (string, error) { + switch { + case p.APIKeyEnv != "": + v, ok := os.LookupEnv(p.APIKeyEnv) + if !ok { + return "", fmt.Errorf("proxy api_key_env %q is not set", p.APIKeyEnv) + } + return v, nil + case p.APIKeyFile != "": + b, err := os.ReadFile(p.APIKeyFile) + if err != nil { + return "", fmt.Errorf("proxy api_key_file: %w", err) + } + return strings.TrimSpace(string(b)), nil + } + return "", nil +} + // IsCloudProxyBackendPassthrough reports whether this model uses the // cloud-proxy gRPC backend in passthrough mode. Empty Mode counts as // passthrough (SetDefaults normalises it, but Validate accepts empty diff --git a/core/config/model_config_failover_test.go b/core/config/model_config_failover_test.go index e79ed5f45..fe34c5c66 100644 --- a/core/config/model_config_failover_test.go +++ b/core/config/model_config_failover_test.go @@ -1,6 +1,8 @@ package config import ( + "os" + "path/filepath" "time" . "github.com/onsi/ginkgo/v2" @@ -66,3 +68,22 @@ failover: Entry("no name", func(c *ModelConfig) { c.Name = "" }, "requires a name"), ) }) + +var _ = Describe("ProxyConfig.ResolveAPIKey", func() { + It("reads the env var", func() { + GinkgoT().Setenv("FAILOVER_TEST_KEY", "k1") + Expect(ProxyConfig{APIKeyEnv: "FAILOVER_TEST_KEY"}.ResolveAPIKey()).To(Equal("k1")) + }) + It("fails on an unset env var", func() { + _, err := ProxyConfig{APIKeyEnv: "FAILOVER_TEST_UNSET_KEY"}.ResolveAPIKey() + Expect(err).To(HaveOccurred()) + }) + It("reads and trims the key file", func() { + f := filepath.Join(GinkgoT().TempDir(), "key") + Expect(os.WriteFile(f, []byte(" k2\n"), 0o600)).To(Succeed()) + Expect(ProxyConfig{APIKeyFile: f}.ResolveAPIKey()).To(Equal("k2")) + }) + It("returns empty when nothing is set", func() { + Expect(ProxyConfig{}.ResolveAPIKey()).To(Equal("")) + }) +}) diff --git a/core/services/failover/prober.go b/core/services/failover/prober.go new file mode 100644 index 000000000..b3055c17f --- /dev/null +++ b/core/services/failover/prober.go @@ -0,0 +1,283 @@ +package failover + +import ( + "bytes" + "context" + "encoding/binary" + "encoding/json" + "errors" + "fmt" + "io" + "io/fs" + "mime/multipart" + "net/http" + "net/url" + "os" + "path/filepath" + "strings" + + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/pkg/grpc" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" +) + +// LoadFunc returns the backend for a local target, loading it if needed. +type LoadFunc func(ctx context.Context, cfg config.ModelConfig) (grpc.Backend, error) + +// DefaultProber probes remote targets over the upstream's OpenAI-compatible +// API and local targets through their gRPC backend. +type DefaultProber struct { + HTTP *http.Client + Load LoadFunc + ModelPath string +} + +func NewProber(load LoadFunc, modelPath string) *DefaultProber { + return &DefaultProber{HTTP: &http.Client{}, Load: load, ModelPath: modelPath} +} + +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) + } + return p.coldLiveness(cfg) +} + +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 +} + +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) localHealth(ctx context.Context, cfg config.ModelConfig) error { + if p.Load == nil { + return errors.New("failover: no backend loader configured") + } + // Load returns the running backend, or starts it again after a crash. + b, err := p.Load(ctx, cfg) + if err != nil { + return err + } + 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 { + if p.Load == nil { + return errors.New("failover: no backend loader configured") + } + b, err := p.Load(ctx, cfg) + if err != nil { + return err + } + 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 p.localHealth(ctx, cfg) +} + +// coldLiveness checks the model file without loading the model. +func (p *DefaultProber) coldLiveness(cfg config.ModelConfig) error { + f := cfg.Model + if f == "" || p.ModelPath == "" || strings.Contains(f, "://") { + return nil + } + path := f + if !filepath.IsAbs(path) { + path = filepath.Join(p.ModelPath, f) + } + if _, err := os.Stat(path); err != nil { + if errors.Is(err, fs.ErrNotExist) && filepath.Ext(f) == "" { + return nil // a repository id, downloaded on demand + } + return fmt.Errorf("model file %s: %w", f, err) + } + return nil +} diff --git a/core/services/failover/prober_test.go b/core/services/failover/prober_test.go new file mode 100644 index 000000000..2b8b7f42d --- /dev/null +++ b/core/services/failover/prober_test.go @@ -0,0 +1,169 @@ +package failover + +import ( + "context" + "encoding/json" + "errors" + "io" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "sync" + + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/pkg/grpc" + pb "github.com/mudler/LocalAI/pkg/grpc/proto" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + ggrpc "google.golang.org/grpc" +) + +type fakeUpstream struct { + mu sync.Mutex + srv *httptest.Server + models []string + status int + paths []string + auth string +} + +func newFakeUpstream() *fakeUpstream { + u := &fakeUpstream{status: http.StatusOK} + u.srv = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + u.mu.Lock() + u.paths = append(u.paths, r.Method+" "+r.URL.Path) + u.auth = r.Header.Get("Authorization") + status, models := u.status, u.models + u.mu.Unlock() + _, _ = io.Copy(io.Discard, r.Body) + if status != http.StatusOK { + w.WriteHeader(status) + return + } + if r.URL.Path == "/v1/models" { + var data []map[string]string + for _, m := range models { + data = append(data, map[string]string{"id": m}) + } + _ = json.NewEncoder(w).Encode(map[string]any{"data": data}) + return + } + _, _ = w.Write([]byte(`{}`)) + })) + return u +} + +type fakeBackend struct { + grpc.Backend + healthy bool + predictErr error + predicted bool +} + +func (b *fakeBackend) HealthCheck(context.Context) (bool, error) { return b.healthy, nil } +func (b *fakeBackend) Predict(context.Context, *pb.PredictOptions, ...ggrpc.CallOption) (*pb.Reply, error) { + b.predicted = true + return &pb.Reply{}, b.predictErr +} + +var _ = Describe("DefaultProber", func() { + var ( + up *fakeUpstream + p *DefaultProber + ctx = context.Background() + ) + + BeforeEach(func() { + up = newFakeUpstream() + DeferCleanup(up.srv.Close) + p = NewProber(nil, "") + }) + + proxied := func(name, upstreamModel string, usecases ...string) config.ModelConfig { + c := config.ModelConfig{Name: name, Backend: "cloud-proxy", KnownUsecaseStrings: usecases} + c.KnownUsecases = config.GetUsecasesFromYAML(usecases) + c.Proxy.UpstreamURL = up.srv.URL + "/v1/chat/completions" + c.Proxy.UpstreamModel = upstreamModel + return c + } + + DescribeTable("UpstreamBase", + func(in, want string) { + got, err := UpstreamBase(in) + Expect(err).ToNot(HaveOccurred()) + Expect(got).To(Equal(want)) + }, + Entry("full endpoint", "https://h:8080/v1/chat/completions", "https://h:8080"), + Entry("path prefix", "https://h/api/v1/chat/completions", "https://h/api"), + Entry("bare host", "https://h", "https://h"), + Entry("bare host slash", "https://h/", "https://h"), + ) + + It("passes liveness when the upstream lists the model", func() { + up.models = []string{"big-llm"} + Expect(p.Liveness(ctx, proxied("argus-llm", "big-llm"), KindRemote, false)).To(Succeed()) + Expect(up.paths).To(ContainElement("GET /v1/models")) + }) + + It("uses the target name when upstream_model is empty", func() { + up.models = []string{"argus-llm"} + Expect(p.Liveness(ctx, proxied("argus-llm", ""), KindRemote, false)).To(Succeed()) + }) + + It("fails liveness when the model is not listed or the upstream errors", func() { + up.models = []string{"other"} + Expect(p.Liveness(ctx, proxied("argus-llm", ""), KindRemote, false)).To(MatchError(ContainSubstring("does not list"))) + up.status = http.StatusServiceUnavailable + Expect(p.Liveness(ctx, proxied("argus-llm", ""), KindRemote, false)).To(MatchError(ContainSubstring("503"))) + }) + + It("sends the API key as a bearer token", func() { + GinkgoT().Setenv("FAILOVER_PROBE_KEY", "sekret") + up.models = []string{"argus-llm"} + c := proxied("argus-llm", "") + c.Proxy.APIKeyEnv = "FAILOVER_PROBE_KEY" + Expect(p.Liveness(ctx, c, KindRemote, false)).To(Succeed()) + Expect(up.auth).To(Equal("Bearer sekret")) + }) + + DescribeTable("remote inference hits the usecase endpoint", + func(usecase, path string) { + Expect(p.Inference(ctx, proxied("m", "", usecase), KindRemote, false)).To(Succeed()) + Expect(up.paths).To(ContainElement("POST " + path)) + }, + Entry("chat", "chat", "/v1/chat/completions"), + Entry("embeddings", "embeddings", "/v1/embeddings"), + Entry("transcription", "transcript", "/v1/audio/transcriptions"), + Entry("tts", "tts", "/v1/audio/speech"), + ) + + It("uses HealthCheck for warm local liveness and Predict for local chat inference", func() { + b := &fakeBackend{healthy: true} + p = NewProber(func(context.Context, config.ModelConfig) (grpc.Backend, error) { return b, nil }, "") + c := config.ModelConfig{Name: "gemma", Backend: "llama-cpp", KnownUsecaseStrings: []string{"chat"}} + c.KnownUsecases = config.GetUsecasesFromYAML(c.KnownUsecaseStrings) + Expect(p.Liveness(ctx, c, KindLocal, true)).To(Succeed()) + b.healthy = false + Expect(p.Liveness(ctx, c, KindLocal, true)).To(HaveOccurred()) + Expect(p.Inference(ctx, c, KindLocal, true)).To(Succeed()) + Expect(b.predicted).To(BeTrue()) + b.predictErr = errors.New("boom") + Expect(p.Inference(ctx, c, KindLocal, true)).To(HaveOccurred()) + }) + + It("checks the model file for cold local liveness without loading", func() { + dir := GinkgoT().TempDir() + p = NewProber(func(context.Context, config.ModelConfig) (grpc.Backend, error) { + Fail("cold liveness must not load the model") + return nil, nil + }, dir) + c := config.ModelConfig{Name: "cold", Backend: "llama-cpp"} + c.Model = "weights.gguf" + Expect(p.Liveness(ctx, c, KindLocal, false)).To(HaveOccurred()) + Expect(os.WriteFile(filepath.Join(dir, "weights.gguf"), []byte("x"), 0o600)).To(Succeed()) + Expect(p.Liveness(ctx, c, KindLocal, false)).To(Succeed()) + c.Model = "org/some-hf-repo" // no extension: downloaded on demand + Expect(p.Liveness(ctx, c, KindLocal, false)).To(Succeed()) + }) +}) diff --git a/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md b/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md index c332082bf..11ef80e4f 100644 --- a/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md +++ b/docs/superpowers/specs/2026-09-26-model-failover-chains-design.md @@ -152,7 +152,7 @@ Chain states: |---|---|---| | remote | `GET /v1/models` returns 2xx and lists the upstream model. `` is the scheme and host of `proxy.upstream_url` plus any path prefix before `/v1`. The upstream model is `proxy.upstream_model`, or the target name when it is empty. `/v1/models` works on any OpenAI-compatible upstream, and `/readyz` exists only on LocalAI. | one minimal real request, chosen by usecase | | local, `warm: true` | gRPC `HealthCheck` on the loaded backend. If the backend is not loaded (it crashed), a reload is the recovery attempt. | chat and completion: `Predict` with 1 token; embeddings: `Embedding` of `"ping"`; other usecases: `HealthCheck`. A local backend process that answers `HealthCheck` rarely fails only for TTS or transcription. | -| local, cold | the config and model files exist and the backend is installed. The model is never loaded only to probe it. | none. After a trip, the target returns to `healthy` when `min_dwell` has passed. The next real request is the test. | +| local, cold | the model file exists (skipped for URLs and repository ids). The model is never loaded only to probe it. | none. After a trip, the target returns to `healthy` when `min_dwell` has passed. The next real request is the test. | Minimal requests by usecase: