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

103 lines
3.6 KiB
Go

package ai
import (
"encoding/json"
"os"
"strings"
"testing"
"git.ipao.vip/rogee/go-sip/internal/configread"
)
func currentFixture(t *testing.T, mode string) (configread.Task, map[string]configread.Provider) {
t.Helper()
name := "config-read-task-full.json"
if mode == "asr_only" {
name = "config-read-task-asr.json"
}
raw, err := os.ReadFile("../../contracts/schema/examples/" + name)
if err != nil {
t.Fatal(err)
}
var task configread.Task
if err := json.Unmarshal(raw, &task); err != nil {
t.Fatal(err)
}
providers := map[string]configread.Provider{}
for _, code := range []string{"volcengine", "openai_compatible", "ali_bailian"} {
providers[code] = configread.Provider{ProviderRef: "catalog-" + code, Code: code, Name: "Example only", Endpoint: "https://example.invalid", WSEndpoint: "wss://example.invalid", Credential: "example-only-not-a-real-secret"}
}
task = changeCurrentAgent(t, task, func(a map[string]any) {
a["asr"] = map[string]any{"provider_ref": "volcengine", "model": "example-asr", "params": map[string]any{"result_type": "full", "enable_nonstream": false, "enable_itn": false}}
if mode == "full_ai" {
a["llm"] = map[string]any{"provider_ref": "openai_compatible", "model": "example-chat", "params": map[string]any{"temperature": 0, "max_tokens": 256}}
a["tts"] = map[string]any{"provider_ref": "ali_bailian", "model": "qwen3-tts-flash", "voice": "Cherry", "params": map[string]any{"format": "pcm", "sample_rate": 16000, "rate": 1, "language": "Chinese"}}
a["prompt"].(map[string]any)["text"] = "Example only"
}
})
return task, providers
}
func changeCurrentAgent(t *testing.T, task configread.Task, change func(map[string]any)) configread.Task {
t.Helper()
var body map[string]any
if err := json.Unmarshal(task.Raw, &body); err != nil {
t.Fatal(err)
}
change(body["agent"].(map[string]any))
raw, err := json.Marshal(body)
if err != nil {
t.Fatal(err)
}
var updated configread.Task
if err := json.Unmarshal(raw, &updated); err != nil {
t.Fatal(err)
}
return updated
}
func TestBindCurrentFullAIUsesApprovedSDKFields(t *testing.T) {
task, ps := currentFixture(t, "full_ai")
b, err := Bind(task, ps)
if err != nil {
t.Fatal(err)
}
if b.LLM == nil || b.TTS == nil || b.ASR.Provider.Code != "volcengine" || b.LLM.Model != "example-chat" || b.TTS.Voice != "Cherry" || b.TTS.Protocol != TTSProtocolDashScopeTask || string(b.LLM.Params["temperature"]) != "0" {
t.Fatal("approved settings changed")
}
if len(b.HangupKeywords) != 1 || b.HangupKeywords[0].ClosingRemark != "好的,祝您生活愉快。" || b.Prompt != "Example only" || b.Opening != "Example greeting" {
t.Fatal("conversation settings changed")
}
}
func TestBindCurrentASROnlyDoesNotBindOtherProviders(t *testing.T) {
task, ps := currentFixture(t, "asr_only")
delete(ps, "openai_compatible")
delete(ps, "ali_bailian")
b, err := Bind(task, ps)
if err != nil || b.LLM != nil || b.TTS != nil {
t.Fatal("ASR-only inherited another provider", err)
}
}
func TestBindCurrentRejectsUnauthorizedProviderBeforeCall(t *testing.T) {
for _, kind := range []string{"missing", "code", "endpoint", "credential"} {
t.Run(kind, func(t *testing.T) {
task, ps := currentFixture(t, "full_ai")
p := ps["volcengine"]
switch kind {
case "missing":
delete(ps, "volcengine")
case "code":
p.Code = "unknown"
ps["volcengine"] = p
case "endpoint":
p.WSEndpoint = ""
ps["volcengine"] = p
case "credential":
p.Credential = ""
ps["volcengine"] = p
}
if _, err := Bind(task, ps); err == nil || strings.Contains(err.Error(), "example-only-not-a-real-secret") {
t.Fatal("missing/redacted provider validation", err)
}
})
}
}