feat: better language selection
This commit is contained in:
@@ -7,11 +7,11 @@ import (
|
||||
"log"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
"github.com/leonardotrapani/hyprvoice/internal/language"
|
||||
"github.com/leonardotrapani/hyprvoice/internal/provider"
|
||||
)
|
||||
|
||||
@@ -21,6 +21,7 @@ type DeepgramAdapter struct {
|
||||
apiKey string
|
||||
model string
|
||||
language string
|
||||
keywords []string
|
||||
conn *websocket.Conn
|
||||
resultsCh chan TranscriptionResult
|
||||
mu sync.Mutex
|
||||
@@ -82,13 +83,14 @@ type deepgramError struct {
|
||||
// endpoint: the WebSocket endpoint config (e.g., wss://api.deepgram.com, /v1/listen)
|
||||
// apiKey: Deepgram API key
|
||||
// model: model ID (e.g., "nova-3")
|
||||
// lang: canonical language code (will be converted to provider format)
|
||||
func NewDeepgramAdapter(endpoint *provider.EndpointConfig, apiKey, model, lang string) *DeepgramAdapter {
|
||||
// lang: provider language code
|
||||
func NewDeepgramAdapter(endpoint *provider.EndpointConfig, apiKey, model, lang string, keywords []string) *DeepgramAdapter {
|
||||
return &DeepgramAdapter{
|
||||
endpoint: endpoint,
|
||||
apiKey: apiKey,
|
||||
model: model,
|
||||
language: lang,
|
||||
keywords: keywords,
|
||||
resultsCh: make(chan TranscriptionResult, 100),
|
||||
maxRetries: 3,
|
||||
retryDelays: defaultRetryDelays,
|
||||
@@ -228,9 +230,12 @@ func (a *DeepgramAdapter) buildURL() (string, error) {
|
||||
q.Set("punctuate", "true")
|
||||
|
||||
// add language if specified
|
||||
providerLang := language.ToProviderFormat(a.language, "deepgram")
|
||||
if providerLang != "" {
|
||||
q.Set("language", providerLang)
|
||||
if a.language != "" {
|
||||
q.Set("language", a.language)
|
||||
}
|
||||
|
||||
if len(a.keywords) > 0 {
|
||||
q.Set("keywords", strings.Join(a.keywords, ","))
|
||||
}
|
||||
|
||||
u.RawQuery = q.Encode()
|
||||
|
||||
@@ -8,8 +8,8 @@ import (
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
|
||||
"github.com/leonardotrapani/hyprvoice/internal/language"
|
||||
"github.com/leonardotrapani/hyprvoice/internal/provider"
|
||||
)
|
||||
|
||||
@@ -19,6 +19,7 @@ type DeepgramBatchAdapter struct {
|
||||
apiKey string
|
||||
model string
|
||||
language string
|
||||
keywords []string
|
||||
}
|
||||
|
||||
// deepgramBatchResponse is the response from the pre-recorded API
|
||||
@@ -36,12 +37,13 @@ type deepgramBatchChannel struct {
|
||||
}
|
||||
|
||||
// NewDeepgramBatchAdapter creates a new batch adapter for Deepgram
|
||||
func NewDeepgramBatchAdapter(endpoint *provider.EndpointConfig, apiKey, model, lang string) *DeepgramBatchAdapter {
|
||||
func NewDeepgramBatchAdapter(endpoint *provider.EndpointConfig, apiKey, model, lang string, keywords []string) *DeepgramBatchAdapter {
|
||||
return &DeepgramBatchAdapter{
|
||||
endpoint: endpoint,
|
||||
apiKey: apiKey,
|
||||
model: model,
|
||||
language: lang,
|
||||
keywords: keywords,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -116,9 +118,12 @@ func (a *DeepgramBatchAdapter) buildURL() (string, error) {
|
||||
q.Set("punctuate", "true")
|
||||
|
||||
// add language if specified
|
||||
providerLang := language.ToProviderFormat(a.language, "deepgram")
|
||||
if providerLang != "" {
|
||||
q.Set("language", providerLang)
|
||||
if a.language != "" {
|
||||
q.Set("language", a.language)
|
||||
}
|
||||
|
||||
if len(a.keywords) > 0 {
|
||||
q.Set("keywords", strings.Join(a.keywords, ","))
|
||||
}
|
||||
|
||||
u.RawQuery = q.Encode()
|
||||
|
||||
@@ -23,7 +23,7 @@ func TestDeepgramAdapter_Creation(t *testing.T) {
|
||||
BaseURL: "wss://api.deepgram.com",
|
||||
Path: "/v1/listen",
|
||||
}
|
||||
adapter := NewDeepgramAdapter(endpoint, "test-api-key", "nova-3", "en")
|
||||
adapter := NewDeepgramAdapter(endpoint, "test-api-key", "nova-3", "en", nil)
|
||||
|
||||
if adapter.apiKey != "test-api-key" {
|
||||
t.Errorf("apiKey = %q, want %q", adapter.apiKey, "test-api-key")
|
||||
@@ -72,7 +72,7 @@ func TestDeepgramAdapter_BuildURL(t *testing.T) {
|
||||
BaseURL: "wss://api.deepgram.com",
|
||||
Path: "/v1/listen",
|
||||
}
|
||||
adapter := NewDeepgramAdapter(endpoint, "test-key", tt.model, tt.language)
|
||||
adapter := NewDeepgramAdapter(endpoint, "test-key", tt.model, tt.language, nil)
|
||||
|
||||
url, err := adapter.buildURL()
|
||||
if err != nil {
|
||||
@@ -93,7 +93,7 @@ func TestDeepgramAdapter_SendChunkNotStarted(t *testing.T) {
|
||||
BaseURL: "wss://api.deepgram.com",
|
||||
Path: "/v1/listen",
|
||||
}
|
||||
adapter := NewDeepgramAdapter(endpoint, "test-key", "nova-3", "en")
|
||||
adapter := NewDeepgramAdapter(endpoint, "test-key", "nova-3", "en", nil)
|
||||
|
||||
err := adapter.SendChunk([]byte("audio data"))
|
||||
if err == nil {
|
||||
@@ -109,7 +109,7 @@ func TestDeepgramAdapter_CloseNotStarted(t *testing.T) {
|
||||
BaseURL: "wss://api.deepgram.com",
|
||||
Path: "/v1/listen",
|
||||
}
|
||||
adapter := NewDeepgramAdapter(endpoint, "test-key", "nova-3", "en")
|
||||
adapter := NewDeepgramAdapter(endpoint, "test-key", "nova-3", "en", nil)
|
||||
|
||||
// closing not-started adapter should not error
|
||||
err := adapter.Close()
|
||||
@@ -173,7 +173,7 @@ func TestDeepgramAdapter_StartAndClose(t *testing.T) {
|
||||
Path: "",
|
||||
}
|
||||
|
||||
adapter := NewDeepgramAdapter(endpoint, "test-api-key", "nova-3", "en")
|
||||
adapter := NewDeepgramAdapter(endpoint, "test-api-key", "nova-3", "en", nil)
|
||||
|
||||
ctx := context.Background()
|
||||
if err := adapter.Start(ctx, ""); err != nil {
|
||||
@@ -229,7 +229,7 @@ func TestDeepgramAdapter_ReceivesResults(t *testing.T) {
|
||||
|
||||
wsURL := "ws" + strings.TrimPrefix(server.URL, "http")
|
||||
endpoint := &provider.EndpointConfig{BaseURL: wsURL, Path: ""}
|
||||
adapter := NewDeepgramAdapter(endpoint, "test-api-key", "nova-3", "en")
|
||||
adapter := NewDeepgramAdapter(endpoint, "test-api-key", "nova-3", "en", nil)
|
||||
|
||||
ctx := context.Background()
|
||||
if err := adapter.Start(ctx, ""); err != nil {
|
||||
@@ -312,7 +312,7 @@ func TestDeepgramAdapter_SendsRawBinaryAudio(t *testing.T) {
|
||||
|
||||
wsURL := "ws" + strings.TrimPrefix(server.URL, "http")
|
||||
endpoint := &provider.EndpointConfig{BaseURL: wsURL, Path: ""}
|
||||
adapter := NewDeepgramAdapter(endpoint, "test-api-key", "nova-3", "en")
|
||||
adapter := NewDeepgramAdapter(endpoint, "test-api-key", "nova-3", "en", nil)
|
||||
|
||||
ctx := context.Background()
|
||||
if err := adapter.Start(ctx, ""); err != nil {
|
||||
@@ -357,7 +357,7 @@ func TestDeepgramAdapter_HandlesError(t *testing.T) {
|
||||
|
||||
wsURL := "ws" + strings.TrimPrefix(server.URL, "http")
|
||||
endpoint := &provider.EndpointConfig{BaseURL: wsURL, Path: ""}
|
||||
adapter := NewDeepgramAdapter(endpoint, "test-api-key", "nova-3", "en")
|
||||
adapter := NewDeepgramAdapter(endpoint, "test-api-key", "nova-3", "en", nil)
|
||||
|
||||
ctx := context.Background()
|
||||
if err := adapter.Start(ctx, ""); err != nil {
|
||||
@@ -398,7 +398,7 @@ func TestDeepgramAdapter_ContextCancellation(t *testing.T) {
|
||||
|
||||
wsURL := "ws" + strings.TrimPrefix(server.URL, "http")
|
||||
endpoint := &provider.EndpointConfig{BaseURL: wsURL, Path: ""}
|
||||
adapter := NewDeepgramAdapter(endpoint, "test-api-key", "nova-3", "en")
|
||||
adapter := NewDeepgramAdapter(endpoint, "test-api-key", "nova-3", "en", nil)
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
if err := adapter.Start(ctx, ""); err != nil {
|
||||
|
||||
@@ -11,7 +11,6 @@ import (
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/leonardotrapani/hyprvoice/internal/language"
|
||||
"github.com/leonardotrapani/hyprvoice/internal/provider"
|
||||
)
|
||||
|
||||
@@ -22,6 +21,7 @@ type ElevenLabsAdapter struct {
|
||||
apiKey string
|
||||
model string
|
||||
language string
|
||||
keywords []string
|
||||
}
|
||||
|
||||
// ElevenLabsResponse represents the API response
|
||||
@@ -33,14 +33,15 @@ type ElevenLabsResponse struct {
|
||||
// endpoint: the endpoint config (BaseURL + Path)
|
||||
// apiKey: ElevenLabs API key
|
||||
// model: model ID (e.g., "scribe_v1")
|
||||
// lang: canonical language code (will be converted to provider format)
|
||||
func NewElevenLabsAdapter(endpoint *provider.EndpointConfig, apiKey, model, lang string) *ElevenLabsAdapter {
|
||||
// lang: provider language code
|
||||
func NewElevenLabsAdapter(endpoint *provider.EndpointConfig, apiKey, model, lang string, keywords []string) *ElevenLabsAdapter {
|
||||
return &ElevenLabsAdapter{
|
||||
client: &http.Client{Timeout: 30 * time.Second},
|
||||
endpoint: endpoint,
|
||||
apiKey: apiKey,
|
||||
model: model,
|
||||
language: lang,
|
||||
keywords: keywords,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -52,6 +53,7 @@ func NewElevenLabsAdapterFromConfig(config Config) *ElevenLabsAdapter {
|
||||
config.APIKey,
|
||||
config.Model,
|
||||
config.Language,
|
||||
config.Keywords,
|
||||
)
|
||||
}
|
||||
|
||||
@@ -85,14 +87,23 @@ func (a *ElevenLabsAdapter) Transcribe(ctx context.Context, audioData []byte) (s
|
||||
return "", fmt.Errorf("write model_id: %w", err)
|
||||
}
|
||||
|
||||
// Add language_code if specified (convert to provider format)
|
||||
providerLang := language.ToProviderFormat(a.language, "elevenlabs")
|
||||
if providerLang != "" {
|
||||
if err := writer.WriteField("language_code", providerLang); err != nil {
|
||||
// Add language_code if specified
|
||||
if a.language != "" {
|
||||
if err := writer.WriteField("language_code", a.language); err != nil {
|
||||
return "", fmt.Errorf("write language_code: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
if len(a.keywords) > 0 {
|
||||
keytermsJSON, err := json.Marshal(a.keywords)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("marshal keyterms: %w", err)
|
||||
}
|
||||
if err := writer.WriteField("keyterms", string(keytermsJSON)); err != nil {
|
||||
return "", fmt.Errorf("write keyterms: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
if err := writer.Close(); err != nil {
|
||||
return "", fmt.Errorf("close writer: %w", err)
|
||||
}
|
||||
|
||||
@@ -8,11 +8,11 @@ import (
|
||||
"log"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/gorilla/websocket"
|
||||
"github.com/leonardotrapani/hyprvoice/internal/language"
|
||||
"github.com/leonardotrapani/hyprvoice/internal/provider"
|
||||
)
|
||||
|
||||
@@ -25,6 +25,7 @@ type ElevenLabsStreamingAdapter struct {
|
||||
apiKey string
|
||||
model string
|
||||
language string
|
||||
keywords []string
|
||||
conn *websocket.Conn
|
||||
resultsCh chan TranscriptionResult
|
||||
mu sync.Mutex
|
||||
@@ -38,15 +39,17 @@ type ElevenLabsStreamingAdapter struct {
|
||||
retryDelays []time.Duration
|
||||
|
||||
// finalization signaling
|
||||
commitDone chan struct{}
|
||||
commitDone chan struct{}
|
||||
contextSent 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"`
|
||||
MessageType string `json:"message_type"`
|
||||
AudioBase64 string `json:"audio_base_64"`
|
||||
Commit bool `json:"commit"`
|
||||
SampleRate int `json:"sample_rate"`
|
||||
PreviousText string `json:"previous_text,omitempty"`
|
||||
}
|
||||
|
||||
// ElevenLabs WebSocket response types (incoming)
|
||||
@@ -62,13 +65,14 @@ type elevenLabsWSMessage struct {
|
||||
// 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 {
|
||||
// lang: provider language code
|
||||
func NewElevenLabsStreamingAdapter(endpoint *provider.EndpointConfig, apiKey, model, lang string, keywords []string) *ElevenLabsStreamingAdapter {
|
||||
return &ElevenLabsStreamingAdapter{
|
||||
endpoint: endpoint,
|
||||
apiKey: apiKey,
|
||||
model: model,
|
||||
language: lang,
|
||||
keywords: keywords,
|
||||
resultsCh: make(chan TranscriptionResult, 100),
|
||||
maxRetries: 3,
|
||||
retryDelays: defaultRetryDelays,
|
||||
@@ -126,6 +130,7 @@ func (a *ElevenLabsStreamingAdapter) connectLocked() error {
|
||||
return fmt.Errorf("websocket dial: %w", err)
|
||||
}
|
||||
a.conn = conn
|
||||
a.contextSent = false
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -199,9 +204,8 @@ func (a *ElevenLabsStreamingAdapter) buildURL() (string, error) {
|
||||
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)
|
||||
if a.language != "" {
|
||||
q.Set("language_code", a.language)
|
||||
}
|
||||
|
||||
// use VAD for automatic commit (easier for real-time use)
|
||||
@@ -334,6 +338,13 @@ func (a *ElevenLabsStreamingAdapter) SendChunk(audio []byte) error {
|
||||
SampleRate: 16000,
|
||||
}
|
||||
|
||||
a.mu.Lock()
|
||||
if !a.contextSent && len(a.keywords) > 0 {
|
||||
msg.PreviousText = strings.Join(a.keywords, ", ")
|
||||
a.contextSent = true
|
||||
}
|
||||
a.mu.Unlock()
|
||||
|
||||
// send as JSON
|
||||
a.mu.Lock()
|
||||
err := a.conn.WriteJSON(msg)
|
||||
|
||||
@@ -72,6 +72,7 @@ func TestElevenLabsStreamingAdapter_Start(t *testing.T) {
|
||||
"test-api-key",
|
||||
"scribe_v1",
|
||||
"en",
|
||||
nil,
|
||||
)
|
||||
|
||||
ctx := context.Background()
|
||||
@@ -126,6 +127,7 @@ func TestElevenLabsStreamingAdapter_SendChunk(t *testing.T) {
|
||||
"test-api-key",
|
||||
"scribe_v1",
|
||||
"en",
|
||||
nil,
|
||||
)
|
||||
|
||||
ctx := context.Background()
|
||||
@@ -188,6 +190,7 @@ func TestElevenLabsStreamingAdapter_Results(t *testing.T) {
|
||||
"test-api-key",
|
||||
"scribe_v1",
|
||||
"en",
|
||||
nil,
|
||||
)
|
||||
|
||||
ctx := context.Background()
|
||||
@@ -262,6 +265,7 @@ func TestElevenLabsStreamingAdapter_ErrorMessages(t *testing.T) {
|
||||
"test-api-key",
|
||||
"scribe_v1",
|
||||
"en",
|
||||
nil,
|
||||
)
|
||||
|
||||
ctx := context.Background()
|
||||
@@ -325,6 +329,7 @@ func TestElevenLabsStreamingAdapter_LanguageConversion(t *testing.T) {
|
||||
"test-api-key",
|
||||
"scribe_v1",
|
||||
"es", // Spanish
|
||||
nil,
|
||||
)
|
||||
|
||||
ctx := context.Background()
|
||||
@@ -372,6 +377,7 @@ func TestElevenLabsStreamingAdapter_Close(t *testing.T) {
|
||||
"test-api-key",
|
||||
"scribe_v1",
|
||||
"en",
|
||||
nil,
|
||||
)
|
||||
|
||||
ctx := context.Background()
|
||||
@@ -408,6 +414,7 @@ func TestElevenLabsStreamingAdapter_NotStarted(t *testing.T) {
|
||||
"test-api-key",
|
||||
"scribe_v1",
|
||||
"en",
|
||||
nil,
|
||||
)
|
||||
|
||||
// SendChunk should fail when not started
|
||||
@@ -470,6 +477,7 @@ func TestElevenLabsStreamingAdapter_ReconnectOnReadError(t *testing.T) {
|
||||
"test-api-key",
|
||||
"scribe_v1",
|
||||
"en",
|
||||
nil,
|
||||
)
|
||||
// use very short delays for testing
|
||||
adapter.retryDelays = []time.Duration{10 * time.Millisecond, 20 * time.Millisecond, 40 * time.Millisecond}
|
||||
@@ -547,6 +555,7 @@ func TestElevenLabsStreamingAdapter_ReconnectNotifiesClient(t *testing.T) {
|
||||
"test-api-key",
|
||||
"scribe_v1",
|
||||
"en",
|
||||
nil,
|
||||
)
|
||||
adapter.retryDelays = []time.Duration{10 * time.Millisecond, 20 * time.Millisecond, 40 * time.Millisecond}
|
||||
|
||||
@@ -633,6 +642,7 @@ func TestElevenLabsStreamingAdapter_MaxRetriesExhausted(t *testing.T) {
|
||||
"test-api-key",
|
||||
"scribe_v1",
|
||||
"en",
|
||||
nil,
|
||||
)
|
||||
adapter.retryDelays = []time.Duration{5 * time.Millisecond, 10 * time.Millisecond, 15 * time.Millisecond}
|
||||
adapter.maxRetries = 2
|
||||
@@ -709,6 +719,7 @@ func TestElevenLabsStreamingAdapter_ReconnectExponentialBackoff(t *testing.T) {
|
||||
"test-api-key",
|
||||
"scribe_v1",
|
||||
"en",
|
||||
nil,
|
||||
)
|
||||
// use measurable delays
|
||||
adapter.retryDelays = []time.Duration{50 * time.Millisecond, 100 * time.Millisecond, 200 * time.Millisecond}
|
||||
|
||||
@@ -13,7 +13,7 @@ func TestNewElevenLabsAdapter(t *testing.T) {
|
||||
Path: "/v1/speech-to-text",
|
||||
}
|
||||
|
||||
adapter := NewElevenLabsAdapter(endpoint, "test-api-key", "scribe_v1", "en")
|
||||
adapter := NewElevenLabsAdapter(endpoint, "test-api-key", "scribe_v1", "en", nil)
|
||||
|
||||
if adapter == nil {
|
||||
t.Fatalf("NewElevenLabsAdapter() returned nil")
|
||||
@@ -74,7 +74,7 @@ func TestElevenLabsAdapter_Transcribe_EmptyAudio(t *testing.T) {
|
||||
Path: "/v1/speech-to-text",
|
||||
}
|
||||
|
||||
adapter := NewElevenLabsAdapter(endpoint, "test-key", "scribe_v1", "")
|
||||
adapter := NewElevenLabsAdapter(endpoint, "test-key", "scribe_v1", "", nil)
|
||||
ctx := context.Background()
|
||||
|
||||
result, err := adapter.Transcribe(ctx, []byte{})
|
||||
@@ -94,7 +94,7 @@ func TestElevenLabsAdapter_Transcribe_ValidAudio(t *testing.T) {
|
||||
Path: "/v1/speech-to-text",
|
||||
}
|
||||
|
||||
adapter := NewElevenLabsAdapter(endpoint, "test-key", "scribe_v1", "en")
|
||||
adapter := NewElevenLabsAdapter(endpoint, "test-key", "scribe_v1", "en", nil)
|
||||
|
||||
if adapter == nil {
|
||||
t.Fatal("NewElevenLabsAdapter() returned nil")
|
||||
|
||||
@@ -1,69 +0,0 @@
|
||||
package transcriber
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"fmt"
|
||||
"log"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/sashabaranov/go-openai"
|
||||
)
|
||||
|
||||
// GroqTranslationAdapter implements BatchAdapter for Groq Translation API
|
||||
// Translates audio to English text. The Language field in config hints at the source language.
|
||||
type GroqTranslationAdapter struct {
|
||||
client *openai.Client
|
||||
config Config
|
||||
}
|
||||
|
||||
func NewGroqTranslationAdapter(config Config) *GroqTranslationAdapter {
|
||||
clientConfig := openai.DefaultConfig(config.APIKey)
|
||||
clientConfig.BaseURL = "https://api.groq.com/openai/v1"
|
||||
client := openai.NewClientWithConfig(clientConfig)
|
||||
|
||||
return &GroqTranslationAdapter{
|
||||
client: client,
|
||||
config: config,
|
||||
}
|
||||
}
|
||||
|
||||
func (a *GroqTranslationAdapter) Transcribe(ctx context.Context, audioData []byte) (string, error) {
|
||||
if len(audioData) == 0 {
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// Convert raw PCM to WAV format
|
||||
wavData, err := convertToWAV(audioData)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("convert to WAV: %w", err)
|
||||
}
|
||||
|
||||
// Create translation request
|
||||
// Note: Translation always outputs English, regardless of target language
|
||||
// The Language field in the request hints at the source audio language for better accuracy
|
||||
req := openai.AudioRequest{
|
||||
Model: a.config.Model,
|
||||
Reader: bytes.NewReader(wavData),
|
||||
FilePath: "audio.wav",
|
||||
Language: a.config.Language, // Source language hint
|
||||
}
|
||||
|
||||
// Add keywords as prompt to help with spelling hints
|
||||
if len(a.config.Keywords) > 0 {
|
||||
req.Prompt = strings.Join(a.config.Keywords, ", ")
|
||||
}
|
||||
|
||||
start := time.Now()
|
||||
resp, err := a.client.CreateTranslation(ctx, req)
|
||||
duration := time.Since(start)
|
||||
|
||||
if err != nil {
|
||||
log.Printf("groq-translation-adapter: API call failed after %v: %v", duration, err)
|
||||
return "", fmt.Errorf("groq translation: %w", err)
|
||||
}
|
||||
|
||||
log.Printf("groq-translation-adapter: translated %d bytes in %v: %q", len(audioData), duration, resp.Text)
|
||||
return resp.Text, nil
|
||||
}
|
||||
@@ -8,7 +8,6 @@ import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/leonardotrapani/hyprvoice/internal/language"
|
||||
"github.com/leonardotrapani/hyprvoice/internal/provider"
|
||||
"github.com/sashabaranov/go-openai"
|
||||
)
|
||||
@@ -27,7 +26,7 @@ type OpenAIAdapter struct {
|
||||
// endpoint: the BaseURL for the API (e.g., "https://api.openai.com", "https://api.groq.com/openai")
|
||||
// apiKey: the API key for authentication
|
||||
// model: model ID to use
|
||||
// lang: canonical language code (will be converted to provider format)
|
||||
// lang: provider language code
|
||||
// keywords: optional spelling hints
|
||||
// providerName: used for logging and language format conversion
|
||||
func NewOpenAIAdapter(endpoint *provider.EndpointConfig, apiKey, model, lang string, keywords []string, providerName string) *OpenAIAdapter {
|
||||
@@ -69,15 +68,12 @@ func (a *OpenAIAdapter) Transcribe(ctx context.Context, audioData []byte) (strin
|
||||
return "", fmt.Errorf("convert to WAV: %w", err)
|
||||
}
|
||||
|
||||
// Convert language code to provider format
|
||||
providerLang := language.ToProviderFormat(a.language, a.providerName)
|
||||
|
||||
// Create transcription request
|
||||
req := openai.AudioRequest{
|
||||
Model: a.model,
|
||||
Reader: bytes.NewReader(wavData),
|
||||
FilePath: "audio.wav",
|
||||
Language: providerLang,
|
||||
Language: a.language,
|
||||
}
|
||||
|
||||
// Add keywords as initial_prompt to help with spelling hints
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"log"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
@@ -21,6 +22,7 @@ type OpenAIRealtimeAdapter struct {
|
||||
apiKey string
|
||||
model string
|
||||
language string
|
||||
keywords []string
|
||||
conn *websocket.Conn
|
||||
resultsCh chan TranscriptionResult
|
||||
mu sync.Mutex
|
||||
@@ -56,6 +58,7 @@ type openaiRealtimeSessionConfig struct {
|
||||
type openaiRealtimeTranscription struct {
|
||||
Model string `json:"model,omitempty"`
|
||||
Language string `json:"language,omitempty"`
|
||||
Prompt string `json:"prompt,omitempty"`
|
||||
}
|
||||
|
||||
type openaiRealtimeTurnDetection struct {
|
||||
@@ -104,12 +107,13 @@ type openaiRealtimeError struct {
|
||||
// apiKey: OpenAI API key
|
||||
// model: model ID (e.g., "gpt-4o-realtime-preview")
|
||||
// lang: canonical language code (will be used for transcription config)
|
||||
func NewOpenAIRealtimeAdapter(endpoint *provider.EndpointConfig, apiKey, model, lang string) *OpenAIRealtimeAdapter {
|
||||
func NewOpenAIRealtimeAdapter(endpoint *provider.EndpointConfig, apiKey, model, lang string, keywords []string) *OpenAIRealtimeAdapter {
|
||||
return &OpenAIRealtimeAdapter{
|
||||
endpoint: endpoint,
|
||||
apiKey: apiKey,
|
||||
model: model,
|
||||
language: lang,
|
||||
keywords: keywords,
|
||||
resultsCh: make(chan TranscriptionResult, 100),
|
||||
maxRetries: 3,
|
||||
retryDelays: defaultRetryDelays,
|
||||
@@ -206,6 +210,10 @@ func (a *OpenAIRealtimeAdapter) configureSession() error {
|
||||
sessionUpdate.Session.InputAudioTranscription.Language = a.language
|
||||
}
|
||||
|
||||
if len(a.keywords) > 0 {
|
||||
sessionUpdate.Session.InputAudioTranscription.Prompt = strings.Join(a.keywords, ", ")
|
||||
}
|
||||
|
||||
return a.conn.WriteJSON(sessionUpdate)
|
||||
}
|
||||
|
||||
|
||||
@@ -118,7 +118,7 @@ func TestOpenAIRealtimeAdapter_Start(t *testing.T) {
|
||||
Path: "",
|
||||
}
|
||||
|
||||
adapter := NewOpenAIRealtimeAdapter(endpoint, "sk-test-key", "gpt-4o-realtime-preview", "en")
|
||||
adapter := NewOpenAIRealtimeAdapter(endpoint, "sk-test-key", "gpt-4o-realtime-preview", "en", nil)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||
defer cancel()
|
||||
@@ -198,7 +198,7 @@ func TestOpenAIRealtimeAdapter_SendChunk(t *testing.T) {
|
||||
|
||||
wsURL := "ws" + strings.TrimPrefix(server.URL, "http")
|
||||
endpoint := &provider.EndpointConfig{BaseURL: wsURL, Path: ""}
|
||||
adapter := NewOpenAIRealtimeAdapter(endpoint, "sk-test", "gpt-4o-realtime-preview", "")
|
||||
adapter := NewOpenAIRealtimeAdapter(endpoint, "sk-test", "gpt-4o-realtime-preview", "", nil)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||
defer cancel()
|
||||
@@ -282,7 +282,7 @@ func TestOpenAIRealtimeAdapter_TranscriptionResults(t *testing.T) {
|
||||
|
||||
wsURL := "ws" + strings.TrimPrefix(server.URL, "http")
|
||||
endpoint := &provider.EndpointConfig{BaseURL: wsURL, Path: ""}
|
||||
adapter := NewOpenAIRealtimeAdapter(endpoint, "sk-test", "gpt-4o-realtime-preview", "en")
|
||||
adapter := NewOpenAIRealtimeAdapter(endpoint, "sk-test", "gpt-4o-realtime-preview", "en", nil)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
|
||||
defer cancel()
|
||||
@@ -369,7 +369,7 @@ func TestOpenAIRealtimeAdapter_ErrorHandling(t *testing.T) {
|
||||
|
||||
wsURL := "ws" + strings.TrimPrefix(server.URL, "http")
|
||||
endpoint := &provider.EndpointConfig{BaseURL: wsURL, Path: ""}
|
||||
adapter := NewOpenAIRealtimeAdapter(endpoint, "sk-test", "gpt-4o-realtime-preview", "")
|
||||
adapter := NewOpenAIRealtimeAdapter(endpoint, "sk-test", "gpt-4o-realtime-preview", "", nil)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
||||
defer cancel()
|
||||
@@ -433,7 +433,7 @@ func TestOpenAIRealtimeAdapter_Reconnection(t *testing.T) {
|
||||
|
||||
wsURL := "ws" + strings.TrimPrefix(server.URL, "http")
|
||||
endpoint := &provider.EndpointConfig{BaseURL: wsURL, Path: ""}
|
||||
adapter := NewOpenAIRealtimeAdapter(endpoint, "sk-test", "gpt-4o-realtime-preview", "")
|
||||
adapter := NewOpenAIRealtimeAdapter(endpoint, "sk-test", "gpt-4o-realtime-preview", "", nil)
|
||||
adapter.retryDelays = []time.Duration{10 * time.Millisecond, 20 * time.Millisecond, 40 * time.Millisecond}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
@@ -473,7 +473,7 @@ func TestOpenAIRealtimeAdapter_Close(t *testing.T) {
|
||||
|
||||
wsURL := "ws" + strings.TrimPrefix(server.URL, "http")
|
||||
endpoint := &provider.EndpointConfig{BaseURL: wsURL, Path: ""}
|
||||
adapter := NewOpenAIRealtimeAdapter(endpoint, "sk-test", "gpt-4o-realtime-preview", "")
|
||||
adapter := NewOpenAIRealtimeAdapter(endpoint, "sk-test", "gpt-4o-realtime-preview", "", nil)
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
|
||||
@@ -10,8 +10,6 @@ import (
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/leonardotrapani/hyprvoice/internal/language"
|
||||
)
|
||||
|
||||
// WhisperCppAdapter implements BatchAdapter for local whisper-cpp transcription
|
||||
@@ -23,7 +21,7 @@ type WhisperCppAdapter struct {
|
||||
|
||||
// NewWhisperCppAdapter creates a new whisper-cpp adapter
|
||||
// modelPath: full path to the model file (e.g., ~/.local/share/hyprvoice/models/whisper/ggml-base.en.bin)
|
||||
// lang: canonical language code (will be converted to whisper-cpp format)
|
||||
// lang: whisper-cpp language code
|
||||
// threads: number of CPU threads (0 for auto)
|
||||
func NewWhisperCppAdapter(modelPath, lang string, threads int) *WhisperCppAdapter {
|
||||
return &WhisperCppAdapter{
|
||||
@@ -63,8 +61,11 @@ func (a *WhisperCppAdapter) Transcribe(ctx context.Context, audioData []byte) (s
|
||||
}
|
||||
defer os.Remove(tmpFile)
|
||||
|
||||
// convert language to whisper-cpp format
|
||||
lang := language.ToProviderFormat(a.language, "whisper-cpp")
|
||||
// use whisper-cpp auto if unspecified
|
||||
lang := a.language
|
||||
if lang == "" {
|
||||
lang = "auto"
|
||||
}
|
||||
|
||||
// build command args
|
||||
args := []string{
|
||||
|
||||
@@ -8,7 +8,6 @@ import (
|
||||
"golang.org/x/text/cases"
|
||||
"golang.org/x/text/language"
|
||||
|
||||
lang "github.com/leonardotrapani/hyprvoice/internal/language"
|
||||
"github.com/leonardotrapani/hyprvoice/internal/models/whisper"
|
||||
"github.com/leonardotrapani/hyprvoice/internal/provider"
|
||||
"github.com/leonardotrapani/hyprvoice/internal/recording"
|
||||
@@ -43,15 +42,6 @@ func NewTranscriber(config Config) (Transcriber, error) {
|
||||
return nil, fmt.Errorf("provider is required")
|
||||
}
|
||||
|
||||
// special case: groq-translation uses CreateTranslation API (different from transcription)
|
||||
if config.Provider == provider.ConfigProviderGroqTranslation {
|
||||
if config.APIKey == "" {
|
||||
return nil, fmt.Errorf("Groq API key required")
|
||||
}
|
||||
adapter := NewGroqTranslationAdapter(config)
|
||||
return NewSimpleTranscriber(config, adapter), nil
|
||||
}
|
||||
|
||||
// map config provider name to registry provider name
|
||||
registryProvider := provider.BaseProviderName(config.Provider)
|
||||
|
||||
@@ -89,8 +79,7 @@ func NewTranscriber(config Config) (Transcriber, error) {
|
||||
// runtime language-model compatibility check with fallback
|
||||
// primary validation happens at config time (hard error), this is a safety net
|
||||
if config.Language != "" && !model.SupportsLanguage(config.Language) {
|
||||
langName := lang.FromCode(config.Language).Name
|
||||
log.Printf("warning: model %s does not support language %s, falling back to auto-detect", model.ID, langName)
|
||||
log.Printf("warning: model %s does not support language %s, falling back to auto-detect", model.ID, config.Language)
|
||||
config.Language = ""
|
||||
}
|
||||
|
||||
@@ -119,11 +108,11 @@ func NewTranscriber(config Config) (Transcriber, error) {
|
||||
var streamingAdapter StreamingAdapter
|
||||
switch adapterType {
|
||||
case provider.AdapterElevenLabsStream:
|
||||
streamingAdapter = NewElevenLabsStreamingAdapter(endpoint, config.APIKey, model.ID, config.Language)
|
||||
streamingAdapter = NewElevenLabsStreamingAdapter(endpoint, config.APIKey, model.ID, config.Language, config.Keywords)
|
||||
case provider.AdapterDeepgram:
|
||||
streamingAdapter = NewDeepgramAdapter(endpoint, config.APIKey, model.ID, config.Language)
|
||||
streamingAdapter = NewDeepgramAdapter(endpoint, config.APIKey, model.ID, config.Language, config.Keywords)
|
||||
case provider.AdapterOpenAIRealtime:
|
||||
streamingAdapter = NewOpenAIRealtimeAdapter(endpoint, config.APIKey, model.ID, config.Language)
|
||||
streamingAdapter = NewOpenAIRealtimeAdapter(endpoint, config.APIKey, model.ID, config.Language, config.Keywords)
|
||||
default:
|
||||
return nil, fmt.Errorf("unsupported streaming adapter type: %s", adapterType)
|
||||
}
|
||||
@@ -136,9 +125,9 @@ func NewTranscriber(config Config) (Transcriber, error) {
|
||||
case provider.AdapterOpenAI:
|
||||
adapter = NewOpenAIAdapter(model.Endpoint, config.APIKey, model.ID, config.Language, config.Keywords, registryProvider)
|
||||
case provider.AdapterElevenLabs:
|
||||
adapter = NewElevenLabsAdapter(model.Endpoint, config.APIKey, model.ID, config.Language)
|
||||
adapter = NewElevenLabsAdapter(model.Endpoint, config.APIKey, model.ID, config.Language, config.Keywords)
|
||||
case provider.AdapterDeepgram:
|
||||
adapter = NewDeepgramBatchAdapter(model.Endpoint, config.APIKey, model.ID, config.Language)
|
||||
adapter = NewDeepgramBatchAdapter(model.Endpoint, config.APIKey, model.ID, config.Language, config.Keywords)
|
||||
case provider.AdapterWhisperCpp:
|
||||
modelPath := whisper.GetModelPath(config.Model)
|
||||
if modelPath == "" {
|
||||
|
||||
@@ -56,26 +56,6 @@ func TestNewTranscriber(t *testing.T) {
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "valid groq-translation config",
|
||||
config: Config{
|
||||
Provider: "groq-translation",
|
||||
APIKey: "gsk-test-key",
|
||||
Language: "es",
|
||||
Model: "whisper-large-v3-turbo",
|
||||
},
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "groq-translation config without api key",
|
||||
config: Config{
|
||||
Provider: "groq-translation",
|
||||
APIKey: "",
|
||||
Language: "es",
|
||||
Model: "whisper-large-v3-turbo",
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "valid mistral-transcription config",
|
||||
config: Config{
|
||||
|
||||
Reference in New Issue
Block a user