diff --git a/core/backend/preload.go b/core/backend/preload.go index 3525da299..ab725bff3 100644 --- a/core/backend/preload.go +++ b/core/backend/preload.go @@ -66,12 +66,12 @@ func pipelineStages(cl *config.ModelConfigLoader, p *config.Pipeline, modelPath } var stages []PreloadStage for _, s := range []struct{ role, name string }{ - {"vad", p.VAD}, - {"transcription", p.Transcription}, - {"llm", p.LLM}, - {"tts", p.TTS}, - {"sound_detection", p.SoundDetection}, - {"voice_recognition", voiceRec}, + {config.PipelineStageVAD, p.VAD}, + {config.PipelineStageTranscription, p.Transcription}, + {config.PipelineStageLLM, p.LLM}, + {config.PipelineStageTTS, p.TTS}, + {config.PipelineStageSoundDetection, p.SoundDetection}, + {config.PipelineStageVoiceRecognition, voiceRec}, } { if s.name == "" { continue diff --git a/core/config/model_config.go b/core/config/model_config.go index 14510d0eb..3c8987920 100644 --- a/core/config/model_config.go +++ b/core/config/model_config.go @@ -800,6 +800,18 @@ type MCPSTDIOServer struct { Command string `json:"command,omitempty"` } +// Pipeline stage names. They match the Pipeline yaml keys and are the stage +// identifiers the realtime endpoint routes by (failover chains per stage, +// model_failover events, preload roles), so every user shares one spelling. +const ( + PipelineStageVAD = "vad" + PipelineStageTranscription = "transcription" + PipelineStageLLM = "llm" + PipelineStageTTS = "tts" + PipelineStageSoundDetection = "sound_detection" + PipelineStageVoiceRecognition = "voice_recognition" +) + // @Description Pipeline defines other models to use for audio-to-audio type Pipeline struct { TTS string `yaml:"tts,omitempty" json:"tts,omitempty"` diff --git a/core/http/endpoints/openai/realtime.go b/core/http/endpoints/openai/realtime.go index 8c8988371..11d1653cb 100644 --- a/core/http/endpoints/openai/realtime.go +++ b/core/http/endpoints/openai/realtime.go @@ -711,7 +711,7 @@ func runRealtimeSession(application *application.Application, t Transport, model var gateErr error if session.voiceGate != nil { _, gateErr = backend.PreloadStages(context.Background(), application.ModelLoader(), application.ApplicationConfig(), []backend.PreloadStage{ - {Role: "voice_recognition", Cfg: session.voiceGate.recCfg}, + {Role: config.PipelineStageVoiceRecognition, Cfg: session.voiceGate.recCfg}, }) } if err := errors.Join(<-warmErr, gateErr); err != nil { diff --git a/core/http/endpoints/openai/realtime_failover.go b/core/http/endpoints/openai/realtime_failover.go index c47c705ce..c998e05c2 100644 --- a/core/http/endpoints/openai/realtime_failover.go +++ b/core/http/endpoints/openai/realtime_failover.go @@ -19,7 +19,7 @@ import ( // 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. + // chain; stageChains maps a stage (config.PipelineStage*) 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 diff --git a/core/http/endpoints/openai/realtime_failover_test.go b/core/http/endpoints/openai/realtime_failover_test.go index 996fa1eb6..06dfda2b3 100644 --- a/core/http/endpoints/openai/realtime_failover_test.go +++ b/core/http/endpoints/openai/realtime_failover_test.go @@ -39,14 +39,14 @@ var _ = Describe("realtime failover", func() { }) chainModel := func() *wrappedModel { - return &wrappedModel{stageRouter: stageRouter{failover: fm, stageChains: map[string]string{"tts": "chain"}, + return &wrappedModel{stageRouter: stageRouter{failover: fm, stageChains: map[string]string{config.PipelineStageTTS: "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 { + err := m.stageCall(context.Background(), config.PipelineStageTTS, nil, func(cfg *config.ModelConfig, _ func()) error { tried = append(tried, cfg.Name) if cfg.Name == "a" { return errors.New("dial tcp: refused") @@ -60,7 +60,7 @@ var _ = Describe("realtime failover", func() { 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 { + err := m.stageCall(context.Background(), config.PipelineStageTTS, nil, func(cfg *config.ModelConfig, commit func()) error { tried = append(tried, cfg.Name) commit() return errors.New("dial tcp: refused") @@ -73,7 +73,7 @@ var _ = Describe("realtime failover", func() { m := &wrappedModel{} base := &config.ModelConfig{Name: "plain"} calls := 0 - err := m.stageCall(context.Background(), "tts", base, func(cfg *config.ModelConfig, _ func()) error { + err := m.stageCall(context.Background(), config.PipelineStageTTS, base, func(cfg *config.ModelConfig, _ func()) error { calls++ Expect(cfg).To(BeIdenticalTo(base)) return nil @@ -93,7 +93,7 @@ var _ = Describe("realtime failover", func() { } return out } - stop := startFailoverEvents(t, fm, map[string]string{"llm": "chain"}) + stop := startFailoverEvents(t, fm, map[string]string{config.PipelineStageLLM: "chain"}) Eventually(failoverEvents).Should(ContainElement(And( HaveField("Reason", "initial"), HaveField("To", "a"), HaveField("Stage", "llm")))) fm.ReportFailure("a", errors.New("dial tcp: refused")) @@ -150,7 +150,7 @@ var _ = Describe("realtime failover in transcription-only and sound-only session 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"})) + Expect(tm.stageChains).To(Equal(map[string]string{config.PipelineStageSoundDetection: "sound-chain"})) var tried []string tm.stageTargetConfig = func(name string) (*config.ModelConfig, error) { @@ -176,7 +176,7 @@ var _ = Describe("realtime failover in transcription-only and sound-only session 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"})) + Expect(tm.stageChains).To(Equal(map[string]string{config.PipelineStageTranscription: "stt-chain"})) var tried []string tm.stageTargetConfig = func(name string) (*config.ModelConfig, error) { diff --git a/core/http/endpoints/openai/realtime_model.go b/core/http/endpoints/openai/realtime_model.go index f94bb510b..bf4bdb6d2 100644 --- a/core/http/endpoints/openai/realtime_model.go +++ b/core/http/endpoints/openai/realtime_model.go @@ -118,7 +118,7 @@ type transcriptOnlyModel struct { func (m *transcriptOnlyModel) VAD(ctx context.Context, request *schema.VADRequest) (*schema.VADResponse, error) { var res *schema.VADResponse - err := m.stageCall(ctx, "vad", m.VADConfig, func(cfg *config.ModelConfig, _ func()) error { + err := m.stageCall(ctx, config.PipelineStageVAD, m.VADConfig, func(cfg *config.ModelConfig, _ func()) error { var err error res, err = backend.VAD(request, ctx, m.modelLoader, m.appConfig, *cfg) return err @@ -128,7 +128,7 @@ func (m *transcriptOnlyModel) VAD(ctx context.Context, request *schema.VADReques func (m *transcriptOnlyModel) Transcribe(ctx context.Context, audio, language string, translate bool, diarize bool, prompt string) (*schema.TranscriptionResult, error) { var res *schema.TranscriptionResult - err := m.stageCall(ctx, "transcription", m.TranscriptionConfig, func(cfg *config.ModelConfig, _ func()) error { + err := m.stageCall(ctx, config.PipelineStageTranscription, 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 @@ -138,7 +138,7 @@ func (m *transcriptOnlyModel) Transcribe(ctx context.Context, audio, language st func (m *transcriptOnlyModel) SoundDetection(ctx context.Context, audio string, topK int, threshold float32) (*schema.SoundClassificationResult, error) { var res *schema.SoundClassificationResult - err := m.stageCall(ctx, "sound_detection", m.SoundDetectionConfig, func(cfg *config.ModelConfig, _ func()) error { + err := m.stageCall(ctx, config.PipelineStageSoundDetection, 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 @@ -171,7 +171,7 @@ 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) { var res *schema.TranscriptionResult - err := m.stageCall(ctx, "transcription", m.TranscriptionConfig, func(cfg *config.ModelConfig, commit func()) error { + err := m.stageCall(ctx, config.PipelineStageTranscription, 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() @@ -185,7 +185,7 @@ func (m *transcriptOnlyModel) TranscribeStream(ctx context.Context, audio, langu func (m *transcriptOnlyModel) TranscribeLive(ctx context.Context, language string, onEvent func(backend.LiveTranscriptionEvent)) (backend.LiveTranscriptionSession, error) { 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 { + err := m.stageCall(ctx, config.PipelineStageTranscription, 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 @@ -199,15 +199,15 @@ func (m *transcriptOnlyModel) PredictConfig() *config.ModelConfig { func (m *transcriptOnlyModel) Warmup(ctx context.Context) error { 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}, + {Role: config.PipelineStageVAD, Cfg: m.VADConfig}, + {Role: config.PipelineStageTranscription, Cfg: m.TranscriptionConfig}, + {Role: config.PipelineStageSoundDetection, Cfg: m.SoundDetectionConfig}, }) } func (m *wrappedModel) VAD(ctx context.Context, request *schema.VADRequest) (*schema.VADResponse, error) { var res *schema.VADResponse - err := m.stageCall(ctx, "vad", m.VADConfig, func(cfg *config.ModelConfig, _ func()) error { + err := m.stageCall(ctx, config.PipelineStageVAD, m.VADConfig, func(cfg *config.ModelConfig, _ func()) error { var err error res, err = backend.VAD(request, ctx, m.modelLoader, m.appConfig, *cfg) return err @@ -217,7 +217,7 @@ func (m *wrappedModel) VAD(ctx context.Context, request *schema.VADRequest) (*sc func (m *wrappedModel) Transcribe(ctx context.Context, audio, language string, translate bool, diarize bool, prompt string) (*schema.TranscriptionResult, error) { var res *schema.TranscriptionResult - err := m.stageCall(ctx, "transcription", m.TranscriptionConfig, func(cfg *config.ModelConfig, _ func()) error { + err := m.stageCall(ctx, config.PipelineStageTranscription, 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 @@ -227,7 +227,7 @@ func (m *wrappedModel) Transcribe(ctx context.Context, audio, language string, t func (m *wrappedModel) SoundDetection(ctx context.Context, audio string, topK int, threshold float32) (*schema.SoundClassificationResult, error) { var res *schema.SoundClassificationResult - err := m.stageCall(ctx, "sound_detection", m.SoundDetectionConfig, func(cfg *config.ModelConfig, _ func()) error { + err := m.stageCall(ctx, config.PipelineStageSoundDetection, 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 @@ -269,13 +269,13 @@ func (m *wrappedModel) Predict(ctx context.Context, messages schema.Messages, im // A routed turn dispatches to the router's pick: chains as router // candidates are not resolved here. - if routed || !m.isChainStage("llm") { + if routed || !m.isChainStage(config.PipelineStageLLM) { 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 { + err := m.stageCall(ctx, config.PipelineStageLLM, turnCfg, func(cfg *config.ModelConfig, commit func()) error { if m.tuneLLM != nil { m.tuneLLM(cfg) } @@ -513,7 +513,7 @@ func (m *wrappedModel) TTS(ctx context.Context, text, voice, language string) (s out string res *proto.Result ) - err := m.stageCall(ctx, "tts", m.TTSConfig, func(cfg *config.ModelConfig, _ func()) error { + err := m.stageCall(ctx, config.PipelineStageTTS, 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 @@ -526,7 +526,7 @@ 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 m.stageCall(ctx, "tts", m.TTSConfig, func(cfg *config.ModelConfig, commit func()) error { + return m.stageCall(ctx, config.PipelineStageTTS, 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 { @@ -562,7 +562,7 @@ 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) { var res *schema.TranscriptionResult - err := m.stageCall(ctx, "transcription", m.TranscriptionConfig, func(cfg *config.ModelConfig, commit func()) error { + err := m.stageCall(ctx, config.PipelineStageTranscription, 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() @@ -577,7 +577,7 @@ func (m *wrappedModel) TranscribeLive(ctx context.Context, language string, onEv 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 { + err := m.stageCall(ctx, config.PipelineStageTranscription, 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 @@ -830,11 +830,11 @@ func (m *wrappedModel) FillToolArguments(ctx context.Context, messages schema.Me func (m *wrappedModel) Warmup(ctx context.Context) error { stages := []backend.PreloadStage{ - {Role: "vad", Cfg: m.VADConfig}, - {Role: "transcription", Cfg: m.TranscriptionConfig}, - {Role: "llm", Cfg: m.LLMConfig}, - {Role: "tts", Cfg: m.TTSConfig}, - {Role: "sound_detection", Cfg: m.SoundDetectionConfig}, + {Role: config.PipelineStageVAD, Cfg: m.VADConfig}, + {Role: config.PipelineStageTranscription, Cfg: m.TranscriptionConfig}, + {Role: config.PipelineStageLLM, Cfg: m.LLMConfig}, + {Role: config.PipelineStageTTS, Cfg: m.TTSConfig}, + {Role: config.PipelineStageSoundDetection, Cfg: m.SoundDetectionConfig}, } // The scoring model is a separate stage only when it isn't the LLM. if m.ScoreConfig != nil && m.ScoreConfig != m.LLMConfig { @@ -937,7 +937,7 @@ func newTranscriptionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfig 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) + cfgVAD, err = sr.resolveStage(config.PipelineStageVAD, cfgVAD) } if err != nil { @@ -950,7 +950,7 @@ 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) + cfgSST, err = sr.resolveStage(config.PipelineStageTranscription, cfgSST) } if err != nil { @@ -963,7 +963,7 @@ 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) + cfgSound, err = sr.resolveStage(config.PipelineStageSoundDetection, cfgSound) } if err != nil { return nil, nil, err @@ -990,7 +990,7 @@ func newSoundDetectionOnlyModel(pipeline *config.Pipeline, cl *config.ModelConfi sr := newStageRouter(fm, cl, ml, appConfig) cfgSound, err := loadSoundDetectionConfig(pipeline, cl, ml, appConfig) if err == nil { - cfgSound, err = sr.resolveStage("sound_detection", cfgSound) + cfgSound, err = sr.resolveStage(config.PipelineStageSoundDetection, cfgSound) } if err != nil { return nil, err @@ -1059,7 +1059,7 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model cfgVAD, err := cl.LoadResolvedModelConfig(pipeline.VAD, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) if err == nil { - cfgVAD, err = resolveStage("vad", cfgVAD) + cfgVAD, err = resolveStage(config.PipelineStageVAD, cfgVAD) } if err != nil { @@ -1073,7 +1073,7 @@ 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) + cfgSST, err = resolveStage(config.PipelineStageTranscription, cfgSST) } if err != nil { @@ -1108,7 +1108,7 @@ 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) + cfgLLM, err = resolveStage(config.PipelineStageLLM, cfgLLM) } if err != nil { @@ -1130,7 +1130,7 @@ func newModel(pipeline *config.Pipeline, cl *config.ModelConfigLoader, ml *model cfgTTS, err := cl.LoadResolvedModelConfig(pipeline.TTS, ml.ModelPath, appConfig.ToConfigLoaderOptions()...) if err == nil { - cfgTTS, err = resolveStage("tts", cfgTTS) + cfgTTS, err = resolveStage(config.PipelineStageTTS, cfgTTS) } if err != nil { @@ -1143,7 +1143,7 @@ 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) + cfgSound, err = resolveStage(config.PipelineStageSoundDetection, cfgSound) } if err != nil { return nil, err