diff options
| -rw-r--r-- | internal/audio/gemini_provider.go | 58 | ||||
| -rw-r--r-- | internal/audio/gemini_provider_test.go | 34 | ||||
| -rw-r--r-- | internal/audio/openai_provider.go | 50 | ||||
| -rw-r--r-- | internal/audio/openai_provider_test.go | 44 | ||||
| -rw-r--r-- | internal/audio/provider.go | 64 | ||||
| -rw-r--r-- | internal/audio/sidecar.go | 25 |
6 files changed, 145 insertions, 130 deletions
diff --git a/internal/audio/gemini_provider.go b/internal/audio/gemini_provider.go index 2893278..bda01f5 100644 --- a/internal/audio/gemini_provider.go +++ b/internal/audio/gemini_provider.go @@ -28,22 +28,24 @@ var execLookPath = exec.LookPath var execCommand = exec.Command // GeminiProvider implements Provider interface for Gemini TTS. +// It stores only the Gemini-specific sub-config so it never sees OpenAI fields. type GeminiProvider struct { client *genai.Client - config *Config + config GeminiAudioConfig } var _ Provider = (*GeminiProvider)(nil) -// NewGeminiProvider creates a new Gemini TTS provider. -func NewGeminiProvider(config *Config) (Provider, error) { - normalized := normalizeGeminiConfig(config) - if normalized.GoogleAPIKey == "" { +// NewGeminiProvider creates a new Gemini TTS provider from the Gemini-specific +// sub-config. Callers that have a flat Config should use NewProvider instead. +func NewGeminiProvider(config GeminiAudioConfig, outputFormat string) (Provider, error) { + normalized := normalizeGeminiAudioConfig(config) + if normalized.APIKey == "" { return nil, errors.New("google API key is required") } client, err := genai.NewClient(context.Background(), &genai.ClientConfig{ - APIKey: normalized.GoogleAPIKey, + APIKey: normalized.APIKey, Backend: genai.BackendGeminiAPI, }) if err != nil { @@ -61,7 +63,7 @@ func (p *GeminiProvider) GenerateAudio(ctx context.Context, text string, outputF if err := ValidateBulgarianText(text); err != nil { return err } - if p == nil || p.client == nil || p.config == nil { + if p == nil || p.client == nil { return errors.New("gemini client not initialized") } @@ -71,7 +73,7 @@ func (p *GeminiProvider) GenerateAudio(ctx context.Context, text string, outputF SpeechConfig: p.speechConfig(), } - response, err := p.client.Models.GenerateContent(ctx, p.config.GeminiTTSModel, []*genai.Content{ + response, err := p.client.Models.GenerateContent(ctx, p.config.TTSModel, []*genai.Content{ genai.NewContentFromText(prompt, genai.RoleUser), }, req) if err != nil { @@ -97,7 +99,7 @@ func (p *GeminiProvider) Name() string { // 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) == "" { + if p == nil || strings.TrimSpace(p.config.APIKey) == "" { return errors.New("google API key not configured") } @@ -105,13 +107,8 @@ func (p *GeminiProvider) IsAvailable() error { } func (p *GeminiProvider) buildPrompt(text string) string { - config := p.config - if config == nil { - config = &Config{} - } - var prompt strings.Builder - prompt.WriteString(geminiPromptInstruction(config)) + prompt.WriteString(geminiPromptInstruction(p.config)) prompt.WriteString("\n") prompt.WriteString(strings.TrimSpace(text)) @@ -119,16 +116,11 @@ func (p *GeminiProvider) buildPrompt(text string) 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 != "" { + if voice := strings.TrimSpace(p.config.Voice); voice != "" { speechConfig.VoiceConfig = &genai.VoiceConfig{ PrebuiltVoiceConfig: &genai.PrebuiltVoiceConfig{ VoiceName: voice, @@ -139,24 +131,20 @@ func (p *GeminiProvider) speechConfig() *genai.SpeechConfig { 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) +// normalizeGeminiAudioConfig applies defaults and trims whitespace from a GeminiAudioConfig. +func normalizeGeminiAudioConfig(config GeminiAudioConfig) GeminiAudioConfig { + config.APIKey = strings.TrimSpace(config.APIKey) + config.TTSModel = strings.TrimSpace(config.TTSModel) + config.Voice = strings.TrimSpace(config.Voice) - if normalized.GeminiTTSModel == "" { - normalized.GeminiTTSModel = defaultGeminiTTSModel + if config.TTSModel == "" { + config.TTSModel = defaultGeminiTTSModel } - if normalized.GeminiSpeed <= 0 { - normalized.GeminiSpeed = 1.0 + if config.Speed <= 0 { + config.Speed = 1.0 } - return normalized + return config } func extractAudioData(response *genai.GenerateContentResponse) ([]byte, string, error) { diff --git a/internal/audio/gemini_provider_test.go b/internal/audio/gemini_provider_test.go index 66f0a74..64b43b1 100644 --- a/internal/audio/gemini_provider_test.go +++ b/internal/audio/gemini_provider_test.go @@ -12,21 +12,19 @@ import ( func TestNewGeminiProvider(t *testing.T) { tests := []struct { name string - config *Config + config GeminiAudioConfig wantErr bool wantModel string wantSpeed float64 }{ { name: "missing google api key", - config: &Config{}, + config: GeminiAudioConfig{}, wantErr: true, }, { - name: "valid config", - config: &Config{ - GoogleAPIKey: "test-key", - }, + name: "valid config", + config: GeminiAudioConfig{APIKey: "test-key"}, wantErr: false, wantModel: defaultGeminiTTSModel, wantSpeed: 1.0, @@ -35,7 +33,7 @@ func TestNewGeminiProvider(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - provider, err := NewGeminiProvider(tt.config) + provider, err := NewGeminiProvider(tt.config, "mp3") if (err != nil) != tt.wantErr { t.Fatalf("NewGeminiProvider() error = %v, wantErr %v", err, tt.wantErr) } @@ -51,11 +49,11 @@ func TestNewGeminiProvider(t *testing.T) { if !ok { t.Fatalf("NewGeminiProvider() returned %T, want *GeminiProvider", provider) } - if geminiProvider.config.GeminiTTSModel != tt.wantModel { - t.Fatalf("GeminiTTSModel = %q, want %q", geminiProvider.config.GeminiTTSModel, tt.wantModel) + if geminiProvider.config.TTSModel != tt.wantModel { + t.Fatalf("TTSModel = %q, want %q", geminiProvider.config.TTSModel, tt.wantModel) } - if geminiProvider.config.GeminiSpeed != tt.wantSpeed { - t.Fatalf("GeminiSpeed = %v, want %v", geminiProvider.config.GeminiSpeed, tt.wantSpeed) + if geminiProvider.config.Speed != tt.wantSpeed { + t.Fatalf("Speed = %v, want %v", geminiProvider.config.Speed, tt.wantSpeed) } }) } @@ -64,17 +62,17 @@ func TestNewGeminiProvider(t *testing.T) { func TestGeminiProviderIsAvailable(t *testing.T) { tests := []struct { name string - config *Config + config GeminiAudioConfig wantErr bool }{ { name: "with API key", - config: &Config{GoogleAPIKey: "test-key"}, + config: GeminiAudioConfig{APIKey: "test-key"}, wantErr: false, }, { name: "without API key", - config: &Config{}, + config: GeminiAudioConfig{}, wantErr: true, }, } @@ -92,9 +90,9 @@ func TestGeminiProviderIsAvailable(t *testing.T) { func TestGeminiProviderBuildPrompt(t *testing.T) { provider := &GeminiProvider{ - config: &Config{ - GeminiVoice: "Kore", - GeminiSpeed: 0.92, + config: GeminiAudioConfig{ + Voice: "Kore", + Speed: 0.92, }, } @@ -277,7 +275,7 @@ func TestNewGeminiProviderWithGoogleAPIKey(t *testing.T) { t.Skip("Skipping smoke test: GOOGLE_API_KEY not set") } - provider, err := NewGeminiProvider(&Config{GoogleAPIKey: apiKey}) + provider, err := NewGeminiProvider(GeminiAudioConfig{APIKey: apiKey}, "mp3") if err != nil { t.Fatalf("NewGeminiProvider() unexpected error: %v", err) } diff --git a/internal/audio/openai_provider.go b/internal/audio/openai_provider.go index bf5ac58..9c6d3a4 100644 --- a/internal/audio/openai_provider.go +++ b/internal/audio/openai_provider.go @@ -15,26 +15,26 @@ import ( // Compile-time check that OpenAIProvider implements the Provider interface. var _ Provider = (*OpenAIProvider)(nil) -// OpenAIProvider implements Provider interface for OpenAI TTS +// OpenAIProvider implements Provider interface for OpenAI TTS. +// It stores only the OpenAI-specific sub-config so it never sees Gemini fields. type OpenAIProvider struct { - client *openai.Client - config *Config + client *openai.Client + config OpenAIAudioConfig + outputFormat string } -// NewOpenAIProvider creates a new OpenAI TTS provider -func NewOpenAIProvider(config *Config) (Provider, error) { - if config.OpenAIKey == "" { +// NewOpenAIProvider creates a new OpenAI TTS provider from the OpenAI-specific +// sub-config. Callers that have a flat Config should use NewProvider instead. +func NewOpenAIProvider(config OpenAIAudioConfig, outputFormat string) (Provider, error) { + if config.Key == "" { return nil, errors.New("OpenAI API key is required") } - client := openai.NewClient(config.OpenAIKey) - - provider := &OpenAIProvider{ - client: client, - config: config, - } - - return provider, nil + return &OpenAIProvider{ + client: openai.NewClient(config.Key), + config: config, + outputFormat: outputFormat, + }, nil } // GenerateAudio generates audio using OpenAI TTS @@ -49,22 +49,22 @@ func (p *OpenAIProvider) GenerateAudio(ctx context.Context, text string, outputF // Prepare the TTS request // OpenAI TTS will automatically detect and pronounce Bulgarian text - fmt.Printf("OpenAI TTS: Using model '%s' with voice '%s' at speed %.2f\n", p.config.OpenAIModel, p.config.OpenAIVoice, p.config.OpenAISpeed) - if p.config.OpenAIInstruction != "" && (p.config.OpenAIModel == "gpt-4o-mini-tts" || p.config.OpenAIModel == "gpt-4o-mini-audio-preview") { - fmt.Printf("OpenAI TTS Instruction: '%s'\n", p.config.OpenAIInstruction) + fmt.Printf("OpenAI TTS: Using model '%s' with voice '%s' at speed %.2f\n", p.config.Model, p.config.Voice, p.config.Speed) + if p.config.Instruction != "" && (p.config.Model == "gpt-4o-mini-tts" || p.config.Model == "gpt-4o-mini-audio-preview") { + fmt.Printf("OpenAI TTS Instruction: '%s'\n", p.config.Instruction) } fmt.Printf("OpenAI TTS Input: '%s'\n", processedText) req := openai.CreateSpeechRequest{ - Model: openai.SpeechModel(p.config.OpenAIModel), + Model: openai.SpeechModel(p.config.Model), Input: processedText, - Voice: openai.SpeechVoice(p.config.OpenAIVoice), - Speed: p.config.OpenAISpeed, + Voice: openai.SpeechVoice(p.config.Voice), + Speed: p.config.Speed, } // Add instructions for gpt-4o-mini-tts model - if p.config.OpenAIInstruction != "" && (p.config.OpenAIModel == "gpt-4o-mini-tts" || p.config.OpenAIModel == "gpt-4o-mini-audio-preview") { - req.Instructions = p.config.OpenAIInstruction + if p.config.Instruction != "" && (p.config.Model == "gpt-4o-mini-tts" || p.config.Model == "gpt-4o-mini-audio-preview") { + req.Instructions = p.config.Instruction } // Determine response format based on output file extension @@ -92,8 +92,8 @@ func (p *OpenAIProvider) GenerateAudio(ctx context.Context, text string, outputF if err != nil { // Check if it's a model access error errStr := err.Error() - if strings.Contains(errStr, "does not have access to model") && (p.config.OpenAIModel == "gpt-4o-mini-tts" || p.config.OpenAIModel == "gpt-4o-mini-audio-preview") { - return fmt.Errorf("OpenAI TTS API error: %w\nNote: The %s model requires access. Try using --openai-model tts-1-hd instead", err, p.config.OpenAIModel) + if strings.Contains(errStr, "does not have access to model") && (p.config.Model == "gpt-4o-mini-tts" || p.config.Model == "gpt-4o-mini-audio-preview") { + return fmt.Errorf("OpenAI TTS API error: %w\nNote: The %s model requires access. Try using --openai-model tts-1-hd instead", err, p.config.Model) } return err } @@ -138,7 +138,7 @@ func (p *OpenAIProvider) Name() string { // IsAvailable checks if the OpenAI API is accessible func (p *OpenAIProvider) IsAvailable() error { - if p.config.OpenAIKey == "" { + if p.config.Key == "" { return errors.New("OpenAI API key not configured") } diff --git a/internal/audio/openai_provider_test.go b/internal/audio/openai_provider_test.go index ae16ca0..09f4c2f 100644 --- a/internal/audio/openai_provider_test.go +++ b/internal/audio/openai_provider_test.go @@ -9,30 +9,26 @@ import ( func TestNewOpenAIProvider(t *testing.T) { tests := []struct { name string - config *Config + config OpenAIAudioConfig wantErr bool errMsg string }{ { - name: "missing API key", - config: &Config{ - OpenAIKey: "", - }, + name: "missing API key", + config: OpenAIAudioConfig{Key: ""}, wantErr: true, errMsg: "OpenAI API key is required", }, { - name: "valid config", - config: &Config{ - OpenAIKey: "test-key", - }, + name: "valid config", + config: OpenAIAudioConfig{Key: "test-key"}, wantErr: false, }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - provider, err := NewOpenAIProvider(tt.config) + provider, err := NewOpenAIProvider(tt.config, "mp3") if (err != nil) != tt.wantErr { t.Errorf("NewOpenAIProvider() error = %v, wantErr %v", err, tt.wantErr) } @@ -53,30 +49,24 @@ func TestNewOpenAIProvider(t *testing.T) { func TestOpenAIProviderIsAvailable(t *testing.T) { tests := []struct { name string - config *Config + config OpenAIAudioConfig wantErr bool }{ { - name: "with API key", - config: &Config{ - OpenAIKey: "test-key", - }, + name: "with API key", + config: OpenAIAudioConfig{Key: "test-key"}, wantErr: false, }, { - name: "without API key", - config: &Config{ - OpenAIKey: "", - }, + name: "without API key", + config: OpenAIAudioConfig{Key: ""}, wantErr: true, }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - provider := &OpenAIProvider{ - config: tt.config, - } + provider := &OpenAIProvider{config: tt.config} err := provider.IsAvailable() if (err != nil) != tt.wantErr { t.Errorf("IsAvailable() error = %v, wantErr %v", err, tt.wantErr) @@ -86,9 +76,7 @@ func TestOpenAIProviderIsAvailable(t *testing.T) { } func TestPreprocessBulgarianText(t *testing.T) { - provider := &OpenAIProvider{ - config: &Config{}, - } + provider := &OpenAIProvider{} tests := []struct { name string @@ -133,11 +121,7 @@ func TestPreprocessBulgarianText(t *testing.T) { } func TestGenerateAudioValidation(t *testing.T) { - provider := &OpenAIProvider{ - config: &Config{ - OpenAIKey: "test-key", - }, - } + provider := &OpenAIProvider{config: OpenAIAudioConfig{Key: "test-key"}} ctx := context.Background() diff --git a/internal/audio/provider.go b/internal/audio/provider.go index 0e863bf..fac0cd9 100644 --- a/internal/audio/provider.go +++ b/internal/audio/provider.go @@ -17,26 +17,76 @@ type Provider interface { IsAvailable() error } -// Config holds common configuration for audio providers. +// OpenAIAudioConfig holds settings specific to the OpenAI TTS backend. +// Callers that only use Gemini never need to populate these fields. +type OpenAIAudioConfig struct { + Key string + Model string // "tts-1", "tts-1-hd", or "gpt-4o-mini-tts" + Voice string // One of OpenAIVoices. + Speed float64 // 0.25 to 4.0 + Instruction string // Voice instructions for gpt-4o-mini-tts model +} + +// GeminiAudioConfig holds settings specific to the Gemini TTS backend. +// Callers that only use OpenAI never need to populate these fields. +type GeminiAudioConfig struct { + APIKey string + TTSModel string // "gemini-2.5-flash-preview-tts" + Voice string // One of GeminiVoices; empty lets the caller choose a random voice. + Speed float64 // Prompt hint for desired speech speed +} + +// Config holds common configuration for audio providers. Provider-specific +// settings are grouped into OpenAI and Gemini sub-configs so callers and +// implementations only see the fields relevant to their backend. type Config struct { Provider string // Provider name: "openai" or "gemini" OutputDir string // Directory for output files OutputFormat string // Output format: "mp3" or "wav" - // OpenAI-specific settings + // OpenAI-specific settings — ignored when Provider == "gemini". OpenAIKey string OpenAIModel string // "tts-1", "tts-1-hd", or "gpt-4o-mini-tts" OpenAIVoice string // One of OpenAIVoices. OpenAISpeed float64 // 0.25 to 4.0 OpenAIInstruction string // Voice instructions for gpt-4o-mini-tts model - // Gemini-specific settings + // Gemini-specific settings — ignored when Provider == "openai". GoogleAPIKey string GeminiTTSModel string // "gemini-2.5-flash-preview-tts" GeminiVoice string // One of GeminiVoices; empty lets the caller choose a random voice. GeminiSpeed float64 // Prompt hint for desired speech speed } +// openAIAudioConfigFrom extracts the OpenAI-specific sub-config from the flat Config. +// A nil Config produces a zero-value OpenAIAudioConfig. +func openAIAudioConfigFrom(c *Config) OpenAIAudioConfig { + if c == nil { + return OpenAIAudioConfig{} + } + return OpenAIAudioConfig{ + Key: c.OpenAIKey, + Model: c.OpenAIModel, + Voice: c.OpenAIVoice, + Speed: c.OpenAISpeed, + Instruction: c.OpenAIInstruction, + } +} + +// geminiAudioConfigFrom extracts the Gemini-specific sub-config from the flat Config. +// A nil Config produces a zero-value GeminiAudioConfig. +func geminiAudioConfigFrom(c *Config) GeminiAudioConfig { + if c == nil { + return GeminiAudioConfig{} + } + return GeminiAudioConfig{ + APIKey: c.GoogleAPIKey, + TTSModel: c.GeminiTTSModel, + Voice: c.GeminiVoice, + Speed: c.GeminiSpeed, + } +} + // DefaultConfig returns default configuration func DefaultProviderConfig() *Config { return &Config{ @@ -53,7 +103,9 @@ func DefaultProviderConfig() *Config { } } -// NewProvider creates the appropriate audio provider based on configuration +// NewProvider creates the appropriate audio provider based on configuration. +// It extracts provider-specific sub-configs so each implementation only +// receives the fields it needs (ISP). func NewProvider(config *Config) (Provider, error) { if config == nil { config = DefaultProviderConfig() @@ -64,12 +116,12 @@ func NewProvider(config *Config) (Provider, error) { if config.OpenAIKey == "" { return nil, fmt.Errorf("OpenAI API key is required") } - return NewOpenAIProvider(config) + return NewOpenAIProvider(openAIAudioConfigFrom(config), config.OutputFormat) case "gemini": if config.GoogleAPIKey == "" { return nil, fmt.Errorf("google API key is required") } - return NewGeminiProvider(config) + return NewGeminiProvider(geminiAudioConfigFrom(config), config.OutputFormat) default: return nil, fmt.Errorf("unknown audio provider: %s", config.Provider) } diff --git a/internal/audio/sidecar.go b/internal/audio/sidecar.go index 554afc6..dbe3064 100644 --- a/internal/audio/sidecar.go +++ b/internal/audio/sidecar.go @@ -48,12 +48,13 @@ func ProcessedTextForProvider(provider, text string) string { } // InstructionForProvider returns the provider-specific instruction semantics written to attribution files. +// It accepts the flat Config and extracts the relevant sub-config internally. func InstructionForProvider(provider string, config *Config) string { switch strings.ToLower(strings.TrimSpace(provider)) { case "openai": - return openAIInstructionForAttribution(config) + return openAIInstructionForAttribution(openAIAudioConfigFrom(config)) case "gemini": - return geminiPromptInstruction(config) + return geminiPromptInstruction(geminiAudioConfigFrom(config)) default: return "" } @@ -69,16 +70,12 @@ func openAIProcessedText(text string) string { return strings.TrimSpace(cleanedText) } -func openAIInstructionForAttribution(config *Config) string { - if config == nil { +func openAIInstructionForAttribution(config OpenAIAudioConfig) string { + if !openAIModelUsesInstructions(config.Model) { return "" } - if !openAIModelUsesInstructions(config.OpenAIModel) { - return "" - } - - return strings.TrimSpace(config.OpenAIInstruction) + return strings.TrimSpace(config.Instruction) } func openAIModelUsesInstructions(model string) bool { @@ -90,23 +87,19 @@ func openAIModelUsesInstructions(model string) bool { } } -func geminiPromptInstruction(config *Config) string { - if config == nil { - config = &Config{} - } - +func geminiPromptInstruction(config GeminiAudioConfig) string { 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 != "" { + if speedHint := geminiSpeedHint(config.Speed); speedHint != "" { prompt.WriteString(" ") prompt.WriteString(speedHint) } prompt.WriteString("\n\nSpeak the following Bulgarian text:") - if voice := strings.TrimSpace(config.GeminiVoice); voice != "" { + if voice := strings.TrimSpace(config.Voice); voice != "" { prompt.WriteString("\n\nUse a clear, natural delivery that matches the voice named ") prompt.WriteString(voice) prompt.WriteString(".") |
