fix(audio): resolve provider options and honor STT timeout

This commit is contained in:
ntduc committed 2026-07-22 20:48:12 +07:00
1 parent b49c6abcc7
commit bb58e03352
5 files changed
+255 -65

No files matched your search

+44 -18
View File
@@ -41,9 +41,9 @@ type Manager struct {
musicProviders map[string]MusicProvider
sfxProviders map[string]SFXProvider
primary string // primary TTS provider
sttChain []string // STT fallback order (Phase 4)
musicChain []string // Music fallback order (Phase 3)
primary string // primary TTS provider
sttChain []string // STT fallback order (Phase 4)
musicChain []string // Music fallback order (Phase 3)
channelSTTOverrides map[string][]string // channel → provider key list (Phase 4)
auto AutoMode
@@ -298,33 +298,59 @@ func (m *Manager) SynthesizeWithFallback(ctx context.Context, text string, opts
// SynthesizeWithFallbackAdapted is like SynthesizeWithFallback but applies
// AdaptAgentParams(genericAgentParams, providerName) per-attempt before
// synthesizing. This is the Finding #1 fix: each fallback attempt receives
// provider-native params rather than the primary's adapted keys.
// synthesizing, so each fallback attempt receives provider-native params
// rather than the primary's adapted keys.
//
// genericAgentParams must use the generic allow-list keys (speed, emotion, style).
// Passing nil is safe and produces the same behaviour as SynthesizeWithFallback.
func (m *Manager) SynthesizeWithFallbackAdapted(ctx context.Context, text string, opts TTSOptions, genericAgentParams map[string]any) (*SynthResult, error) {
var providerErrs []error
if p, ok := m.ttsProviders[m.primary]; ok {
attemptOpts := m.withAdaptedParams(opts, m.primary, genericAgentParams)
if result, err := p.Synthesize(ctx, text, attemptOpts); err == nil {
return result, nil
} else {
slog.Warn("tts primary provider failed, trying fallback", "provider", m.primary, "error", err)
providerErrs = append(providerErrs, fmt.Errorf("%s: %w", m.primary, err))
return m.SynthesizeWithFallbackResolved(ctx, text, m.primary, func(string) TTSOptions {
return opts
}, genericAgentParams)
}
// SynthesizeWithFallbackResolved tries preferredProvider first, then the
// manager primary and remaining registered providers. resolveOpts runs for
// every actual attempt so provider-specific voice, model, and format defaults
// never leak into another provider's request.
func (m *Manager) SynthesizeWithFallbackResolved(
ctx context.Context,
text string,
preferredProvider string,
resolveOpts func(providerName string) TTSOptions,
genericAgentParams map[string]any,
) (*SynthResult, error) {
if preferredProvider == "" {
preferredProvider = m.primary
}
providerNames := make([]string, 0, len(m.ttsProviders))
if _, ok := m.ttsProviders[preferredProvider]; ok {
providerNames = append(providerNames, preferredProvider)
}
if m.primary != preferredProvider {
if _, ok := m.ttsProviders[m.primary]; ok {
providerNames = append(providerNames, m.primary)
}
}
for name, p := range m.ttsProviders {
if name == m.primary {
continue
for name := range m.ttsProviders {
if name != preferredProvider && name != m.primary {
providerNames = append(providerNames, name)
}
}
var providerErrs []error
for _, name := range providerNames {
opts := TTSOptions{}
if resolveOpts != nil {
opts = resolveOpts(name)
}
attemptOpts := m.withAdaptedParams(opts, name, genericAgentParams)
p := m.ttsProviders[name]
result, err := p.Synthesize(ctx, text, attemptOpts)
if err == nil {
slog.Info("tts fallback succeeded", "provider", name)
return result, nil
}
slog.Warn("tts fallback provider failed", "provider", name, "error", err)
slog.Warn("tts provider failed, trying fallback", "provider", name, "error", err)
providerErrs = append(providerErrs, fmt.Errorf("%s: %w", name, err))
}
if len(providerErrs) == 0 {
+5 -10
View File
@@ -16,15 +16,10 @@ import (
"github.com/nextlevelbuilder/goclaw/internal/audio"
)
const (
// sttMaxBytes matches the OpenAI upload limit, which compatible engines
// generally adopt. Enforced before the request so an oversized file fails
// locally instead of after a full upload.
sttMaxBytes = 25 << 20 // 25 MB
// sttDefaultTimeout is the floor for transcription, which is slower than
// synthesis and routinely exceeds the shared 30 s default.
sttDefaultTimeout = 120 * time.Second
)
// sttMaxBytes matches the OpenAI upload limit, which compatible engines
// generally adopt. Enforced before the request so an oversized file fails
// locally instead of after a full upload.
const sttMaxBytes = 25 << 20 // 25 MB
// STTProvider transcribes audio via an OpenAI-compatible
// POST /audio/transcriptions.
@@ -67,7 +62,7 @@ func (p *STTProvider) Transcribe(ctx context.Context, in audio.STTInput, opts au
req.Header.Set("Authorization", "Bearer "+p.cfg.APIKey)
}
timeout := sttDefaultTimeout
timeout := time.Duration(p.cfg.TimeoutMs) * time.Millisecond
if opts.TimeoutMs > 0 {
timeout = time.Duration(opts.TimeoutMs) * time.Millisecond
}
+95
View File
@@ -8,6 +8,7 @@ import (
"os"
"path/filepath"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -202,6 +203,100 @@ func TestSTTOptionsModelOverridesConfig(t *testing.T) {
assert.Equal(t, "opts-model", got.model)
}
func TestSTTConfiguredTimeoutAppliesWithoutPerCallOverride(t *testing.T) {
requestStarted := make(chan struct{})
releaseServer := make(chan struct{})
srv := httptest.NewServer(http.HandlerFunc(func(_ http.ResponseWriter, r *http.Request) {
close(requestStarted)
select {
case <-releaseServer:
case <-r.Context().Done():
}
}))
t.Cleanup(srv.Close)
t.Cleanup(func() { close(releaseServer) })
provider, err := openaicompat.NewSTTProvider(openaicompat.Config{
APIBase: srv.URL,
TimeoutMs: 25,
})
require.NoError(t, err)
ctx, cancel := context.WithCancel(context.Background())
t.Cleanup(cancel)
result := make(chan error, 1)
go func() {
_, err := provider.Transcribe(ctx,
audio.STTInput{Bytes: []byte("x"), MimeType: "audio/wav"},
audio.STTOptions{})
result <- err
}()
select {
case <-requestStarted:
case <-time.After(time.Second):
t.Fatal("transcription request did not reach the test server")
}
select {
case err := <-result:
require.Error(t, err)
assert.ErrorIs(t, err, context.DeadlineExceeded)
case <-time.After(500 * time.Millisecond):
cancel()
<-result
t.Fatal("Transcribe did not honor Config.TimeoutMs")
}
}
func TestSTTOptionsTimeoutOverridesConfig(t *testing.T) {
requestStarted := make(chan struct{})
releaseResponse := make(chan struct{})
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
close(requestStarted)
select {
case <-releaseResponse:
_, _ = io.WriteString(w, `{"text":"ok"}`)
case <-r.Context().Done():
}
}))
t.Cleanup(srv.Close)
provider, err := openaicompat.NewSTTProvider(openaicompat.Config{
APIBase: srv.URL,
TimeoutMs: 25,
})
require.NoError(t, err)
result := make(chan error, 1)
go func() {
_, err := provider.Transcribe(context.Background(),
audio.STTInput{Bytes: []byte("x"), MimeType: "audio/wav"},
audio.STTOptions{TimeoutMs: 1000})
result <- err
}()
select {
case <-requestStarted:
case <-time.After(time.Second):
t.Fatal("transcription request did not reach the test server")
}
select {
case err := <-result:
t.Fatalf("Transcribe used Config.TimeoutMs instead of opts.TimeoutMs: %v", err)
case <-time.After(100 * time.Millisecond):
}
close(releaseResponse)
select {
case err := <-result:
require.NoError(t, err)
case <-time.After(time.Second):
t.Fatal("Transcribe did not return after the test server responded")
}
}
// TestSTTLanguageFallsBackToHint covers a plain (non-verbose) json response,
// which carries no language field.
func TestSTTLanguageFallsBackToHint(t *testing.T) {
+17 -35
View File
@@ -235,8 +235,8 @@ func (t *TtsTool) resolveTenantProvider(ctx context.Context, mgr *tts.Manager, r
// resolveAgentGenericTTSParams reads the per-agent TTSParams generic map from
// the dispatcher-injected AgentAudioSnapshot. Returns nil when no snapshot
// is present or no tts_params are configured. The caller is responsible for
// calling audio.AdaptAgentParams(generic, providerName) PER-ATTEMPT to convert
// generic keys to provider-specific keys (Finding #1 CRITICAL).
// calling audio.AdaptAgentParams(generic, providerName) for each provider attempt
// so generic keys are converted to the correct provider-specific keys.
func (t *TtsTool) resolveAgentGenericTTSParams(ctx context.Context) map[string]any {
snap, ok := store.AgentAudioFromCtx(ctx)
if !ok || len(snap.OtherConfig) == 0 {
@@ -275,7 +275,7 @@ func (t *TtsTool) Execute(ctx context.Context, args map[string]any) *Result {
argModel, _ := args["model"].(string)
providerName, _ := args["provider"].(string)
// Read generic agent TTS params once; adapt PER-ATTEMPT below (Finding #1 CRITICAL).
// Read generic agent TTS params once; adapt them for each provider attempt below.
// Storing generic keys here so each fallback provider gets its own adapted copy.
genericAgentParams := t.resolveAgentGenericTTSParams(ctx)
@@ -284,17 +284,16 @@ func (t *TtsTool) Execute(ctx context.Context, args map[string]any) *Result {
mgr := t.manager
t.mu.RUnlock()
effectiveProvider := providerName
if effectiveProvider == "" {
effectiveProvider = t.resolvePrimary(ctx, mgr)
}
voice, model, voiceFromAgent := t.resolveVoiceAndModel(ctx, effectiveProvider, argVoice, argModel)
// Determine format based on channel (read from ctx — thread-safe)
channel := ToolChannelFromCtx(ctx)
opts := tts.Options{Voice: voice, Model: model}
if channel == "telegram" {
opts.Format = "opus"
resolveProviderOpts := func(name string) tts.Options {
voice, model, voiceFromAgent := t.resolveVoiceAndModel(ctx, name, argVoice, argModel)
opts := tts.Options{Voice: voice, Model: model}
if channel == "telegram" {
opts.Format = "opus"
}
opts.Voice = applyVoiceCompat(name, opts.Voice, voiceFromAgent)
return opts
}
var result *tts.SynthResult
@@ -312,7 +311,7 @@ func (t *TtsTool) Execute(ctx context.Context, args map[string]any) *Result {
}
providerName = tenantName
}
opts.Voice = applyVoiceCompat(providerName, opts.Voice, voiceFromAgent)
opts := resolveProviderOpts(providerName)
if adapted := audio.AdaptAgentParams(genericAgentParams, providerName); len(adapted) > 0 {
opts.Params = mergeParams(opts.Params, adapted)
}
@@ -320,37 +319,20 @@ func (t *TtsTool) Execute(ctx context.Context, args map[string]any) *Result {
} else {
// Prefer the DB-configured tenant provider used by /tts and auto-TTS.
if p, tenantName, ok := t.resolveTenantProvider(ctx, mgr, ""); ok {
tenantOpts := opts
tenantOpts.Voice = applyVoiceCompat(tenantName, tenantOpts.Voice, voiceFromAgent)
tenantOpts := resolveProviderOpts(tenantName)
if adapted := audio.AdaptAgentParams(genericAgentParams, tenantName); len(adapted) > 0 {
tenantOpts.Params = mergeParams(opts.Params, adapted)
tenantOpts.Params = mergeParams(tenantOpts.Params, adapted)
}
result, err = p.Synthesize(ctx, text, tenantOpts)
if err != nil {
slog.Warn("tts tenant provider failed, trying fallback", "provider", tenantName, "error", err)
result, err = mgr.SynthesizeWithFallbackAdapted(ctx, text, opts, genericAgentParams)
primary := t.resolvePrimary(ctx, mgr)
result, err = mgr.SynthesizeWithFallbackResolved(ctx, text, primary, resolveProviderOpts, genericAgentParams)
}
} else {
// Resolve primary from tenant settings or default.
primary := t.resolvePrimary(ctx, mgr)
if p, ok := mgr.GetProvider(primary); ok {
// Adapt for the primary provider attempt specifically.
primaryOpts := opts
primaryOpts.Voice = applyVoiceCompat(primary, primaryOpts.Voice, voiceFromAgent)
if adapted := audio.AdaptAgentParams(genericAgentParams, primary); len(adapted) > 0 {
primaryOpts.Params = mergeParams(opts.Params, adapted)
}
result, err = p.Synthesize(ctx, text, primaryOpts)
if err != nil {
slog.Warn("tts primary provider failed, trying fallback", "provider", primary, "error", err)
// SynthesizeWithFallbackAdapted adapts genericAgentParams per-attempt
// (Finding #1 CRITICAL): each fallback provider receives its own
// provider-native keys, not the primary's adapted map.
result, err = mgr.SynthesizeWithFallbackAdapted(ctx, text, opts, genericAgentParams)
}
} else {
result, err = mgr.SynthesizeWithFallbackAdapted(ctx, text, opts, genericAgentParams)
}
result, err = mgr.SynthesizeWithFallbackResolved(ctx, text, primary, resolveProviderOpts, genericAgentParams)
}
}
+94 -2
View File
@@ -14,18 +14,57 @@ type tenantRoutingTTSProvider struct {
name string
calls int
err error
voice string
model string
}
func (p *tenantRoutingTTSProvider) Name() string { return p.name }
func (p *tenantRoutingTTSProvider) Synthesize(_ context.Context, _ string, _ tts.Options) (*tts.SynthResult, error) {
func (p *tenantRoutingTTSProvider) Synthesize(_ context.Context, _ string, opts tts.Options) (*tts.SynthResult, error) {
p.calls++
p.voice = opts.Voice
p.model = opts.Model
if p.err != nil {
return nil, p.err
}
return &tts.SynthResult{Audio: []byte("audio"), Extension: "mp3", MimeType: "audio/mpeg"}, nil
}
func TestTtsTool_UsesSystemConfigForResolvedTenantProvider(t *testing.T) {
t.Parallel()
globalProvider := &tenantRoutingTTSProvider{name: "openai"}
tenantProvider := &tenantRoutingTTSProvider{name: "edge"}
mgr := newTenantRoutingManager("openai", globalProvider)
mgr.SetTenantResolver(func(context.Context) (audio.TTSProvider, string, audio.AutoMode, error) {
return tenantProvider, "edge", audio.AutoOff, nil
})
tool := NewTtsTool(mgr)
tool.SetSystemConfigStore(&fakeSystemConfigStore{data: map[string]string{
"tts.openai.voice": "alloy",
"tts.openai.model": "openai-model",
"tts.edge.voice": "vi-VN-HoaiMyNeural",
"tts.edge.model": "edge-model",
}})
result := tool.Execute(WithToolWorkspace(context.Background(), t.TempDir()), map[string]any{
"text": "hello from tenant provider",
})
if result.IsError {
t.Fatalf("unexpected error: %s", result.ForLLM)
}
if tenantProvider.voice != "vi-VN-HoaiMyNeural" {
t.Fatalf("tenant voice = %q, want %q", tenantProvider.voice, "vi-VN-HoaiMyNeural")
}
if tenantProvider.model != "edge-model" {
t.Fatalf("tenant model = %q, want %q", tenantProvider.model, "edge-model")
}
if globalProvider.calls != 0 {
t.Fatalf("global provider calls = %d, want 0", globalProvider.calls)
}
}
func newTenantRoutingManager(primary string, providers ...*tenantRoutingTTSProvider) *tts.Manager {
mgr := tts.NewManager(tts.ManagerConfig{Primary: primary})
for _, provider := range providers {
@@ -98,7 +137,13 @@ func TestTtsTool_TenantProviderFailureFallsBackWhenProviderOmitted(t *testing.T)
})
tool := NewTtsTool(mgr)
result := tool.Execute(context.Background(), map[string]any{
tool.SetSystemConfigStore(&fakeSystemConfigStore{data: map[string]string{
"tts.edge.voice": "vi-VN-HoaiMyNeural",
"tts.edge.model": "edge-model",
"tts.gemini.voice": "Kore",
"tts.gemini.model": "gemini-model",
}})
result := tool.Execute(WithToolWorkspace(context.Background(), t.TempDir()), map[string]any{
"text": "hello with tenant fallback",
})
@@ -111,6 +156,53 @@ func TestTtsTool_TenantProviderFailureFallsBackWhenProviderOmitted(t *testing.T)
if edgeProvider.calls != 1 {
t.Fatalf("global edge calls = %d, want 1", edgeProvider.calls)
}
if geminiProvider.voice != "Kore" || geminiProvider.model != "gemini-model" {
t.Fatalf("tenant options = voice %q model %q, want tenant system config", geminiProvider.voice, geminiProvider.model)
}
if edgeProvider.voice != "vi-VN-HoaiMyNeural" || edgeProvider.model != "edge-model" {
t.Fatalf("fallback options = voice %q model %q, want global system config", edgeProvider.voice, edgeProvider.model)
}
}
func TestTtsTool_FallbackResolvesOptionsForEveryActualProvider(t *testing.T) {
t.Parallel()
managerPrimary := &tenantRoutingTTSProvider{name: "edge"}
preferred := &tenantRoutingTTSProvider{name: "openai", err: errors.New("openai unavailable")}
tenantProvider := &tenantRoutingTTSProvider{name: "gemini", err: errors.New("tenant provider unavailable")}
mgr := newTenantRoutingManager("edge", managerPrimary, preferred)
mgr.SetTenantResolver(func(context.Context) (audio.TTSProvider, string, audio.AutoMode, error) {
return tenantProvider, "gemini", audio.AutoOff, nil
})
tool := NewTtsTool(mgr)
tool.SetSystemConfigStore(&fakeSystemConfigStore{data: map[string]string{
"tts.gemini.voice": "Kore",
"tts.gemini.model": "gemini-model",
"tts.openai.voice": "alloy",
"tts.openai.model": "openai-model",
"tts.edge.voice": "vi-VN-HoaiMyNeural",
"tts.edge.model": "edge-model",
}})
ctx := ctxWithTTSSettings(t, ttsOverride{Primary: "openai"})
ctx = WithToolWorkspace(ctx, t.TempDir())
result := tool.Execute(ctx, map[string]any{"text": "fallback across providers"})
if result.IsError {
t.Fatalf("unexpected error: %s", result.ForLLM)
}
if tenantProvider.voice != "Kore" || tenantProvider.model != "gemini-model" {
t.Fatalf("tenant options = voice %q model %q, want gemini config", tenantProvider.voice, tenantProvider.model)
}
if preferred.voice != "alloy" || preferred.model != "openai-model" {
t.Fatalf("preferred options = voice %q model %q, want openai config", preferred.voice, preferred.model)
}
if managerPrimary.voice != "vi-VN-HoaiMyNeural" || managerPrimary.model != "edge-model" {
t.Fatalf("secondary options = voice %q model %q, want edge config", managerPrimary.voice, managerPrimary.model)
}
if preferred.calls != 1 || managerPrimary.calls != 1 {
t.Fatalf("fallback calls = preferred %d secondary %d, want one each", preferred.calls, managerPrimary.calls)
}
}
func TestTtsTool_ExplicitProviderMismatchStillErrors(t *testing.T) {