summaryrefslogtreecommitdiff
path: root/internal/audio
diff options
context:
space:
mode:
authorPaul Buetow <paul@buetow.org>2026-04-01 13:45:50 +0300
committerPaul Buetow <paul@buetow.org>2026-04-01 13:45:50 +0300
commit155b3e78e184b8f1c2193d875d4db8a56ea36d34 (patch)
treea97dc41297c51ba029890cf7947008c9b863796e /internal/audio
parentbe758529bd22fd1f0d43a8b9f8a55197db0f1f60 (diff)
zf: fix Gemini provider routing and output validation
Diffstat (limited to 'internal/audio')
-rw-r--r--internal/audio/gemini_provider.go18
-rw-r--r--internal/audio/gemini_provider_test.go18
-rw-r--r--internal/audio/provider.go5
-rw-r--r--internal/audio/provider_test.go28
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)
+ }
+ }
})
}
}