mirror of
https://github.com/mudler/LocalAI.git
synced 2026-09-29 17:44:30 -04:00
fix(failover): resolve chains in transcription and sound-only sessions
Transcription-only and sound-detection-only realtime sessions passed a chain config straight to the model loader. It has no backend, so the loader fell back to greedy backend auto-detection: slow, and ending in an unhelpful error. Sound-only sessions are a main use of chains. The stage routing of the full pipeline moves into a stageRouter that both realtime model kinds embed. Every stage resolves to the chain's active target at build time and goes through the failover plan per call. The session sends failover events for any model with chain stages, and restarts them when a transcription session.update swaps the model. Assisted-by: Claude:claude-opus-5-5
This commit is contained in:
1 parent
e183ea11d1
commit
1bd8e70173
6 files changed
+320
-95
No files matched your search
@@ -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
|
||||
}
|
||||
|
||||
@@ -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() {
|
||||
|
||||
@@ -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())
|
||||
})
|
||||
})
|
||||
@@ -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
|
||||
}
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in new issue
Block a user