add arbitrary whisper transcriptor support

This commit is contained in:
Huaaudio 2026-03-30 18:19:09 +02:00
parent c9ec8fa2c2
commit c3b85c7d55
7 changed files with 492 additions and 211 deletions

View file

@ -48,7 +48,11 @@ model_list:
PicoClaw's `DetectTranscriber` function will attempt to detect the appropriate Transcriber in the following order:
1. **Targeted Selection**: Standard matching via `cfg.Voice.ModelName`.
- If the protocol matches `elevenlabs/`, the ElevenLabs transcriber is initiated.
- If the protocol supports general OpenAI-compatible audio transcription endpoints (e.g., `openai`, `azure`, `groq`, `deepseek`), `AudioModelTranscriber` is leveraged.
2. **Fallback Scanning**: If no `model_name` is selected, it scans `model_list` specifically looking for `elevenlabs/` protocol models or `groq/` provider formats (e.g. for Whisper fallback).
1. **Targeted Selection**: Resolve `cfg.Voice.ModelName` against `model_list`, then create the transcriber from that resolved model entry.
- This means aliases such as `my-asr-model` are the primary ASR contract.
- If the resolved model uses `elevenlabs/...`, the ElevenLabs transcriber is initiated.
- If the resolved model uses an OpenAI-compatible Whisper model name such as `openai/whisper-1` or `groq/whisper-large-v3`, the Whisper transcriber is initiated.
- If the resolved model uses an OpenAI-compatible audio-capable provider (for example `openai`, `azure`, `gemini`, `deepseek`), `AudioModelTranscriber` is leveraged.
2. **Fallback Scanning**: If no `model_name` is selected, PicoClaw performs a compatibility scan through `model_list` for legacy auto-detected ASR providers such as `elevenlabs/...` entries and OpenAI-compatible Whisper models.
Fallback scanning exists for compatibility, but the recommended configuration is to set `voice.model_name` to a named `model_list` entry such as `my-asr-model`.

View file

@ -41,31 +41,90 @@ func supportsAudioTranscription(model string) bool {
}
}
func supportsWhisperTranscription(model string) bool {
protocol, _ := providers.ExtractProtocol(model)
switch protocol {
case "openai", "litellm", "openrouter", "groq", "zhipu", "gemini", "nvidia",
"ollama", "moonshot", "shengsuanyun", "deepseek", "cerebras",
"vivgrid", "volcengine", "vllm", "qwen", "qwen-intl", "qwen-international", "dashscope-intl",
"qwen-us", "dashscope-us", "mistral", "avian", "minimax", "longcat", "modelscope", "novita",
"coding-plan", "alibaba-coding", "qwen-coding", "mimo":
return true
default:
return false
}
}
func whisperModelID(modelCfg *config.ModelConfig) string {
if modelCfg == nil || modelCfg.APIKey() == "" {
return ""
}
if !supportsWhisperTranscription(modelCfg.Model) {
return ""
}
_, modelID := providers.ExtractProtocol(strings.TrimSpace(modelCfg.Model))
if strings.Contains(strings.ToLower(modelID), "whisper") {
return modelID
}
return ""
}
func transcriberFromModelConfig(modelCfg *config.ModelConfig) Transcriber {
if modelCfg == nil {
return nil
}
protocol, _ := providers.ExtractProtocol(modelCfg.Model)
if protocol == "elevenlabs" && modelCfg.APIKey() != "" {
return NewElevenLabsTranscriber(modelCfg.APIKey(), modelCfg.APIBase)
}
if modelID := whisperModelID(modelCfg); modelID != "" {
return NewWhisperTranscriber(modelCfg)
}
if supportsAudioTranscription(modelCfg.Model) {
return NewAudioModelTranscriber(modelCfg)
}
return nil
}
func fallbackTranscriberFromModelConfig(modelCfg *config.ModelConfig) Transcriber {
if modelCfg == nil {
return nil
}
protocol, _ := providers.ExtractProtocol(modelCfg.Model)
if protocol == "elevenlabs" && modelCfg.APIKey() != "" {
return NewElevenLabsTranscriber(modelCfg.APIKey(), modelCfg.APIBase)
}
if modelID := whisperModelID(modelCfg); modelID != "" {
return NewWhisperTranscriber(modelCfg)
}
return nil
}
// DetectTranscriber inspects cfg and returns the appropriate Transcriber, or
// nil if no supported transcription provider is configured.
func DetectTranscriber(cfg *config.Config) Transcriber {
if cfg == nil {
return nil
}
if modelName := strings.TrimSpace(cfg.Voice.ModelName); modelName != "" {
modelCfg, err := cfg.GetModelConfig(modelName)
if err == nil {
protocol, _ := providers.ExtractProtocol(modelCfg.Model)
if protocol == "elevenlabs" && modelCfg.APIKey() != "" {
return NewElevenLabsTranscriber(modelCfg.APIKey(), modelCfg.APIBase)
}
if supportsAudioTranscription(modelCfg.Model) {
return NewAudioModelTranscriber(modelCfg)
if tr := transcriberFromModelConfig(modelCfg); tr != nil {
return tr
}
}
}
// Fall back to scanning ModelList for suitable ASR providers
// Fall back to compatibility scanning for legacy auto-detected ASR providers.
for _, mc := range cfg.ModelList {
protocol, _ := providers.ExtractProtocol(mc.Model)
if protocol == "elevenlabs" && mc.APIKey() != "" {
return NewElevenLabsTranscriber(mc.APIKey(), mc.APIBase)
}
if (strings.HasPrefix(mc.Model, "groq/") || mc.ModelName == "groq" || mc.Model == "whisper-large-v3-turbo") &&
mc.APIKey() != "" {
return NewGroqTranscriber(mc.APIKey())
if tr := fallbackTranscriberFromModelConfig(mc); tr != nil {
return tr
}
}
return nil

View file

@ -33,26 +33,68 @@ func TestDetectTranscriber(t *testing.T) {
wantName: "audio-model",
},
{
name: "groq via model list",
name: "voice model name alias selects elevenlabs transcriber",
cfg: &config.Config{
Voice: config.VoiceConfig{ModelName: "my-asr-model"},
ModelList: []*config.ModelConfig{
{
ModelName: "my-asr-model",
Model: "elevenlabs/scribe_v1",
APIKeys: config.SimpleSecureStrings("sk_elevenlabs_test"),
},
},
},
wantName: "elevenlabs",
},
{
name: "voice model name alias selects whisper transcriber for groq",
cfg: &config.Config{
Voice: config.VoiceConfig{ModelName: "my-asr-model"},
ModelList: []*config.ModelConfig{
{
ModelName: "my-asr-model",
Model: "groq/whisper-large-v3",
APIKeys: config.SimpleSecureStrings("sk-groq-model"),
},
},
},
wantName: "whisper",
},
{
name: "openai whisper alias selects whisper transcriber",
cfg: &config.Config{
Voice: config.VoiceConfig{ModelName: "my-asr-model"},
ModelList: []*config.ModelConfig{
{
ModelName: "my-asr-model",
Model: "openai/whisper-1",
APIKeys: config.SimpleSecureStrings("sk-openai-model"),
},
},
},
wantName: "whisper",
},
{
name: "whisper via model list fallback",
cfg: &config.Config{
ModelList: []*config.ModelConfig{
{ModelName: "openai", Model: "openai/gpt-4o", APIKeys: config.SimpleSecureStrings("sk-openai")},
{
ModelName: "groq",
Model: "groq/llama-3.3-70b",
Model: "groq/whisper-large-v3-turbo",
APIKeys: config.SimpleSecureStrings("sk-groq-model"),
},
},
},
wantName: "groq",
wantName: "whisper",
},
{
name: "voice model name selects non-gemini audio model transcriber",
name: "voice model name alias selects non-gemini audio model transcriber",
cfg: &config.Config{
Voice: config.VoiceConfig{ModelName: "voice-openai-audio"},
Voice: config.VoiceConfig{ModelName: "my-asr-model"},
ModelList: []*config.ModelConfig{
{
ModelName: "voice-openai-audio",
ModelName: "my-asr-model",
Model: "openai/gpt-4o-audio-preview",
APIKeys: config.SimpleSecureStrings("sk-openai"),
},
@ -92,7 +134,7 @@ func TestDetectTranscriber(t *testing.T) {
name: "groq model list entry without key is skipped",
cfg: &config.Config{
ModelList: []*config.ModelConfig{
{Model: "groq/llama-3.3-70b"},
{Model: "groq/whisper-large-v3"},
},
},
wantNil: true,
@ -103,12 +145,12 @@ func TestDetectTranscriber(t *testing.T) {
ModelList: []*config.ModelConfig{
{
ModelName: "groq",
Model: "groq/llama-3.3-70b",
Model: "groq/whisper-large-v3",
APIKeys: config.SimpleSecureStrings("sk-groq-model"),
},
},
},
wantName: "groq",
wantName: "whisper",
},
{
name: "missing voice model name config returns nil",

View file

@ -1,184 +0,0 @@
package asr
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"mime/multipart"
"net/http"
"os"
"path/filepath"
"time"
"github.com/sipeed/picoclaw/pkg/logger"
"github.com/sipeed/picoclaw/pkg/utils"
)
type GroqTranscriber struct {
apiKey string
apiBase string
httpClient *http.Client
}
func NewGroqTranscriber(apiKey string) *GroqTranscriber {
logger.DebugCF("voice", "Creating Groq transcriber", map[string]any{"has_api_key": apiKey != ""})
apiBase := "https://api.groq.com/openai/v1"
return &GroqTranscriber{
apiKey: apiKey,
apiBase: apiBase,
httpClient: &http.Client{
Timeout: 60 * time.Second,
},
}
}
func (t *GroqTranscriber) TranscribeData(
ctx context.Context,
data []byte,
filename string,
) (*TranscriptionResponse, error) {
logger.InfoCF("voice", "Starting memory transcription", map[string]any{"filename": filename, "bytes": len(data)})
var requestBody bytes.Buffer
writer := multipart.NewWriter(&requestBody)
part, err := writer.CreateFormFile("file", filename)
if err != nil {
logger.ErrorCF("voice", "Failed to create form file", map[string]any{"error": err})
return nil, fmt.Errorf("failed to create form file: %w", err)
}
if _, copyErr := io.Copy(part, bytes.NewReader(data)); copyErr != nil {
logger.ErrorCF("voice", "Failed to copy file content", map[string]any{"error": copyErr})
return nil, fmt.Errorf("failed to copy file content: %w", copyErr)
}
if err = writer.WriteField("model", "whisper-large-v3-turbo"); err != nil {
logger.ErrorCF("voice", "Failed to write model field", map[string]any{"error": err})
return nil, fmt.Errorf("failed to write model field: %w", err)
}
if err = writer.WriteField("response_format", "json"); err != nil {
logger.ErrorCF("voice", "Failed to write response_format field", map[string]any{"error": err})
return nil, fmt.Errorf("failed to write response_format field: %w", err)
}
if err = writer.Close(); err != nil {
logger.ErrorCF("voice", "Failed to close multipart writer", map[string]any{"error": err})
return nil, fmt.Errorf("failed to close multipart writer: %w", err)
}
return t.doRequest(ctx, &requestBody, writer.FormDataContentType(), int64(len(data)))
}
func (t *GroqTranscriber) Transcribe(ctx context.Context, audioFilePath string) (*TranscriptionResponse, error) {
logger.InfoCF("voice", "Starting transcription", map[string]any{"audio_file": audioFilePath})
audioFile, err := os.Open(audioFilePath)
if err != nil {
return nil, fmt.Errorf("failed to open audio file %s: %w", audioFilePath, err)
}
defer audioFile.Close()
fileInfo, err := audioFile.Stat()
if err != nil {
return nil, fmt.Errorf("failed to stat audio file %s: %w", audioFilePath, err)
}
var requestBody bytes.Buffer
writer := multipart.NewWriter(&requestBody)
part, err := writer.CreateFormFile("file", filepath.Base(audioFilePath))
if err != nil {
return nil, fmt.Errorf("failed to create form file: %w", err)
}
if _, copyErr := io.Copy(part, audioFile); copyErr != nil {
return nil, fmt.Errorf("failed to copy audio data: %w", copyErr)
}
if err = writer.WriteField("model", "whisper-large-v3-turbo"); err != nil {
return nil, fmt.Errorf("failed to write model field: %w", err)
}
if err = writer.WriteField("response_format", "json"); err != nil {
return nil, fmt.Errorf("failed to write response_format field: %w", err)
}
if err = writer.Close(); err != nil {
return nil, fmt.Errorf("failed to close multipart writer: %w", err)
}
return t.doRequest(ctx, &requestBody, writer.FormDataContentType(), fileInfo.Size())
}
func (t *GroqTranscriber) doRequest(
ctx context.Context,
requestBody *bytes.Buffer,
contentType string,
fileSize int64,
) (*TranscriptionResponse, error) {
url := t.apiBase + "/audio/transcriptions"
req, err := http.NewRequestWithContext(ctx, "POST", url, requestBody)
if err != nil {
logger.ErrorCF("voice", "Failed to create request", map[string]any{"error": err})
return nil, fmt.Errorf("failed to create request: %w", err)
}
req.Header.Set("Content-Type", contentType)
req.Header.Set("Authorization", "Bearer "+t.apiKey)
logger.DebugCF("voice", "Sending transcription request to Groq API", map[string]any{
"url": url,
"request_size_bytes": requestBody.Len(),
"file_size_bytes": fileSize,
})
resp, err := t.httpClient.Do(req)
if err != nil {
logger.ErrorCF("voice", "Failed to send request", map[string]any{"error": err})
return nil, fmt.Errorf("failed to send request: %w", err)
}
defer resp.Body.Close()
body, err := io.ReadAll(resp.Body)
if err != nil {
logger.ErrorCF("voice", "Failed to read response", map[string]any{"error": err})
return nil, fmt.Errorf("failed to read response: %w", err)
}
if resp.StatusCode != http.StatusOK {
logger.ErrorCF("voice", "API error", map[string]any{
"status_code": resp.StatusCode,
"response": string(body),
})
return nil, fmt.Errorf("API error (status %d): %s", resp.StatusCode, string(body))
}
logger.DebugCF("voice", "Received response from Groq API", map[string]any{
"status_code": resp.StatusCode,
"response_size_bytes": len(body),
})
var result TranscriptionResponse
if err := json.Unmarshal(body, &result); err != nil {
logger.ErrorCF("voice", "Failed to unmarshal response", map[string]any{"error": err})
return nil, fmt.Errorf("failed to unmarshal response: %w", err)
}
logger.InfoCF("voice", "Transcription completed successfully", map[string]any{
"text_length": len(result.Text),
"language": result.Language,
"duration_seconds": result.Duration,
"transcription_preview": utils.Truncate(result.Text, 50),
})
return &result, nil
}
func (t *GroqTranscriber) Name() string {
return "groq"
}

View file

@ -0,0 +1,245 @@
package asr
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"mime/multipart"
"net/http"
"os"
"path/filepath"
"strings"
"time"
"github.com/sipeed/picoclaw/pkg/config"
"github.com/sipeed/picoclaw/pkg/logger"
"github.com/sipeed/picoclaw/pkg/providers"
"github.com/sipeed/picoclaw/pkg/utils"
)
type WhisperTranscriber struct {
apiKey string
apiBase string
modelID string
providerName string
httpClient *http.Client
}
func NewWhisperTranscriber(modelCfg *config.ModelConfig) *WhisperTranscriber {
if modelCfg == nil {
return nil
}
protocol, modelID := providers.ExtractProtocol(modelCfg.Model)
if modelID == "" {
modelID = strings.TrimSpace(modelCfg.Model)
}
tr := newWhisperTranscriber(
modelCfg.APIKey(),
providers.ResolveAPIBase(modelCfg),
modelID,
protocol,
)
if tr == nil {
return nil
}
logger.DebugCF("voice", "Creating whisper transcriber", map[string]any{
"api_base": tr.apiBase,
"has_key": tr.apiKey != "",
"model": tr.modelID,
"provider": tr.providerName,
})
return tr
}
func NewGroqTranscriber(apiKey, modelID string) *WhisperTranscriber {
return newWhisperTranscriber(apiKey, "https://api.groq.com/openai/v1", modelID, "groq")
}
func newWhisperTranscriber(apiKey, apiBase, modelID, providerName string) *WhisperTranscriber {
if modelID == "" {
return nil
}
if providerName == "" {
providerName = "whisper"
}
return &WhisperTranscriber{
apiKey: apiKey,
apiBase: strings.TrimRight(apiBase, "/"),
modelID: modelID,
providerName: providerName,
httpClient: &http.Client{
Timeout: 60 * time.Second,
},
}
}
func (t *WhisperTranscriber) transcriptionURL() string {
base := strings.TrimRight(t.apiBase, "/")
if strings.HasSuffix(base, "/audio/transcriptions") {
return base
}
return base + "/audio/transcriptions"
}
func (t *WhisperTranscriber) TranscribeData(
ctx context.Context,
data []byte,
filename string,
) (*TranscriptionResponse, error) {
logger.InfoCF("voice", "Starting whisper transcription from memory", map[string]any{
"bytes": len(data),
"filename": filename,
"model": t.modelID,
"provider": t.providerName,
})
var requestBody bytes.Buffer
writer := multipart.NewWriter(&requestBody)
part, err := writer.CreateFormFile("file", filename)
if err != nil {
logger.ErrorCF("voice", "Failed to create whisper form file", map[string]any{"error": err})
return nil, fmt.Errorf("failed to create form file: %w", err)
}
if _, copyErr := io.Copy(part, bytes.NewReader(data)); copyErr != nil {
logger.ErrorCF("voice", "Failed to copy whisper file content", map[string]any{"error": copyErr})
return nil, fmt.Errorf("failed to copy file content: %w", copyErr)
}
if err = writer.WriteField("model", t.modelID); err != nil {
logger.ErrorCF("voice", "Failed to write whisper model field", map[string]any{"error": err})
return nil, fmt.Errorf("failed to write model field: %w", err)
}
if err = writer.WriteField("response_format", "json"); err != nil {
logger.ErrorCF("voice", "Failed to write whisper response_format field", map[string]any{"error": err})
return nil, fmt.Errorf("failed to write response_format field: %w", err)
}
if err = writer.Close(); err != nil {
logger.ErrorCF("voice", "Failed to close whisper multipart writer", map[string]any{"error": err})
return nil, fmt.Errorf("failed to close multipart writer: %w", err)
}
return t.doRequest(ctx, &requestBody, writer.FormDataContentType(), int64(len(data)))
}
func (t *WhisperTranscriber) Transcribe(ctx context.Context, audioFilePath string) (*TranscriptionResponse, error) {
logger.InfoCF("voice", "Starting whisper transcription", map[string]any{
"audio_file": audioFilePath,
"model": t.modelID,
"provider": t.providerName,
})
audioFile, err := os.Open(audioFilePath)
if err != nil {
return nil, fmt.Errorf("failed to open audio file %s: %w", audioFilePath, err)
}
defer audioFile.Close()
fileInfo, err := audioFile.Stat()
if err != nil {
return nil, fmt.Errorf("failed to stat audio file %s: %w", audioFilePath, err)
}
var requestBody bytes.Buffer
writer := multipart.NewWriter(&requestBody)
part, err := writer.CreateFormFile("file", filepath.Base(audioFilePath))
if err != nil {
return nil, fmt.Errorf("failed to create form file: %w", err)
}
if _, copyErr := io.Copy(part, audioFile); copyErr != nil {
return nil, fmt.Errorf("failed to copy audio data: %w", copyErr)
}
if err = writer.WriteField("model", t.modelID); err != nil {
return nil, fmt.Errorf("failed to write model field: %w", err)
}
if err = writer.WriteField("response_format", "json"); err != nil {
return nil, fmt.Errorf("failed to write response_format field: %w", err)
}
if err = writer.Close(); err != nil {
return nil, fmt.Errorf("failed to close multipart writer: %w", err)
}
return t.doRequest(ctx, &requestBody, writer.FormDataContentType(), fileInfo.Size())
}
func (t *WhisperTranscriber) doRequest(
ctx context.Context,
requestBody *bytes.Buffer,
contentType string,
fileSize int64,
) (*TranscriptionResponse, error) {
url := t.transcriptionURL()
req, err := http.NewRequestWithContext(ctx, "POST", url, requestBody)
if err != nil {
logger.ErrorCF("voice", "Failed to create whisper request", map[string]any{"error": err})
return nil, fmt.Errorf("failed to create request: %w", err)
}
req.Header.Set("Content-Type", contentType)
if t.apiKey != "" {
req.Header.Set("Authorization", "Bearer "+t.apiKey)
}
logger.DebugCF("voice", "Sending whisper transcription request", map[string]any{
"file_size_bytes": fileSize,
"model": t.modelID,
"provider": t.providerName,
"request_size_bytes": requestBody.Len(),
"url": url,
})
resp, err := t.httpClient.Do(req)
if err != nil {
logger.ErrorCF("voice", "Failed to send whisper request", map[string]any{"error": err})
return nil, fmt.Errorf("failed to send request: %w", err)
}
defer resp.Body.Close()
body, err := io.ReadAll(resp.Body)
if err != nil {
logger.ErrorCF("voice", "Failed to read whisper response", map[string]any{"error": err})
return nil, fmt.Errorf("failed to read response: %w", err)
}
if resp.StatusCode != http.StatusOK {
logger.ErrorCF("voice", "Whisper API error", map[string]any{
"provider": t.providerName,
"response": string(body),
"status_code": resp.StatusCode,
})
return nil, fmt.Errorf("API error (status %d): %s", resp.StatusCode, string(body))
}
var result TranscriptionResponse
if err := json.Unmarshal(body, &result); err != nil {
logger.ErrorCF("voice", "Failed to unmarshal whisper response", map[string]any{"error": err})
return nil, fmt.Errorf("failed to unmarshal response: %w", err)
}
logger.InfoCF("voice", "Whisper transcription completed successfully", map[string]any{
"duration_seconds": result.Duration,
"language": result.Language,
"provider": t.providerName,
"text_length": len(result.Text),
"transcription_preview": utils.Truncate(result.Text, 50),
})
return &result, nil
}
func (t *WhisperTranscriber) Name() string {
return "whisper"
}

View file

@ -0,0 +1,102 @@
package asr
import (
"context"
"encoding/json"
"io"
"net/http"
"net/http/httptest"
"testing"
"github.com/sipeed/picoclaw/pkg/config"
)
func TestWhisperTranscriberTranscribeDataUsesConfiguredModel(t *testing.T) {
var gotModel string
var gotPath string
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
gotPath = r.URL.Path
if got := r.Header.Get("Authorization"); got != "Bearer sk-openai-test" {
t.Errorf("Authorization = %q, want %q", got, "Bearer sk-openai-test")
}
reader, err := r.MultipartReader()
if err != nil {
t.Fatalf("MultipartReader() error: %v", err)
}
for {
part, err := reader.NextPart()
if err == io.EOF {
break
}
if err != nil {
t.Fatalf("NextPart() error: %v", err)
}
data, err := io.ReadAll(part)
if err != nil {
t.Fatalf("ReadAll() error: %v", err)
}
if part.FormName() == "model" {
gotModel = string(data)
}
}
w.Header().Set("Content-Type", "application/json")
if err := json.NewEncoder(w).Encode(TranscriptionResponse{Text: "hello from whisper"}); err != nil {
t.Fatalf("Encode() error: %v", err)
}
}))
defer server.Close()
tr := NewWhisperTranscriber(&config.ModelConfig{
Model: "openai/whisper-1",
APIBase: server.URL,
APIKeys: config.SimpleSecureStrings("sk-openai-test"),
})
tr.httpClient = server.Client()
resp, err := tr.TranscribeData(context.Background(), []byte("audio"), "clip.ogg")
if err != nil {
t.Fatalf("TranscribeData() error: %v", err)
}
if resp.Text != "hello from whisper" {
t.Errorf("Text = %q, want %q", resp.Text, "hello from whisper")
}
if gotModel != "whisper-1" {
t.Errorf("model field = %q, want %q", gotModel, "whisper-1")
}
if gotPath != "/audio/transcriptions" {
t.Errorf("path = %q, want %q", gotPath, "/audio/transcriptions")
}
}
func TestWhisperTranscriberUsesEndpointAPIBaseWithoutDoubleAppend(t *testing.T) {
var gotPath string
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
gotPath = r.URL.Path
w.Header().Set("Content-Type", "application/json")
if err := json.NewEncoder(w).Encode(TranscriptionResponse{Text: "ok"}); err != nil {
t.Fatalf("Encode() error: %v", err)
}
}))
defer server.Close()
tr := NewWhisperTranscriber(&config.ModelConfig{
Model: "groq/whisper-large-v3",
APIBase: server.URL + "/audio/transcriptions",
APIKeys: config.SimpleSecureStrings("sk-groq-test"),
})
tr.httpClient = server.Client()
if _, err := tr.TranscribeData(context.Background(), []byte("audio"), "clip.ogg"); err != nil {
t.Fatalf("TranscribeData() error: %v", err)
}
if gotPath != "/audio/transcriptions" {
t.Errorf("path = %q, want %q", gotPath, "/audio/transcriptions")
}
}

View file

@ -56,6 +56,19 @@ func ExtractProtocol(model string) (protocol, modelID string) {
return protocol, modelID
}
// ResolveAPIBase returns the configured API base, or the protocol default when
// the model uses an HTTP-based provider family with a known default endpoint.
func ResolveAPIBase(cfg *config.ModelConfig) string {
if cfg == nil {
return ""
}
if apiBase := strings.TrimSpace(cfg.APIBase); apiBase != "" {
return strings.TrimRight(apiBase, "/")
}
protocol, _ := ExtractProtocol(cfg.Model)
return strings.TrimRight(getDefaultAPIBase(protocol), "/")
}
// CreateProviderFromConfig creates a provider based on the ModelConfig.
// It uses the protocol prefix in the Model field to determine which provider to create.
// Supported protocol families include OpenAI-compatible prefixes (e.g., openai, openrouter, groq, gemini),