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.
This commit is contained in:
viettranx
2026-04-15 11:24:57 +07:00
parent 96f11720e8
commit 37224937c6
2 changed files with 484 additions and 0 deletions
+207
View File
@@ -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"
}
}
+277
View File
@@ -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
}