diff --git a/core/http/endpoints/openai/realtime.go b/core/http/endpoints/openai/realtime.go index 9db5f899f..fe2c15068 100644 --- a/core/http/endpoints/openai/realtime.go +++ b/core/http/endpoints/openai/realtime.go @@ -754,6 +754,14 @@ func runRealtimeSession(application *application.Application, t Transport, model Session: session.ToServer(), }) + // 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() + } + var ( msg []byte wg sync.WaitGroup diff --git a/core/http/endpoints/openai/realtime_classifier_test.go b/core/http/endpoints/openai/realtime_classifier_test.go index b8630ca81..224808350 100644 --- a/core/http/endpoints/openai/realtime_classifier_test.go +++ b/core/http/endpoints/openai/realtime_classifier_test.go @@ -45,7 +45,7 @@ var classifierTestHistory = schema.Messages{ func classifierResultEvents(t *fakeTransport) []types.ClassifierResultEvent { var out []types.ClassifierResultEvent - for _, e := range t.events { + for _, e := range t.events() { if ev, ok := e.(types.ClassifierResultEvent); ok { out = append(out, ev) } @@ -57,7 +57,7 @@ func classifierResultEvents(t *fakeTransport) []types.ClassifierResultEvent { // item — what a classifier response actually "spoke". func replyTexts(t *fakeTransport) []string { var out []string - for _, e := range t.events { + for _, e := range t.events() { if ev, ok := e.(types.ResponseOutputTextDoneEvent); ok { out = append(out, ev.Text) } @@ -277,7 +277,7 @@ var _ = Describe("classifierRespond", func() { Expect(t.countEvents(types.ServerEventTypeResponseOutputTextDone)).To(Equal(1)) Expect(t.countEvents(types.ServerEventTypeResponseFunctionCallArgumentsDone)).To(Equal(1)) var fcArgs string - for _, e := range t.events { + for _, e := range t.events() { if done, ok := e.(types.ResponseFunctionCallArgumentsDoneEvent); ok { fcArgs = done.Arguments } @@ -656,7 +656,7 @@ var _ = Describe("classifierRespond slot filling", func() { Expect(results[0].Arguments).To(MatchJSON(`{"direction":"up","distance":3,"units":"meters"}`)) var fcArgs string - for _, e := range t.events { + for _, e := range t.events() { if done, ok := e.(types.ResponseFunctionCallArgumentsDoneEvent); ok { fcArgs = done.Arguments } @@ -695,7 +695,7 @@ var _ = Describe("classifierRespond slot filling", func() { Expect(handled).To(BeTrue()) var fcArgs string - for _, e := range t.events { + for _, e := range t.events() { if done, ok := e.(types.ResponseFunctionCallArgumentsDoneEvent); ok { fcArgs = done.Arguments } diff --git a/core/http/endpoints/openai/realtime_doubles_test.go b/core/http/endpoints/openai/realtime_doubles_test.go index a2c104b3c..e82a4a130 100644 --- a/core/http/endpoints/openai/realtime_doubles_test.go +++ b/core/http/endpoints/openai/realtime_doubles_test.go @@ -17,8 +17,10 @@ import ( // so streaming behaviour can be asserted without a real WebSocket/WebRTC peer. // It is not a *WebRTCTransport, so handler code takes the WebSocket path. type fakeTransport struct { - events []types.ServerEvent - audio []fakeAudioChunk + // mu guards sent: some specs send from a background goroutine. + mu sync.Mutex + sent []types.ServerEvent + audio []fakeAudioChunk } type fakeAudioChunk struct { @@ -27,10 +29,19 @@ type fakeAudioChunk struct { } func (f *fakeTransport) SendEvent(e types.ServerEvent) error { - f.events = append(f.events, e) + f.mu.Lock() + defer f.mu.Unlock() + f.sent = append(f.sent, e) return nil } +// events returns a copy of the server events sent so far. +func (f *fakeTransport) events() []types.ServerEvent { + f.mu.Lock() + defer f.mu.Unlock() + return append([]types.ServerEvent(nil), f.sent...) +} + func (f *fakeTransport) ReadEvent() ([]byte, error) { return nil, nil } func (f *fakeTransport) SendAudio(_ context.Context, pcm []byte, sampleRate int) error { @@ -43,7 +54,7 @@ func (f *fakeTransport) Close() error { return nil } // countEvents returns how many recorded events have the given type. func (f *fakeTransport) countEvents(et types.ServerEventType) int { n := 0 - for _, e := range f.events { + for _, e := range f.events() { if e.ServerEventType() == et { n++ } @@ -55,7 +66,7 @@ func (f *fakeTransport) countEvents(et types.ServerEventType) int { // delta event — i.e. the text streamed to the client as it is generated. func (f *fakeTransport) transcriptDeltaText() string { var b strings.Builder - for _, e := range f.events { + for _, e := range f.events() { if d, ok := e.(types.ResponseOutputAudioTranscriptDeltaEvent); ok { b.WriteString(d.Delta) } diff --git a/core/http/endpoints/openai/realtime_failover.go b/core/http/endpoints/openai/realtime_failover.go new file mode 100644 index 000000000..d58d52f33 --- /dev/null +++ b/core/http/endpoints/openai/realtime_failover.go @@ -0,0 +1,70 @@ +package openai + +import ( + "context" + "sort" + + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/http/endpoints/openai/types" + "github.com/mudler/LocalAI/core/services/failover" +) + +// isChainStage reports whether stage names a failover chain. +func (m *wrappedModel) isChainStage(stage string) bool { + _, ok := m.stageChains[stage] + return ok && m.failover != nil +} + +// 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) { + 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) + if err != nil { + return err + } + err = fn(cfg, commit) + if err != nil { + failover.RecordAttemptTrace(m.appTracing, chain, target, err) + } + return err + }) +} + +// 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() { + // Subscribe before reading the status, so a switch that lands in + // between is still delivered. + events, cancel := fm.Subscribe(16) + stages := make([]string, 0, len(stageChains)) + for s := range stageChains { + stages = append(stages, s) + } + sort.Strings(stages) + for _, stage := range stages { + chain := stageChains[stage] + if st, ok := fm.ChainStatus(chain); ok { + sendEvent(t, types.ModelFailoverEvent{Chain: chain, Stage: stage, To: st.Active, State: string(st.State), Reason: string(failover.ReasonInitial)}) + } + } + go func() { + for ev := range events { + if ev.Type != failover.EventChainSwitched { + continue + } + for _, stage := range stages { + if stageChains[stage] == ev.Chain { + sendEvent(t, types.ModelFailoverEvent{Chain: ev.Chain, Stage: stage, From: ev.From, To: ev.To, State: ev.State, Reason: string(ev.Reason)}) + } + } + } + }() + return cancel +} diff --git a/core/http/endpoints/openai/realtime_failover_test.go b/core/http/endpoints/openai/realtime_failover_test.go new file mode 100644 index 000000000..48737630c --- /dev/null +++ b/core/http/endpoints/openai/realtime_failover_test.go @@ -0,0 +1,99 @@ +package openai + +import ( + "context" + "errors" + + "github.com/mudler/LocalAI/core/config" + "github.com/mudler/LocalAI/core/http/endpoints/openai/types" + "github.com/mudler/LocalAI/core/services/failover" + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" +) + +type rtSource map[string]config.ModelConfig + +func (s rtSource) GetModelConfig(n string) (config.ModelConfig, bool) { c, ok := s[n]; return c, ok } +func (s rtSource) GetAllModelsConfigs() []config.ModelConfig { + var out []config.ModelConfig + for _, c := range s { + out = append(out, c) + } + return out +} + +var _ = Describe("realtime failover", func() { + var fm *failover.Manager + + BeforeEach(func() { + fm = failover.New(rtSource{ + "a": {Name: "a", Backend: "cloud-proxy"}, + "b": {Name: "b", Backend: "llama-cpp"}, + "chain": {Name: "chain", Failover: &config.FailoverConfig{Targets: []config.FailoverTarget{{Model: "a"}, {Model: "b"}}}}, + }) + }) + + 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 }} + } + + It("routes a chain stage through the plan and retries before commit", func() { + m := chainModel() + var tried []string + err := m.stageCall(context.Background(), "tts", nil, func(cfg *config.ModelConfig, _ func()) error { + tried = append(tried, cfg.Name) + if cfg.Name == "a" { + return errors.New("dial tcp: refused") + } + return nil + }) + Expect(err).ToNot(HaveOccurred()) + Expect(tried).To(Equal([]string{"a", "b"})) + }) + + It("does not retry a chain stage once output was committed", func() { + m := chainModel() + var tried []string + err := m.stageCall(context.Background(), "tts", nil, func(cfg *config.ModelConfig, commit func()) error { + tried = append(tried, cfg.Name) + commit() + return errors.New("dial tcp: refused") + }) + Expect(err).To(HaveOccurred()) + Expect(tried).To(Equal([]string{"a"})) + }) + + It("calls a plain stage once with its own config", func() { + m := &wrappedModel{} + base := &config.ModelConfig{Name: "plain"} + calls := 0 + err := m.stageCall(context.Background(), "tts", base, func(cfg *config.ModelConfig, _ func()) error { + calls++ + Expect(cfg).To(BeIdenticalTo(base)) + return nil + }) + Expect(err).ToNot(HaveOccurred()) + Expect(calls).To(Equal(1)) + }) + + It("sends initial events, then switch events, and stops on cancel", func() { + t := &fakeTransport{} + failoverEvents := 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 + } + stop := startFailoverEvents(t, fm, map[string]string{"llm": "chain"}) + Eventually(failoverEvents).Should(ContainElement(And( + HaveField("Reason", "initial"), HaveField("To", "a"), HaveField("Stage", "llm")))) + fm.ReportFailure("a", errors.New("dial tcp: refused")) + Eventually(failoverEvents).Should(ContainElement(And( + HaveField("Reason", "trip"), HaveField("From", "a"), HaveField("To", "b"), HaveField("Chain", "chain")))) + stop() + }) +}) diff --git a/core/http/endpoints/openai/realtime_model.go b/core/http/endpoints/openai/realtime_model.go index b525eee26..a90a0e4e5 100644 --- a/core/http/endpoints/openai/realtime_model.go +++ b/core/http/endpoints/openai/realtime_model.go @@ -19,6 +19,7 @@ import ( "github.com/mudler/LocalAI/core/http/endpoints/openai/types" "github.com/mudler/LocalAI/core/http/middleware" "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" @@ -84,6 +85,18 @@ type wrappedModel struct { routerStore router.DecisionStore 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) + // 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 } // anyToAnyModel represent a model which supports Any-to-Any operations @@ -165,15 +178,33 @@ func (m *transcriptOnlyModel) Warmup(ctx context.Context) error { } func (m *wrappedModel) 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 *wrappedModel) 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 *wrappedModel) 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 *wrappedModel) 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) { @@ -181,11 +212,22 @@ func (m *wrappedModel) Predict(ctx context.Context, messages schema.Messages, im Messages: messages, } + toolsJSON, toolChoiceJSON := realtimeToolsJSON(tools, toolChoice) + + // infer renders the prompt for cfg and starts inference on it. Everything + // that reads the LLM config lives here, so a chain stage can run it again + // against the next target. + infer := func(cfg *config.ModelConfig, cb func(string, backend.TokenUsage) bool) (func() (backend.LLMResponse, error), error) { + predInput := m.renderPredictPrompt(input, cfg, tools, toolChoice) + return backend.ModelInference(ctx, predInput, messages, images, videos, audios, m.modelLoader, cfg, m.confLoader, m.appConfig, cb, toolsJSON, toolChoiceJSON, logprobs, topLogprobs, logitBias, nil) + } + // Per-turn routing: when the session's LLMConfig is a router, swap // to the candidate the classifier picks for this turn's prompt. // LLMConfig itself is held by value (we never mutate it) — turnCfg // is the config we dispatch against. turnCfg := m.LLMConfig + routed := false if m.LLMConfig.HasRouter() && m.routerDeps != nil { chosen, err := m.routeTurn(ctx, &input) if err != nil { @@ -193,9 +235,47 @@ func (m *wrappedModel) Predict(ctx context.Context, messages schema.Messages, im "router_model", m.LLMConfig.Name, "error", err) } else if chosen != nil { turnCfg = chosen + routed = true } } + // A routed turn dispatches to the router's pick: chains as router + // candidates are not resolved here. + if routed || !m.isChainStage("llm") { + return infer(turnCfg, tokenCallback) + } + + return func() (backend.LLMResponse, error) { + var resp backend.LLMResponse + err := m.stageCall(ctx, "llm", turnCfg, func(cfg *config.ModelConfig, commit func()) error { + if m.tuneLLM != nil { + m.tuneLLM(cfg) + } + // Without a callback nothing reaches the client before the + // reply is complete, so every failure can still be retried. + var cb func(string, backend.TokenUsage) bool + if tokenCallback != nil { + cb = func(s string, u backend.TokenUsage) bool { + commit() + return tokenCallback(s, u) + } + } + predict, err := infer(cfg, cb) + if err != nil { + return err + } + resp, err = predict() + return err + }) + return resp, err + }, nil +} + +// renderPredictPrompt templates the turn's prompt for cfg. It also applies +// the turn's tool choice and function-calling grammar to cfg, which the +// backend reads when inference starts. The prompt is empty for models that +// use the tokenizer's template. +func (m *wrappedModel) renderPredictPrompt(input schema.OpenAIRequest, turnCfg *config.ModelConfig, tools []types.ToolUnion, toolChoice *types.ToolChoiceUnion) string { // Surface the resolved reasoning effort to the Go-side template path too // (jinja models get it via backend metadata in gRPCPredictOpts; Go-templated // models like gpt-oss read it from the template's .ReasoningEffort). @@ -303,6 +383,12 @@ func (m *wrappedModel) Predict(ctx context.Context, messages schema.Messages, im } } + return predInput +} + +// realtimeToolsJSON serializes the turn's tools and tool choice the way the +// backends expect them. Neither depends on the LLM config. +func realtimeToolsJSON(tools []types.ToolUnion, toolChoice *types.ToolChoiceUnion) (string, string) { var toolsJSON string if len(tools) > 0 { // Convert tools to OpenAI Chat Completions format (nested) @@ -348,7 +434,7 @@ func (m *wrappedModel) Predict(ctx context.Context, messages schema.Messages, im toolChoiceJSON = string(b) } - return backend.ModelInference(ctx, predInput, messages, images, videos, audios, m.modelLoader, turnCfg, m.confLoader, m.appConfig, tokenCallback, toolsJSON, toolChoiceJSON, logprobs, topLogprobs, logitBias, nil) + return toolsJSON, toolChoiceJSON } // routeTurn classifies this turn's prompt against the session's router @@ -395,7 +481,16 @@ func newRealtimeDecisionID() string { } func (m *wrappedModel) TTS(ctx context.Context, text, voice, language string) (string, *proto.Result, error) { - return backend.ModelTTS(ctx, text, voice, language, "", maps.Clone(m.ttsParams), m.modelLoader, m.appConfig, *m.TTSConfig) + var ( + out string + res *proto.Result + ) + err := m.stageCall(ctx, "tts", m.TTSConfig, func(cfg *config.ModelConfig, _ func()) error { + var err error + out, res, err = backend.ModelTTS(ctx, text, voice, language, "", maps.Clone(m.ttsParams), m.modelLoader, m.appConfig, *cfg) + return err + }) + return out, res, err } func (m *wrappedModel) setTTSParams(params map[string]string) { @@ -403,7 +498,14 @@ func (m *wrappedModel) setTTSParams(params map[string]string) { } func (m *wrappedModel) TTSStream(ctx context.Context, text, voice, language string, onAudio func(pcm []byte, sampleRate int) error) error { - return ttsStream(ctx, m.modelLoader, m.appConfig, *m.TTSConfig, text, voice, language, maps.Clone(m.ttsParams), onAudio) + return m.stageCall(ctx, "tts", m.TTSConfig, func(cfg *config.ModelConfig, commit func()) error { + // Audio that reached the client cannot be taken back, so the first + // chunk ends the retries. + return ttsStream(ctx, m.modelLoader, m.appConfig, *cfg, text, voice, language, maps.Clone(m.ttsParams), func(pcm []byte, sr int) error { + commit() + return onAudio(pcm, sr) + }) + }) } func resolveRealtimeVoice(ctx context.Context, configuredVoice string, ttsConfig *config.ModelConfig, profiles *voiceprofile.Store) (string, map[string]string, func(), error) { @@ -431,11 +533,28 @@ func resolveRealtimeVoice(ctx context.Context, configuredVoice string, ttsConfig } func (m *wrappedModel) 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 *wrappedModel) 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: once it is + // open, events flow to the client for the rest of the utterance. + 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 *wrappedModel) PredictConfig() *config.ModelConfig { @@ -693,8 +812,32 @@ 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}) } - _, err := backend.PreloadStages(ctx, m.modelLoader, m.appConfig, stages) - return err + // 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)...) } // wavStreamHeaderBytes is the size of the WAV header that backend.ModelTTSStream @@ -854,6 +997,8 @@ type RealtimeRoutingContext struct { Store router.DecisionStore SessionID string UserID string + // Failover resolves pipeline stages that name a failover chain. + Failover *failover.Manager } // buildRealtimeRoutingContext assembles the routing dependencies the @@ -875,6 +1020,7 @@ func buildRealtimeRoutingContext(a *application.Application, sessionID string) * Store: a.RouterDecisions(), SessionID: sessionID, UserID: userID, + Failover: a.FailoverManager(), } } @@ -882,7 +1028,29 @@ func buildRealtimeRoutingContext(a *application.Application, sessionID string) * func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model.ModelLoader, appConfig *config.ApplicationConfig, evaluator *templates.Evaluator, routing *RealtimeRoutingContext) (Model, error) { xlog.Debug("Creating new model pipeline model", "pipeline", pipeline) + // 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{} + 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 cl.LoadResolvedModelConfig(st.Active, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) + } + cfgVAD, err := cl.LoadResolvedModelConfig(pipeline.VAD, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) + if err == nil { + cfgVAD, err = resolveStage("vad", cfgVAD) + } if err != nil { return nil, fmt.Errorf("failed to load backend config: %w", err) @@ -894,6 +1062,9 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model // TODO: Do we always need a transcription model? It can be disabled. Note that any-to-any instruction following models don't transcribe as such, so if transcription is required it is a separate process cfgSST, err := cl.LoadResolvedModelConfig(pipeline.Transcription, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) + if err == nil { + cfgSST, err = resolveStage("transcription", cfgSST) + } if err != nil { return nil, fmt.Errorf("failed to load backend config: %w", err) @@ -926,6 +1097,9 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model // Otherwise we want to return a wrapped model, which is a "virtual" model that re-uses other models to perform operations cfgLLM, err := cl.LoadResolvedModelConfig(pipeline.LLM, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) + if err == nil { + cfgLLM, err = resolveStage("llm", cfgLLM) + } if err != nil { return nil, fmt.Errorf("failed to load backend config: %w", err) @@ -937,10 +1111,17 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model // Let the pipeline set the LLM's reasoning effort and force thinking off // (cfgLLM is a per-session copy). disable_thinking applies after the effort. - applyPipelineReasoning(cfgLLM, *pipeline) - applyPipelineThinking(cfgLLM, *pipeline) + pipelineCopy := *pipeline + tuneLLM := func(cfg *config.ModelConfig) { + applyPipelineReasoning(cfg, pipelineCopy) + applyPipelineThinking(cfg, pipelineCopy) + } + tuneLLM(cfgLLM) cfgTTS, err := cl.LoadResolvedModelConfig(pipeline.TTS, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) + if err == nil { + cfgTTS, err = resolveStage("tts", cfgTTS) + } if err != nil { return nil, fmt.Errorf("failed to load backend config: %w", err) @@ -951,6 +1132,9 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model } cfgSound, err := loadSoundDetectionConfig(pipeline, cl, ml, appConfig) + if err == nil { + cfgSound, err = resolveStage("sound_detection", cfgSound) + } if err != nil { return nil, err } @@ -1000,12 +1184,20 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model modelLoader: ml, appConfig: appConfig, evaluator: evaluator, + + stageChains: stageChains, + stageTargetConfig: func(name string) (*config.ModelConfig, error) { + return cl.LoadResolvedModelConfig(name, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) + }, + tuneLLM: tuneLLM, + appTracing: appConfig.EnableTracing, } 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/core/http/endpoints/openai/realtime_semantic_vad_test.go b/core/http/endpoints/openai/realtime_semantic_vad_test.go index c36e13563..b1107c1f6 100644 --- a/core/http/endpoints/openai/realtime_semantic_vad_test.go +++ b/core/http/endpoints/openai/realtime_semantic_vad_test.go @@ -269,7 +269,7 @@ var _ = Describe("liveTurnState", func() { lts.drainEvents(1.0) var got []types.ConversationItemInputAudioTranscriptionDeltaEvent - for _, e := range ftr.events { + for _, e := range ftr.events() { if d, ok := e.(types.ConversationItemInputAudioTranscriptionDeltaEvent); ok { got = append(got, d) } @@ -335,7 +335,7 @@ var _ = Describe("commitUtteranceWithTranscript", func() { Expect(tr.countEvents(types.ServerEventTypeConversationItemInputAudioTranscriptionCompleted)).To(Equal(1)) var completed types.ConversationItemInputAudioTranscriptionCompletedEvent - for _, e := range tr.events { + for _, e := range tr.events() { if c, ok := e.(types.ConversationItemInputAudioTranscriptionCompletedEvent); ok { completed = c } @@ -407,7 +407,7 @@ var _ = Describe("emitPrecomputedTranscription", func() { Expect(tr.countEvents(types.ServerEventTypeConversationItemInputAudioTranscriptionDelta)).To(Equal(2), "empty deltas skipped") Expect(tr.countEvents(types.ServerEventTypeConversationItemInputAudioTranscriptionCompleted)).To(Equal(1)) - for _, e := range tr.events { + for _, e := range tr.events() { switch ev := e.(type) { case types.ConversationItemInputAudioTranscriptionDeltaEvent: Expect(ev.ItemID).To(Equal("item42")) diff --git a/core/http/endpoints/openai/realtime_sound_detection_test.go b/core/http/endpoints/openai/realtime_sound_detection_test.go index e440e80c3..058c74076 100644 --- a/core/http/endpoints/openai/realtime_sound_detection_test.go +++ b/core/http/endpoints/openai/realtime_sound_detection_test.go @@ -38,7 +38,7 @@ var _ = Describe("emitSoundDetection", func() { Expect(err).ToNot(HaveOccurred()) Expect(t.countEvents(types.ServerEventTypeConversationItemSoundDetection)).To(Equal(1)) - ev, ok := t.events[0].(types.ConversationItemSoundDetectionEvent) + ev, ok := t.events()[0].(types.ConversationItemSoundDetectionEvent) Expect(ok).To(BeTrue()) Expect(ev.ItemID).To(Equal("item1")) Expect(ev.ContentIndex).To(Equal(0)) @@ -62,7 +62,7 @@ var _ = Describe("emitSoundDetection", func() { Expect(err).ToNot(HaveOccurred()) Expect(t.countEvents(types.ServerEventTypeConversationItemSoundDetection)).To(Equal(1)) - ev, ok := t.events[0].(types.ConversationItemSoundDetectionEvent) + ev, ok := t.events()[0].(types.ConversationItemSoundDetectionEvent) Expect(ok).To(BeTrue()) Expect(ev.Detections).To(BeEmpty()) }) diff --git a/core/http/endpoints/openai/realtime_stream_test.go b/core/http/endpoints/openai/realtime_stream_test.go index 2d5d7d7a1..5f160027e 100644 --- a/core/http/endpoints/openai/realtime_stream_test.go +++ b/core/http/endpoints/openai/realtime_stream_test.go @@ -250,8 +250,8 @@ var _ = Describe("triggerResponse", func() { // The single terminal carries the produced output item and the usage — // both empty in the legacy code. var done *types.ResponseDoneEvent - for i := range t.events { - if d, ok := t.events[i].(types.ResponseDoneEvent); ok { + for i := range t.events() { + if d, ok := t.events()[i].(types.ResponseDoneEvent); ok { done = &d } } @@ -287,8 +287,8 @@ var _ = Describe("triggerResponse", func() { var created *types.ResponseCreatedEvent var done *types.ResponseDoneEvent - for i := range t.events { - switch e := t.events[i].(type) { + for i := range t.events() { + switch e := t.events()[i].(type) { case types.ResponseCreatedEvent: created = &e case types.ResponseDoneEvent: @@ -317,8 +317,8 @@ var _ = Describe("triggerResponse", func() { triggerResponse(context.Background(), session, &Conversation{}, t, nil) - for i := range t.events { - if d, ok := t.events[i].(types.ResponseDoneEvent); ok { + for i := range t.events() { + if d, ok := t.events()[i].(types.ResponseDoneEvent); ok { Expect(d.Response.Metadata).To(BeEmpty()) } } diff --git a/core/http/endpoints/openai/realtime_voicegate_integration_test.go b/core/http/endpoints/openai/realtime_voicegate_integration_test.go index b0f7f0b49..4da774c76 100644 --- a/core/http/endpoints/openai/realtime_voicegate_integration_test.go +++ b/core/http/endpoints/openai/realtime_voicegate_integration_test.go @@ -67,7 +67,7 @@ func itSession(gate *voiceGate) (*Session, *fakeModel) { // hasSpeakerNotAuthorized reports whether a speaker_not_authorized error event // was emitted to the client. func hasSpeakerNotAuthorized(tr *fakeTransport) bool { - for _, e := range tr.events { + for _, e := range tr.events() { if ev, ok := e.(types.ErrorEvent); ok && ev.Error.Code == "speaker_not_authorized" { return true } diff --git a/core/http/endpoints/openai/types/failover.go b/core/http/endpoints/openai/types/failover.go new file mode 100644 index 000000000..69c048207 --- /dev/null +++ b/core/http/endpoints/openai/types/failover.go @@ -0,0 +1,46 @@ +package types + +import "encoding/json" + +// ModelFailoverEvent is a LocalAI extension server event +// (localai.model.failover). It tells a client which target serves a +// pipeline stage that names a failover chain: once per chain stage at +// session start (reason "initial"), then on every switch of that chain. +type ModelFailoverEvent struct { + ServerEventBase + + // The failover chain the stage names. + Chain string `json:"chain"` + + // The pipeline stage: vad, transcription, llm, tts or sound_detection. + Stage string `json:"stage"` + + // The target that served the stage before the switch; "" at session start. + From string `json:"from"` + + // The target that serves the stage now. + To string `json:"to"` + + // The chain state: primary, fallback or degraded. + State string `json:"state"` + + // Why the chain switched, or "initial" at session start. + Reason string `json:"reason"` +} + +func (m ModelFailoverEvent) ServerEventType() ServerEventType { + return ServerEventTypeModelFailover +} + +func (m ModelFailoverEvent) MarshalJSON() ([]byte, error) { + type typeAlias ModelFailoverEvent + type typeWrapper struct { + typeAlias + Type ServerEventType `json:"type"` + } + shadow := typeWrapper{ + typeAlias: typeAlias(m), + Type: m.ServerEventType(), + } + return json.Marshal(shadow) +} diff --git a/core/http/endpoints/openai/types/server_events.go b/core/http/endpoints/openai/types/server_events.go index b847a35a7..114a7065a 100644 --- a/core/http/endpoints/openai/types/server_events.go +++ b/core/http/endpoints/openai/types/server_events.go @@ -27,7 +27,11 @@ const ( // ServerEventTypeClassifierResult is a LocalAI extension: it carries the // classifier-mode score distribution and decision for a response. OpenAI // clients ignore it. - ServerEventTypeClassifierResult ServerEventType = "localai.classifier.result" + ServerEventTypeClassifierResult ServerEventType = "localai.classifier.result" + // ServerEventTypeModelFailover is a LocalAI extension: it names the target + // that serves a pipeline stage backed by a failover chain, at session start + // and on every chain switch. OpenAI clients ignore it. + ServerEventTypeModelFailover ServerEventType = "localai.model.failover" ServerEventTypeInputAudioBufferCommitted ServerEventType = "input_audio_buffer.committed" ServerEventTypeInputAudioBufferCleared ServerEventType = "input_audio_buffer.cleared" ServerEventTypeInputAudioBufferSpeechStarted ServerEventType = "input_audio_buffer.speech_started" diff --git a/tests/e2e/e2e_suite_test.go b/tests/e2e/e2e_suite_test.go index a89318dfc..e7bb844d5 100644 --- a/tests/e2e/e2e_suite_test.go +++ b/tests/e2e/e2e_suite_test.go @@ -273,6 +273,44 @@ var _ = BeforeSuite(func() { Expect(err).ToNot(HaveOccurred()) Expect(os.WriteFile(filepath.Join(modelsPath, "realtime-pipeline.yaml"), pipelineData, 0644)).To(Succeed()) + // Realtime pipelines whose LLM is a failover chain: target 0 always fails + // to load, target 1 is mock-llm. rt-failover skips the warm-up so the + // switch happens on the first turn, mid-session; rt-failover-warm keeps it + // so the switch happens while the session starts. Each has its own chain + // because chain state is shared across sessions. + for _, rt := range []struct { + name, suffix string + disableWarmup bool + }{{"rt-failover", "rt", true}, {"rt-failover-warm", "rt-warm", false}} { + for _, cfg := range []map[string]any{ + { + "name": "fail-" + rt.suffix, + "backend": "mock-backend", + "parameters": map[string]any{"model": "fail-load-" + rt.suffix}, + }, + { + "name": "chain-" + rt.suffix, + "failover": map[string]any{ + "targets": []map[string]any{{"model": "fail-" + rt.suffix}, {"model": "mock-llm"}}, + }, + }, + { + "name": rt.name, + "pipeline": map[string]any{ + "vad": "mock-vad", + "transcription": "mock-stt", + "llm": "chain-" + rt.suffix, + "tts": "mock-tts", + "disable_warmup": rt.disableWarmup, + }, + }, + } { + data, err := yaml.Marshal(cfg) + Expect(err).ToNot(HaveOccurred()) + Expect(os.WriteFile(filepath.Join(modelsPath, cfg["name"].(string)+".yaml"), data, 0644)).To(Succeed()) + } + } + // Classifier-mode pipeline (LocalAI extension): responses are // prefill-scored against the option list via the mock backend's // ROUTE_HINT-driven Score instead of being generated. Threshold 0.6: diff --git a/tests/e2e/realtime_ws_test.go b/tests/e2e/realtime_ws_test.go index 6aeb3fcd4..1aa2552a1 100644 --- a/tests/e2e/realtime_ws_test.go +++ b/tests/e2e/realtime_ws_test.go @@ -192,6 +192,106 @@ var _ = Describe("Realtime WebSocket API", Label("Realtime"), func() { }) }) + Context("Failover chain stage", Label("failover"), func() { + // userTurn adds a user text item, asks for a response and reads until + // response.done. It returns the user item id, the response.done event + // and the localai.model.failover events seen on the way. + userTurn := func(conn *websocket.Conn, text string) (string, map[string]any, []map[string]any) { + sendClientEvent(conn, map[string]any{ + "type": "conversation.item.create", + "item": map[string]any{ + "type": "message", + "role": "user", + "content": []map[string]any{{"type": "input_text", "text": text}}, + }, + }) + added := drainUntil(conn, "conversation.item.added", 10*time.Second) + item, _ := added["item"].(map[string]any) + userID, _ := item["id"].(string) + ExpectWithOffset(1, userID).ToNot(BeEmpty()) + + sendClientEvent(conn, map[string]any{"type": "response.create"}) + var failovers []map[string]any + deadline := time.Now().Add(60 * time.Second) + for time.Now().Before(deadline) { + evt := readServerEvent(conn, time.Until(deadline)) + switch evt["type"] { + case "localai.model.failover": + failovers = append(failovers, evt) + case "error": + Fail(fmt.Sprintf("unexpected error event: %v", evt)) + case "response.done": + return userID, evt, failovers + } + } + Fail("timed out waiting for response.done") + return "", nil, nil + } + + retrieveItem := func(conn *websocket.Conn, id string) map[string]any { + sendClientEvent(conn, map[string]any{"type": "conversation.item.retrieve", "item_id": id}) + evt := drainUntil(conn, "conversation.item.retrieved", 10*time.Second) + item, _ := evt["item"].(map[string]any) + return item + } + + It("switches the LLM mid-session and keeps the conversation", func() { + conn := connectWS("rt-failover") + defer conn.Close() + + Expect(readServerEvent(conn, 30*time.Second)["type"]).To(Equal("session.created")) + initial := drainUntil(conn, "localai.model.failover", 10*time.Second) + Expect(initial).To(HaveKeyWithValue("stage", "llm")) + Expect(initial).To(HaveKeyWithValue("chain", "chain-rt")) + Expect(initial).To(HaveKeyWithValue("reason", "initial")) + Expect(initial).To(HaveKeyWithValue("to", "fail-rt")) + + sendClientEvent(conn, disableVADEvent()) + drainUntil(conn, "session.updated", 10*time.Second) + + firstID, done, failovers := userTurn(conn, "Hello, how are you?") + Expect(failovers).To(ContainElement(And( + HaveKeyWithValue("stage", "llm"), + HaveKeyWithValue("from", "fail-rt"), + HaveKeyWithValue("to", "mock-llm"), + HaveKeyWithValue("reason", "trip"), + ))) + resp, _ := done["response"].(map[string]any) + Expect(resp).To(HaveKeyWithValue("status", "completed")) + output, _ := resp["output"].([]any) + Expect(output).ToNot(BeEmpty()) + firstReply, _ := output[0].(map[string]any) + firstReplyID, _ := firstReply["id"].(string) + Expect(firstReplyID).ToNot(BeEmpty()) + + _, done, _ = userTurn(conn, "And now?") + resp, _ = done["response"].(map[string]any) + Expect(resp).To(HaveKeyWithValue("status", "completed")) + + // The switch kept the session: the first turn is still in it. + Expect(retrieveItem(conn, firstID)).To(HaveKeyWithValue("id", firstID)) + Expect(retrieveItem(conn, firstReplyID)).To(HaveKeyWithValue("id", firstReplyID)) + }) + + 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() + + Expect(readServerEvent(conn, 30*time.Second)["type"]).To(Equal("session.created")) + initial := drainUntil(conn, "localai.model.failover", 10*time.Second) + Expect(initial).To(HaveKeyWithValue("chain", "chain-rt-warm")) + Expect(initial).To(HaveKeyWithValue("reason", "initial")) + Expect(initial).To(HaveKeyWithValue("to", "mock-llm")) + + sendClientEvent(conn, disableVADEvent()) + drainUntil(conn, "session.updated", 10*time.Second) + + _, done, _ := userTurn(conn, "Hello?") + resp, _ := done["response"].(map[string]any) + Expect(resp).To(HaveKeyWithValue("status", "completed")) + }) + }) + Context("Manual audio commit", func() { It("should produce a response with audio when audio is committed", func() { conn := connectWS(pipelineModel())