80 lines
2.7 KiB
Go
80 lines
2.7 KiB
Go
package ai
|
|
|
|
import (
|
|
"git.ipao.vip/rogee/go-sip/internal/configread"
|
|
"testing"
|
|
)
|
|
|
|
func protocolTask(t *testing.T, protocol string) (configread.Task, map[string]configread.Provider) {
|
|
t.Helper()
|
|
task, providers := currentFixture(t, "full_ai")
|
|
task = changeCurrentAgent(t, task, func(a map[string]any) {
|
|
asr := a["asr"].(map[string]any)
|
|
delete(asr, "interim")
|
|
asr["model"] = "unlisted-asr-model-2029"
|
|
asr["language"] = "en-US"
|
|
asr["input"].(map[string]any)["sample_rate_hz"] = 16000
|
|
llm := a["llm"].(map[string]any)
|
|
llm["model"] = "unlisted-chat-model-2029"
|
|
tts := a["tts"].(map[string]any)
|
|
tts["protocol"] = protocol
|
|
tts["model"] = "unlisted-tts-model-2029"
|
|
tts["voice"] = "unlisted-voice"
|
|
tts["language_type"] = "English"
|
|
tts["speed"] = 1.0
|
|
if protocol == "dashscope_task_websocket" {
|
|
tts["speed"] = 1.75
|
|
}
|
|
})
|
|
asr := providers["asr-example"]
|
|
asr.Code = "ali_bailian"
|
|
asr.WSEndpoint = "wss://speech.example.invalid/inference"
|
|
providers[asr.ProviderRef] = asr
|
|
tts := providers["tts-example"]
|
|
tts.WSEndpoint = asr.WSEndpoint
|
|
providers[tts.ProviderRef] = tts
|
|
return task, providers
|
|
}
|
|
|
|
func TestModelsAndVoicesAreDataWithinExplicitTTSProtocol(t *testing.T) {
|
|
for _, protocol := range []string{"dashscope_tts_http", "dashscope_task_websocket"} {
|
|
t.Run(protocol, func(t *testing.T) {
|
|
task, providers := protocolTask(t, protocol)
|
|
bound, err := Bind(task, providers)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if bound.ASR.Model != "unlisted-asr-model-2029" || bound.ASR.Request.SampleRate != 16000 || bound.ASR.Language != "en-US" || bound.LLM.Model != "unlisted-chat-model-2029" || bound.TTS.Model != "unlisted-tts-model-2029" || bound.TTS.Voice != "unlisted-voice" || bound.TTS.LanguageType != "English" {
|
|
t.Fatal("task-selected model/voice/language changed")
|
|
}
|
|
expected := providers["tts-example"].Endpoint
|
|
if protocol == "dashscope_task_websocket" {
|
|
expected = providers["tts-example"].WSEndpoint
|
|
}
|
|
if bound.TTS.Provider.Endpoint != expected {
|
|
t.Fatal("connection selected from model name instead of explicit protocol")
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestMissingOrUnsupportedProtocolNeverFallsBack(t *testing.T) {
|
|
for _, protocol := range []string{"", "unknown-protocol"} {
|
|
t.Run(protocol, func(t *testing.T) {
|
|
task, providers := protocolTask(t, "dashscope_task_websocket")
|
|
task = changeCurrentAgent(t, task, func(a map[string]any) {
|
|
tts := a["tts"].(map[string]any)
|
|
if protocol == "" {
|
|
delete(tts, "protocol")
|
|
} else {
|
|
tts["protocol"] = protocol
|
|
}
|
|
tts["model"] = "cosyvoice-v3-flash"
|
|
})
|
|
if _, err := Bind(task, providers); err == nil {
|
|
t.Fatal("protocol guessed from a familiar model")
|
|
}
|
|
})
|
|
}
|
|
}
|