mirror of
https://github.com/mudler/LocalAI.git
synced 2026-09-29 01:25:03 -04:00
feat(failover): switch realtime pipeline stages per call
A stage that names a chain is resolved on every call, so a switch keeps the session and its conversation. Clients get localai.model.failover events at session start and on every switch. Assisted-by: Claude:claude-opus-5-5 Signed-off-by: Ettore Di Giacinto <mudler@localai.io>
This commit is contained in:
1 parent
3e8ca016ee
commit
ea55e2ffc3
14 files changed
+603
-35
No files matched your search
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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()
|
||||
})
|
||||
})
|
||||
@@ -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
|
||||
}
|
||||
@@ -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"))
|
||||
|
||||
@@ -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())
|
||||
})
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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"
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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())
|
||||
|
||||
Reference in new issue
Block a user