summaryrefslogtreecommitdiff
path: root/internal/tts/gemini.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/tts/gemini.go')
-rw-r--r--internal/tts/gemini.go140
1 files changed, 140 insertions, 0 deletions
diff --git a/internal/tts/gemini.go b/internal/tts/gemini.go
new file mode 100644
index 0000000..1836be4
--- /dev/null
+++ b/internal/tts/gemini.go
@@ -0,0 +1,140 @@
+package tts
+
+import (
+ "context"
+ "fmt"
+ "os"
+ "strings"
+
+ "google.golang.org/genai"
+
+ "codeberg.org/snonux/comicforge/internal/provider"
+)
+
+const (
+ // DefaultModel is the Gemini TTS model used for narration.
+ DefaultModel = "gemini-2.5-flash-preview-tts"
+)
+
+// GeminiConfig configures the Gemini TTS provider.
+type GeminiConfig struct {
+ APIKey string
+ Model string
+ Voice string
+}
+
+// GeminiProvider generates MP3 audio with Gemini TTS.
+type GeminiProvider struct {
+ client *genai.Client
+ model string
+ voice string
+ err error
+}
+
+var _ provider.TTSProvider = (*GeminiProvider)(nil)
+
+// NewGeminiProvider creates a Gemini TTS provider.
+func NewGeminiProvider(cfg *GeminiConfig) *GeminiProvider {
+ g := &GeminiProvider{model: DefaultModel}
+ if cfg == nil {
+ g.err = fmt.Errorf("tts config is required")
+ return g
+ }
+ g.model = defaultOr(cfg.Model, DefaultModel)
+ g.voice = cfg.Voice
+ if strings.TrimSpace(cfg.APIKey) == "" {
+ g.err = fmt.Errorf("Google API key is required for TTS")
+ return g
+ }
+ client, err := genai.NewClient(context.Background(), &genai.ClientConfig{
+ APIKey: cfg.APIKey,
+ Backend: genai.BackendGeminiAPI,
+ })
+ if err != nil {
+ g.err = fmt.Errorf("create Gemini client: %w", err)
+ return g
+ }
+ g.client = client
+ return g
+}
+
+// Name returns the provider name.
+func (g *GeminiProvider) Name() string { return "gemini" }
+
+// IsAvailable reports whether the provider was initialized successfully.
+func (g *GeminiProvider) IsAvailable() error {
+ if g == nil {
+ return fmt.Errorf("tts provider is nil")
+ }
+ return g.err
+}
+
+// GenerateAudio writes MP3 audio for the provided text to outputFile.
+func (g *GeminiProvider) GenerateAudio(ctx context.Context, text, outputFile string) error {
+ if g == nil {
+ return fmt.Errorf("tts provider is nil")
+ }
+ if ctx == nil {
+ ctx = context.Background()
+ }
+ if g.err != nil {
+ return g.err
+ }
+ if strings.TrimSpace(text) == "" {
+ return fmt.Errorf("text is required")
+ }
+ if strings.TrimSpace(outputFile) == "" {
+ return fmt.Errorf("output file is required")
+ }
+ voiceName := g.voice
+ if strings.TrimSpace(voiceName) == "" {
+ voiceName = "Aoede"
+ }
+
+ speechCfg := &genai.SpeechConfig{
+ VoiceConfig: &genai.VoiceConfig{
+ PrebuiltVoiceConfig: &genai.PrebuiltVoiceConfig{VoiceName: voiceName},
+ },
+ LanguageCode: "bg-BG",
+ }
+ resp, err := g.client.Models.GenerateContent(ctx, g.model, genai.Text(text), &genai.GenerateContentConfig{
+ ResponseModalities: []string{"AUDIO"},
+ SpeechConfig: speechCfg,
+ })
+ if err != nil {
+ return fmt.Errorf("generate audio: %w", err)
+ }
+ data, mimeType, err := extractAudio(resp)
+ if err != nil {
+ return err
+ }
+ if strings.TrimSpace(outputFile) == "" {
+ return fmt.Errorf("output file is required")
+ }
+ if err := os.WriteFile(outputFile, data, 0o644); err != nil {
+ return fmt.Errorf("write audio: %w", err)
+ }
+ if mimeType != "" && !strings.HasPrefix(mimeType, "audio/") {
+ return fmt.Errorf("unexpected audio mime type %q", mimeType)
+ }
+ return nil
+}
+
+func extractAudio(resp *genai.GenerateContentResponse) ([]byte, string, error) {
+ if resp == nil || len(resp.Candidates) == 0 || resp.Candidates[0] == nil || resp.Candidates[0].Content == nil {
+ return nil, "", fmt.Errorf("no audio returned")
+ }
+ for _, part := range resp.Candidates[0].Content.Parts {
+ if part != nil && part.InlineData != nil && len(part.InlineData.Data) > 0 {
+ return part.InlineData.Data, part.InlineData.MIMEType, nil
+ }
+ }
+ return nil, "", fmt.Errorf("audio payload missing")
+}
+
+func defaultOr(value, fallback string) string {
+ if strings.TrimSpace(value) != "" {
+ return value
+ }
+ return fallback
+}