diff options
Diffstat (limited to 'internal/audio')
| -rw-r--r-- | internal/audio/gemini_provider.go | 63 | ||||
| -rw-r--r-- | internal/audio/gemini_provider_test.go | 40 | ||||
| -rw-r--r-- | internal/audio/provider.go | 4 | ||||
| -rw-r--r-- | internal/audio/provider_test.go | 8 | ||||
| -rw-r--r-- | internal/audio/voices.go | 26 | ||||
| -rw-r--r-- | internal/audio/voices_test.go | 23 |
6 files changed, 146 insertions, 18 deletions
diff --git a/internal/audio/gemini_provider.go b/internal/audio/gemini_provider.go index 5662454..2893278 100644 --- a/internal/audio/gemini_provider.go +++ b/internal/audio/gemini_provider.go @@ -7,6 +7,7 @@ import ( "errors" "fmt" "os" + "os/exec" "path/filepath" "strings" @@ -21,6 +22,11 @@ const ( geminiTTSBitsPerSample = 16 ) +var ErrGeminiNoAudioData = errors.New("no audio data returned from Gemini") + +var execLookPath = exec.LookPath +var execCommand = exec.Command + // GeminiProvider implements Provider interface for Gemini TTS. type GeminiProvider struct { client *genai.Client @@ -173,7 +179,12 @@ func extractAudioData(response *genai.GenerateContentResponse) ([]byte, string, } } - return nil, "", errors.New("no audio data returned from Gemini") + return nil, "", ErrGeminiNoAudioData +} + +// IsGeminiNoAudioDataError reports whether the error means Gemini returned no audio payload. +func IsGeminiNoAudioDataError(err error) bool { + return errors.Is(err, ErrGeminiNoAudioData) } func writeGeminiAudioFile(outputFile string, audioData []byte, mimeType string) error { @@ -182,20 +193,22 @@ func writeGeminiAudioFile(outputFile string, audioData []byte, mimeType string) } ext := strings.ToLower(filepath.Ext(outputFile)) - if ext != ".wav" { - return fmt.Errorf("gemini TTS only supports .wav output files, got %q", outputFile) - } - encoded, err := encodePCMAsWAV(audioData) if err != nil { return err } - if err := os.WriteFile(outputFile, encoded, 0644); err != nil { - return fmt.Errorf("failed to write output file: %w", err) + switch ext { + case ".wav": + if err := os.WriteFile(outputFile, encoded, 0644); err != nil { + return fmt.Errorf("failed to write output file: %w", err) + } + return nil + case ".mp3": + return transcodeWAVToMP3(encoded, outputFile) + default: + return fmt.Errorf("gemini TTS only supports .wav and .mp3 output files, got %q", outputFile) } - - return nil } func ensureOutputDirectory(outputFile string) error { @@ -263,3 +276,35 @@ func encodePCMAsWAV(pcmData []byte) ([]byte, error) { return buffer.Bytes(), nil } + +func transcodeWAVToMP3(wavData []byte, outputFile string) error { + ffmpegPath, err := execLookPath("ffmpeg") + if err != nil { + return fmt.Errorf("ffmpeg is required to convert Gemini audio to mp3: %w", err) + } + + cmd := execCommand( + ffmpegPath, + "-nostdin", + "-hide_banner", + "-loglevel", "error", + "-y", + "-f", "wav", + "-i", "pipe:0", + "-codec:a", "libmp3lame", + "-q:a", "4", + outputFile, + ) + cmd.Stdin = bytes.NewReader(wavData) + + output, err := cmd.CombinedOutput() + if err != nil { + message := strings.TrimSpace(string(output)) + if message == "" { + message = err.Error() + } + return fmt.Errorf("failed to convert Gemini audio to mp3: %s", message) + } + + return nil +} diff --git a/internal/audio/gemini_provider_test.go b/internal/audio/gemini_provider_test.go index ff0b2f5..66f0a74 100644 --- a/internal/audio/gemini_provider_test.go +++ b/internal/audio/gemini_provider_test.go @@ -219,16 +219,50 @@ func TestWriteGeminiAudioFileWritesWAV(t *testing.T) { } } -func TestWriteGeminiAudioFileRejectsUnsupportedFormats(t *testing.T) { +func TestWriteGeminiAudioFileWritesMP3ViaFFmpeg(t *testing.T) { dir := t.TempDir() outputFile := filepath.Join(dir, "output.mp3") + ffmpegScript := filepath.Join(dir, "ffmpeg") + script := "#!/bin/sh\nout=\"\"\nfor arg in \"$@\"; do out=\"$arg\"; done\ncat >/dev/null\nprintf 'mp3' > \"$out\"\n" + if err := os.WriteFile(ffmpegScript, []byte(script), 0755); err != nil { + t.Fatalf("failed to write fake ffmpeg script: %v", err) + } + + originalLookPath := execLookPath + execLookPath = func(file string) (string, error) { + if file == "ffmpeg" { + return ffmpegScript, nil + } + return originalLookPath(file) + } + t.Cleanup(func() { + execLookPath = originalLookPath + }) + + if err := writeGeminiAudioFile(outputFile, []byte{0x11, 0x22}, "audio/pcm"); err != nil { + t.Fatalf("writeGeminiAudioFile() unexpected error: %v", err) + } + + fileData, err := os.ReadFile(outputFile) + if err != nil { + t.Fatalf("ReadFile() unexpected error: %v", err) + } + if string(fileData) != "mp3" { + t.Fatalf("output file = %q, want fake mp3 payload", string(fileData)) + } +} + +func TestWriteGeminiAudioFileRejectsUnsupportedFormats(t *testing.T) { + dir := t.TempDir() + outputFile := filepath.Join(dir, "output.flac") + err := writeGeminiAudioFile(outputFile, []byte{0x11, 0x22}, "audio/pcm") if err == nil { - t.Fatal("writeGeminiAudioFile() expected error for non-wav output") + t.Fatal("writeGeminiAudioFile() expected error for unsupported output") } - if !strings.Contains(err.Error(), "only supports .wav output files") { + if !strings.Contains(err.Error(), "only supports .wav and .mp3 output files") { t.Fatalf("writeGeminiAudioFile() error = %v, want unsupported-format message", err) } diff --git a/internal/audio/provider.go b/internal/audio/provider.go index b7f6bd9..4fdccfb 100644 --- a/internal/audio/provider.go +++ b/internal/audio/provider.go @@ -33,7 +33,7 @@ type Config struct { // Gemini-specific settings GoogleAPIKey string GeminiTTSModel string // "gemini-2.5-flash-preview-tts" - GeminiVoice string // One of GeminiVoices, or empty for the model default. + GeminiVoice string // One of GeminiVoices; empty lets the caller choose a random voice. GeminiSpeed float64 // Prompt hint for desired speech speed } @@ -42,7 +42,7 @@ func DefaultProviderConfig() *Config { return &Config{ Provider: "gemini", OutputDir: "./", - OutputFormat: "wav", + OutputFormat: "mp3", OpenAIModel: "gpt-4o-mini-tts", // New model with voice instructions support OpenAIVoice: "alloy", OpenAISpeed: 1.0, diff --git a/internal/audio/provider_test.go b/internal/audio/provider_test.go index c13e823..a08b7a6 100644 --- a/internal/audio/provider_test.go +++ b/internal/audio/provider_test.go @@ -36,8 +36,8 @@ func TestDefaultProviderConfig(t *testing.T) { t.Errorf("Expected provider 'gemini', got '%s'", config.Provider) } - if config.OutputFormat != "wav" { - t.Errorf("Expected output format 'wav', got '%s'", config.OutputFormat) + if config.OutputFormat != "mp3" { + t.Errorf("Expected output format 'mp3', got '%s'", config.OutputFormat) } if config.OpenAIModel != "gpt-4o-mini-tts" { @@ -73,8 +73,8 @@ func TestDefaultProviderConfigIsGeminiCompatible(t *testing.T) { } outputFile := filepath.Join(t.TempDir(), "audio."+config.OutputFormat) - if filepath.Ext(outputFile) != ".wav" { - t.Fatalf("DefaultProviderConfig() output file %q is incompatible with Gemini TTS", outputFile) + if filepath.Ext(outputFile) != ".mp3" { + t.Fatalf("DefaultProviderConfig() output file %q does not use the default mp3 extension", outputFile) } if !strings.HasSuffix(config.GeminiTTSModel, "-tts") { diff --git a/internal/audio/voices.go b/internal/audio/voices.go index 9a96c76..2f6b5aa 100644 --- a/internal/audio/voices.go +++ b/internal/audio/voices.go @@ -1,5 +1,7 @@ package audio +import "strings" + // OpenAIVoices lists the OpenAI voices supported by the app. var OpenAIVoices = []string{ "alloy", @@ -44,3 +46,27 @@ var GeminiVoices = []string{ "Vindemiatrix", "Zubenelgenubi", } + +// GeminiVoiceFallbacks returns the selected voice first, followed by the remaining known Gemini voices. +func GeminiVoiceFallbacks(selected string) []string { + selected = strings.TrimSpace(selected) + if selected == "" { + return append([]string(nil), GeminiVoices...) + } + + fallbacks := []string{selected} + seen := map[string]struct{}{selected: {}} + for _, voice := range GeminiVoices { + voice = strings.TrimSpace(voice) + if voice == "" { + continue + } + if _, ok := seen[voice]; ok { + continue + } + fallbacks = append(fallbacks, voice) + seen[voice] = struct{}{} + } + + return fallbacks +} diff --git a/internal/audio/voices_test.go b/internal/audio/voices_test.go index 121f328..ea0f797 100644 --- a/internal/audio/voices_test.go +++ b/internal/audio/voices_test.go @@ -34,3 +34,26 @@ func TestVoiceLists(t *testing.T) { }) } } + +func TestGeminiVoiceFallbacks(t *testing.T) { + t.Parallel() + + t.Run("selected voice comes first", func(t *testing.T) { + t.Parallel() + + got := GeminiVoiceFallbacks("Kore") + wantPrefix := []string{"Kore", "Zephyr", "Puck", "Charon"} + if !reflect.DeepEqual(got[:len(wantPrefix)], wantPrefix) { + t.Fatalf("GeminiVoiceFallbacks() prefix mismatch\nwant: %#v\ngot: %#v", wantPrefix, got[:len(wantPrefix)]) + } + }) + + t.Run("empty selection returns known voices", func(t *testing.T) { + t.Parallel() + + got := GeminiVoiceFallbacks("") + if !reflect.DeepEqual(got, GeminiVoices) { + t.Fatalf("GeminiVoiceFallbacks() mismatch\nwant: %#v\ngot: %#v", GeminiVoices, got) + } + }) +} |
