diff --git a/core/application/startup.go b/core/application/startup.go index 4c1ddd4e9..db9e1caaf 100644 --- a/core/application/startup.go +++ b/core/application/startup.go @@ -255,7 +255,7 @@ func New(opts ...config.AppOption) (*Application, error) { // chain. WithOnWarmChanged pins and preloads warm local targets so a // switch to them does not wait for a cold load. application.failoverManager = failover.New(application.ModelConfigLoader(), - failover.WithProber(failover.NewProber(failoverLoadedBackend(application.ModelLoader()))), + failover.WithProber(failover.NewProber(failoverLoadedBackend(application.ModelLoader()), options.ProxyAPIKeyEnvLookup)), failover.WithOnWarmChanged(application.applyFailoverWarmTargets), ) // The assistant client was built in start() (above), before this diff --git a/core/cli/run.go b/core/cli/run.go index 3dc165ecc..de6621ba9 100644 --- a/core/cli/run.go +++ b/core/cli/run.go @@ -291,6 +291,7 @@ func (r *RunCMD) Run(ctx *cliContext.Context) error { } opts := []config.AppOption{ + config.WithProxyAPIKeyEnvLookup(os.Getenv), config.WithContext(context.Background()), config.WithArtifactDownloadConcurrency(r.ArtifactDownloadConcurrency), config.WithModelArtifactMaterializer(modelartifacts.NewDefaultManager( diff --git a/core/config/application_config.go b/core/config/application_config.go index b83e8fe58..24e1867c4 100644 --- a/core/config/application_config.go +++ b/core/config/application_config.go @@ -14,6 +14,9 @@ import ( ) type ApplicationConfig struct { + // ProxyAPIKeyEnvLookup resolves upstream credentials at the CLI boundary. + ProxyAPIKeyEnvLookup func(string) string `json:"-" yaml:"-"` + Context context.Context ConfigFile string SystemState *system.SystemState @@ -271,6 +274,10 @@ type AgentPoolConfig struct { type AppOption func(*ApplicationConfig) +func WithProxyAPIKeyEnvLookup(lookup func(string) string) AppOption { + return func(o *ApplicationConfig) { o.ProxyAPIKeyEnvLookup = lookup } +} + func NewApplicationConfig(o ...AppOption) *ApplicationConfig { opt := &ApplicationConfig{ Context: context.Background(), diff --git a/core/config/model_config.go b/core/config/model_config.go index 9c38f178d..7b67acb2a 100644 --- a/core/config/model_config.go +++ b/core/config/model_config.go @@ -306,10 +306,13 @@ const ( // "" when neither is set. Mirrored (not imported, to keep backends independent // of core's package layout) by resolveAPIKey in backend/go/cloud-proxy/proxy.go // — keep the two in sync, empty-value handling included. -func (p ProxyConfig) ResolveAPIKey() (string, error) { +func (p ProxyConfig) ResolveAPIKey(envLookup func(string) string) (string, error) { switch { case p.APIKeyEnv != "": - v := os.Getenv(p.APIKeyEnv) + var v string + if envLookup != nil { + v = envLookup(p.APIKeyEnv) + } if v == "" { return "", fmt.Errorf("proxy api_key_env %q is unset", p.APIKeyEnv) } diff --git a/core/config/model_config_failover_test.go b/core/config/model_config_failover_test.go index 9e82868e9..6a276f056 100644 --- a/core/config/model_config_failover_test.go +++ b/core/config/model_config_failover_test.go @@ -70,25 +70,32 @@ failover: }) 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("reads the key through the application lookup", func() { + app := NewApplicationConfig(WithProxyAPIKeyEnvLookup(func(name string) string { + Expect(name).To(Equal("FAILOVER_TEST_KEY")) + return "k1" + })) + Expect(ProxyConfig{APIKeyEnv: "FAILOVER_TEST_KEY"}.ResolveAPIKey(app.ProxyAPIKeyEnvLookup)).To(Equal("k1")) + }) + It("fails without an environment lookup", func() { + _, err := ProxyConfig{APIKeyEnv: "FAILOVER_TEST_KEY"}.ResolveAPIKey(nil) + Expect(err).To(HaveOccurred()) }) It("fails on an unset env var", func() { - _, err := ProxyConfig{APIKeyEnv: "FAILOVER_TEST_UNSET_KEY"}.ResolveAPIKey() + _, err := ProxyConfig{APIKeyEnv: "FAILOVER_TEST_UNSET_KEY"}.ResolveAPIKey(os.Getenv) Expect(err).To(HaveOccurred()) }) It("fails on a set-but-empty env var", func() { GinkgoT().Setenv("FAILOVER_TEST_EMPTY_KEY", "") - _, err := ProxyConfig{APIKeyEnv: "FAILOVER_TEST_EMPTY_KEY"}.ResolveAPIKey() + _, err := ProxyConfig{APIKeyEnv: "FAILOVER_TEST_EMPTY_KEY"}.ResolveAPIKey(os.Getenv) 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")) + Expect(ProxyConfig{APIKeyFile: f}.ResolveAPIKey(os.Getenv)).To(Equal("k2")) }) It("returns empty when nothing is set", func() { - Expect(ProxyConfig{}.ResolveAPIKey()).To(Equal("")) + Expect(ProxyConfig{}.ResolveAPIKey(os.Getenv)).To(Equal("")) }) }) diff --git a/core/http/endpoints/localai/failover_test.go b/core/http/endpoints/localai/failover_test.go index fd3b6a733..77fb48678 100644 --- a/core/http/endpoints/localai/failover_test.go +++ b/core/http/endpoints/localai/failover_test.go @@ -91,7 +91,7 @@ var _ = Describe("failover endpoints", func() { req, _ := http.NewRequestWithContext(ctx, http.MethodGet, srv.URL+"/api/failover/events", nil) resp, err := http.DefaultClient.Do(req) Expect(err).ToNot(HaveOccurred()) - defer resp.Body.Close() + defer func() { _ = resp.Body.Close() }() Expect(resp.Header.Get("Content-Type")).To(HavePrefix("text/event-stream")) r := bufio.NewReader(resp.Body) next := func() string { diff --git a/core/services/failover/prober.go b/core/services/failover/prober.go index 9a407e8a0..6b27e4f8a 100644 --- a/core/services/failover/prober.go +++ b/core/services/failover/prober.go @@ -32,17 +32,18 @@ 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 + HTTP *http.Client + Loaded LoadedFunc + EnvLookup func(string) string } -func NewProber(loaded LoadedFunc) *DefaultProber { +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} + }, Loaded: loaded, EnvLookup: envLookup} } func (p *DefaultProber) Liveness(ctx context.Context, cfg config.ModelConfig, kind Kind, warm bool) error { @@ -102,7 +103,7 @@ func PrepareTarget(cfg *config.ModelConfig) { } func (p *DefaultProber) authorize(req *http.Request, cfg config.ModelConfig) error { - key, err := cfg.Proxy.ResolveAPIKey() + key, err := cfg.Proxy.ResolveAPIKey(p.EnvLookup) if err != nil || key == "" { return err } @@ -124,7 +125,7 @@ func (p *DefaultProber) do(req *http.Request, cfg config.ModelConfig) (*http.Res return nil, err } if resp.StatusCode/100 != 2 { - resp.Body.Close() + _ = resp.Body.Close() return nil, fmt.Errorf("upstream %s: HTTP %d", req.URL.Path, resp.StatusCode) } return resp, nil @@ -143,7 +144,7 @@ func (p *DefaultProber) remoteLiveness(ctx context.Context, cfg config.ModelConf if err != nil { return err } - defer resp.Body.Close() + defer func() { _ = resp.Body.Close() }() var list struct { Data []struct { ID string `json:"id"` diff --git a/core/services/failover/prober_test.go b/core/services/failover/prober_test.go index 842934876..fd22bd1e3 100644 --- a/core/services/failover/prober_test.go +++ b/core/services/failover/prober_test.go @@ -7,6 +7,7 @@ import ( "io" "net/http" "net/http/httptest" + "os" "sync" "github.com/mudler/LocalAI/core/config" @@ -75,7 +76,7 @@ var _ = Describe("DefaultProber", func() { BeforeEach(func() { up = newFakeUpstream() DeferCleanup(up.srv.Close) - p = NewProber(nil) + p = NewProber(nil, os.Getenv) }) proxied := func(name, upstreamModel string, usecases ...string) config.ModelConfig { @@ -155,7 +156,7 @@ var _ = Describe("DefaultProber", func() { It("uses HealthCheck for warm local liveness and Predict for local chat inference", func() { b := &fakeBackend{healthy: true} - p = NewProber(func(config.ModelConfig) grpc.Backend { return b }) + p = NewProber(func(config.ModelConfig) grpc.Backend { 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()) @@ -171,7 +172,7 @@ var _ = Describe("DefaultProber", func() { // The warm preload loads the model; a probe that loaded it too would // block until the load finished and then judge it on an expired ctx. asked := 0 - p = NewProber(func(config.ModelConfig) grpc.Backend { asked++; return nil }) + p = NewProber(func(config.ModelConfig) grpc.Backend { asked++; return nil }, nil) for _, uc := range []string{"chat", "tts"} { c := config.ModelConfig{Name: "gemma", Backend: "llama-cpp", KnownUsecaseStrings: []string{uc}} c.KnownUsecases = config.GetUsecasesFromYAML(c.KnownUsecaseStrings) @@ -185,7 +186,7 @@ var _ = Describe("DefaultProber", func() { p = NewProber(func(config.ModelConfig) grpc.Backend { Fail("cold liveness must not look up the backend") return nil - }) + }, nil) // None of these files exist: a missing file says nothing about whether // the target can serve (download on first use, dotted names, backends // that need no file). Only a real request may trip a cold target. diff --git a/tests/e2e/realtime_ws_test.go b/tests/e2e/realtime_ws_test.go index 1aa2552a1..68181ef35 100644 --- a/tests/e2e/realtime_ws_test.go +++ b/tests/e2e/realtime_ws_test.go @@ -237,7 +237,7 @@ var _ = Describe("Realtime WebSocket API", Label("Realtime"), func() { It("switches the LLM mid-session and keeps the conversation", func() { conn := connectWS("rt-failover") - defer conn.Close() + defer func() { _ = conn.Close() }() Expect(readServerEvent(conn, 30*time.Second)["type"]).To(Equal("session.created")) initial := drainUntil(conn, "localai.model.failover", 10*time.Second) @@ -275,7 +275,7 @@ var _ = Describe("Realtime WebSocket API", Label("Realtime"), func() { It("starts the session on the next target when the active one fails to warm up", func() { conn := connectWS("rt-failover-warm") - defer conn.Close() + defer func() { _ = conn.Close() }() Expect(readServerEvent(conn, 30*time.Second)["type"]).To(Equal("session.created")) initial := drainUntil(conn, "localai.model.failover", 10*time.Second)