From 155b3e78e184b8f1c2193d875d4db8a56ea36d34 Mon Sep 17 00:00:00 2001 From: Paul Buetow Date: Wed, 1 Apr 2026 13:45:50 +0300 Subject: zf: fix Gemini provider routing and output validation --- internal/audio/gemini_provider.go | 18 +++++++----------- internal/audio/gemini_provider_test.go | 18 ++++++++++++++++++ internal/audio/provider.go | 5 +++++ internal/audio/provider_test.go | 28 +++++++++++++++++++++++----- 4 files changed, 53 insertions(+), 16 deletions(-) diff --git a/internal/audio/gemini_provider.go b/internal/audio/gemini_provider.go index 23c4f4b..8d0c517 100644 --- a/internal/audio/gemini_provider.go +++ b/internal/audio/gemini_provider.go @@ -206,20 +206,16 @@ func writeGeminiAudioFile(outputFile string, audioData []byte, mimeType string) } ext := strings.ToLower(filepath.Ext(outputFile)) - mimeType = strings.ToLower(mimeType) - if ext == ".wav" || (ext == "" && (mimeType == "" || strings.Contains(mimeType, "pcm"))) { - encoded, err := encodePCMAsWAV(audioData) - if err != nil { - return err - } + if ext != ".wav" { + return fmt.Errorf("Gemini TTS only supports .wav output files, got %q", outputFile) + } - if err := os.WriteFile(outputFile, encoded, 0644); err != nil { - return fmt.Errorf("failed to write output file: %w", err) - } - return nil + encoded, err := encodePCMAsWAV(audioData) + if err != nil { + return err } - if err := os.WriteFile(outputFile, audioData, 0644); err != nil { + if err := os.WriteFile(outputFile, encoded, 0644); err != nil { return fmt.Errorf("failed to write output file: %w", err) } diff --git a/internal/audio/gemini_provider_test.go b/internal/audio/gemini_provider_test.go index 1c04f55..b8ab4c9 100644 --- a/internal/audio/gemini_provider_test.go +++ b/internal/audio/gemini_provider_test.go @@ -141,3 +141,21 @@ func TestWriteGeminiAudioFileWritesWAV(t *testing.T) { t.Fatalf("len(output) = %d, want %d", got, want) } } + +func TestWriteGeminiAudioFileRejectsUnsupportedFormats(t *testing.T) { + dir := t.TempDir() + outputFile := filepath.Join(dir, "output.mp3") + + err := writeGeminiAudioFile(outputFile, []byte{0x11, 0x22}, "audio/pcm") + if err == nil { + t.Fatal("writeGeminiAudioFile() expected error for non-wav output") + } + + if !strings.Contains(err.Error(), "only supports .wav output files") { + t.Fatalf("writeGeminiAudioFile() error = %v, want unsupported-format message", err) + } + + if _, statErr := os.Stat(outputFile); !os.IsNotExist(statErr) { + t.Fatalf("expected no output file to be written, statErr=%v", statErr) + } +} diff --git a/internal/audio/provider.go b/internal/audio/provider.go index a4379f1..7ca4282 100644 --- a/internal/audio/provider.go +++ b/internal/audio/provider.go @@ -65,6 +65,11 @@ func NewProvider(config *Config) (Provider, error) { return nil, fmt.Errorf("OpenAI API key is required") } return NewOpenAIProvider(config) + case "gemini": + if config.GoogleAPIKey == "" { + return nil, fmt.Errorf("Google API key is required") + } + return NewGeminiProvider(config) default: return nil, fmt.Errorf("unknown audio provider: %s", config.Provider) diff --git a/internal/audio/provider_test.go b/internal/audio/provider_test.go index 97d3e9c..24d2729 100644 --- a/internal/audio/provider_test.go +++ b/internal/audio/provider_test.go @@ -54,10 +54,11 @@ func TestDefaultProviderConfig(t *testing.T) { func TestNewProvider(t *testing.T) { tests := []struct { - name string - config *Config - wantErr bool - errMsg string + name string + config *Config + wantErr bool + errMsg string + wantProvider string }{ { name: "nil config uses defaults", @@ -81,17 +82,34 @@ func TestNewProvider(t *testing.T) { wantErr: true, errMsg: "unknown audio provider: unknown", }, + { + name: "gemini provider with key", + config: &Config{ + Provider: "gemini", + GoogleAPIKey: "test-google-key", + }, + wantErr: false, + wantProvider: "gemini", + }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - _, err := NewProvider(tt.config) + provider, err := NewProvider(tt.config) if (err != nil) != tt.wantErr { t.Errorf("NewProvider() error = %v, wantErr %v", err, tt.wantErr) } if tt.wantErr && err != nil && err.Error() != tt.errMsg { t.Errorf("NewProvider() error = %v, want %v", err.Error(), tt.errMsg) } + if !tt.wantErr && tt.wantProvider != "" { + if provider == nil { + t.Fatalf("NewProvider() returned nil provider") + } + if provider.Name() != tt.wantProvider { + t.Fatalf("NewProvider() Name() = %v, want %v", provider.Name(), tt.wantProvider) + } + } }) } } -- cgit v1.2.3