diff options
| author | Paul Buetow <paul@buetow.org> | 2026-04-01 13:45:50 +0300 |
|---|---|---|
| committer | Paul Buetow <paul@buetow.org> | 2026-04-01 13:45:50 +0300 |
| commit | 155b3e78e184b8f1c2193d875d4db8a56ea36d34 (patch) | |
| tree | a97dc41297c51ba029890cf7947008c9b863796e /internal/audio | |
| parent | be758529bd22fd1f0d43a8b9f8a55197db0f1f60 (diff) | |
zf: fix Gemini provider routing and output validation
Diffstat (limited to 'internal/audio')
| -rw-r--r-- | internal/audio/gemini_provider.go | 18 | ||||
| -rw-r--r-- | internal/audio/gemini_provider_test.go | 18 | ||||
| -rw-r--r-- | internal/audio/provider.go | 5 | ||||
| -rw-r--r-- | 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) + } + } }) } } |
