From 37224937c616ff84cedbb44c8d7edab86ce12b68 Mon Sep 17 00:00:00 2001 From: viettranx Date: Wed, 15 Apr 2026 07:06:00 +0700 Subject: [PATCH] feat(audio): ElevenLabs Scribe STT provider Implement native ElevenLabs Scribe STT provider with POST /v1/speech-to-text endpoint. Supports 20MB audio cap, multipart form submission, configurable model and language. Includes comprehensive unit tests for happy path, audio size validation, API errors, and edge cases. --- internal/audio/elevenlabs/stt.go | 207 +++++++++++++++++++ internal/audio/elevenlabs/stt_test.go | 277 ++++++++++++++++++++++++++ 2 files changed, 484 insertions(+) create mode 100644 internal/audio/elevenlabs/stt.go create mode 100644 internal/audio/elevenlabs/stt_test.go diff --git a/internal/audio/elevenlabs/stt.go b/internal/audio/elevenlabs/stt.go new file mode 100644 index 00000000..37c430f4 --- /dev/null +++ b/internal/audio/elevenlabs/stt.go @@ -0,0 +1,207 @@ +package elevenlabs + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io" + "mime/multipart" + "net/http" + "os" + "path/filepath" + "strings" + "time" + + "github.com/nextlevelbuilder/goclaw/internal/audio" +) + +const ( + sttMaxBytes = 20 << 20 // 20 MB safety cap (conservative; ElevenLabs allows 3 GB) + sttDefaultModelID = "scribe_v1" + sttDefaultTimeout = 60 * time.Second +) + +// STTProvider transcribes audio via ElevenLabs Scribe. +// Endpoint: POST /v1/speech-to-text (xi-api-key auth). +type STTProvider struct { + cfg Config + c *client +} + +// NewSTTProvider returns an ElevenLabs STT (Scribe) provider. +func NewSTTProvider(cfg Config) *STTProvider { + cfg = cfg.withDefaults() + timeoutMs := max(cfg.TimeoutMs, int(sttDefaultTimeout.Milliseconds())) + return &STTProvider{ + cfg: cfg, + c: newClient(cfg.APIKey, cfg.BaseURL, timeoutMs), + } +} + +// Name returns the stable provider identifier. +func (p *STTProvider) Name() string { return "elevenlabs" } + +// Transcribe converts audio to text via Scribe. FilePath is preferred over +// Bytes to avoid buffering large files in memory. 20 MB cap enforced before POST. +func (p *STTProvider) Transcribe(ctx context.Context, in audio.STTInput, opts audio.STTOptions) (*audio.TranscriptResult, error) { + // Resolve file path or temp file from bytes. + filePath, cleanup, err := resolveFilePath(in) + if err != nil { + return nil, fmt.Errorf("elevenlabs stt: resolve input: %w", err) + } + if cleanup != nil { + defer cleanup() + } + + // Size cap before multipart assembly. + if err := checkFileSize(filePath, sttMaxBytes); err != nil { + return nil, fmt.Errorf("elevenlabs stt: %w", err) + } + + // Build multipart body. + var buf bytes.Buffer + mw := multipart.NewWriter(&buf) + + modelID := opts.ModelID + if modelID == "" { + modelID = sttDefaultModelID + } + if err := mw.WriteField("model_id", modelID); err != nil { + return nil, fmt.Errorf("elevenlabs stt: write model_id field: %w", err) + } + + if opts.Language != "" { + if err := mw.WriteField("language_code", opts.Language); err != nil { + return nil, fmt.Errorf("elevenlabs stt: write language_code field: %w", err) + } + } + + if opts.Diarize { + if err := mw.WriteField("diarize", "true"); err != nil { + return nil, fmt.Errorf("elevenlabs stt: write diarize field: %w", err) + } + } + + // Attach file. + filename := in.Filename + if filename == "" { + filename = filepath.Base(filePath) + } + fw, err := mw.CreateFormFile("file", filename) + if err != nil { + return nil, fmt.Errorf("elevenlabs stt: create form file: %w", err) + } + f, err := os.Open(filePath) + if err != nil { + return nil, fmt.Errorf("elevenlabs stt: open file: %w", err) + } + defer f.Close() + if _, err := io.Copy(fw, f); err != nil { + return nil, fmt.Errorf("elevenlabs stt: write file bytes: %w", err) + } + if err := mw.Close(); err != nil { + return nil, fmt.Errorf("elevenlabs stt: close multipart writer: %w", err) + } + + // Build HTTP request. + url := strings.TrimRight(p.cfg.BaseURL, "/") + "/v1/speech-to-text" + req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, &buf) + if err != nil { + return nil, fmt.Errorf("elevenlabs stt: create request: %w", err) + } + req.Header.Set("Content-Type", mw.FormDataContentType()) + req.Header.Set("xi-api-key", p.cfg.APIKey) + + timeout := sttDefaultTimeout + if opts.TimeoutMs > 0 { + timeout = time.Duration(opts.TimeoutMs) * time.Millisecond + } + hc := &http.Client{Timeout: timeout} + resp, err := hc.Do(req) + if err != nil { + return nil, fmt.Errorf("elevenlabs stt: http request: %w", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + body, _ := io.ReadAll(io.LimitReader(resp.Body, 512)) + return nil, fmt.Errorf("elevenlabs stt: API error %d: %s", resp.StatusCode, string(body)) + } + + var result struct { + Text string `json:"text"` + LanguageCode string `json:"language_code"` + Duration float64 `json:"audio_duration_secs"` + } + if err := json.NewDecoder(resp.Body).Decode(&result); err != nil { + return nil, fmt.Errorf("elevenlabs stt: parse response: %w", err) + } + + return &audio.TranscriptResult{ + Text: result.Text, + Language: result.LanguageCode, + Duration: result.Duration, + Provider: "elevenlabs", + }, nil +} + +// resolveFilePath returns a usable file path. When only Bytes is set, writes a +// temp file (0600) and returns a cleanup func to remove it. +func resolveFilePath(in audio.STTInput) (path string, cleanup func(), err error) { + if in.FilePath != "" { + return in.FilePath, nil, nil + } + if len(in.Bytes) == 0 { + return "", nil, fmt.Errorf("neither FilePath nor Bytes provided") + } + ext := extFromMime(in.MimeType) + f, err := os.CreateTemp("", "stt-*"+ext) + if err != nil { + return "", nil, fmt.Errorf("create temp file: %w", err) + } + if err := os.Chmod(f.Name(), 0600); err != nil { + f.Close() + os.Remove(f.Name()) + return "", nil, fmt.Errorf("chmod temp file: %w", err) + } + if _, err := f.Write(in.Bytes); err != nil { + f.Close() + os.Remove(f.Name()) + return "", nil, fmt.Errorf("write temp file: %w", err) + } + f.Close() + return f.Name(), func() { os.Remove(f.Name()) }, nil +} + +// checkFileSize returns an error if the file exceeds maxBytes. +func checkFileSize(path string, maxBytes int64) error { + info, err := os.Stat(path) + if err != nil { + return fmt.Errorf("stat file: %w", err) + } + if info.Size() > maxBytes { + return fmt.Errorf("file too large (%d bytes, max %d)", info.Size(), maxBytes) + } + return nil +} + +// extFromMime returns a file extension for a MIME type. +func extFromMime(mime string) string { + switch mime { + case "audio/ogg", "audio/ogg; codecs=opus": + return ".ogg" + case "audio/mpeg", "audio/mp3": + return ".mp3" + case "audio/wav", "audio/wave": + return ".wav" + case "audio/mp4", "audio/m4a": + return ".m4a" + case "audio/webm": + return ".webm" + case "audio/flac": + return ".flac" + default: + return ".bin" + } +} diff --git a/internal/audio/elevenlabs/stt_test.go b/internal/audio/elevenlabs/stt_test.go new file mode 100644 index 00000000..9d87267e --- /dev/null +++ b/internal/audio/elevenlabs/stt_test.go @@ -0,0 +1,277 @@ +package elevenlabs + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "os" + "testing" + + "github.com/nextlevelbuilder/goclaw/internal/audio" +) + +// sttHappyResponse matches ElevenLabs Scribe response schema. +type sttHappyResponse struct { + Text string `json:"text"` + LanguageCode string `json:"language_code"` + DurationSecs float64 `json:"audio_duration_secs"` +} + +func newTestSTTServer(t *testing.T, handler http.HandlerFunc) (*httptest.Server, *STTProvider) { + t.Helper() + srv := httptest.NewServer(handler) + p := NewSTTProvider(Config{ + APIKey: "test-key", + BaseURL: srv.URL, + }) + return srv, p +} + +func writeTempAudioFile(t *testing.T, content string) string { + t.Helper() + f, err := os.CreateTemp("", "stt_elabs_test_*.ogg") + if err != nil { + t.Fatalf("create temp: %v", err) + } + f.WriteString(content) + f.Close() + return f.Name() +} + +// Case 1: happy path — multipart upload returns transcript. +func TestSTTProvider_HappyPath(t *testing.T) { + audioFile := writeTempAudioFile(t, "fake-ogg") + defer os.Remove(audioFile) + + srv, p := newTestSTTServer(t, func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/v1/speech-to-text" { + t.Errorf("unexpected path: %s", r.URL.Path) + } + if r.Header.Get("xi-api-key") == "" { + t.Error("missing xi-api-key header") + } + if err := r.ParseMultipartForm(1 << 20); err != nil { + t.Errorf("parse multipart: %v", err) + } + if _, _, err := r.FormFile("file"); err != nil { + t.Errorf("expected 'file' field: %v", err) + } + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(sttHappyResponse{ + Text: "hello world", + LanguageCode: "en", + DurationSecs: 5.0, + }) + }) + defer srv.Close() + + res, err := p.Transcribe(context.Background(), audio.STTInput{FilePath: audioFile}, audio.STTOptions{}) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if res.Text != "hello world" { + t.Errorf("expected 'hello world', got %q", res.Text) + } + if res.Provider != "elevenlabs" { + t.Errorf("expected provider 'elevenlabs', got %q", res.Provider) + } + if res.Duration != 5.0 { + t.Errorf("expected duration 5.0, got %f", res.Duration) + } +} + +// Case 2: 401 surfaces error with status code. +func TestSTTProvider_Unauthorized(t *testing.T) { + audioFile := writeTempAudioFile(t, "fake-ogg") + defer os.Remove(audioFile) + + srv, p := newTestSTTServer(t, func(w http.ResponseWriter, r *http.Request) { + http.Error(w, `{"detail":"invalid api key"}`, http.StatusUnauthorized) + }) + defer srv.Close() + + _, err := p.Transcribe(context.Background(), audio.STTInput{FilePath: audioFile}, audio.STTOptions{}) + if err == nil { + t.Fatal("expected error for 401, got nil") + } + // Should contain status code. + if err.Error() == "" { + t.Error("error message is empty") + } +} + +// Case 3: language_code passthrough in multipart. +func TestSTTProvider_LanguagePassthrough(t *testing.T) { + audioFile := writeTempAudioFile(t, "fake-ogg") + defer os.Remove(audioFile) + + var gotLang string + srv, p := newTestSTTServer(t, func(w http.ResponseWriter, r *http.Request) { + r.ParseMultipartForm(1 << 20) + gotLang = r.FormValue("language_code") + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(sttHappyResponse{Text: "ok", LanguageCode: "vi"}) + }) + defer srv.Close() + + p.Transcribe(context.Background(), audio.STTInput{FilePath: audioFile}, audio.STTOptions{Language: "vi"}) + if gotLang != "vi" { + t.Errorf("expected language_code 'vi', got %q", gotLang) + } +} + +// Case 4: diarize=true sets multipart field. +func TestSTTProvider_DiarizeField(t *testing.T) { + audioFile := writeTempAudioFile(t, "fake-ogg") + defer os.Remove(audioFile) + + var gotDiarize string + srv, p := newTestSTTServer(t, func(w http.ResponseWriter, r *http.Request) { + r.ParseMultipartForm(1 << 20) + gotDiarize = r.FormValue("diarize") + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(sttHappyResponse{Text: "ok"}) + }) + defer srv.Close() + + p.Transcribe(context.Background(), audio.STTInput{FilePath: audioFile}, audio.STTOptions{Diarize: true}) + if gotDiarize != "true" { + t.Errorf("expected diarize 'true', got %q", gotDiarize) + } +} + +// Case 5: FilePath preferred over Bytes when both set. +func TestSTTProvider_FilePathPreferredOverBytes(t *testing.T) { + audioFile := writeTempAudioFile(t, "file-content") + defer os.Remove(audioFile) + + var receivedSize int64 + srv, p := newTestSTTServer(t, func(w http.ResponseWriter, r *http.Request) { + r.ParseMultipartForm(1 << 20) + f, _, _ := r.FormFile("file") + if f != nil { + buf := make([]byte, 100) + n, _ := f.Read(buf) + receivedSize = int64(n) + } + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(sttHappyResponse{Text: "ok"}) + }) + defer srv.Close() + + // When both set, FilePath wins (Bytes ignored by resolveFilePath when FilePath != ""). + in := audio.STTInput{ + FilePath: audioFile, + Bytes: []byte("different-bytes"), + } + _, err := p.Transcribe(context.Background(), in, audio.STTOptions{}) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + // file-content is 12 bytes; if Bytes were used it would be 15 bytes. + if receivedSize != 12 { + t.Errorf("expected 12 bytes from FilePath, got %d (Bytes may have been used instead)", receivedSize) + } +} + +// Case 6: Bytes-only writes temp file with 0600 perm and cleans up. +func TestSTTProvider_BytesWritesTempFileAndCleanup(t *testing.T) { + var capturedTempPath string + + srv, p := newTestSTTServer(t, func(w http.ResponseWriter, r *http.Request) { + // We can't easily intercept the temp path here; just respond OK. + w.Header().Set("Content-Type", "application/json") + json.NewEncoder(w).Encode(sttHappyResponse{Text: "from bytes"}) + }) + defer srv.Close() + + // Inject a spy by using a real file that we can track. + // Instead, verify behavior: temp file cleaned up after call. + tmpDir := os.TempDir() + beforeFiles := countTempSTTFiles(tmpDir) + + in := audio.STTInput{ + Bytes: []byte("fake-audio-bytes"), + MimeType: "audio/ogg", + } + res, err := p.Transcribe(context.Background(), in, audio.STTOptions{}) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if res.Text != "from bytes" { + t.Errorf("expected 'from bytes', got %q", res.Text) + } + + afterFiles := countTempSTTFiles(tmpDir) + if afterFiles > beforeFiles { + t.Errorf("temp file not cleaned up: before=%d, after=%d, path=%s", beforeFiles, afterFiles, capturedTempPath) + } +} + +// Case 7: oversized file rejected before POST. +func TestSTTProvider_OversizedFileRejected(t *testing.T) { + // Create a file that appears >20MB via stat (we mock size check by writing enough bytes). + // To avoid writing 20MB in tests, we do a stat-based check — we'll write a file + // and manually verify the error path by making the maxBytes tiny. + // Instead: create a real oversized content using a temp file trick. + // We'll write a marker file and test the checkFileSize logic directly. + + f, err := os.CreateTemp("", "stt_big_*.ogg") + if err != nil { + t.Fatalf("create temp: %v", err) + } + defer os.Remove(f.Name()) + + // Write 21MB of zeros. + chunk := make([]byte, 1<<20) // 1 MB + for range 21 { + f.Write(chunk) + } + f.Close() + + // Server should NEVER be called. + srv, p := newTestSTTServer(t, func(w http.ResponseWriter, r *http.Request) { + t.Error("unexpected HTTP call for oversized file") + }) + defer srv.Close() + + _, err = p.Transcribe(context.Background(), audio.STTInput{FilePath: f.Name()}, audio.STTOptions{}) + if err == nil { + t.Fatal("expected error for oversized file, got nil") + } +} + +// Case 8: context cancellation aborts request. +func TestSTTProvider_ContextCancellation(t *testing.T) { + audioFile := writeTempAudioFile(t, "fake-ogg") + defer os.Remove(audioFile) + + srv, p := newTestSTTServer(t, func(w http.ResponseWriter, r *http.Request) { + <-r.Context().Done() + }) + defer srv.Close() + + ctx, cancel := context.WithCancel(context.Background()) + cancel() // cancel immediately + + _, err := p.Transcribe(ctx, audio.STTInput{FilePath: audioFile}, audio.STTOptions{}) + if err == nil { + t.Fatal("expected error for cancelled context, got nil") + } +} + +// countTempSTTFiles counts stt-* temp files in dir (for cleanup verification). +func countTempSTTFiles(dir string) int { + entries, err := os.ReadDir(dir) + if err != nil { + return 0 + } + count := 0 + for _, e := range entries { + if len(e.Name()) > 4 && e.Name()[:4] == "stt-" { + count++ + } + } + return count +}