diff --git a/core/http/endpoints/openai/realtime.go b/core/http/endpoints/openai/realtime.go index fe2c15068..8c8988371 100644 --- a/core/http/endpoints/openai/realtime.go +++ b/core/http/endpoints/openai/realtime.go @@ -32,6 +32,7 @@ import ( "github.com/mudler/LocalAI/core/http/endpoints/openai/turncoord" "github.com/mudler/LocalAI/core/http/endpoints/openai/types" "github.com/mudler/LocalAI/core/schema" + "github.com/mudler/LocalAI/core/services/failover" "github.com/mudler/LocalAI/core/services/routing/router" "github.com/mudler/LocalAI/core/services/voiceprofile" "github.com/mudler/LocalAI/core/templates" @@ -638,6 +639,7 @@ func runRealtimeSession(application *application.Application, t Transport, model application.ModelConfigLoader(), application.ModelLoader(), application.ApplicationConfig(), + application.FailoverManager(), ) } else { m, err = newModel( @@ -756,11 +758,10 @@ func runRealtimeSession(application *application.Application, t Transport, model // Sent after session.created, which clients expect as the first event. // This function runs until the connection closes, so the defer stops the - // events at session end. - if wrapped, ok := m.(*wrappedModel); ok && wrapped.failover != nil && len(wrapped.stageChains) > 0 { - stopFailoverEvents := startFailoverEvents(t, wrapped.failover, wrapped.stageChains) - defer stopFailoverEvents() - } + // events at session end. A transcription session.update that swaps the + // model restarts them for the new model's chains. + stopFailoverEvents := startModelFailoverEvents(t, m) + defer func() { stopFailoverEvents() }() var ( msg []byte @@ -832,12 +833,14 @@ func runRealtimeSession(application *application.Application, t Transport, model // Handle transcription session update if e.Session.Transcription != nil { + prevModel := session.ModelInterface if err := updateTransSession( session, &e.Session, application.ModelConfigLoader(), application.ModelLoader(), application.ApplicationConfig(), + application.FailoverManager(), ); err != nil { xlog.Error("failed to update session", "error", err) // The cause is validation feedback on the client's own @@ -854,6 +857,10 @@ func runRealtimeSession(application *application.Application, t Transport, model }, Session: session.ToServer(), }) + if session.ModelInterface != prevModel { + stopFailoverEvents() + stopFailoverEvents = startModelFailoverEvents(t, session.ModelInterface) + } } // Handle realtime session update @@ -1151,7 +1158,7 @@ func sendTestTone(t Transport) { } } -func updateTransSession(session *Session, update *types.SessionUnion, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig) error { +func updateTransSession(session *Session, update *types.SessionUnion, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig, fm *failover.Manager) error { sessionLock.Lock() defer sessionLock.Unlock() @@ -1174,7 +1181,7 @@ func updateTransSession(session *Session, update *types.SessionUnion, cl *config return fmt.Errorf("model is not a valid pipeline model: %s", trUpd.Model) } - m, cfg, err := newTranscriptionOnlyModel(&cfg.Pipeline, cl, ml, appConfig) + m, cfg, err := newTranscriptionOnlyModel(&cfg.Pipeline, cl, ml, appConfig, fm) if err != nil { return err } diff --git a/core/http/endpoints/openai/realtime_failover.go b/core/http/endpoints/openai/realtime_failover.go index d58d52f33..c47c705ce 100644 --- a/core/http/endpoints/openai/realtime_failover.go +++ b/core/http/endpoints/openai/realtime_failover.go @@ -2,41 +2,151 @@ package openai import ( "context" + "errors" + "fmt" "sort" + "sync" + "github.com/mudler/LocalAI/core/backend" "github.com/mudler/LocalAI/core/config" "github.com/mudler/LocalAI/core/http/endpoints/openai/types" "github.com/mudler/LocalAI/core/services/failover" + "github.com/mudler/LocalAI/pkg/model" ) -// isChainStage reports whether stage names a failover chain. -func (m *wrappedModel) isChainStage(stage string) bool { - _, ok := m.stageChains[stage] - return ok && m.failover != nil +// stageRouter routes realtime pipeline stages that name a failover chain. +// Every realtime model kind (full pipeline, transcription-only, sound-only) +// embeds one, so a chain resolves the same way whatever the session does. +type stageRouter struct { + // failover and stageChains route pipeline stages that name a failover + // chain; stageChains maps a stage ("llm", "tts", ...) to its chain. + // The model's *Config fields then hold the target that was active at + // session start, for the checks that run once (voice, templates). + failover *failover.Manager + stageChains map[string]string + stageTargetConfig func(name string) (*config.ModelConfig, error) + appTracing bool } +func newStageRouter(fm *failover.Manager, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig) stageRouter { + return stageRouter{ + failover: fm, + stageChains: map[string]string{}, + stageTargetConfig: func(name string) (*config.ModelConfig, error) { + cfg, err := cl.LoadResolvedModelConfig(name, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) + if err != nil { + return nil, err + } + failover.PrepareTarget(cfg) + return cfg, nil + }, + appTracing: appConfig.EnableTracing, + } +} + +// resolveStage records a stage that names a chain and returns the chain's +// active target, so everything that inspects stage configs at session start +// sees a real model. Any other config is returned as is. A chain config +// reaching the model loader would have no backend and trigger backend +// auto-detection. +func (r *stageRouter) resolveStage(stage string, cfg *config.ModelConfig) (*config.ModelConfig, error) { + if cfg == nil || !cfg.IsFailover() { + return cfg, nil + } + if r.failover == nil { + return nil, fmt.Errorf("pipeline %s stage %q is a failover chain, but failover is not running", stage, cfg.Name) + } + st, ok := r.failover.ChainStatus(cfg.Name) + if !ok { + return nil, fmt.Errorf("failover chain %q not found", cfg.Name) + } + r.stageChains[stage] = cfg.Name + return r.stageTargetConfig(st.Active) +} + +// isChainStage reports whether stage names a failover chain. +func (r *stageRouter) isChainStage(stage string) bool { + _, ok := r.stageChains[stage] + return ok && r.failover != nil +} + +// hasChainStages reports whether any stage names a chain. +func (r *stageRouter) hasChainStages() bool { + return r.failover != nil && len(r.stageChains) > 0 +} + +func (r *stageRouter) router() *stageRouter { return r } + +// stageRouted is implemented by every realtime model that embeds a +// stageRouter; the session uses it to send failover events. +type stageRouted interface{ router() *stageRouter } + // stageCall runs fn with the config that should serve stage now. A plain // stage uses base. A chain stage goes through the failover plan on every // call, so a switch takes effect on the next call without rebuilding the // session, and fn is retried on the next target until it calls commit. -func (m *wrappedModel) stageCall(ctx context.Context, stage string, base *config.ModelConfig, fn func(cfg *config.ModelConfig, commit func()) error) error { - if !m.isChainStage(stage) { +func (r *stageRouter) stageCall(ctx context.Context, stage string, base *config.ModelConfig, fn func(cfg *config.ModelConfig, commit func()) error) error { + if !r.isChainStage(stage) { return fn(base, func() {}) } - chain := m.stageChains[stage] - return m.failover.Do(ctx, chain, func(ctx context.Context, target string, commit func()) error { - cfg, err := m.stageTargetConfig(target) + chain := r.stageChains[stage] + return r.failover.Do(ctx, chain, func(ctx context.Context, target string, commit func()) error { + cfg, err := r.stageTargetConfig(target) if err != nil { return err } err = fn(cfg, commit) if err != nil { - failover.RecordAttemptTrace(m.appTracing, chain, target, err) + failover.RecordAttemptTrace(r.appTracing, chain, target, err) } return err }) } +// warmStages preloads the stages. A chain stage warms through its failover +// plan: a target that fails to load moves the stage to the next one instead +// of failing the session. +func (r *stageRouter) warmStages(ctx context.Context, ml *model.ModelLoader, appConfig *config.ApplicationConfig, stages []backend.PreloadStage) error { + var ( + plain []backend.PreloadStage + wg sync.WaitGroup + mu sync.Mutex + errs []error + ) + for _, s := range stages { + if !r.isChainStage(s.Role) { + plain = append(plain, s) + continue + } + wg.Go(func() { + err := r.stageCall(ctx, s.Role, s.Cfg, func(cfg *config.ModelConfig, _ func()) error { + _, err := backend.PreloadStages(ctx, ml, appConfig, []backend.PreloadStage{{Role: s.Role, Cfg: cfg}}) + return err + }) + mu.Lock() + errs = append(errs, err) + mu.Unlock() + }) + } + _, err := backend.PreloadStages(ctx, ml, appConfig, plain) + wg.Wait() + return errors.Join(append(errs, err)...) +} + +// startModelFailoverEvents starts failover events for m when it has chain +// stages. The returned func stops them and is never nil. +func startModelFailoverEvents(t Transport, m Model) func() { + sr, ok := m.(stageRouted) + if !ok { + return func() {} + } + r := sr.router() + if !r.hasChainStages() { + return func() {} + } + return startFailoverEvents(t, r.failover, r.stageChains) +} + // startFailoverEvents tells the client which target serves each chain stage // now, and again whenever a chain switches. The returned func stops it. func startFailoverEvents(t Transport, fm *failover.Manager, stageChains map[string]string) func() { diff --git a/core/http/endpoints/openai/realtime_failover_test.go b/core/http/endpoints/openai/realtime_failover_test.go index 48737630c..996fa1eb6 100644 --- a/core/http/endpoints/openai/realtime_failover_test.go +++ b/core/http/endpoints/openai/realtime_failover_test.go @@ -3,10 +3,15 @@ package openai import ( "context" "errors" + "os" + "path/filepath" + "time" "github.com/mudler/LocalAI/core/config" "github.com/mudler/LocalAI/core/http/endpoints/openai/types" "github.com/mudler/LocalAI/core/services/failover" + "github.com/mudler/LocalAI/pkg/model" + "github.com/mudler/LocalAI/pkg/system" . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" ) @@ -34,8 +39,8 @@ var _ = Describe("realtime failover", func() { }) chainModel := func() *wrappedModel { - return &wrappedModel{failover: fm, stageChains: map[string]string{"tts": "chain"}, - stageTargetConfig: func(name string) (*config.ModelConfig, error) { return &config.ModelConfig{Name: name}, nil }} + return &wrappedModel{stageRouter: stageRouter{failover: fm, stageChains: map[string]string{"tts": "chain"}, + stageTargetConfig: func(name string) (*config.ModelConfig, error) { return &config.ModelConfig{Name: name}, nil }}} } It("routes a chain stage through the plan and retries before commit", func() { @@ -97,3 +102,105 @@ var _ = Describe("realtime failover", func() { stop() }) }) + +var _ = Describe("realtime failover in transcription-only and sound-only sessions", func() { + var ( + cl *config.ModelConfigLoader + ml *model.ModelLoader + appConfig *config.ApplicationConfig + fm *failover.Manager + ) + + BeforeEach(func() { + dir := GinkgoT().TempDir() + write := func(name, body string) { + Expect(os.WriteFile(filepath.Join(dir, name+".yaml"), []byte(body), 0o600)).To(Succeed()) + } + write("vad", "name: vad\nbackend: silero-vad\n") + write("stt-a", "name: stt-a\nbackend: fake-stt-a\n") + write("stt-b", "name: stt-b\nbackend: fake-stt-b\n") + write("stt-chain", "name: stt-chain\nfailover:\n targets:\n - model: stt-a\n - model: stt-b\n") + write("sound-a", "name: sound-a\nbackend: fake-sound-a\n") + write("sound-b", "name: sound-b\nbackend: fake-sound-b\n") + write("sound-chain", "name: sound-chain\nfailover:\n targets:\n - model: sound-a\n - model: sound-b\n") + ss := &system.SystemState{Model: system.Model{ModelsPath: dir}} + appConfig = config.NewApplicationConfig() + appConfig.SystemState = ss + cl = config.NewModelConfigLoader(dir) + Expect(cl.LoadModelConfigsFromPath(dir)).To(Succeed()) + ml = model.NewModelLoader(ss) + fm = failover.New(cl) + }) + + // failoverEvents collects the failover events a transport received. + failoverEvents := func(t *fakeTransport) func() []types.ModelFailoverEvent { + return func() []types.ModelFailoverEvent { + var out []types.ModelFailoverEvent + for _, e := range t.events() { + if fe, ok := e.(types.ModelFailoverEvent); ok { + out = append(out, fe) + } + } + return out + } + } + + It("resolves a sound_detection chain and routes each call through it", func() { + m, err := newSoundDetectionOnlyModel(&config.Pipeline{SoundDetection: "sound-chain"}, cl, ml, appConfig, fm) + Expect(err).ToNot(HaveOccurred()) + tm := m.(*transcriptOnlyModel) + Expect(tm.SoundDetectionConfig.Name).To(Equal("sound-a")) + Expect(tm.stageChains).To(Equal(map[string]string{"sound_detection": "sound-chain"})) + + var tried []string + tm.stageTargetConfig = func(name string) (*config.ModelConfig, error) { + tried = append(tried, name) + return nil, errors.New("dial tcp: refused") + } + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + _, err = tm.SoundDetection(ctx, "a.wav", 3, 0.1) + Expect(err).To(HaveOccurred()) + Expect(tried).To(Equal([]string{"sound-a", "sound-b"})) + + t := &fakeTransport{} + stop := startModelFailoverEvents(t, m) + defer stop() + Eventually(failoverEvents(t)).Should(ContainElement(And( + HaveField("Stage", "sound_detection"), HaveField("Chain", "sound-chain"), HaveField("Reason", "initial")))) + }) + + It("resolves a transcription chain in a transcription-only session", func() { + m, cfg, err := newTranscriptionOnlyModel(&config.Pipeline{VAD: "vad", Transcription: "stt-chain"}, cl, ml, appConfig, fm) + Expect(err).ToNot(HaveOccurred()) + Expect(cfg.Name).To(Equal("stt-a")) + tm := m.(*transcriptOnlyModel) + Expect(tm.VADConfig.Name).To(Equal("vad")) + Expect(tm.stageChains).To(Equal(map[string]string{"transcription": "stt-chain"})) + + var tried []string + tm.stageTargetConfig = func(name string) (*config.ModelConfig, error) { + tried = append(tried, name) + return nil, errors.New("dial tcp: refused") + } + _, err = tm.Transcribe(context.Background(), "a.wav", "", false, false, "") + Expect(err).To(HaveOccurred()) + Expect(tried).To(Equal([]string{"stt-a", "stt-b"})) + }) + + It("fails before touching a backend when failover is not running", func() { + _, err := newSoundDetectionOnlyModel(&config.Pipeline{SoundDetection: "sound-chain"}, cl, ml, appConfig, nil) + Expect(err).To(MatchError(ContainSubstring("failover is not running"))) + _, _, err = newTranscriptionOnlyModel(&config.Pipeline{VAD: "vad", Transcription: "stt-chain"}, cl, ml, appConfig, nil) + Expect(err).To(MatchError(ContainSubstring("failover is not running"))) + Expect(ml.ListLoadedModels()).To(BeEmpty()) + }) + + It("sends no failover events for a session without chains", func() { + m, err := newSoundDetectionOnlyModel(&config.Pipeline{SoundDetection: "sound-a"}, cl, ml, appConfig, fm) + Expect(err).ToNot(HaveOccurred()) + t := &fakeTransport{} + startModelFailoverEvents(t, m)() + Consistently(failoverEvents(t), 200*time.Millisecond).Should(BeEmpty()) + }) +}) diff --git a/core/http/endpoints/openai/realtime_model.go b/core/http/endpoints/openai/realtime_model.go index ee42bde2d..f94bb510b 100644 --- a/core/http/endpoints/openai/realtime_model.go +++ b/core/http/endpoints/openai/realtime_model.go @@ -86,17 +86,10 @@ type wrappedModel struct { routerSessionID string routerUserID string - // failover and stageChains route pipeline stages that name a failover - // chain; stageChains maps a stage ("llm", "tts", ...) to its chain. - // The *Config fields above then hold the target that was active at - // session start, for the checks that run once (voice, templates). - failover *failover.Manager - stageChains map[string]string - stageTargetConfig func(name string) (*config.ModelConfig, error) + stageRouter // tuneLLM applies the pipeline's LLM overrides (reasoning effort, // disable_thinking) to a chain target loaded per call. - tuneLLM func(cfg *config.ModelConfig) - appTracing bool + tuneLLM func(cfg *config.ModelConfig) } // anyToAnyModel represent a model which supports Any-to-Any operations @@ -119,18 +112,38 @@ type transcriptOnlyModel struct { appConfig *config.ApplicationConfig modelLoader *model.ModelLoader confLoader *config.ModelConfigLoader + + stageRouter } func (m *transcriptOnlyModel) VAD(ctx context.Context, request *schema.VADRequest) (*schema.VADResponse, error) { - return backend.VAD(request, ctx, m.modelLoader, m.appConfig, *m.VADConfig) + var res *schema.VADResponse + err := m.stageCall(ctx, "vad", m.VADConfig, func(cfg *config.ModelConfig, _ func()) error { + var err error + res, err = backend.VAD(request, ctx, m.modelLoader, m.appConfig, *cfg) + return err + }) + return res, err } func (m *transcriptOnlyModel) Transcribe(ctx context.Context, audio, language string, translate bool, diarize bool, prompt string) (*schema.TranscriptionResult, error) { - return backend.ModelTranscription(ctx, audio, language, translate, diarize, prompt, m.modelLoader, *m.TranscriptionConfig, m.appConfig) + var res *schema.TranscriptionResult + err := m.stageCall(ctx, "transcription", m.TranscriptionConfig, func(cfg *config.ModelConfig, _ func()) error { + var err error + res, err = backend.ModelTranscription(ctx, audio, language, translate, diarize, prompt, m.modelLoader, *cfg, m.appConfig) + return err + }) + return res, err } func (m *transcriptOnlyModel) SoundDetection(ctx context.Context, audio string, topK int, threshold float32) (*schema.SoundClassificationResult, error) { - return modelSoundDetection(ctx, m.modelLoader, m.appConfig, m.SoundDetectionConfig, audio, topK, threshold) + var res *schema.SoundClassificationResult + err := m.stageCall(ctx, "sound_detection", m.SoundDetectionConfig, func(cfg *config.ModelConfig, _ func()) error { + var err error + res, err = modelSoundDetection(ctx, m.modelLoader, m.appConfig, cfg, audio, topK, threshold) + return err + }) + return res, err } func (m *transcriptOnlyModel) Predict(ctx context.Context, messages schema.Messages, images, videos, audios []string, tokenCallback func(string, backend.TokenUsage) bool, tools []types.ToolUnion, toolChoice *types.ToolChoiceUnion, logprobs *int, topLogprobs *int, logitBias map[string]float64) (func() (backend.LLMResponse, error), error) { @@ -157,11 +170,27 @@ func (m *transcriptOnlyModel) TTSStream(ctx context.Context, text, voice, langua } func (m *transcriptOnlyModel) TranscribeStream(ctx context.Context, audio, language string, translate, diarize bool, prompt string, onDelta func(text string)) (*schema.TranscriptionResult, error) { - return transcribeStream(ctx, m.modelLoader, *m.TranscriptionConfig, m.appConfig, audio, language, translate, diarize, prompt, onDelta) + var res *schema.TranscriptionResult + err := m.stageCall(ctx, "transcription", m.TranscriptionConfig, func(cfg *config.ModelConfig, commit func()) error { + var err error + res, err = transcribeStream(ctx, m.modelLoader, *cfg, m.appConfig, audio, language, translate, diarize, prompt, func(s string) { + commit() + onDelta(s) + }) + return err + }) + return res, err } func (m *transcriptOnlyModel) TranscribeLive(ctx context.Context, language string, onEvent func(backend.LiveTranscriptionEvent)) (backend.LiveTranscriptionSession, error) { - return backend.ModelTranscriptionLive(ctx, language, m.modelLoader, *m.TranscriptionConfig, m.appConfig, onEvent) + var live backend.LiveTranscriptionSession + // Only opening the live session can move to the next target. + err := m.stageCall(ctx, "transcription", m.TranscriptionConfig, func(cfg *config.ModelConfig, _ func()) error { + var err error + live, err = backend.ModelTranscriptionLive(ctx, language, m.modelLoader, *cfg, m.appConfig, onEvent) + return err + }) + return live, err } func (m *transcriptOnlyModel) PredictConfig() *config.ModelConfig { @@ -169,12 +198,11 @@ func (m *transcriptOnlyModel) PredictConfig() *config.ModelConfig { } func (m *transcriptOnlyModel) Warmup(ctx context.Context) error { - _, err := backend.PreloadStages(ctx, m.modelLoader, m.appConfig, []backend.PreloadStage{ + return m.warmStages(ctx, m.modelLoader, m.appConfig, []backend.PreloadStage{ {Role: "vad", Cfg: m.VADConfig}, {Role: "transcription", Cfg: m.TranscriptionConfig}, {Role: "sound_detection", Cfg: m.SoundDetectionConfig}, }) - return err } func (m *wrappedModel) VAD(ctx context.Context, request *schema.VADRequest) (*schema.VADResponse, error) { @@ -812,32 +840,7 @@ func (m *wrappedModel) Warmup(ctx context.Context) error { if m.ScoreConfig != nil && m.ScoreConfig != m.LLMConfig { stages = append(stages, backend.PreloadStage{Role: "classifier", Cfg: m.ScoreConfig}) } - // A chain stage warms through its failover plan: a target that fails to - // load moves the stage to the next one instead of failing the session. - var ( - plain []backend.PreloadStage - wg sync.WaitGroup - mu sync.Mutex - errs []error - ) - for _, s := range stages { - if !m.isChainStage(s.Role) { - plain = append(plain, s) - continue - } - wg.Go(func() { - err := m.stageCall(ctx, s.Role, s.Cfg, func(cfg *config.ModelConfig, _ func()) error { - _, err := backend.PreloadStages(ctx, m.modelLoader, m.appConfig, []backend.PreloadStage{{Role: s.Role, Cfg: cfg}}) - return err - }) - mu.Lock() - errs = append(errs, err) - mu.Unlock() - }) - } - _, err := backend.PreloadStages(ctx, m.modelLoader, m.appConfig, plain) - wg.Wait() - return errors.Join(append(errs, err)...) + return m.warmStages(ctx, m.modelLoader, m.appConfig, stages) } // wavStreamHeaderBytes is the size of the WAV header that backend.ModelTTSStream @@ -930,8 +933,12 @@ func loadSoundDetectionConfig(pipeline *config.Pipeline, cl *config.ModelConfigL return cfg, nil } -func newTranscriptionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig) (Model, *config.ModelConfig, error) { +func newTranscriptionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig, fm *failover.Manager) (Model, *config.ModelConfig, error) { + sr := newStageRouter(fm, cl, ml, appConfig) cfgVAD, err := cl.LoadResolvedModelConfig(pipeline.VAD, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) + if err == nil { + cfgVAD, err = sr.resolveStage("vad", cfgVAD) + } if err != nil { return nil, nil, fmt.Errorf("failed to load backend config: %w", err) @@ -942,6 +949,9 @@ func newTranscriptionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfig } cfgSST, err := cl.LoadResolvedModelConfig(pipeline.Transcription, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) + if err == nil { + cfgSST, err = sr.resolveStage("transcription", cfgSST) + } if err != nil { return nil, nil, fmt.Errorf("failed to load backend config: %w", err) @@ -952,6 +962,9 @@ func newTranscriptionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfig } cfgSound, err := loadSoundDetectionConfig(pipeline, cl, ml, appConfig) + if err == nil { + cfgSound, err = sr.resolveStage("sound_detection", cfgSound) + } if err != nil { return nil, nil, err } @@ -964,6 +977,7 @@ func newTranscriptionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfig confLoader: cl, modelLoader: ml, appConfig: appConfig, + stageRouter: sr, }, cfgSST, nil } @@ -972,8 +986,12 @@ func newTranscriptionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfig // a sound-detection-only realtime session, which activates on sounds (not // speech) and is driven by client-side windowing (turn_detection none + // input_audio_buffer.commit) rather than the voice VAD loop. -func newSoundDetectionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig) (Model, error) { +func newSoundDetectionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig, fm *failover.Manager) (Model, error) { + sr := newStageRouter(fm, cl, ml, appConfig) cfgSound, err := loadSoundDetectionConfig(pipeline, cl, ml, appConfig) + if err == nil { + cfgSound, err = sr.resolveStage("sound_detection", cfgSound) + } if err != nil { return nil, err } @@ -985,6 +1003,7 @@ func newSoundDetectionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfi confLoader: cl, modelLoader: ml, appConfig: appConfig, + stageRouter: sr, }, nil } @@ -1031,29 +1050,12 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model // A stage that names a failover chain is resolved on every call. Here it // takes the chain's active target, so everything that inspects stage // configs at session start (voice, reasoning, templates) sees a real model. - stageChains := map[string]string{} - loadTarget := func(name string) (*config.ModelConfig, error) { - cfg, err := cl.LoadResolvedModelConfig(name, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) - if err != nil { - return nil, err - } - failover.PrepareTarget(cfg) - return cfg, nil - } - resolveStage := func(stage string, cfg *config.ModelConfig) (*config.ModelConfig, error) { - if cfg == nil || !cfg.IsFailover() { - return cfg, nil - } - if routing == nil || routing.Failover == nil { - return nil, fmt.Errorf("pipeline %s stage %q is a failover chain, but failover is not running", stage, cfg.Name) - } - st, ok := routing.Failover.ChainStatus(cfg.Name) - if !ok { - return nil, fmt.Errorf("failover chain %q not found", cfg.Name) - } - stageChains[stage] = cfg.Name - return loadTarget(st.Active) + var fm *failover.Manager + if routing != nil { + fm = routing.Failover } + sr := newStageRouter(fm, cl, ml, appConfig) + resolveStage := sr.resolveStage cfgVAD, err := cl.LoadResolvedModelConfig(pipeline.VAD, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) if err == nil { @@ -1193,17 +1195,14 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model appConfig: appConfig, evaluator: evaluator, - stageChains: stageChains, - stageTargetConfig: loadTarget, - tuneLLM: tuneLLM, - appTracing: appConfig.EnableTracing, + stageRouter: sr, + tuneLLM: tuneLLM, } if routing != nil { wm.routerDeps = routing.Deps wm.routerStore = routing.Store wm.routerSessionID = routing.SessionID wm.routerUserID = routing.UserID - wm.failover = routing.Failover } return wm, nil } diff --git a/docs/content/features/model-failover.md b/docs/content/features/model-failover.md index c7fa1039a..6bc5089e7 100644 --- a/docs/content/features/model-failover.md +++ b/docs/content/features/model-failover.md @@ -120,7 +120,8 @@ pipeline: tts: voice-chain ``` -LocalAI resolves the chain for every call of the stage. When a chain switches, +LocalAI resolves the chain for every call of the stage, in full pipelines and +in transcription-only and sound-detection-only sessions. When a chain switches, the session stays open and keeps its conversation. The next turn uses the new target. @@ -134,8 +135,6 @@ it starts (`reason: initial`) and each time a chain switches: Limits: -- Chains are resolved only in full realtime pipelines. A transcription-only or - sound-detection-only session does not resolve chains yet. - After a `session.update` that changes the pipeline, `localai.model.failover` events keep describing the chains from session start. - A chain used as a router candidate, or as the classifier-mode scoring model, 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 57bb7e5b0..326cfb45a 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 @@ -278,8 +278,11 @@ Plain HTTP clients can see failover without subscribing to events. ## Request path (realtime) - In `core/http/endpoints/openai/realtime_model.go`, a pipeline stage that names - a chain is resolved **for each call** of `wrappedModel`, not once at session - start. A helper, `mgr.Do(ctx, chain, func(cfg *config.ModelConfig) error)`, + a chain is resolved **for each call**, not once at session start. This holds + for the full pipeline (`wrappedModel`) and for transcription-only and + sound-detection-only sessions (`transcriptOnlyModel`); both embed the same + stage router. A chain config never reaches the model loader: it has no + backend and would start backend auto-detection. A helper, `mgr.Do(ctx, chain, func(cfg *config.ModelConfig) error)`, goes through the plan with the same classification as HTTP. - Streaming stages (`Predict` with a token callback, `TTSStream`, `TranscribeStream`) wrap the callback. A retry is allowed only until the first