diff --git a/core/config/model_config_loader.go b/core/config/model_config_loader.go index 7733a7392..363fcc652 100644 --- a/core/config/model_config_loader.go +++ b/core/config/model_config_loader.go @@ -501,6 +501,80 @@ func (bcl *ModelConfigLoader) ValidateAliasTarget(cfg *ModelConfig) error { return nil } +// failoverUsecases are the single usecases a chain can share. Checking one +// flag at a time avoids treating "chat+tts" and "tts" as unrelated. +var failoverUsecases = []ModelConfigUsecase{ + FLAG_CHAT, FLAG_COMPLETION, FLAG_EMBEDDINGS, FLAG_RERANK, FLAG_IMAGE, + FLAG_TRANSCRIPT, FLAG_TTS, FLAG_SOUND_GENERATION, FLAG_VAD, FLAG_VIDEO, + FLAG_SOUND_CLASSIFICATION, +} + +// ValidateFailoverTargets checks that every target of a chain exists and is +// not itself a chain. Alias targets are allowed and resolve one hop. +func (bcl *ModelConfigLoader) ValidateFailoverTargets(cfg *ModelConfig) error { + return validateFailoverTargets(cfg, bcl.GetModelConfig) +} + +// FailoverTargetsShareUsecase reports whether all targets of a chain have at +// least one usecase in common. A false result is only a warning: usecases are +// often inferred. +func (bcl *ModelConfigLoader) FailoverTargetsShareUsecase(cfg *ModelConfig) bool { + return failoverTargetsShareUsecase(cfg, bcl.GetModelConfig) +} + +func validateFailoverTargets(cfg *ModelConfig, lookup func(string) (ModelConfig, bool)) error { + if cfg == nil || !cfg.IsFailover() { + return nil + } + for _, t := range cfg.Failover.Targets { + target, ok := lookup(t.Model) + if !ok { + return fmt.Errorf("failover chain %q: target %q does not exist", cfg.Name, t.Model) + } + if target.IsAlias() { + if resolved, ok := lookup(target.Alias); ok { + target = resolved + } + } + if target.IsFailover() { + return fmt.Errorf("failover chain %q: target %q is a chain (chains do not nest)", cfg.Name, t.Model) + } + } + return nil +} + +func failoverTargetsShareUsecase(cfg *ModelConfig, lookup func(string) (ModelConfig, bool)) bool { + if cfg == nil || !cfg.IsFailover() { + return true + } + var targets []ModelConfig + for _, t := range cfg.Failover.Targets { + target, ok := lookup(t.Model) + if !ok { + return true // missing targets are reported by validateFailoverTargets + } + if target.IsAlias() { + if resolved, ok := lookup(target.Alias); ok { + target = resolved + } + } + targets = append(targets, target) + } + for _, u := range failoverUsecases { + all := true + for i := range targets { + if !targets[i].HasUsecases(u) { + all = false + break + } + } + if all { + return true + } + } + return false +} + type preloadWork struct { key string config ModelConfig @@ -837,6 +911,28 @@ func (bcl *ModelConfigLoader) loadModelConfigsFromPath(path string, strict bool, } } + // Reject failover chains whose targets are missing or are themselves + // chains. bcl.Lock() is held here, so look up configs directly rather + // than through GetModelConfig, which would deadlock on the same mutex. + lookup := func(n string) (ModelConfig, bool) { c, ok := bcl.configs[n]; return c, ok } + for name, cfg := range bcl.configs { + if !cfg.IsFailover() { + continue + } + c := cfg + if err := validateFailoverTargets(&c, lookup); err != nil { + if strict { + return fmt.Errorf("invalid model config %q: %w", name, err) + } + xlog.Error("skipping invalid failover chain", "model", name, "error", err) + delete(bcl.configs, name) + continue + } + if !failoverTargetsShareUsecase(&c, lookup) { + xlog.Warn("failover chain targets share no known usecase", "model", name) + } + } + return nil } diff --git a/core/config/model_config_loader_test.go b/core/config/model_config_loader_test.go index 1a3e9b03a..9912a365b 100644 --- a/core/config/model_config_loader_test.go +++ b/core/config/model_config_loader_test.go @@ -402,3 +402,43 @@ var _ = Describe("ModelConfigLoader ResolveAliasName", func() { Expect(target).To(BeEmpty()) }) }) + +var _ = Describe("ModelConfigLoader failover validation", func() { + var loader *ModelConfigLoader + chain := func(targets ...string) *ModelConfig { + c := &ModelConfig{Name: "chain", Failover: &FailoverConfig{}} + for _, t := range targets { + c.Failover.Targets = append(c.Failover.Targets, FailoverTarget{Model: t}) + } + return c + } + + BeforeEach(func() { + loader = NewModelConfigLoader("") + loader.configs["a"] = ModelConfig{Name: "a", Backend: "llama-cpp", KnownUsecaseStrings: []string{"chat"}} + loader.configs["b"] = ModelConfig{Name: "b", Backend: "llama-cpp", KnownUsecaseStrings: []string{"chat"}} + loader.configs["tts"] = ModelConfig{Name: "tts", Backend: "piper", KnownUsecaseStrings: []string{"tts"}} + loader.configs["alias-b"] = ModelConfig{Name: "alias-b", Alias: "b"} + loader.configs["other-chain"] = *chain("a", "b") + loader.configs["alias-chain"] = ModelConfig{Name: "alias-chain", Alias: "other-chain"} + for k, c := range loader.configs { + c.KnownUsecases = GetUsecasesFromYAML(c.KnownUsecaseStrings) + loader.configs[k] = c + } + }) + + It("accepts existing targets and alias targets", func() { + Expect(loader.ValidateFailoverTargets(chain("a", "alias-b"))).To(Succeed()) + }) + It("rejects a missing target", func() { + Expect(loader.ValidateFailoverTargets(chain("a", "nope"))).To(MatchError(ContainSubstring("does not exist"))) + }) + It("rejects a nested chain, directly or through an alias", func() { + Expect(loader.ValidateFailoverTargets(chain("a", "other-chain"))).To(MatchError(ContainSubstring("chains do not nest"))) + Expect(loader.ValidateFailoverTargets(chain("a", "alias-chain"))).To(MatchError(ContainSubstring("chains do not nest"))) + }) + It("reports whether targets share a usecase", func() { + Expect(loader.FailoverTargetsShareUsecase(chain("a", "b"))).To(BeTrue()) + Expect(loader.FailoverTargetsShareUsecase(chain("a", "tts"))).To(BeFalse()) + }) +}) diff --git a/core/http/endpoints/localai/import_model.go b/core/http/endpoints/localai/import_model.go index b3cf491eb..73448430a 100644 --- a/core/http/endpoints/localai/import_model.go +++ b/core/http/endpoints/localai/import_model.go @@ -187,6 +187,12 @@ func ImportModelEndpoint(cl *config.ModelConfigLoader, gs *galleryop.GalleryServ return c.JSON(http.StatusBadRequest, ModelResponse{Success: false, Error: err.Error()}) } + // Reject failover chains whose targets are missing or are themselves + // chains, for the same reason. + if err := cl.ValidateFailoverTargets(&modelConfig); err != nil { + return c.JSON(http.StatusBadRequest, ModelResponse{Success: false, Error: err.Error()}) + } + // Create the configuration file configPath := filepath.Join(appConfig.SystemState.Model.ModelsPath, modelConfig.Name+".yaml") if err := utils.VerifyPath(modelConfig.Name+".yaml", appConfig.SystemState.Model.ModelsPath); err != nil { diff --git a/core/services/modeladmin/config.go b/core/services/modeladmin/config.go index 2cadfcaed..463ffe3f3 100644 --- a/core/services/modeladmin/config.go +++ b/core/services/modeladmin/config.go @@ -169,6 +169,9 @@ func (s *ConfigService) patchConfig(ctx context.Context, name string, patch map[ if err := s.Loader.ValidateAliasTarget(&updated); err != nil { return nil, fmt.Errorf("%w: %v", ErrInvalidConfig, err) } + if err := s.Loader.ValidateFailoverTargets(&updated); err != nil { + return nil, fmt.Errorf("%w: %v", ErrInvalidConfig, err) + } var result *PatchResult err = s.withMutationRollback([]string{configPath}, func() error { if err := writeFileAtomic(configPath, yamlData, 0644); err != nil { @@ -286,6 +289,9 @@ func (s *ConfigService) editYAML(ctx context.Context, name string, body []byte) if err := s.Loader.ValidateAliasTarget(&req); err != nil { return nil, fmt.Errorf("%w: %v", ErrInvalidConfig, err) } + if err := s.Loader.ValidateFailoverTargets(&req); err != nil { + return nil, fmt.Errorf("%w: %v", ErrInvalidConfig, err) + } configPath := existing.GetModelConfigFile() modelsPath := s.modelsPath()