diff options
| author | Paul Buetow <paul@buetow.org> | 2026-04-01 13:40:50 +0300 |
|---|---|---|
| committer | Paul Buetow <paul@buetow.org> | 2026-04-01 13:40:50 +0300 |
| commit | be758529bd22fd1f0d43a8b9f8a55197db0f1f60 (patch) | |
| tree | 603b07c6d80b1b71e8c726c5a33046ad6a3b3510 /internal/audio | |
| parent | 7435d08d70107d20e13a0a725c9e8645ac3f7c57 (diff) | |
zf: add Gemini TTS provider
Diffstat (limited to 'internal/audio')
| -rw-r--r-- | internal/audio/doc.go | 4 | ||||
| -rw-r--r-- | internal/audio/gemini_provider.go | 293 | ||||
| -rw-r--r-- | internal/audio/gemini_provider_test.go | 143 | ||||
| -rw-r--r-- | internal/audio/provider.go | 9 |
4 files changed, 446 insertions, 3 deletions
diff --git a/internal/audio/doc.go b/internal/audio/doc.go index c8a5ce4..ba12f56 100644 --- a/internal/audio/doc.go +++ b/internal/audio/doc.go @@ -1,3 +1,3 @@ -// Package audio provides audio generation functionality using OpenAI TTS -// for Bulgarian text-to-speech conversion. +// Package audio provides audio generation functionality using OpenAI and +// Gemini TTS for Bulgarian text-to-speech conversion. package audio diff --git a/internal/audio/gemini_provider.go b/internal/audio/gemini_provider.go new file mode 100644 index 0000000..23c4f4b --- /dev/null +++ b/internal/audio/gemini_provider.go @@ -0,0 +1,293 @@ +package audio + +import ( + "bytes" + "context" + "encoding/binary" + "errors" + "fmt" + "os" + "path/filepath" + "strings" + + "google.golang.org/genai" +) + +const ( + defaultGeminiTTSModel = "gemini-2.5-flash" + geminiTTSLanguageCode = "bg" + geminiTTSChannels = 1 + geminiTTSSampleRate = 24000 + geminiTTSBitsPerSample = 16 +) + +// GeminiProvider implements Provider interface for Gemini TTS. +type GeminiProvider struct { + client *genai.Client + config *Config +} + +var _ Provider = (*GeminiProvider)(nil) + +// NewGeminiProvider creates a new Gemini TTS provider. +func NewGeminiProvider(config *Config) (Provider, error) { + normalized := normalizeGeminiConfig(config) + if normalized.GoogleAPIKey == "" { + return nil, errors.New("Google API key is required") + } + + client, err := genai.NewClient(context.Background(), &genai.ClientConfig{ + APIKey: normalized.GoogleAPIKey, + Backend: genai.BackendGeminiAPI, + }) + if err != nil { + return nil, fmt.Errorf("failed to create Gemini client: %w", err) + } + + return &GeminiProvider{ + client: client, + config: normalized, + }, nil +} + +// GenerateAudio generates audio using Gemini TTS and writes it to the output file. +func (p *GeminiProvider) GenerateAudio(ctx context.Context, text string, outputFile string) error { + if err := ValidateBulgarianText(text); err != nil { + return err + } + if p == nil || p.client == nil || p.config == nil { + return errors.New("Gemini client not initialized") + } + + prompt := p.buildPrompt(text) + req := &genai.GenerateContentConfig{ + ResponseModalities: []string{string(genai.ModalityAudio)}, + SpeechConfig: p.speechConfig(), + } + + response, err := p.client.Models.GenerateContent(ctx, p.config.GeminiTTSModel, []*genai.Content{ + genai.NewContentFromText(prompt, genai.RoleUser), + }, req) + if err != nil { + return fmt.Errorf("Gemini API error: %w", err) + } + + audioData, mimeType, err := extractAudioData(response) + if err != nil { + return err + } + + if err := writeGeminiAudioFile(outputFile, audioData, mimeType); err != nil { + return err + } + + return nil +} + +// Name returns the provider name. +func (p *GeminiProvider) Name() string { + return "gemini" +} + +// IsAvailable checks if the Google API key is configured. +func (p *GeminiProvider) IsAvailable() error { + if p == nil || p.config == nil || strings.TrimSpace(p.config.GoogleAPIKey) == "" { + return errors.New("Google API key not configured") + } + + return nil +} + +func (p *GeminiProvider) buildPrompt(text string) string { + config := p.config + if config == nil { + config = &Config{} + } + + var prompt strings.Builder + prompt.WriteString("You are speaking Bulgarian language (български език). ") + prompt.WriteString("Pronounce the Bulgarian text with authentic Bulgarian phonetics, not Russian.") + + if speedHint := geminiSpeedHint(config.GeminiSpeed); speedHint != "" { + prompt.WriteString(" ") + prompt.WriteString(speedHint) + } + + prompt.WriteString("\n\nSpeak the following Bulgarian text:\n") + prompt.WriteString(strings.TrimSpace(text)) + + if voice := strings.TrimSpace(config.GeminiVoice); voice != "" { + prompt.WriteString("\n\nUse a clear, natural delivery that matches the voice named ") + prompt.WriteString(voice) + prompt.WriteString(".") + } + + return prompt.String() +} + +func (p *GeminiProvider) speechConfig() *genai.SpeechConfig { + config := p.config + if config == nil { + config = &Config{} + } + + speechConfig := &genai.SpeechConfig{ + LanguageCode: geminiTTSLanguageCode, + } + + if voice := strings.TrimSpace(config.GeminiVoice); voice != "" { + speechConfig.VoiceConfig = &genai.VoiceConfig{ + PrebuiltVoiceConfig: &genai.PrebuiltVoiceConfig{ + VoiceName: voice, + }, + } + } + + return speechConfig +} + +func normalizeGeminiConfig(config *Config) *Config { + normalized := &Config{} + if config != nil { + *normalized = *config + } + + normalized.GoogleAPIKey = strings.TrimSpace(normalized.GoogleAPIKey) + normalized.GeminiTTSModel = strings.TrimSpace(normalized.GeminiTTSModel) + normalized.GeminiVoice = strings.TrimSpace(normalized.GeminiVoice) + + if normalized.GeminiTTSModel == "" { + normalized.GeminiTTSModel = defaultGeminiTTSModel + } + if normalized.GeminiSpeed <= 0 { + normalized.GeminiSpeed = 1.0 + } + + return normalized +} + +func geminiSpeedHint(speed float64) string { + switch { + case speed < 0.95: + return "Speak slowly and clearly for language learners." + case speed > 1.05: + return "Speak slightly faster than normal while staying clear." + default: + return "Speak at a natural pace." + } +} + +func extractAudioData(response *genai.GenerateContentResponse) ([]byte, string, error) { + if response == nil { + return nil, "", errors.New("no response from Gemini") + } + + for _, candidate := range response.Candidates { + if candidate == nil || candidate.Content == nil { + continue + } + + for _, part := range candidate.Content.Parts { + if part == nil || part.InlineData == nil || len(part.InlineData.Data) == 0 { + continue + } + + audio := append([]byte(nil), part.InlineData.Data...) + return audio, part.InlineData.MIMEType, nil + } + } + + return nil, "", errors.New("no audio data returned from Gemini") +} + +func writeGeminiAudioFile(outputFile string, audioData []byte, mimeType string) error { + if err := ensureOutputDirectory(outputFile); err != nil { + return err + } + + 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 err := os.WriteFile(outputFile, encoded, 0644); err != nil { + return fmt.Errorf("failed to write output file: %w", err) + } + return nil + } + + if err := os.WriteFile(outputFile, audioData, 0644); err != nil { + return fmt.Errorf("failed to write output file: %w", err) + } + + return nil +} + +func ensureOutputDirectory(outputFile string) error { + dir := filepath.Dir(outputFile) + if dir == "" || dir == "." { + return nil + } + + if err := os.MkdirAll(dir, 0755); err != nil { + return fmt.Errorf("failed to create output directory: %w", err) + } + + return nil +} + +func encodePCMAsWAV(pcmData []byte) ([]byte, error) { + var buffer bytes.Buffer + + if _, err := buffer.WriteString("RIFF"); err != nil { + return nil, err + } + if err := binary.Write(&buffer, binary.LittleEndian, uint32(36+len(pcmData))); err != nil { + return nil, err + } + if _, err := buffer.WriteString("WAVE"); err != nil { + return nil, err + } + if _, err := buffer.WriteString("fmt "); err != nil { + return nil, err + } + if err := binary.Write(&buffer, binary.LittleEndian, uint32(16)); err != nil { + return nil, err + } + if err := binary.Write(&buffer, binary.LittleEndian, uint16(1)); err != nil { + return nil, err + } + if err := binary.Write(&buffer, binary.LittleEndian, uint16(geminiTTSChannels)); err != nil { + return nil, err + } + if err := binary.Write(&buffer, binary.LittleEndian, uint32(geminiTTSSampleRate)); err != nil { + return nil, err + } + + byteRate := uint32(geminiTTSSampleRate * geminiTTSChannels * geminiTTSBitsPerSample / 8) + if err := binary.Write(&buffer, binary.LittleEndian, byteRate); err != nil { + return nil, err + } + + blockAlign := uint16(geminiTTSChannels * geminiTTSBitsPerSample / 8) + if err := binary.Write(&buffer, binary.LittleEndian, blockAlign); err != nil { + return nil, err + } + if err := binary.Write(&buffer, binary.LittleEndian, uint16(geminiTTSBitsPerSample)); err != nil { + return nil, err + } + if _, err := buffer.WriteString("data"); err != nil { + return nil, err + } + if err := binary.Write(&buffer, binary.LittleEndian, uint32(len(pcmData))); err != nil { + return nil, err + } + if _, err := buffer.Write(pcmData); err != nil { + return nil, err + } + + return buffer.Bytes(), nil +} diff --git a/internal/audio/gemini_provider_test.go b/internal/audio/gemini_provider_test.go new file mode 100644 index 0000000..1c04f55 --- /dev/null +++ b/internal/audio/gemini_provider_test.go @@ -0,0 +1,143 @@ +package audio + +import ( + "os" + "path/filepath" + "strings" + "testing" + + "google.golang.org/genai" +) + +func TestNewGeminiProvider(t *testing.T) { + tests := []struct { + name string + config *Config + wantErr bool + }{ + { + name: "missing google api key", + config: &Config{}, + wantErr: true, + }, + { + name: "valid config", + config: &Config{ + GoogleAPIKey: "test-key", + GeminiTTSModel: "gemini-2.5-flash", + GeminiVoice: "Kore", + GeminiSpeed: 1.0, + }, + wantErr: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + provider, err := NewGeminiProvider(tt.config) + if (err != nil) != tt.wantErr { + t.Fatalf("NewGeminiProvider() error = %v, wantErr %v", err, tt.wantErr) + } + if err != nil { + return + } + + if provider.Name() != "gemini" { + t.Fatalf("Name() = %q, want %q", provider.Name(), "gemini") + } + }) + } +} + +func TestGeminiProviderIsAvailable(t *testing.T) { + provider := &GeminiProvider{ + config: &Config{GoogleAPIKey: "test-key"}, + } + + if err := provider.IsAvailable(); err != nil { + t.Fatalf("IsAvailable() unexpected error: %v", err) + } + + provider.config.GoogleAPIKey = "" + if err := provider.IsAvailable(); err == nil { + t.Fatal("IsAvailable() expected error when API key is missing") + } +} + +func TestGeminiProviderBuildPrompt(t *testing.T) { + provider := &GeminiProvider{ + config: &Config{ + GeminiVoice: "Kore", + GeminiSpeed: 0.92, + }, + } + + prompt := provider.buildPrompt("ябълка") + + for _, want := range []string{ + "Bulgarian language", + "authentic Bulgarian phonetics", + "Speak slowly and clearly for language learners.", + "ябълка", + "voice named Kore", + } { + if !strings.Contains(prompt, want) { + t.Fatalf("buildPrompt() = %q, missing %q", prompt, want) + } + } +} + +func TestExtractAudioData(t *testing.T) { + response := &genai.GenerateContentResponse{ + Candidates: []*genai.Candidate{ + { + Content: &genai.Content{ + Parts: []*genai.Part{ + { + InlineData: &genai.Blob{ + Data: []byte{0x01, 0x02, 0x03}, + MIMEType: "audio/pcm", + }, + }, + }, + }, + }, + }, + } + + data, mimeType, err := extractAudioData(response) + if err != nil { + t.Fatalf("extractAudioData() unexpected error: %v", err) + } + + if mimeType != "audio/pcm" { + t.Fatalf("extractAudioData() mimeType = %q, want %q", mimeType, "audio/pcm") + } + + if len(data) != 3 || data[0] != 0x01 || data[1] != 0x02 || data[2] != 0x03 { + t.Fatalf("extractAudioData() data = %v, want raw audio bytes", data) + } +} + +func TestWriteGeminiAudioFileWritesWAV(t *testing.T) { + dir := t.TempDir() + outputFile := filepath.Join(dir, "output.wav") + pcmData := []byte{0x11, 0x22, 0x33, 0x44} + + if err := writeGeminiAudioFile(outputFile, pcmData, "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 !strings.HasPrefix(string(fileData[:4]), "RIFF") { + t.Fatalf("output file does not look like WAV data: %q", fileData[:4]) + } + + if got, want := len(fileData), 44+len(pcmData); got != want { + t.Fatalf("len(output) = %d, want %d", got, want) + } +} diff --git a/internal/audio/provider.go b/internal/audio/provider.go index 06894bf..a4379f1 100644 --- a/internal/audio/provider.go +++ b/internal/audio/provider.go @@ -19,7 +19,7 @@ type Provider interface { // Config holds common configuration for audio providers type Config struct { - Provider string // Provider name: "openai" + Provider string // Provider name: "openai" or "gemini" OutputDir string // Directory for output files OutputFormat string // Output format: "mp3" or "wav" @@ -30,6 +30,11 @@ type Config struct { OpenAISpeed float64 // 0.25 to 4.0 OpenAIInstruction string // Voice instructions for gpt-4o-mini-tts model + // Gemini-specific settings + GoogleAPIKey string + GeminiTTSModel string // "gemini-2.5-flash" + GeminiVoice string // Prebuilt Gemini TTS voice name, or empty for the model default + GeminiSpeed float64 // Prompt hint for desired speech speed } // DefaultConfig returns default configuration @@ -43,6 +48,8 @@ func DefaultProviderConfig() *Config { OpenAISpeed: 1.0, // OpenAISpeed: 0.98, // Default speed for clarity OpenAIInstruction: "You are speaking Bulgarian language (български език). Pronounce the Bulgarian text with authentic Bulgarian phonetics, not Russian. Speak slowly and clearly for language learners.", + GeminiTTSModel: "gemini-2.5-flash", + GeminiSpeed: 1.0, } } |
