158 lines
5.9 KiB
Go
158 lines
5.9 KiB
Go
package ai
|
|
|
|
import (
|
|
"encoding/json"
|
|
"os"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
|
|
"git.ipao.vip/rogee/go-sip/internal/configread"
|
|
doubaospeech "github.com/GizClaw/doubao-speech-go"
|
|
)
|
|
|
|
func currentFixture(t *testing.T, mode string) (configread.CurrentTask, map[string]configread.CurrentProvider) {
|
|
t.Helper()
|
|
name := "config-read-task-full.json"
|
|
if mode == "asr_only" {
|
|
name = "config-read-task-asr.json"
|
|
}
|
|
raw, err := os.ReadFile("../../contracts/local/examples/" + name)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
var task configread.CurrentTask
|
|
if err := json.Unmarshal(raw, &task); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
raw, err = os.ReadFile("../../contracts/local/examples/config-read-providers.json")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
var list struct {
|
|
Providers []configread.CurrentProvider `json:"providers"`
|
|
}
|
|
if err := json.Unmarshal(raw, &list); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
providers := make(map[string]configread.CurrentProvider, len(list.Providers))
|
|
for _, p := range list.Providers {
|
|
providers[p.ProviderRef] = p
|
|
}
|
|
return task, providers
|
|
}
|
|
|
|
func changeCurrentAgent(t *testing.T, task configread.CurrentTask, change func(map[string]any)) configread.CurrentTask {
|
|
t.Helper()
|
|
var body map[string]any
|
|
if err := json.Unmarshal(task.Raw, &body); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
agent := body["agent"].(map[string]any)
|
|
change(agent)
|
|
raw, err := json.Marshal(body)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
var updated configread.CurrentTask
|
|
if err := json.Unmarshal(raw, &updated); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
return updated
|
|
}
|
|
|
|
func TestBindCurrentFullAIUsesApprovedSDKFields(t *testing.T) {
|
|
task, providers := currentFixture(t, "full_ai")
|
|
bound, err := BindCurrent(task, providers)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if bound.Mode != "full_ai" || bound.LLM == nil || bound.TTS == nil {
|
|
t.Fatalf("full mode needs all three providers: mode=%q LLM=%t TTS=%t", bound.Mode, bound.LLM != nil, bound.TTS != nil)
|
|
}
|
|
if bound.ASR.Provider.Credential != providers["asr-example"].Credential || bound.ASR.Request.Format != doubaospeech.FormatPCMS16LE || bound.ASR.Request.SampleRate != 16000 || bound.ASR.Request.Channel != 1 || bound.ASR.Request.Bits != 16 || bound.ASR.Request.Language != doubaospeech.LanguageZhCN || bound.ASR.Request.ResultType != "full" || bound.ASR.Timeout != 5*time.Second {
|
|
t.Fatal("ASR approved input/interim/credential/timeout not bound to SDK request")
|
|
}
|
|
if bound.LLM.Provider.Credential != providers["llm-example"].Credential || bound.LLM.Model != "example-chat" || bound.LLM.Temperature == nil || *bound.LLM.Temperature != 0 || bound.LLM.MaxTokens == nil || *bound.LLM.MaxTokens != 256 || bound.LLM.Timeout != 5*time.Second {
|
|
t.Fatal("LLM model/explicit zero/limit/credential/timeout not bound")
|
|
}
|
|
if bound.TTS.Provider.Credential != providers["tts-example"].Credential || bound.TTS.Request.ResourceID != "example-tts" || bound.TTS.Request.Speaker != "example-neutral" || bound.TTS.Request.Format != doubaospeech.FormatPCMS16LE || bound.TTS.Request.SampleRate != 16000 || bound.TTS.Request.SpeechRate != 0 || bound.TTS.Timeout != 5*time.Second {
|
|
t.Fatal("TTS model/voice/speed/format/credential/timeout not bound to SDK request")
|
|
}
|
|
if len(bound.HangupKeywords) != 1 || bound.HangupKeywords[0] != "不用了" || bound.Prompt != "Example only" || bound.Opening != "Example greeting" {
|
|
t.Fatal("immutable prompt and keyword behavior not bound")
|
|
}
|
|
}
|
|
|
|
func TestBindCurrentASROnlyDoesNotBindOtherProviders(t *testing.T) {
|
|
task, providers := currentFixture(t, "asr_only")
|
|
delete(providers, "llm-example")
|
|
delete(providers, "tts-example")
|
|
bound, err := BindCurrent(task, providers)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if bound.Mode != "asr_only" || bound.LLM != nil || bound.TTS != nil || bound.ASR.Request.ResultType != "single" {
|
|
t.Fatal("ASR-only mode must not inherit LLM/TTS configuration or interim results")
|
|
}
|
|
}
|
|
|
|
func TestBindCurrentRejectsSDKUnsupportedTTSWithoutChangingSchema(t *testing.T) {
|
|
for _, tc := range []struct {
|
|
name string
|
|
edit func(map[string]any)
|
|
}{
|
|
{"pcma", func(tts map[string]any) { tts["format"].(map[string]any)["encoding"] = "pcma" }},
|
|
{"speed-below", func(tts map[string]any) { tts["speed"] = 0.25 }},
|
|
{"speed-above", func(tts map[string]any) { tts["speed"] = 3.0 }},
|
|
{"speed-unrepresentable", func(tts map[string]any) { tts["speed"] = 1.005 }},
|
|
{"sample-rate", func(tts map[string]any) { tts["format"].(map[string]any)["sample_rate_hz"] = 12345 }},
|
|
} {
|
|
t.Run(tc.name, func(t *testing.T) {
|
|
task, providers := currentFixture(t, "full_ai")
|
|
task = changeCurrentAgent(t, task, func(agent map[string]any) { tc.edit(agent["tts"].(map[string]any)) })
|
|
_, err := BindCurrent(task, providers)
|
|
if err == nil || !strings.Contains(err.Error(), "TTS") || strings.Contains(err.Error(), providers["tts-example"].Credential) {
|
|
t.Fatalf("expected explicit non-secret TTS capability error, got %v", err)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestBindCurrentRejectsUnauthorizedProviderBeforeCall(t *testing.T) {
|
|
for _, tc := range []struct {
|
|
name string
|
|
mutate func(map[string]configread.CurrentProvider)
|
|
}{
|
|
{"missing", func(ps map[string]configread.CurrentProvider) { delete(ps, "asr-example") }},
|
|
{"disabled", func(ps map[string]configread.CurrentProvider) {
|
|
p := ps["asr-example"]
|
|
p.Enabled = false
|
|
ps[p.ProviderRef] = p
|
|
}},
|
|
{"wrong-role", func(ps map[string]configread.CurrentProvider) {
|
|
p := ps["asr-example"]
|
|
p.Role = "tts"
|
|
ps[p.ProviderRef] = p
|
|
}},
|
|
{"wrong-adapter", func(ps map[string]configread.CurrentProvider) {
|
|
p := ps["asr-example"]
|
|
p.Adapter = "unknown"
|
|
ps[p.ProviderRef] = p
|
|
}},
|
|
{"missing-credential", func(ps map[string]configread.CurrentProvider) {
|
|
p := ps["asr-example"]
|
|
p.Credential = ""
|
|
ps[p.ProviderRef] = p
|
|
}},
|
|
} {
|
|
t.Run(tc.name, func(t *testing.T) {
|
|
task, providers := currentFixture(t, "full_ai")
|
|
tc.mutate(providers)
|
|
if _, err := BindCurrent(task, providers); err == nil {
|
|
t.Fatal("unavailable AI provider cannot authorize execution")
|
|
}
|
|
})
|
|
}
|
|
}
|