feat: refactor
This commit is contained in:
@@ -0,0 +1,126 @@
|
||||
package transcriber
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
|
||||
"github.com/leonardotrapani/hyprvoice/internal/language"
|
||||
"github.com/leonardotrapani/hyprvoice/internal/provider"
|
||||
)
|
||||
|
||||
// DeepgramBatchAdapter implements BatchAdapter for Deepgram pre-recorded transcription
|
||||
type DeepgramBatchAdapter struct {
|
||||
endpoint *provider.EndpointConfig
|
||||
apiKey string
|
||||
model string
|
||||
language string
|
||||
}
|
||||
|
||||
// deepgramBatchResponse is the response from the pre-recorded API
|
||||
type deepgramBatchResponse struct {
|
||||
Results *deepgramBatchResults `json:"results,omitempty"`
|
||||
Error *deepgramError `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
type deepgramBatchResults struct {
|
||||
Channels []deepgramBatchChannel `json:"channels,omitempty"`
|
||||
}
|
||||
|
||||
type deepgramBatchChannel struct {
|
||||
Alternatives []deepgramAlternative `json:"alternatives,omitempty"`
|
||||
}
|
||||
|
||||
// NewDeepgramBatchAdapter creates a new batch adapter for Deepgram
|
||||
func NewDeepgramBatchAdapter(endpoint *provider.EndpointConfig, apiKey, model, lang string) *DeepgramBatchAdapter {
|
||||
return &DeepgramBatchAdapter{
|
||||
endpoint: endpoint,
|
||||
apiKey: apiKey,
|
||||
model: model,
|
||||
language: lang,
|
||||
}
|
||||
}
|
||||
|
||||
// Transcribe sends audio data to Deepgram's pre-recorded API
|
||||
func (a *DeepgramBatchAdapter) Transcribe(ctx context.Context, audioData []byte) (string, error) {
|
||||
// build URL with query parameters
|
||||
apiURL, err := a.buildURL()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("build url: %w", err)
|
||||
}
|
||||
|
||||
// create request with audio data as body
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", apiURL, bytes.NewReader(audioData))
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("create request: %w", err)
|
||||
}
|
||||
|
||||
// set headers
|
||||
req.Header.Set("Authorization", "Token "+a.apiKey)
|
||||
req.Header.Set("Content-Type", "audio/wav") // we send raw PCM wrapped as WAV
|
||||
|
||||
// send request
|
||||
resp, err := http.DefaultClient.Do(req)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("http request: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
// read response
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("read response: %w", err)
|
||||
}
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", fmt.Errorf("deepgram api error (status %d): %s", resp.StatusCode, string(body))
|
||||
}
|
||||
|
||||
// parse response
|
||||
var result deepgramBatchResponse
|
||||
if err := json.Unmarshal(body, &result); err != nil {
|
||||
return "", fmt.Errorf("parse response: %w", err)
|
||||
}
|
||||
|
||||
if result.Error != nil {
|
||||
return "", fmt.Errorf("deepgram error: %s", result.Error.Message)
|
||||
}
|
||||
|
||||
// extract transcript
|
||||
if result.Results == nil || len(result.Results.Channels) == 0 {
|
||||
return "", nil
|
||||
}
|
||||
if len(result.Results.Channels[0].Alternatives) == 0 {
|
||||
return "", nil
|
||||
}
|
||||
|
||||
return result.Results.Channels[0].Alternatives[0].Transcript, nil
|
||||
}
|
||||
|
||||
// buildURL constructs the API URL with query parameters
|
||||
func (a *DeepgramBatchAdapter) buildURL() (string, error) {
|
||||
baseURL := a.endpoint.BaseURL + a.endpoint.Path
|
||||
|
||||
u, err := url.Parse(baseURL)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("parse base url: %w", err)
|
||||
}
|
||||
|
||||
q := u.Query()
|
||||
q.Set("model", a.model)
|
||||
q.Set("smart_format", "true")
|
||||
q.Set("punctuate", "true")
|
||||
|
||||
// add language if specified
|
||||
providerLang := language.ToProviderFormat(a.language, "deepgram")
|
||||
if providerLang != "" {
|
||||
q.Set("language", providerLang)
|
||||
}
|
||||
|
||||
u.RawQuery = q.Encode()
|
||||
return u.String(), nil
|
||||
}
|
||||
@@ -4,11 +4,12 @@ import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log"
|
||||
"strings"
|
||||
|
||||
"github.com/leonardotrapani/hyprvoice/internal/language"
|
||||
"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/notify"
|
||||
"github.com/leonardotrapani/hyprvoice/internal/provider"
|
||||
"github.com/leonardotrapani/hyprvoice/internal/recording"
|
||||
)
|
||||
@@ -27,26 +28,13 @@ type BatchAdapter interface {
|
||||
|
||||
// Configuration for the transcriber
|
||||
type Config struct {
|
||||
Provider string
|
||||
APIKey string
|
||||
Language string
|
||||
Model string
|
||||
Keywords []string
|
||||
Threads int // CPU threads for local transcription (0 = auto)
|
||||
}
|
||||
|
||||
// mapConfigProviderToRegistryName maps config provider names to provider registry names
|
||||
// Config uses names like "groq-transcription", "groq-translation", "mistral-transcription"
|
||||
// Registry uses base names like "groq", "mistral"
|
||||
func mapConfigProviderToRegistryName(configProvider string) string {
|
||||
switch configProvider {
|
||||
case "groq-transcription", "groq-translation":
|
||||
return "groq"
|
||||
case "mistral-transcription":
|
||||
return "mistral"
|
||||
default:
|
||||
return configProvider
|
||||
}
|
||||
Provider string
|
||||
APIKey string
|
||||
Language string
|
||||
Model string
|
||||
Keywords []string
|
||||
Threads int // CPU threads for local transcription (0 = auto)
|
||||
Streaming bool // use streaming mode if model supports it
|
||||
}
|
||||
|
||||
// NewTranscriber creates a new transcriber based on model metadata
|
||||
@@ -56,7 +44,7 @@ func NewTranscriber(config Config) (Transcriber, error) {
|
||||
}
|
||||
|
||||
// special case: groq-translation uses CreateTranslation API (different from transcription)
|
||||
if config.Provider == "groq-translation" {
|
||||
if config.Provider == provider.ConfigProviderGroqTranslation {
|
||||
if config.APIKey == "" {
|
||||
return nil, fmt.Errorf("Groq API key required")
|
||||
}
|
||||
@@ -65,7 +53,7 @@ func NewTranscriber(config Config) (Transcriber, error) {
|
||||
}
|
||||
|
||||
// map config provider name to registry provider name
|
||||
registryProvider := mapConfigProviderToRegistryName(config.Provider)
|
||||
registryProvider := provider.BaseProviderName(config.Provider)
|
||||
|
||||
// lookup provider
|
||||
p := provider.GetProvider(registryProvider)
|
||||
@@ -75,7 +63,7 @@ func NewTranscriber(config Config) (Transcriber, error) {
|
||||
|
||||
// check API key requirement
|
||||
if p.RequiresAPIKey() && config.APIKey == "" {
|
||||
return nil, fmt.Errorf("%s API key required", strings.Title(registryProvider))
|
||||
return nil, fmt.Errorf("%s API key required", cases.Title(language.English).String(registryProvider))
|
||||
}
|
||||
|
||||
// lookup model from provider
|
||||
@@ -101,41 +89,57 @@ 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 := language.FromCode(config.Language).Name
|
||||
langName := lang.FromCode(config.Language).Name
|
||||
log.Printf("warning: model %s does not support language %s, falling back to auto-detect", model.ID, langName)
|
||||
|
||||
// send desktop notification to alert user
|
||||
notifier := notify.NewDesktop(nil)
|
||||
notifier.Error(fmt.Sprintf("Model %s does not support %s. Using auto-detect.", model.Name, langName))
|
||||
|
||||
// override language to auto for this session
|
||||
config.Language = ""
|
||||
}
|
||||
|
||||
// streaming models use StreamingTranscriber
|
||||
if model.Streaming {
|
||||
// determine if we should use streaming mode
|
||||
useStreaming := config.Streaming && model.SupportsStreaming
|
||||
|
||||
// fail if streaming-only model is used without streaming enabled
|
||||
if !useStreaming && !model.SupportsBatch {
|
||||
return nil, fmt.Errorf("model %s requires streaming mode (set streaming = true in config)", model.ID)
|
||||
}
|
||||
|
||||
// streaming mode: use StreamingTranscriber
|
||||
if useStreaming {
|
||||
// pick the right adapter type for streaming
|
||||
adapterType := model.AdapterType
|
||||
if model.StreamingAdapter != "" {
|
||||
adapterType = model.StreamingAdapter
|
||||
}
|
||||
|
||||
// pick the right endpoint for streaming
|
||||
endpoint := model.Endpoint
|
||||
if model.StreamingEndpoint != nil {
|
||||
endpoint = model.StreamingEndpoint
|
||||
}
|
||||
|
||||
var streamingAdapter StreamingAdapter
|
||||
switch model.AdapterType {
|
||||
case "elevenlabs-streaming":
|
||||
streamingAdapter = NewElevenLabsStreamingAdapter(model.Endpoint, config.APIKey, model.ID, config.Language)
|
||||
case "deepgram":
|
||||
streamingAdapter = NewDeepgramAdapter(model.Endpoint, config.APIKey, model.ID, config.Language)
|
||||
case "openai-realtime":
|
||||
streamingAdapter = NewOpenAIRealtimeAdapter(model.Endpoint, config.APIKey, model.ID, config.Language)
|
||||
switch adapterType {
|
||||
case provider.AdapterElevenLabsStream:
|
||||
streamingAdapter = NewElevenLabsStreamingAdapter(endpoint, config.APIKey, model.ID, config.Language)
|
||||
case provider.AdapterDeepgram:
|
||||
streamingAdapter = NewDeepgramAdapter(endpoint, config.APIKey, model.ID, config.Language)
|
||||
case provider.AdapterOpenAIRealtime:
|
||||
streamingAdapter = NewOpenAIRealtimeAdapter(endpoint, config.APIKey, model.ID, config.Language)
|
||||
default:
|
||||
return nil, fmt.Errorf("unsupported streaming adapter type: %s", model.AdapterType)
|
||||
return nil, fmt.Errorf("unsupported streaming adapter type: %s", adapterType)
|
||||
}
|
||||
return NewStreamingTranscriber(streamingAdapter, config.Language), nil
|
||||
}
|
||||
|
||||
// batch models use SimpleTranscriber
|
||||
// batch mode: use SimpleTranscriber
|
||||
var adapter BatchAdapter
|
||||
switch model.AdapterType {
|
||||
case "openai":
|
||||
case provider.AdapterOpenAI:
|
||||
adapter = NewOpenAIAdapter(model.Endpoint, config.APIKey, model.ID, config.Language, config.Keywords, registryProvider)
|
||||
case "elevenlabs":
|
||||
case provider.AdapterElevenLabs:
|
||||
adapter = NewElevenLabsAdapter(model.Endpoint, config.APIKey, model.ID, config.Language)
|
||||
case "whisper-cpp":
|
||||
case provider.AdapterDeepgram:
|
||||
adapter = NewDeepgramBatchAdapter(model.Endpoint, config.APIKey, model.ID, config.Language)
|
||||
case provider.AdapterWhisperCpp:
|
||||
modelPath := whisper.GetModelPath(config.Model)
|
||||
if modelPath == "" {
|
||||
return nil, fmt.Errorf("unknown whisper model: %s", config.Model)
|
||||
|
||||
@@ -157,30 +157,33 @@ func TestNewTranscriber(t *testing.T) {
|
||||
{
|
||||
name: "elevenlabs streaming model creates StreamingTranscriber",
|
||||
config: Config{
|
||||
Provider: "elevenlabs",
|
||||
APIKey: "test-key",
|
||||
Language: "en",
|
||||
Model: "scribe_v1-streaming",
|
||||
},
|
||||
wantErr: false, // streaming is now supported
|
||||
},
|
||||
{
|
||||
name: "deepgram streaming model creates StreamingTranscriber",
|
||||
config: Config{
|
||||
Provider: "deepgram",
|
||||
APIKey: "test-key",
|
||||
Language: "en",
|
||||
Model: "nova-3",
|
||||
Provider: "elevenlabs",
|
||||
APIKey: "test-key",
|
||||
Language: "en",
|
||||
Model: "scribe_v2_realtime",
|
||||
Streaming: true,
|
||||
},
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "openai realtime streaming model creates StreamingTranscriber",
|
||||
name: "deepgram streaming model creates StreamingTranscriber",
|
||||
config: Config{
|
||||
Provider: "openai",
|
||||
APIKey: "test-key",
|
||||
Language: "en",
|
||||
Model: "gpt-4o-realtime-preview",
|
||||
Provider: "deepgram",
|
||||
APIKey: "test-key",
|
||||
Language: "en",
|
||||
Model: "nova-3",
|
||||
Streaming: true,
|
||||
},
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "openai streaming model creates StreamingTranscriber",
|
||||
config: Config{
|
||||
Provider: "openai",
|
||||
APIKey: "test-key",
|
||||
Language: "en",
|
||||
Model: "gpt-4o-transcribe",
|
||||
Streaming: true,
|
||||
},
|
||||
wantErr: false,
|
||||
},
|
||||
@@ -997,12 +1000,11 @@ func TestStreamingTranscriber_GetFinalTranscriptionSafe(t *testing.T) {
|
||||
|
||||
func TestNewTranscriber_LanguageFallback(t *testing.T) {
|
||||
// test that incompatible language falls back to auto-detect (no error)
|
||||
// distil-whisper-large-v3-en only supports English
|
||||
// base.en only supports English
|
||||
config := Config{
|
||||
Provider: "groq-transcription",
|
||||
APIKey: "test-key",
|
||||
Provider: "whisper-cpp",
|
||||
Language: "es", // Spanish not supported by English-only model
|
||||
Model: "distil-whisper-large-v3-en",
|
||||
Model: "base.en",
|
||||
}
|
||||
|
||||
// should succeed (fallback to auto), not error
|
||||
@@ -1020,10 +1022,9 @@ func TestNewTranscriber_LanguageFallback(t *testing.T) {
|
||||
func TestNewTranscriber_AutoLanguageNoFallback(t *testing.T) {
|
||||
// test that auto language never triggers warning/fallback
|
||||
config := Config{
|
||||
Provider: "groq-transcription",
|
||||
APIKey: "test-key",
|
||||
Provider: "whisper-cpp",
|
||||
Language: "", // auto
|
||||
Model: "distil-whisper-large-v3-en",
|
||||
Model: "base.en",
|
||||
}
|
||||
|
||||
transcriber, err := NewTranscriber(config)
|
||||
@@ -1040,10 +1041,9 @@ func TestNewTranscriber_AutoLanguageNoFallback(t *testing.T) {
|
||||
func TestNewTranscriber_CompatibleLanguageNoFallback(t *testing.T) {
|
||||
// test that compatible language works normally
|
||||
config := Config{
|
||||
Provider: "groq-transcription",
|
||||
APIKey: "test-key",
|
||||
Provider: "whisper-cpp",
|
||||
Language: "en", // English supported by English-only model
|
||||
Model: "distil-whisper-large-v3-en",
|
||||
Model: "base.en",
|
||||
}
|
||||
|
||||
transcriber, err := NewTranscriber(config)
|
||||
|
||||
Reference in New Issue
Block a user