feat: generalize STT to support any OpenAI-compatible Whisper endpoint
Extract Transcriber interface and rename GroqTranscriber to OpenAICompatTranscriber. Add stt_model field to agents.defaults for explicit STT model selection via model_list. 3-tier backward-compatible resolution: 1. agents.defaults.stt_model → model_list lookup 2. providers.groq.api_key (legacy) 3. groq/ prefix in model_list (legacy) Also attach transcriber to OneBot channel (was missing).
This commit is contained in:
parent
cb0c8703fb
commit
c05ac2e8cc
11 changed files with 465 additions and 36 deletions
|
|
@ -10,7 +10,6 @@ import (
|
||||||
"os"
|
"os"
|
||||||
"os/signal"
|
"os/signal"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strings"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/sipeed/picoclaw/pkg/agent"
|
"github.com/sipeed/picoclaw/pkg/agent"
|
||||||
|
|
@ -25,7 +24,6 @@ import (
|
||||||
"github.com/sipeed/picoclaw/pkg/providers"
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
"github.com/sipeed/picoclaw/pkg/state"
|
"github.com/sipeed/picoclaw/pkg/state"
|
||||||
"github.com/sipeed/picoclaw/pkg/tools"
|
"github.com/sipeed/picoclaw/pkg/tools"
|
||||||
"github.com/sipeed/picoclaw/pkg/voice"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func gatewayCmd() {
|
func gatewayCmd() {
|
||||||
|
|
@ -121,38 +119,35 @@ func gatewayCmd() {
|
||||||
// Inject channel manager into agent loop for command handling
|
// Inject channel manager into agent loop for command handling
|
||||||
agentLoop.SetChannelManager(channelManager)
|
agentLoop.SetChannelManager(channelManager)
|
||||||
|
|
||||||
var transcriber *voice.GroqTranscriber
|
// STT (Speech-to-Text) transcriber setup
|
||||||
groqAPIKey := cfg.Providers.Groq.APIKey
|
transcriber := resolveSTTTranscriber(cfg)
|
||||||
if groqAPIKey == "" {
|
if transcriber != nil {
|
||||||
for _, mc := range cfg.ModelList {
|
logger.InfoC("voice", "STT voice transcription enabled")
|
||||||
if strings.HasPrefix(mc.Model, "groq/") && mc.APIKey != "" {
|
|
||||||
groqAPIKey = mc.APIKey
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if groqAPIKey != "" {
|
|
||||||
transcriber = voice.NewGroqTranscriber(groqAPIKey)
|
|
||||||
logger.InfoC("voice", "Groq voice transcription enabled")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if transcriber != nil {
|
if transcriber != nil {
|
||||||
if telegramChannel, ok := channelManager.GetChannel("telegram"); ok {
|
if telegramChannel, ok := channelManager.GetChannel("telegram"); ok {
|
||||||
if tc, ok := telegramChannel.(*channels.TelegramChannel); ok {
|
if tc, ok := telegramChannel.(*channels.TelegramChannel); ok {
|
||||||
tc.SetTranscriber(transcriber)
|
tc.SetTranscriber(transcriber)
|
||||||
logger.InfoC("voice", "Groq transcription attached to Telegram channel")
|
logger.InfoC("voice", "STT transcription attached to Telegram channel")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if discordChannel, ok := channelManager.GetChannel("discord"); ok {
|
if discordChannel, ok := channelManager.GetChannel("discord"); ok {
|
||||||
if dc, ok := discordChannel.(*channels.DiscordChannel); ok {
|
if dc, ok := discordChannel.(*channels.DiscordChannel); ok {
|
||||||
dc.SetTranscriber(transcriber)
|
dc.SetTranscriber(transcriber)
|
||||||
logger.InfoC("voice", "Groq transcription attached to Discord channel")
|
logger.InfoC("voice", "STT transcription attached to Discord channel")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if slackChannel, ok := channelManager.GetChannel("slack"); ok {
|
if slackChannel, ok := channelManager.GetChannel("slack"); ok {
|
||||||
if sc, ok := slackChannel.(*channels.SlackChannel); ok {
|
if sc, ok := slackChannel.(*channels.SlackChannel); ok {
|
||||||
sc.SetTranscriber(transcriber)
|
sc.SetTranscriber(transcriber)
|
||||||
logger.InfoC("voice", "Groq transcription attached to Slack channel")
|
logger.InfoC("voice", "STT transcription attached to Slack channel")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if onebotChannel, ok := channelManager.GetChannel("onebot"); ok {
|
||||||
|
if oc, ok := onebotChannel.(*channels.OneBotChannel); ok {
|
||||||
|
oc.SetTranscriber(transcriber)
|
||||||
|
logger.InfoC("voice", "STT transcription attached to OneBot channel")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
@ -246,3 +241,14 @@ func setupCronTool(
|
||||||
|
|
||||||
return cronService
|
return cronService
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func getDefaultSTTBase(protocol string) string {
|
||||||
|
switch protocol {
|
||||||
|
case "openai":
|
||||||
|
return "https://api.openai.com/v1"
|
||||||
|
case "groq":
|
||||||
|
return "https://api.groq.com/openai/v1"
|
||||||
|
default:
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
51
cmd/picoclaw/stt_resolution.go
Normal file
51
cmd/picoclaw/stt_resolution.go
Normal file
|
|
@ -0,0 +1,51 @@
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/providers"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/voice"
|
||||||
|
)
|
||||||
|
|
||||||
|
// resolveSTTTranscriber resolves the STT transcriber using a 3-tier fallback strategy:
|
||||||
|
// 1. agents.defaults.stt_model → model_list lookup
|
||||||
|
// 2. providers.groq.api_key (backward compat)
|
||||||
|
// 3. groq/ prefix in model_list (backward compat)
|
||||||
|
func resolveSTTTranscriber(cfg *config.Config) voice.Transcriber {
|
||||||
|
// 1. Resolve from agents.defaults.stt_model → model_list lookup
|
||||||
|
if cfg.Agents.Defaults.STTModel != "" {
|
||||||
|
for _, mc := range cfg.ModelList {
|
||||||
|
if mc.ModelName == cfg.Agents.Defaults.STTModel && mc.APIKey != "" {
|
||||||
|
protocol, modelID := providers.ExtractProtocol(mc.Model)
|
||||||
|
apiBase := mc.APIBase
|
||||||
|
if apiBase == "" {
|
||||||
|
apiBase = getDefaultSTTBase(protocol)
|
||||||
|
}
|
||||||
|
return voice.NewOpenAICompatTranscriber(mc.APIKey, apiBase, modelID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// 2. Backward compat: providers.groq.api_key
|
||||||
|
if cfg.Providers.Groq.APIKey != "" {
|
||||||
|
return voice.NewOpenAICompatTranscriber(
|
||||||
|
cfg.Providers.Groq.APIKey,
|
||||||
|
"https://api.groq.com/openai/v1",
|
||||||
|
"whisper-large-v3",
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
// 3. Backward compat: groq/ in model_list (no stt_model set)
|
||||||
|
for _, mc := range cfg.ModelList {
|
||||||
|
if strings.HasPrefix(mc.Model, "groq/") && mc.APIKey != "" {
|
||||||
|
return voice.NewOpenAICompatTranscriber(
|
||||||
|
mc.APIKey,
|
||||||
|
"https://api.groq.com/openai/v1",
|
||||||
|
"whisper-large-v3",
|
||||||
|
)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
|
}
|
||||||
241
cmd/picoclaw/stt_resolution_test.go
Normal file
241
cmd/picoclaw/stt_resolution_test.go
Normal file
|
|
@ -0,0 +1,241 @@
|
||||||
|
package main
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/config"
|
||||||
|
"github.com/sipeed/picoclaw/pkg/voice"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestGetDefaultSTTBase(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
protocol string
|
||||||
|
expected string
|
||||||
|
}{
|
||||||
|
{"openai", "https://api.openai.com/v1"},
|
||||||
|
{"groq", "https://api.groq.com/openai/v1"},
|
||||||
|
{"unknown", ""},
|
||||||
|
{"", ""},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.protocol, func(t *testing.T) {
|
||||||
|
if got := getDefaultSTTBase(tt.protocol); got != tt.expected {
|
||||||
|
t.Errorf("getDefaultSTTBase(%q) = %q, want %q", tt.protocol, got, tt.expected)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveSTTTranscriber_STTModel(t *testing.T) {
|
||||||
|
cfg := &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
STTModel: "whisper",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
ModelList: []config.ModelConfig{
|
||||||
|
{
|
||||||
|
ModelName: "whisper",
|
||||||
|
Model: "openai/whisper-1",
|
||||||
|
APIKey: "sk-test",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
tr := resolveSTTTranscriber(cfg)
|
||||||
|
if tr == nil {
|
||||||
|
t.Fatal("expected transcriber, got nil")
|
||||||
|
}
|
||||||
|
if !tr.IsAvailable() {
|
||||||
|
t.Error("expected transcriber to be available")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveSTTTranscriber_STTModelWithAPIBase(t *testing.T) {
|
||||||
|
cfg := &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
STTModel: "whisper",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
ModelList: []config.ModelConfig{
|
||||||
|
{
|
||||||
|
ModelName: "whisper",
|
||||||
|
Model: "openai/whisper-1",
|
||||||
|
APIKey: "sk-test",
|
||||||
|
APIBase: "https://custom.api.com/v1",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
tr := resolveSTTTranscriber(cfg)
|
||||||
|
if tr == nil {
|
||||||
|
t.Fatal("expected transcriber, got nil")
|
||||||
|
}
|
||||||
|
// Verify it's the right type and has the right fields
|
||||||
|
oat, ok := tr.(*voice.OpenAICompatTranscriber)
|
||||||
|
if !ok {
|
||||||
|
t.Fatal("expected *voice.OpenAICompatTranscriber")
|
||||||
|
}
|
||||||
|
_ = oat // Can't access unexported fields from test, but the resolution worked
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveSTTTranscriber_STTModelGroq(t *testing.T) {
|
||||||
|
cfg := &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
STTModel: "whisper",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
ModelList: []config.ModelConfig{
|
||||||
|
{
|
||||||
|
ModelName: "whisper",
|
||||||
|
Model: "groq/whisper-large-v3",
|
||||||
|
APIKey: "gsk-test",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
tr := resolveSTTTranscriber(cfg)
|
||||||
|
if tr == nil {
|
||||||
|
t.Fatal("expected transcriber, got nil")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveSTTTranscriber_BackwardCompatGroqProvider(t *testing.T) {
|
||||||
|
cfg := &config.Config{
|
||||||
|
Providers: config.ProvidersConfig{
|
||||||
|
Groq: config.ProviderConfig{
|
||||||
|
APIKey: "gsk-test-key",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
tr := resolveSTTTranscriber(cfg)
|
||||||
|
if tr == nil {
|
||||||
|
t.Fatal("expected transcriber, got nil")
|
||||||
|
}
|
||||||
|
if !tr.IsAvailable() {
|
||||||
|
t.Error("expected transcriber to be available")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveSTTTranscriber_BackwardCompatGroqModelList(t *testing.T) {
|
||||||
|
cfg := &config.Config{
|
||||||
|
ModelList: []config.ModelConfig{
|
||||||
|
{
|
||||||
|
ModelName: "groq-llama",
|
||||||
|
Model: "groq/llama-3.3-70b",
|
||||||
|
APIKey: "gsk-test-key",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
tr := resolveSTTTranscriber(cfg)
|
||||||
|
if tr == nil {
|
||||||
|
t.Fatal("expected transcriber, got nil")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveSTTTranscriber_NoneAvailable(t *testing.T) {
|
||||||
|
cfg := &config.Config{}
|
||||||
|
tr := resolveSTTTranscriber(cfg)
|
||||||
|
if tr != nil {
|
||||||
|
t.Error("expected nil transcriber when no config available")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveSTTTranscriber_STTModelPriority(t *testing.T) {
|
||||||
|
// stt_model should take priority over providers.groq.api_key
|
||||||
|
cfg := &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
STTModel: "whisper",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
ModelList: []config.ModelConfig{
|
||||||
|
{
|
||||||
|
ModelName: "whisper",
|
||||||
|
Model: "openai/whisper-1",
|
||||||
|
APIKey: "sk-openai-key",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: config.ProvidersConfig{
|
||||||
|
Groq: config.ProviderConfig{
|
||||||
|
APIKey: "gsk-groq-key",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
tr := resolveSTTTranscriber(cfg)
|
||||||
|
if tr == nil {
|
||||||
|
t.Fatal("expected transcriber, got nil")
|
||||||
|
}
|
||||||
|
// The transcriber should be from stt_model (OpenAI), not from Groq
|
||||||
|
// We can verify by checking it's available
|
||||||
|
if !tr.IsAvailable() {
|
||||||
|
t.Error("expected transcriber to be available")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveSTTTranscriber_STTModelNotInModelList(t *testing.T) {
|
||||||
|
// stt_model set but not found in model_list, should fall back
|
||||||
|
cfg := &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
STTModel: "nonexistent",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
Providers: config.ProvidersConfig{
|
||||||
|
Groq: config.ProviderConfig{
|
||||||
|
APIKey: "gsk-fallback",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
tr := resolveSTTTranscriber(cfg)
|
||||||
|
if tr == nil {
|
||||||
|
t.Fatal("expected transcriber from fallback, got nil")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveSTTTranscriber_LLMEntryNotMatchedAsSTT(t *testing.T) {
|
||||||
|
// A non-groq LLM entry should NOT be matched as STT provider
|
||||||
|
cfg := &config.Config{
|
||||||
|
ModelList: []config.ModelConfig{
|
||||||
|
{
|
||||||
|
ModelName: "gpt4",
|
||||||
|
Model: "openai/gpt-4o",
|
||||||
|
APIKey: "sk-test",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
tr := resolveSTTTranscriber(cfg)
|
||||||
|
if tr != nil {
|
||||||
|
t.Error("expected nil - LLM entries should not match as STT providers")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveSTTTranscriber_STTModelNoAPIKey(t *testing.T) {
|
||||||
|
// stt_model found but no API key, should fall back
|
||||||
|
cfg := &config.Config{
|
||||||
|
Agents: config.AgentsConfig{
|
||||||
|
Defaults: config.AgentDefaults{
|
||||||
|
STTModel: "whisper",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
ModelList: []config.ModelConfig{
|
||||||
|
{
|
||||||
|
ModelName: "whisper",
|
||||||
|
Model: "openai/whisper-1",
|
||||||
|
APIKey: "",
|
||||||
|
},
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
tr := resolveSTTTranscriber(cfg)
|
||||||
|
if tr != nil {
|
||||||
|
t.Error("expected nil - model has no API key")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -4,6 +4,7 @@
|
||||||
"workspace": "~/.picoclaw/workspace",
|
"workspace": "~/.picoclaw/workspace",
|
||||||
"restrict_to_workspace": true,
|
"restrict_to_workspace": true,
|
||||||
"model": "gpt4",
|
"model": "gpt4",
|
||||||
|
"stt_model": "",
|
||||||
"max_tokens": 8192,
|
"max_tokens": 8192,
|
||||||
"temperature": 0.7,
|
"temperature": 0.7,
|
||||||
"max_tool_iterations": 20
|
"max_tool_iterations": 20
|
||||||
|
|
@ -43,6 +44,12 @@
|
||||||
"model": "openai/gpt-5.2",
|
"model": "openai/gpt-5.2",
|
||||||
"api_key": "sk-key2",
|
"api_key": "sk-key2",
|
||||||
"api_base": "https://api2.example.com/v1"
|
"api_base": "https://api2.example.com/v1"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"model_name": "whisper",
|
||||||
|
"model": "groq/whisper-large-v3",
|
||||||
|
"api_key": "gsk_xxx",
|
||||||
|
"_comment": "STT model for voice transcription. Set agents.defaults.stt_model to 'whisper' to use."
|
||||||
}
|
}
|
||||||
],
|
],
|
||||||
"channels": {
|
"channels": {
|
||||||
|
|
|
||||||
|
|
@ -26,7 +26,7 @@ type DiscordChannel struct {
|
||||||
*BaseChannel
|
*BaseChannel
|
||||||
session *discordgo.Session
|
session *discordgo.Session
|
||||||
config config.DiscordConfig
|
config config.DiscordConfig
|
||||||
transcriber *voice.GroqTranscriber
|
transcriber voice.Transcriber
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
typingMu sync.Mutex
|
typingMu sync.Mutex
|
||||||
typingStop map[string]chan struct{} // chatID → stop signal
|
typingStop map[string]chan struct{} // chatID → stop signal
|
||||||
|
|
@ -51,7 +51,7 @@ func NewDiscordChannel(cfg config.DiscordConfig, bus *bus.MessageBus) (*DiscordC
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *DiscordChannel) SetTranscriber(transcriber *voice.GroqTranscriber) {
|
func (c *DiscordChannel) SetTranscriber(transcriber voice.Transcriber) {
|
||||||
c.transcriber = transcriber
|
c.transcriber = transcriber
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -35,7 +35,7 @@ type OneBotChannel struct {
|
||||||
selfID int64
|
selfID int64
|
||||||
pending map[string]chan json.RawMessage
|
pending map[string]chan json.RawMessage
|
||||||
pendingMu sync.Mutex
|
pendingMu sync.Mutex
|
||||||
transcriber *voice.GroqTranscriber
|
transcriber voice.Transcriber
|
||||||
lastMessageID sync.Map
|
lastMessageID sync.Map
|
||||||
pendingEmojiMsg sync.Map
|
pendingEmojiMsg sync.Map
|
||||||
}
|
}
|
||||||
|
|
@ -111,7 +111,7 @@ func NewOneBotChannel(cfg config.OneBotConfig, messageBus *bus.MessageBus) (*One
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *OneBotChannel) SetTranscriber(transcriber *voice.GroqTranscriber) {
|
func (c *OneBotChannel) SetTranscriber(transcriber voice.Transcriber) {
|
||||||
c.transcriber = transcriber
|
c.transcriber = transcriber
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -26,7 +26,7 @@ type SlackChannel struct {
|
||||||
socketClient *socketmode.Client
|
socketClient *socketmode.Client
|
||||||
botUserID string
|
botUserID string
|
||||||
teamID string
|
teamID string
|
||||||
transcriber *voice.GroqTranscriber
|
transcriber voice.Transcriber
|
||||||
ctx context.Context
|
ctx context.Context
|
||||||
cancel context.CancelFunc
|
cancel context.CancelFunc
|
||||||
pendingAcks sync.Map
|
pendingAcks sync.Map
|
||||||
|
|
@ -59,7 +59,7 @@ func NewSlackChannel(cfg config.SlackConfig, messageBus *bus.MessageBus) (*Slack
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *SlackChannel) SetTranscriber(transcriber *voice.GroqTranscriber) {
|
func (c *SlackChannel) SetTranscriber(transcriber voice.Transcriber) {
|
||||||
c.transcriber = transcriber
|
c.transcriber = transcriber
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -29,7 +29,7 @@ type TelegramChannel struct {
|
||||||
commands TelegramCommander
|
commands TelegramCommander
|
||||||
config *config.Config
|
config *config.Config
|
||||||
chatIDs map[string]int64
|
chatIDs map[string]int64
|
||||||
transcriber *voice.GroqTranscriber
|
transcriber voice.Transcriber
|
||||||
placeholders sync.Map // chatID -> messageID
|
placeholders sync.Map // chatID -> messageID
|
||||||
stopThinking sync.Map // chatID -> thinkingCancel
|
stopThinking sync.Map // chatID -> thinkingCancel
|
||||||
}
|
}
|
||||||
|
|
@ -86,7 +86,7 @@ func NewTelegramChannel(cfg *config.Config, bus *bus.MessageBus) (*TelegramChann
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *TelegramChannel) SetTranscriber(transcriber *voice.GroqTranscriber) {
|
func (c *TelegramChannel) SetTranscriber(transcriber voice.Transcriber) {
|
||||||
c.transcriber = transcriber
|
c.transcriber = transcriber
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -177,6 +177,7 @@ type AgentDefaults struct {
|
||||||
MaxTokens int `json:"max_tokens" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_TOKENS"`
|
MaxTokens int `json:"max_tokens" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_TOKENS"`
|
||||||
Temperature *float64 `json:"temperature,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_TEMPERATURE"`
|
Temperature *float64 `json:"temperature,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_TEMPERATURE"`
|
||||||
MaxToolIterations int `json:"max_tool_iterations" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_TOOL_ITERATIONS"`
|
MaxToolIterations int `json:"max_tool_iterations" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_TOOL_ITERATIONS"`
|
||||||
|
STTModel string `json:"stt_model,omitempty" env:"PICOCLAW_AGENTS_DEFAULTS_STT_MODEL"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type ChannelsConfig struct {
|
type ChannelsConfig struct {
|
||||||
|
|
|
||||||
|
|
@ -16,9 +16,15 @@ import (
|
||||||
"github.com/sipeed/picoclaw/pkg/utils"
|
"github.com/sipeed/picoclaw/pkg/utils"
|
||||||
)
|
)
|
||||||
|
|
||||||
type GroqTranscriber struct {
|
type Transcriber interface {
|
||||||
|
Transcribe(ctx context.Context, audioFilePath string) (*TranscriptionResponse, error)
|
||||||
|
IsAvailable() bool
|
||||||
|
}
|
||||||
|
|
||||||
|
type OpenAICompatTranscriber struct {
|
||||||
apiKey string
|
apiKey string
|
||||||
apiBase string
|
apiBase string
|
||||||
|
model string
|
||||||
httpClient *http.Client
|
httpClient *http.Client
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -28,20 +34,20 @@ type TranscriptionResponse struct {
|
||||||
Duration float64 `json:"duration,omitempty"`
|
Duration float64 `json:"duration,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewGroqTranscriber(apiKey string) *GroqTranscriber {
|
func NewOpenAICompatTranscriber(apiKey, apiBase, model string) *OpenAICompatTranscriber {
|
||||||
logger.DebugCF("voice", "Creating Groq transcriber", map[string]any{"has_api_key": apiKey != ""})
|
logger.DebugCF("voice", "Creating STT transcriber", map[string]any{"has_api_key": apiKey != ""})
|
||||||
|
|
||||||
apiBase := "https://api.groq.com/openai/v1"
|
return &OpenAICompatTranscriber{
|
||||||
return &GroqTranscriber{
|
|
||||||
apiKey: apiKey,
|
apiKey: apiKey,
|
||||||
apiBase: apiBase,
|
apiBase: apiBase,
|
||||||
|
model: model,
|
||||||
httpClient: &http.Client{
|
httpClient: &http.Client{
|
||||||
Timeout: 60 * time.Second,
|
Timeout: 60 * time.Second,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *GroqTranscriber) Transcribe(ctx context.Context, audioFilePath string) (*TranscriptionResponse, error) {
|
func (t *OpenAICompatTranscriber) Transcribe(ctx context.Context, audioFilePath string) (*TranscriptionResponse, error) {
|
||||||
logger.InfoCF("voice", "Starting transcription", map[string]any{"audio_file": audioFilePath})
|
logger.InfoCF("voice", "Starting transcription", map[string]any{"audio_file": audioFilePath})
|
||||||
|
|
||||||
audioFile, err := os.Open(audioFilePath)
|
audioFile, err := os.Open(audioFilePath)
|
||||||
|
|
@ -79,7 +85,7 @@ func (t *GroqTranscriber) Transcribe(ctx context.Context, audioFilePath string)
|
||||||
|
|
||||||
logger.DebugCF("voice", "File copied to request", map[string]any{"bytes_copied": copied})
|
logger.DebugCF("voice", "File copied to request", map[string]any{"bytes_copied": copied})
|
||||||
|
|
||||||
if err = writer.WriteField("model", "whisper-large-v3"); err != nil {
|
if err = writer.WriteField("model", t.model); err != nil {
|
||||||
logger.ErrorCF("voice", "Failed to write model field", map[string]any{"error": err})
|
logger.ErrorCF("voice", "Failed to write model field", map[string]any{"error": err})
|
||||||
return nil, fmt.Errorf("failed to write model field: %w", err)
|
return nil, fmt.Errorf("failed to write model field: %w", err)
|
||||||
}
|
}
|
||||||
|
|
@ -104,7 +110,7 @@ func (t *GroqTranscriber) Transcribe(ctx context.Context, audioFilePath string)
|
||||||
req.Header.Set("Content-Type", writer.FormDataContentType())
|
req.Header.Set("Content-Type", writer.FormDataContentType())
|
||||||
req.Header.Set("Authorization", "Bearer "+t.apiKey)
|
req.Header.Set("Authorization", "Bearer "+t.apiKey)
|
||||||
|
|
||||||
logger.DebugCF("voice", "Sending transcription request to Groq API", map[string]any{
|
logger.DebugCF("voice", "Sending transcription request to STT API", map[string]any{
|
||||||
"url": url,
|
"url": url,
|
||||||
"request_size_bytes": requestBody.Len(),
|
"request_size_bytes": requestBody.Len(),
|
||||||
"file_size_bytes": fileInfo.Size(),
|
"file_size_bytes": fileInfo.Size(),
|
||||||
|
|
@ -131,7 +137,7 @@ func (t *GroqTranscriber) Transcribe(ctx context.Context, audioFilePath string)
|
||||||
return nil, fmt.Errorf("API error (status %d): %s", resp.StatusCode, 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{
|
logger.DebugCF("voice", "Received response from STT API", map[string]any{
|
||||||
"status_code": resp.StatusCode,
|
"status_code": resp.StatusCode,
|
||||||
"response_size_bytes": len(body),
|
"response_size_bytes": len(body),
|
||||||
})
|
})
|
||||||
|
|
@ -152,7 +158,7 @@ func (t *GroqTranscriber) Transcribe(ctx context.Context, audioFilePath string)
|
||||||
return &result, nil
|
return &result, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (t *GroqTranscriber) IsAvailable() bool {
|
func (t *OpenAICompatTranscriber) IsAvailable() bool {
|
||||||
available := t.apiKey != ""
|
available := t.apiKey != ""
|
||||||
logger.DebugCF("voice", "Checking transcriber availability", map[string]any{"available": available})
|
logger.DebugCF("voice", "Checking transcriber availability", map[string]any{"available": available})
|
||||||
return available
|
return available
|
||||||
|
|
|
||||||
117
pkg/voice/transcriber_test.go
Normal file
117
pkg/voice/transcriber_test.go
Normal file
|
|
@ -0,0 +1,117 @@
|
||||||
|
package voice
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"os"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestNewOpenAICompatTranscriber(t *testing.T) {
|
||||||
|
tr := NewOpenAICompatTranscriber("test-key", "https://api.example.com/v1", "whisper-1")
|
||||||
|
if tr.apiKey != "test-key" {
|
||||||
|
t.Errorf("expected apiKey 'test-key', got %q", tr.apiKey)
|
||||||
|
}
|
||||||
|
if tr.apiBase != "https://api.example.com/v1" {
|
||||||
|
t.Errorf("expected apiBase 'https://api.example.com/v1', got %q", tr.apiBase)
|
||||||
|
}
|
||||||
|
if tr.model != "whisper-1" {
|
||||||
|
t.Errorf("expected model 'whisper-1', got %q", tr.model)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOpenAICompatTranscriber_IsAvailable(t *testing.T) {
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
apiKey string
|
||||||
|
expected bool
|
||||||
|
}{
|
||||||
|
{"with key", "test-key", true},
|
||||||
|
{"empty key", "", false},
|
||||||
|
}
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
tr := NewOpenAICompatTranscriber(tt.apiKey, "https://example.com", "model")
|
||||||
|
if got := tr.IsAvailable(); got != tt.expected {
|
||||||
|
t.Errorf("IsAvailable() = %v, want %v", got, tt.expected)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOpenAICompatTranscriber_ImplementsInterface(t *testing.T) {
|
||||||
|
var _ Transcriber = (*OpenAICompatTranscriber)(nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOpenAICompatTranscriber_Transcribe(t *testing.T) {
|
||||||
|
// Create a mock HTTP server
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
// Verify it hits the right endpoint
|
||||||
|
if r.URL.Path != "/audio/transcriptions" {
|
||||||
|
t.Errorf("expected path /audio/transcriptions, got %s", r.URL.Path)
|
||||||
|
}
|
||||||
|
// Verify auth header
|
||||||
|
if r.Header.Get("Authorization") != "Bearer test-key" {
|
||||||
|
t.Errorf("unexpected auth header: %s", r.Header.Get("Authorization"))
|
||||||
|
}
|
||||||
|
// Verify it's multipart
|
||||||
|
if err := r.ParseMultipartForm(10 << 20); err != nil {
|
||||||
|
t.Errorf("failed to parse multipart form: %v", err)
|
||||||
|
}
|
||||||
|
// Verify model field
|
||||||
|
if model := r.FormValue("model"); model != "whisper-1" {
|
||||||
|
t.Errorf("expected model 'whisper-1', got %q", model)
|
||||||
|
}
|
||||||
|
// Return mock response
|
||||||
|
resp := TranscriptionResponse{
|
||||||
|
Text: "Hello world",
|
||||||
|
Language: "en",
|
||||||
|
Duration: 1.5,
|
||||||
|
}
|
||||||
|
w.Header().Set("Content-Type", "application/json")
|
||||||
|
json.NewEncoder(w).Encode(resp)
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
tr := NewOpenAICompatTranscriber("test-key", server.URL, "whisper-1")
|
||||||
|
|
||||||
|
// Create a temp audio file
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
audioFile := filepath.Join(tmpDir, "test.ogg")
|
||||||
|
if err := os.WriteFile(audioFile, []byte("fake audio data"), 0644); err != nil {
|
||||||
|
t.Fatalf("failed to create temp file: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := tr.Transcribe(context.Background(), audioFile)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Transcribe() error: %v", err)
|
||||||
|
}
|
||||||
|
if result.Text != "Hello world" {
|
||||||
|
t.Errorf("expected text 'Hello world', got %q", result.Text)
|
||||||
|
}
|
||||||
|
if result.Language != "en" {
|
||||||
|
t.Errorf("expected language 'en', got %q", result.Language)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestOpenAICompatTranscriber_TranscribeError(t *testing.T) {
|
||||||
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
w.WriteHeader(http.StatusUnauthorized)
|
||||||
|
w.Write([]byte(`{"error": "invalid api key"}`))
|
||||||
|
}))
|
||||||
|
defer server.Close()
|
||||||
|
|
||||||
|
tr := NewOpenAICompatTranscriber("bad-key", server.URL, "whisper-1")
|
||||||
|
|
||||||
|
tmpDir := t.TempDir()
|
||||||
|
audioFile := filepath.Join(tmpDir, "test.ogg")
|
||||||
|
os.WriteFile(audioFile, []byte("fake audio data"), 0644)
|
||||||
|
|
||||||
|
_, err := tr.Transcribe(context.Background(), audioFile)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("expected error, got nil")
|
||||||
|
}
|
||||||
|
}
|
||||||
Loading…
Add table
Reference in a new issue