Files
go-sip/internal/ai/protocol_models_test.go
T

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")
}
})
}
}