From eb0f16b5832e5144b2c93309f282da6114d72258 Mon Sep 17 00:00:00 2001 From: leonardotrapani Date: Sun, 1 Feb 2026 01:37:03 +0100 Subject: [PATCH] add elevenlabs streaming adapter with websocket support --- go.mod | 1 + go.sum | 2 + .../adapter_elevenlabs_streaming.go | 282 ++++++++++++ .../adapter_elevenlabs_streaming_test.go | 426 ++++++++++++++++++ progress.txt | 16 +- tasks/prd.jsonc | 2 +- 6 files changed, 727 insertions(+), 2 deletions(-) create mode 100644 internal/transcriber/adapter_elevenlabs_streaming.go create mode 100644 internal/transcriber/adapter_elevenlabs_streaming_test.go diff --git a/go.mod b/go.mod index 929464f..5fd4b71 100644 --- a/go.mod +++ b/go.mod @@ -25,6 +25,7 @@ require ( github.com/charmbracelet/x/term v0.2.1 // indirect github.com/dustin/go-humanize v1.0.1 // indirect github.com/erikgeiser/coninput v0.0.0-20211004153227-1c3628e74d0f // indirect + github.com/gorilla/websocket v1.5.3 // indirect github.com/inconshreveable/mousetrap v1.1.0 // indirect github.com/lucasb-eyer/go-colorful v1.2.0 // indirect github.com/mattn/go-isatty v0.0.20 // indirect diff --git a/go.sum b/go.sum index d198054..811fd78 100644 --- a/go.sum +++ b/go.sum @@ -47,6 +47,8 @@ github.com/erikgeiser/coninput v0.0.0-20211004153227-1c3628e74d0f h1:Y/CXytFA4m6 github.com/erikgeiser/coninput v0.0.0-20211004153227-1c3628e74d0f/go.mod h1:vw97MGsxSvLiUE2X8qFplwetxpGLQrlU1Q9AUEIzCaM= github.com/fsnotify/fsnotify v1.9.0 h1:2Ml+OJNzbYCTzsxtv8vKSFD9PbJjmhYF14k/jKC7S9k= github.com/fsnotify/fsnotify v1.9.0/go.mod h1:8jBTzvmWwFyi3Pb8djgCCO5IBqzKJ/Jwo8TRcHyHii0= +github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg= +github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE= github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8= github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw= github.com/lucasb-eyer/go-colorful v1.2.0 h1:1nnpGOrhyZZuNyfu1QjKiUICQ74+3FNCN69Aj6K7nkY= diff --git a/internal/transcriber/adapter_elevenlabs_streaming.go b/internal/transcriber/adapter_elevenlabs_streaming.go new file mode 100644 index 0000000..9199fe4 --- /dev/null +++ b/internal/transcriber/adapter_elevenlabs_streaming.go @@ -0,0 +1,282 @@ +package transcriber + +import ( + "context" + "encoding/base64" + "encoding/json" + "fmt" + "log" + "net/http" + "net/url" + "sync" + + "github.com/gorilla/websocket" + "github.com/leonardotrapani/hyprvoice/internal/language" + "github.com/leonardotrapani/hyprvoice/internal/provider" +) + +// ElevenLabsStreamingAdapter implements StreamingAdapter for ElevenLabs real-time transcription +type ElevenLabsStreamingAdapter struct { + endpoint *provider.EndpointConfig + apiKey string + model string + language string + conn *websocket.Conn + resultsCh chan TranscriptionResult + mu sync.Mutex + ctx context.Context + cancel context.CancelFunc + wg sync.WaitGroup + started bool +} + +// ElevenLabs WebSocket message types (outgoing) +type elevenLabsInputAudioChunk struct { + MessageType string `json:"message_type"` + AudioBase64 string `json:"audio_base_64"` + Commit bool `json:"commit"` + SampleRate int `json:"sample_rate"` +} + +// ElevenLabs WebSocket response types (incoming) +type elevenLabsWSMessage struct { + MessageType string `json:"message_type"` + Text string `json:"text,omitempty"` + Error string `json:"error,omitempty"` + SessionID string `json:"session_id,omitempty"` + LanguageCode string `json:"language_code,omitempty"` +} + +// NewElevenLabsStreamingAdapter creates a new streaming adapter for ElevenLabs +// endpoint: the WebSocket endpoint config (e.g., wss://api.elevenlabs.io, /v1/speech-to-text/realtime) +// apiKey: ElevenLabs API key +// model: model ID (e.g., "scribe_v1") +// lang: canonical language code (will be converted to provider format) +func NewElevenLabsStreamingAdapter(endpoint *provider.EndpointConfig, apiKey, model, lang string) *ElevenLabsStreamingAdapter { + return &ElevenLabsStreamingAdapter{ + endpoint: endpoint, + apiKey: apiKey, + model: model, + language: lang, + resultsCh: make(chan TranscriptionResult, 100), + } +} + +// Start initiates the WebSocket connection to ElevenLabs +func (a *ElevenLabsStreamingAdapter) Start(ctx context.Context, lang string) error { + a.mu.Lock() + defer a.mu.Unlock() + + if a.started { + return fmt.Errorf("adapter already started") + } + + // use lang param if provided, otherwise use constructor lang + if lang != "" { + a.language = lang + } + + // create cancelable context + a.ctx, a.cancel = context.WithCancel(ctx) + + // build WebSocket URL with query params + wsURL, err := a.buildURL() + if err != nil { + return fmt.Errorf("build websocket url: %w", err) + } + + // prepare headers with API key + headers := http.Header{} + headers.Set("xi-api-key", a.apiKey) + + // connect to WebSocket + log.Printf("elevenlabs-streaming: connecting to %s", wsURL) + conn, resp, err := websocket.DefaultDialer.DialContext(a.ctx, wsURL, headers) + if err != nil { + if resp != nil { + log.Printf("elevenlabs-streaming: dial failed with status %d", resp.StatusCode) + } + return fmt.Errorf("websocket dial: %w", err) + } + a.conn = conn + a.started = true + + // start reader goroutine + a.wg.Add(1) + go a.readLoop() + + log.Printf("elevenlabs-streaming: connected, model=%s, language=%s", a.model, a.language) + return nil +} + +// buildURL constructs the WebSocket URL with query parameters +func (a *ElevenLabsStreamingAdapter) buildURL() (string, error) { + // parse base URL and path + baseURL := a.endpoint.BaseURL + a.endpoint.Path + + u, err := url.Parse(baseURL) + if err != nil { + return "", fmt.Errorf("parse base url: %w", err) + } + + // add query parameters + q := u.Query() + q.Set("model_id", a.model) + q.Set("audio_format", "pcm_16000") // we use 16kHz PCM + + // add language if specified + providerLang := language.ToProviderFormat(a.language, "elevenlabs") + if providerLang != "" { + q.Set("language_code", providerLang) + } + + // use VAD for automatic commit (easier for real-time use) + q.Set("commit_strategy", "vad") + + u.RawQuery = q.Encode() + return u.String(), nil +} + +// readLoop reads messages from the WebSocket and sends results to the channel +func (a *ElevenLabsStreamingAdapter) readLoop() { + defer a.wg.Done() + defer close(a.resultsCh) + + for { + select { + case <-a.ctx.Done(): + return + default: + } + + _, message, err := a.conn.ReadMessage() + if err != nil { + // check if context was cancelled (normal shutdown) + select { + case <-a.ctx.Done(): + return + default: + } + + // actual error + log.Printf("elevenlabs-streaming: read error: %v", err) + a.resultsCh <- TranscriptionResult{Error: fmt.Errorf("websocket read: %w", err)} + return + } + + // parse message + var msg elevenLabsWSMessage + if err := json.Unmarshal(message, &msg); err != nil { + log.Printf("elevenlabs-streaming: parse error: %v", err) + continue + } + + // handle different message types + switch msg.MessageType { + case "session_started": + log.Printf("elevenlabs-streaming: session started, id=%s", msg.SessionID) + + case "partial_transcript": + // interim result + if msg.Text != "" { + a.resultsCh <- TranscriptionResult{Text: msg.Text, IsFinal: false} + } + + case "committed_transcript", "committed_transcript_with_timestamps": + // final result + if msg.Text != "" { + log.Printf("elevenlabs-streaming: committed: %q", msg.Text) + a.resultsCh <- TranscriptionResult{Text: msg.Text, IsFinal: true} + } + + case "error", "auth_error", "quota_exceeded", "rate_limited", + "queue_overflow", "resource_exhausted", "session_time_limit_exceeded", + "input_error", "chunk_size_exceeded", "insufficient_audio_activity", + "transcriber_error", "commit_throttled", "unaccepted_terms": + // error message + errMsg := msg.Error + if errMsg == "" { + errMsg = msg.MessageType + } + log.Printf("elevenlabs-streaming: error: %s", errMsg) + a.resultsCh <- TranscriptionResult{Error: fmt.Errorf("elevenlabs: %s", errMsg)} + + default: + log.Printf("elevenlabs-streaming: unknown message type: %s", msg.MessageType) + } + } +} + +// SendChunk sends audio data to the WebSocket +func (a *ElevenLabsStreamingAdapter) SendChunk(audio []byte) error { + a.mu.Lock() + defer a.mu.Unlock() + + if !a.started || a.conn == nil { + return fmt.Errorf("adapter not started") + } + + // check context + select { + case <-a.ctx.Done(): + return a.ctx.Err() + default: + } + + // encode audio as base64 + audioB64 := base64.StdEncoding.EncodeToString(audio) + + // create message + msg := elevenLabsInputAudioChunk{ + MessageType: "input_audio_chunk", + AudioBase64: audioB64, + Commit: false, // let VAD handle commits + SampleRate: 16000, + } + + // send as JSON + if err := a.conn.WriteJSON(msg); err != nil { + return fmt.Errorf("websocket write: %w", err) + } + + return nil +} + +// Results returns the channel for receiving transcription results +func (a *ElevenLabsStreamingAdapter) Results() <-chan TranscriptionResult { + return a.resultsCh +} + +// Close gracefully closes the WebSocket connection +func (a *ElevenLabsStreamingAdapter) Close() error { + a.mu.Lock() + + if !a.started { + a.mu.Unlock() + return nil + } + + // cancel context first to signal reader to stop + if a.cancel != nil { + a.cancel() + } + + // get conn ref while holding lock + conn := a.conn + + a.started = false + a.mu.Unlock() + + // close websocket outside of lock (readLoop may be blocked on read) + if conn != nil { + // send close frame (best effort) + _ = conn.WriteMessage(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseNormalClosure, "")) + conn.Close() + } + + // wait for reader to finish + a.wg.Wait() + + log.Printf("elevenlabs-streaming: closed") + return nil +} diff --git a/internal/transcriber/adapter_elevenlabs_streaming_test.go b/internal/transcriber/adapter_elevenlabs_streaming_test.go new file mode 100644 index 0000000..b5c4248 --- /dev/null +++ b/internal/transcriber/adapter_elevenlabs_streaming_test.go @@ -0,0 +1,426 @@ +package transcriber + +import ( + "context" + "encoding/base64" + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/gorilla/websocket" + "github.com/leonardotrapani/hyprvoice/internal/provider" +) + +var upgrader = websocket.Upgrader{ + CheckOrigin: func(r *http.Request) bool { return true }, +} + +// mockElevenLabsServer creates a test WebSocket server that simulates ElevenLabs +func mockElevenLabsServer(t *testing.T, handler func(*websocket.Conn)) *httptest.Server { + return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // check API key header + apiKey := r.Header.Get("xi-api-key") + if apiKey == "" { + http.Error(w, "missing api key", http.StatusUnauthorized) + return + } + + // upgrade to websocket + conn, err := upgrader.Upgrade(w, r, nil) + if err != nil { + t.Logf("upgrade error: %v", err) + return + } + defer conn.Close() + + handler(conn) + })) +} + +func TestElevenLabsStreamingAdapter_ImplementsInterface(t *testing.T) { + var _ StreamingAdapter = (*ElevenLabsStreamingAdapter)(nil) +} + +func TestElevenLabsStreamingAdapter_Start(t *testing.T) { + server := mockElevenLabsServer(t, func(conn *websocket.Conn) { + // send session started + msg := elevenLabsWSMessage{ + MessageType: "session_started", + SessionID: "test-session-123", + } + conn.WriteJSON(msg) + + // keep connection open + for { + _, _, err := conn.ReadMessage() + if err != nil { + return + } + } + }) + defer server.Close() + + // convert http://... to ws://... + wsURL := "ws" + strings.TrimPrefix(server.URL, "http") + + adapter := NewElevenLabsStreamingAdapter( + &provider.EndpointConfig{BaseURL: wsURL, Path: ""}, + "test-api-key", + "scribe_v1", + "en", + ) + + ctx := context.Background() + err := adapter.Start(ctx, "") + if err != nil { + t.Fatalf("Start() error: %v", err) + } + + // give time for session_started to be received + time.Sleep(50 * time.Millisecond) + + err = adapter.Close() + if err != nil { + t.Errorf("Close() error: %v", err) + } +} + +func TestElevenLabsStreamingAdapter_SendChunk(t *testing.T) { + receivedChunks := make(chan []byte, 10) + + server := mockElevenLabsServer(t, func(conn *websocket.Conn) { + // send session started + conn.WriteJSON(elevenLabsWSMessage{ + MessageType: "session_started", + SessionID: "test-session", + }) + + // read incoming messages + for { + _, message, err := conn.ReadMessage() + if err != nil { + return + } + + var msg elevenLabsInputAudioChunk + if err := json.Unmarshal(message, &msg); err != nil { + continue + } + + if msg.MessageType == "input_audio_chunk" { + decoded, _ := base64.StdEncoding.DecodeString(msg.AudioBase64) + receivedChunks <- decoded + } + } + }) + defer server.Close() + + wsURL := "ws" + strings.TrimPrefix(server.URL, "http") + + adapter := NewElevenLabsStreamingAdapter( + &provider.EndpointConfig{BaseURL: wsURL, Path: ""}, + "test-api-key", + "scribe_v1", + "en", + ) + + ctx := context.Background() + if err := adapter.Start(ctx, ""); err != nil { + t.Fatalf("Start() error: %v", err) + } + defer adapter.Close() + + // send audio chunk + testAudio := []byte{0x01, 0x02, 0x03, 0x04} + if err := adapter.SendChunk(testAudio); err != nil { + t.Fatalf("SendChunk() error: %v", err) + } + + // verify received + select { + case received := <-receivedChunks: + if string(received) != string(testAudio) { + t.Errorf("received audio mismatch: got %v, want %v", received, testAudio) + } + case <-time.After(time.Second): + t.Error("timeout waiting for audio chunk") + } +} + +func TestElevenLabsStreamingAdapter_Results(t *testing.T) { + server := mockElevenLabsServer(t, func(conn *websocket.Conn) { + // send session started + conn.WriteJSON(elevenLabsWSMessage{ + MessageType: "session_started", + SessionID: "test-session", + }) + + // send partial transcript + conn.WriteJSON(elevenLabsWSMessage{ + MessageType: "partial_transcript", + Text: "hello", + }) + + // send committed transcript + conn.WriteJSON(elevenLabsWSMessage{ + MessageType: "committed_transcript", + Text: "hello world", + }) + + // keep connection open + for { + _, _, err := conn.ReadMessage() + if err != nil { + return + } + } + }) + defer server.Close() + + wsURL := "ws" + strings.TrimPrefix(server.URL, "http") + + adapter := NewElevenLabsStreamingAdapter( + &provider.EndpointConfig{BaseURL: wsURL, Path: ""}, + "test-api-key", + "scribe_v1", + "en", + ) + + ctx := context.Background() + if err := adapter.Start(ctx, ""); err != nil { + t.Fatalf("Start() error: %v", err) + } + defer adapter.Close() + + results := adapter.Results() + + // check partial result + select { + case result := <-results: + if result.Error != nil { + t.Fatalf("unexpected error: %v", result.Error) + } + if result.Text != "hello" { + t.Errorf("partial text: got %q, want %q", result.Text, "hello") + } + if result.IsFinal { + t.Error("partial result should not be final") + } + case <-time.After(time.Second): + t.Fatal("timeout waiting for partial result") + } + + // check final result + select { + case result := <-results: + if result.Error != nil { + t.Fatalf("unexpected error: %v", result.Error) + } + if result.Text != "hello world" { + t.Errorf("final text: got %q, want %q", result.Text, "hello world") + } + if !result.IsFinal { + t.Error("committed result should be final") + } + case <-time.After(time.Second): + t.Fatal("timeout waiting for final result") + } +} + +func TestElevenLabsStreamingAdapter_ErrorMessages(t *testing.T) { + server := mockElevenLabsServer(t, func(conn *websocket.Conn) { + // send session started + conn.WriteJSON(elevenLabsWSMessage{ + MessageType: "session_started", + SessionID: "test-session", + }) + + // send error + conn.WriteJSON(elevenLabsWSMessage{ + MessageType: "error", + Error: "test error message", + }) + + // keep connection open + for { + _, _, err := conn.ReadMessage() + if err != nil { + return + } + } + }) + defer server.Close() + + wsURL := "ws" + strings.TrimPrefix(server.URL, "http") + + adapter := NewElevenLabsStreamingAdapter( + &provider.EndpointConfig{BaseURL: wsURL, Path: ""}, + "test-api-key", + "scribe_v1", + "en", + ) + + ctx := context.Background() + if err := adapter.Start(ctx, ""); err != nil { + t.Fatalf("Start() error: %v", err) + } + defer adapter.Close() + + results := adapter.Results() + + // check error result + select { + case result := <-results: + if result.Error == nil { + t.Fatal("expected error, got nil") + } + if !strings.Contains(result.Error.Error(), "test error message") { + t.Errorf("error message: got %q, want to contain %q", result.Error.Error(), "test error message") + } + case <-time.After(time.Second): + t.Fatal("timeout waiting for error result") + } +} + +func TestElevenLabsStreamingAdapter_LanguageConversion(t *testing.T) { + var receivedURL string + + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + receivedURL = r.URL.String() + + // check API key header + if r.Header.Get("xi-api-key") == "" { + http.Error(w, "missing api key", http.StatusUnauthorized) + return + } + + conn, err := upgrader.Upgrade(w, r, nil) + if err != nil { + return + } + defer conn.Close() + + conn.WriteJSON(elevenLabsWSMessage{ + MessageType: "session_started", + SessionID: "test-session", + }) + + for { + _, _, err := conn.ReadMessage() + if err != nil { + return + } + } + })) + defer server.Close() + + wsURL := "ws" + strings.TrimPrefix(server.URL, "http") + + adapter := NewElevenLabsStreamingAdapter( + &provider.EndpointConfig{BaseURL: wsURL, Path: "/v1/speech-to-text/realtime"}, + "test-api-key", + "scribe_v1", + "es", // Spanish + ) + + ctx := context.Background() + if err := adapter.Start(ctx, ""); err != nil { + t.Fatalf("Start() error: %v", err) + } + adapter.Close() + + // verify language_code was set + if !strings.Contains(receivedURL, "language_code=es") { + t.Errorf("URL should contain language_code=es, got: %s", receivedURL) + } + + // verify model_id was set + if !strings.Contains(receivedURL, "model_id=scribe_v1") { + t.Errorf("URL should contain model_id=scribe_v1, got: %s", receivedURL) + } + + // verify audio_format was set + if !strings.Contains(receivedURL, "audio_format=pcm_16000") { + t.Errorf("URL should contain audio_format=pcm_16000, got: %s", receivedURL) + } +} + +func TestElevenLabsStreamingAdapter_Close(t *testing.T) { + server := mockElevenLabsServer(t, func(conn *websocket.Conn) { + conn.WriteJSON(elevenLabsWSMessage{ + MessageType: "session_started", + SessionID: "test-session", + }) + + for { + _, _, err := conn.ReadMessage() + if err != nil { + return + } + } + }) + defer server.Close() + + wsURL := "ws" + strings.TrimPrefix(server.URL, "http") + + adapter := NewElevenLabsStreamingAdapter( + &provider.EndpointConfig{BaseURL: wsURL, Path: ""}, + "test-api-key", + "scribe_v1", + "en", + ) + + ctx := context.Background() + if err := adapter.Start(ctx, ""); err != nil { + t.Fatalf("Start() error: %v", err) + } + + // close should not block + done := make(chan struct{}) + go func() { + adapter.Close() + close(done) + }() + + select { + case <-done: + // ok + case <-time.After(2 * time.Second): + t.Fatal("Close() blocked for too long") + } + + // results channel should be closed + _, ok := <-adapter.Results() + if ok { + // there might be buffered results, drain them + for range adapter.Results() { + } + } +} + +func TestElevenLabsStreamingAdapter_NotStarted(t *testing.T) { + adapter := NewElevenLabsStreamingAdapter( + &provider.EndpointConfig{BaseURL: "wss://api.elevenlabs.io", Path: "/v1/speech-to-text/realtime"}, + "test-api-key", + "scribe_v1", + "en", + ) + + // SendChunk should fail when not started + err := adapter.SendChunk([]byte{0x01, 0x02}) + if err == nil { + t.Error("SendChunk() should fail when adapter not started") + } + if !strings.Contains(err.Error(), "not started") { + t.Errorf("error should mention 'not started', got: %v", err) + } + + // Close should not fail when not started + err = adapter.Close() + if err != nil { + t.Errorf("Close() should not fail when not started: %v", err) + } +} diff --git a/progress.txt b/progress.txt index 25a8e39..94ef7aa 100644 --- a/progress.txt +++ b/progress.txt @@ -306,4 +306,18 @@ Started: Sun Feb 1 12:22:47 AM CET 2026 - Shows confirm dialog "Try again?" - if yes, recursively calls `editTranscription()` to let user fix - Config only saved AFTER validation passes (no save on cancel) - Leverages existing `ValidateModelLanguage` which returns error with supported languages list -- All tests passing, typecheck passes \ No newline at end of file +- All tests passing, typecheck passes + +### Task 31: Create ElevenLabs StreamingAdapter +- Created `internal/transcriber/adapter_elevenlabs_streaming.go` +- Added gorilla/websocket dependency +- ElevenLabsStreamingAdapter struct: endpoint, apiKey, model, language, conn, resultsCh, mutex, ctx/cancel, WaitGroup +- Start(): connects to wss://api.elevenlabs.io/v1/speech-to-text/realtime with xi-api-key header +- Query params: model_id, language_code, audio_format=pcm_16000, commit_strategy=vad +- Language conversion via language.ToProviderFormat(lang, "elevenlabs") +- SendChunk(): sends input_audio_chunk JSON message with base64-encoded audio +- readLoop goroutine: parses session_started, partial_transcript, committed_transcript messages +- Handles all ElevenLabs error types (auth_error, quota_exceeded, rate_limited, etc.) +- Close(): cancels context, sends close frame, waits for reader goroutine +- Comprehensive tests with mock WebSocket server +- All tests passing with -race flag, typecheck passes \ No newline at end of file diff --git a/tasks/prd.jsonc b/tasks/prd.jsonc index 4943e06..14f3da6 100644 --- a/tasks/prd.jsonc +++ b/tasks/prd.jsonc @@ -739,7 +739,7 @@ "Close() terminates cleanly", "Typecheck passes" ], - "passes": false + "passes": true }, { "title": "Add reconnection logic to ElevenLabs StreamingAdapter",