summaryrefslogtreecommitdiff
path: root/internal
diff options
context:
space:
mode:
Diffstat (limited to 'internal')
-rw-r--r--internal/tts/gemini.go87
1 files changed, 79 insertions, 8 deletions
diff --git a/internal/tts/gemini.go b/internal/tts/gemini.go
index 736e1d1..e4cec59 100644
--- a/internal/tts/gemini.go
+++ b/internal/tts/gemini.go
@@ -4,6 +4,9 @@ import (
"context"
"fmt"
"os"
+ "os/exec"
+ "path/filepath"
+ "strconv"
"strings"
"google.golang.org/genai"
@@ -108,14 +111,8 @@ func (g *GeminiProvider) GenerateAudio(ctx context.Context, text, outputFile str
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)
+ if err := writeAudioFile(data, mimeType, outputFile); err != nil {
+ return err
}
return nil
}
@@ -132,6 +129,80 @@ func extractAudio(resp *genai.GenerateContentResponse) ([]byte, string, error) {
return nil, "", fmt.Errorf("audio payload missing")
}
+func writeAudioFile(data []byte, mimeType, outputFile string) error {
+ if strings.TrimSpace(outputFile) == "" {
+ return fmt.Errorf("output file is required")
+ }
+ lowerMime := strings.ToLower(strings.TrimSpace(mimeType))
+ switch {
+ case strings.HasPrefix(lowerMime, "audio/mpeg"), strings.HasPrefix(lowerMime, "audio/mp3"):
+ if err := os.WriteFile(outputFile, data, 0o644); err != nil {
+ return fmt.Errorf("write audio: %w", err)
+ }
+ return nil
+ case strings.Contains(lowerMime, "audio/l16"), strings.Contains(lowerMime, "codec=pcm"), lowerMime == "":
+ return encodePCMToMP3(data, lowerMime, outputFile)
+ default:
+ return fmt.Errorf("unsupported audio mime type %q", mimeType)
+ }
+}
+
+func encodePCMToMP3(data []byte, mimeType, outputFile string) error {
+ ffmpegPath, err := exec.LookPath("ffmpeg")
+ if err != nil {
+ return fmt.Errorf("ffmpeg not found for audio conversion: %w", err)
+ }
+
+ rate := 24000
+ if idx := strings.Index(mimeType, "rate="); idx >= 0 {
+ value := mimeType[idx+len("rate="):]
+ for i, r := range value {
+ if r < '0' || r > '9' {
+ value = value[:i]
+ break
+ }
+ }
+ if parsed, parseErr := strconv.Atoi(value); parseErr == nil && parsed > 0 {
+ rate = parsed
+ }
+ }
+
+ tmpDir := filepath.Dir(outputFile)
+ rawFile, err := os.CreateTemp(tmpDir, "comicforge-tts-*.pcm")
+ if err != nil {
+ return fmt.Errorf("create temporary audio file: %w", err)
+ }
+ rawPath := rawFile.Name()
+ if _, err := rawFile.Write(data); err != nil {
+ _ = rawFile.Close()
+ _ = os.Remove(rawPath)
+ return fmt.Errorf("write temporary audio file: %w", err)
+ }
+ if err := rawFile.Close(); err != nil {
+ _ = os.Remove(rawPath)
+ return fmt.Errorf("close temporary audio file: %w", err)
+ }
+ defer func() {
+ _ = os.Remove(rawPath)
+ }()
+
+ cmd := exec.Command(ffmpegPath,
+ "-nostdin", "-hide_banner", "-loglevel", "error", "-y",
+ "-f", "s16le",
+ "-ar", fmt.Sprintf("%d", rate),
+ "-ac", "1",
+ "-i", rawPath,
+ "-codec:a", "libmp3lame",
+ "-q:a", "2",
+ outputFile,
+ )
+ out, err := cmd.CombinedOutput()
+ if err != nil {
+ return fmt.Errorf("convert PCM audio to mp3: %w\n%s", err, strings.TrimSpace(string(out)))
+ }
+ return nil
+}
+
func defaultOr(value, fallback string) string {
if strings.TrimSpace(value) != "" {
return value