summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorPaul Buetow <paul@buetow.org>2026-04-01 13:40:50 +0300
committerPaul Buetow <paul@buetow.org>2026-04-01 13:40:50 +0300
commitbe758529bd22fd1f0d43a8b9f8a55197db0f1f60 (patch)
tree603b07c6d80b1b71e8c726c5a33046ad6a3b3510
parent7435d08d70107d20e13a0a725c9e8645ac3f7c57 (diff)
zf: add Gemini TTS provider
-rw-r--r--internal/audio/doc.go4
-rw-r--r--internal/audio/gemini_provider.go293
-rw-r--r--internal/audio/gemini_provider_test.go143
-rw-r--r--internal/audio/provider.go9
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,
}
}