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

158 lines
5.8 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.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/local/examples/" + name)
if err != nil {
t.Fatal(err)
}
var task configread.Task
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.Provider `json:"providers"`
}
if err := json.Unmarshal(raw, &list); err != nil {
t.Fatal(err)
}
providers := make(map[string]configread.Provider, len(list.Providers))
for _, p := range list.Providers {
providers[p.ProviderRef] = p
}
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)
}
agent := body["agent"].(map[string]any)
change(agent)
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, providers := currentFixture(t, "full_ai")
bound, err := Bind(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 := Bind(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 := Bind(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.Provider)
}{
{"missing", func(ps map[string]configread.Provider) { delete(ps, "asr-example") }},
{"disabled", func(ps map[string]configread.Provider) {
p := ps["asr-example"]
p.Enabled = false
ps[p.ProviderRef] = p
}},
{"wrong-role", func(ps map[string]configread.Provider) {
p := ps["asr-example"]
p.Role = "tts"
ps[p.ProviderRef] = p
}},
{"wrong-adapter", func(ps map[string]configread.Provider) {
p := ps["asr-example"]
p.Adapter = "unknown"
ps[p.ProviderRef] = p
}},
{"missing-credential", func(ps map[string]configread.Provider) {
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 := Bind(task, providers); err == nil {
t.Fatal("unavailable AI provider cannot authorize execution")
}
})
}
}