add StreamingTranscriber wrapper for real-time transcription
This commit is contained in:
@@ -0,0 +1,127 @@
|
||||
package transcriber
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/leonardotrapani/hyprvoice/internal/recording"
|
||||
)
|
||||
|
||||
// StreamingTranscriber wraps a StreamingAdapter and implements the Transcriber interface.
|
||||
// It streams audio chunks to the adapter in real-time and accumulates transcription results.
|
||||
type StreamingTranscriber struct {
|
||||
adapter StreamingAdapter
|
||||
language string
|
||||
|
||||
// accumulated final text
|
||||
finalText strings.Builder
|
||||
mu sync.Mutex
|
||||
|
||||
// coordination
|
||||
ctx context.Context
|
||||
cancel context.CancelFunc
|
||||
wg sync.WaitGroup
|
||||
}
|
||||
|
||||
func NewStreamingTranscriber(adapter StreamingAdapter, language string) *StreamingTranscriber {
|
||||
return &StreamingTranscriber{
|
||||
adapter: adapter,
|
||||
language: language,
|
||||
}
|
||||
}
|
||||
|
||||
func (t *StreamingTranscriber) Start(ctx context.Context, frameCh <-chan recording.AudioFrame) (<-chan error, error) {
|
||||
t.ctx, t.cancel = context.WithCancel(ctx)
|
||||
|
||||
if err := t.adapter.Start(t.ctx, t.language); err != nil {
|
||||
t.cancel()
|
||||
return nil, err
|
||||
}
|
||||
|
||||
errCh := make(chan error, 2)
|
||||
|
||||
// goroutine 1: read audio frames and send to adapter
|
||||
t.wg.Add(1)
|
||||
go t.sendAudio(frameCh, errCh)
|
||||
|
||||
// goroutine 2: read results from adapter and accumulate
|
||||
t.wg.Add(1)
|
||||
go t.receiveResults(errCh)
|
||||
|
||||
return errCh, nil
|
||||
}
|
||||
|
||||
func (t *StreamingTranscriber) sendAudio(frameCh <-chan recording.AudioFrame, errCh chan<- error) {
|
||||
defer t.wg.Done()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-t.ctx.Done():
|
||||
return
|
||||
case frame, ok := <-frameCh:
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if err := t.adapter.SendChunk(frame.Data); err != nil {
|
||||
select {
|
||||
case errCh <- err:
|
||||
default:
|
||||
}
|
||||
// don't treat send errors as fatal - adapter may handle reconnection
|
||||
log.Printf("streaming transcriber: send error: %v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (t *StreamingTranscriber) receiveResults(errCh chan<- error) {
|
||||
defer t.wg.Done()
|
||||
|
||||
resultsCh := t.adapter.Results()
|
||||
for {
|
||||
select {
|
||||
case <-t.ctx.Done():
|
||||
return
|
||||
case result, ok := <-resultsCh:
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
if result.Error != nil {
|
||||
select {
|
||||
case errCh <- result.Error:
|
||||
default:
|
||||
}
|
||||
log.Printf("streaming transcriber: result error: %v", result.Error)
|
||||
continue
|
||||
}
|
||||
if result.IsFinal && result.Text != "" {
|
||||
t.mu.Lock()
|
||||
if t.finalText.Len() > 0 {
|
||||
t.finalText.WriteString(" ")
|
||||
}
|
||||
t.finalText.WriteString(result.Text)
|
||||
t.mu.Unlock()
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (t *StreamingTranscriber) Stop(ctx context.Context) error {
|
||||
if t.cancel != nil {
|
||||
t.cancel()
|
||||
}
|
||||
|
||||
// wait for goroutines to finish
|
||||
t.wg.Wait()
|
||||
|
||||
// close the adapter
|
||||
return t.adapter.Close()
|
||||
}
|
||||
|
||||
func (t *StreamingTranscriber) GetFinalTranscription() (string, error) {
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
return t.finalText.String(), nil
|
||||
}
|
||||
@@ -497,3 +497,318 @@ func TestTranscriptionAdapter(t *testing.T) {
|
||||
t.Errorf("Transcribe() = %q, want %q", result, "test result")
|
||||
}
|
||||
}
|
||||
|
||||
// MockStreamingAdapter implements StreamingAdapter for testing
|
||||
type MockStreamingAdapter struct {
|
||||
StartFunc func(ctx context.Context, language string) error
|
||||
SendChunkFunc func(audio []byte) error
|
||||
ResultsFunc func() <-chan TranscriptionResult
|
||||
CloseFunc func() error
|
||||
|
||||
resultsCh chan TranscriptionResult
|
||||
}
|
||||
|
||||
func NewMockStreamingAdapter() *MockStreamingAdapter {
|
||||
return &MockStreamingAdapter{
|
||||
resultsCh: make(chan TranscriptionResult, 10),
|
||||
}
|
||||
}
|
||||
|
||||
func (m *MockStreamingAdapter) Start(ctx context.Context, language string) error {
|
||||
if m.StartFunc != nil {
|
||||
return m.StartFunc(ctx, language)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *MockStreamingAdapter) SendChunk(audio []byte) error {
|
||||
if m.SendChunkFunc != nil {
|
||||
return m.SendChunkFunc(audio)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *MockStreamingAdapter) Results() <-chan TranscriptionResult {
|
||||
if m.ResultsFunc != nil {
|
||||
return m.ResultsFunc()
|
||||
}
|
||||
return m.resultsCh
|
||||
}
|
||||
|
||||
func (m *MockStreamingAdapter) Close() error {
|
||||
if m.CloseFunc != nil {
|
||||
return m.CloseFunc()
|
||||
}
|
||||
close(m.resultsCh)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *MockStreamingAdapter) SendResult(result TranscriptionResult) {
|
||||
m.resultsCh <- result
|
||||
}
|
||||
|
||||
func TestStreamingTranscriber_Start(t *testing.T) {
|
||||
adapter := NewMockStreamingAdapter()
|
||||
transcriber := NewStreamingTranscriber(adapter, "en")
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
frameCh := make(chan recording.AudioFrame, 10)
|
||||
|
||||
errCh, err := transcriber.Start(ctx, frameCh)
|
||||
if err != nil {
|
||||
t.Errorf("Start() error = %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
if errCh == nil {
|
||||
t.Errorf("Start() returned nil error channel")
|
||||
}
|
||||
|
||||
close(frameCh)
|
||||
err = transcriber.Stop(ctx)
|
||||
if err != nil {
|
||||
t.Errorf("Stop() error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStreamingTranscriber_StartError(t *testing.T) {
|
||||
adapter := NewMockStreamingAdapter()
|
||||
adapter.StartFunc = func(ctx context.Context, language string) error {
|
||||
return fmt.Errorf("connection failed")
|
||||
}
|
||||
|
||||
transcriber := NewStreamingTranscriber(adapter, "en")
|
||||
|
||||
ctx := context.Background()
|
||||
frameCh := make(chan recording.AudioFrame, 10)
|
||||
|
||||
_, err := transcriber.Start(ctx, frameCh)
|
||||
if err == nil {
|
||||
t.Errorf("Start() should fail when adapter.Start fails")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStreamingTranscriber_AccumulatesResults(t *testing.T) {
|
||||
adapter := NewMockStreamingAdapter()
|
||||
transcriber := NewStreamingTranscriber(adapter, "en")
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
frameCh := make(chan recording.AudioFrame, 10)
|
||||
|
||||
_, err := transcriber.Start(ctx, frameCh)
|
||||
if err != nil {
|
||||
t.Errorf("Start() error = %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
// send some final results
|
||||
adapter.SendResult(TranscriptionResult{Text: "hello", IsFinal: true})
|
||||
adapter.SendResult(TranscriptionResult{Text: "world", IsFinal: true})
|
||||
|
||||
// give time for results to be processed
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
close(frameCh)
|
||||
err = transcriber.Stop(ctx)
|
||||
if err != nil {
|
||||
t.Errorf("Stop() error = %v", err)
|
||||
}
|
||||
|
||||
result, err := transcriber.GetFinalTranscription()
|
||||
if err != nil {
|
||||
t.Errorf("GetFinalTranscription() error = %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
if result != "hello world" {
|
||||
t.Errorf("GetFinalTranscription() = %q, want %q", result, "hello world")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStreamingTranscriber_IgnoresPartialResults(t *testing.T) {
|
||||
adapter := NewMockStreamingAdapter()
|
||||
transcriber := NewStreamingTranscriber(adapter, "en")
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
frameCh := make(chan recording.AudioFrame, 10)
|
||||
|
||||
_, err := transcriber.Start(ctx, frameCh)
|
||||
if err != nil {
|
||||
t.Errorf("Start() error = %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
// partial results should be ignored
|
||||
adapter.SendResult(TranscriptionResult{Text: "hel", IsFinal: false})
|
||||
adapter.SendResult(TranscriptionResult{Text: "hello", IsFinal: true})
|
||||
adapter.SendResult(TranscriptionResult{Text: "hello wor", IsFinal: false})
|
||||
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
close(frameCh)
|
||||
err = transcriber.Stop(ctx)
|
||||
if err != nil {
|
||||
t.Errorf("Stop() error = %v", err)
|
||||
}
|
||||
|
||||
result, err := transcriber.GetFinalTranscription()
|
||||
if err != nil {
|
||||
t.Errorf("GetFinalTranscription() error = %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
if result != "hello" {
|
||||
t.Errorf("GetFinalTranscription() = %q, want %q", result, "hello")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStreamingTranscriber_HandlesErrors(t *testing.T) {
|
||||
adapter := NewMockStreamingAdapter()
|
||||
transcriber := NewStreamingTranscriber(adapter, "en")
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
frameCh := make(chan recording.AudioFrame, 10)
|
||||
|
||||
errCh, err := transcriber.Start(ctx, frameCh)
|
||||
if err != nil {
|
||||
t.Errorf("Start() error = %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
// send an error result
|
||||
adapter.SendResult(TranscriptionResult{Error: fmt.Errorf("transcription error")})
|
||||
|
||||
// error should be received on errCh
|
||||
select {
|
||||
case e := <-errCh:
|
||||
if e == nil {
|
||||
t.Errorf("expected error on errCh")
|
||||
}
|
||||
case <-time.After(100 * time.Millisecond):
|
||||
t.Errorf("timeout waiting for error on errCh")
|
||||
}
|
||||
|
||||
close(frameCh)
|
||||
_ = transcriber.Stop(ctx)
|
||||
}
|
||||
|
||||
func TestStreamingTranscriber_SendsAudioChunks(t *testing.T) {
|
||||
var receivedChunks [][]byte
|
||||
adapter := NewMockStreamingAdapter()
|
||||
adapter.SendChunkFunc = func(audio []byte) error {
|
||||
chunk := make([]byte, len(audio))
|
||||
copy(chunk, audio)
|
||||
receivedChunks = append(receivedChunks, chunk)
|
||||
return nil
|
||||
}
|
||||
|
||||
transcriber := NewStreamingTranscriber(adapter, "en")
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
frameCh := make(chan recording.AudioFrame, 10)
|
||||
|
||||
_, err := transcriber.Start(ctx, frameCh)
|
||||
if err != nil {
|
||||
t.Errorf("Start() error = %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
// send audio frames
|
||||
frameCh <- recording.AudioFrame{Data: []byte{1, 2, 3, 4}}
|
||||
frameCh <- recording.AudioFrame{Data: []byte{5, 6, 7, 8}}
|
||||
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
|
||||
close(frameCh)
|
||||
err = transcriber.Stop(ctx)
|
||||
if err != nil {
|
||||
t.Errorf("Stop() error = %v", err)
|
||||
}
|
||||
|
||||
if len(receivedChunks) != 2 {
|
||||
t.Errorf("expected 2 chunks, got %d", len(receivedChunks))
|
||||
}
|
||||
}
|
||||
|
||||
func TestStreamingTranscriber_ContextCancellation(t *testing.T) {
|
||||
adapter := NewMockStreamingAdapter()
|
||||
transcriber := NewStreamingTranscriber(adapter, "en")
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
defer cancel()
|
||||
|
||||
frameCh := make(chan recording.AudioFrame, 10)
|
||||
|
||||
_, err := transcriber.Start(ctx, frameCh)
|
||||
if err != nil {
|
||||
t.Errorf("Start() error = %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
// cancel context
|
||||
cancel()
|
||||
|
||||
// stop should complete without hanging
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
_ = transcriber.Stop(context.Background())
|
||||
close(done)
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-done:
|
||||
// success
|
||||
case <-time.After(2 * time.Second):
|
||||
t.Errorf("Stop() timed out after context cancellation")
|
||||
}
|
||||
}
|
||||
|
||||
func TestStreamingTranscriber_GetFinalTranscriptionSafe(t *testing.T) {
|
||||
adapter := NewMockStreamingAdapter()
|
||||
transcriber := NewStreamingTranscriber(adapter, "en")
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
frameCh := make(chan recording.AudioFrame, 10)
|
||||
|
||||
_, err := transcriber.Start(ctx, frameCh)
|
||||
if err != nil {
|
||||
t.Errorf("Start() error = %v", err)
|
||||
return
|
||||
}
|
||||
|
||||
// call GetFinalTranscription concurrently while results are being added
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
for i := 0; i < 100; i++ {
|
||||
_, _ = transcriber.GetFinalTranscription()
|
||||
time.Sleep(time.Millisecond)
|
||||
}
|
||||
close(done)
|
||||
}()
|
||||
|
||||
// send results concurrently
|
||||
for i := 0; i < 10; i++ {
|
||||
adapter.SendResult(TranscriptionResult{Text: "word", IsFinal: true})
|
||||
time.Sleep(5 * time.Millisecond)
|
||||
}
|
||||
|
||||
<-done
|
||||
|
||||
close(frameCh)
|
||||
err = transcriber.Stop(ctx)
|
||||
if err != nil {
|
||||
t.Errorf("Stop() error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -50,3 +50,15 @@ Started: Sun Feb 1 12:22:47 AM CET 2026
|
||||
- `TranscriptionResult` struct: Text, IsFinal, Error fields
|
||||
- `StreamingAdapter` interface: Start, SendChunk, Results, Close methods
|
||||
- All tests passing, typecheck passes
|
||||
|
||||
### Task 6: Create StreamingTranscriber wrapper
|
||||
- Created `internal/transcriber/streaming_transcriber.go`
|
||||
- StreamingTranscriber struct with: adapter, language, finalText builder, mutex, ctx/cancel, WaitGroup
|
||||
- Start() creates cancelable context, starts adapter, spawns 2 goroutines
|
||||
- Goroutine 1: reads frames from channel, calls adapter.SendChunk()
|
||||
- Goroutine 2: reads from adapter.Results(), accumulates final results with space separator
|
||||
- Stop() cancels context, waits for goroutines, closes adapter
|
||||
- GetFinalTranscription() returns accumulated text with mutex protection
|
||||
- Added MockStreamingAdapter and comprehensive tests
|
||||
- Tests verify: start/stop, result accumulation, partial result filtering, error handling, concurrent access
|
||||
- All tests passing with -race flag, typecheck passes
|
||||
|
||||
+1
-1
@@ -154,7 +154,7 @@
|
||||
"No race conditions (run with -race flag)",
|
||||
"Typecheck passes"
|
||||
],
|
||||
"passes": false
|
||||
"passes": true
|
||||
},
|
||||
{
|
||||
"title": "Write tests for Model, Provider, and interfaces",
|
||||
|
||||
Reference in New Issue
Block a user