103 lines
3.6 KiB
Go
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)
|
|
}
|
|
})
|
|
}
|
|
}
|