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

213 lines
7.7 KiB
Go

package ai
import (
"bytes"
"context"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"strings"
"sync/atomic"
"testing"
"time"
"git.ipao.vip/rogee/go-sip/internal/media"
doubaospeech "github.com/GizClaw/doubao-speech-go"
"github.com/coder/websocket"
)
func TestTaskTTSProtocolForwardsOpaqueModelVoiceSpeedAndLanguage(t *testing.T) {
for _, language := range []string{"English", "fr-FR", "Auto"} {
t.Run(language, func(t *testing.T) {
task, providers := protocolTask(t, TTSProtocolDashScopeTask)
bound, err := Bind(task, providers)
if err != nil {
t.Fatal(err)
}
bound.TTS.LanguageType = language
var calls atomic.Int32
expected := []byte{1, 0, 2, 0}
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
calls.Add(1)
c, err := websocket.Accept(w, r, nil)
if err != nil {
t.Error(err)
return
}
defer c.CloseNow()
ctx, cancel := context.WithTimeout(r.Context(), 2*time.Second)
defer cancel()
cmd := mockRead(t, c, ctx)
p := cmd.Payload["parameters"].(map[string]any)
if cmd.Payload["model"] != "unlisted-tts-model-2029" || p["voice"] != "unlisted-voice" || p["rate"] != 1.75 || p["sample_rate"] != float64(16000) || p["format"] != "pcm" {
t.Errorf("approved parameters changed: %+v", cmd.Payload)
}
if language == "Auto" {
if _, ok := p["language_hints"]; ok {
t.Error("explicit Auto must not insert a language hint")
}
} else {
expectedHint := "en"
if language == "fr-FR" {
expectedHint = "fr"
}
hints, ok := p["language_hints"].([]any)
if !ok || len(hints) != 1 || hints[0] != expectedHint {
t.Error("language was replaced with a model-specific default")
}
}
mockEvent(c, ctx, cmd.Header.TaskID, "task-started", nil)
text := mockRead(t, c, ctx)
finish := mockRead(t, c, ctx)
if text.Payload["input"].(map[string]any)["text"] != "Protocol fixture." || finish.Header.Action != "finish-task" {
t.Error("TTS text/lifecycle changed")
}
c.Write(ctx, websocket.MessageBinary, expected)
mockEvent(c, ctx, cmd.Header.TaskID, "task-finished", nil)
}))
defer server.Close()
bound.TTS.Provider.WSEndpoint = "ws" + strings.TrimPrefix(server.URL, "http")
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
audio, err := bound.Synthesize(ctx, "Protocol fixture.")
if err != nil || !bytes.Equal(audio, expected) || calls.Load() != 1 {
t.Fatalf("task TTS: audio=%d err=%v requests=%d", len(audio), err, calls.Load())
}
})
}
}
func TestTaskASRProtocolForwardsOpaqueModelAndSelectedSampleRate(t *testing.T) {
for _, rate := range []int{8000, 16000} {
t.Run(fmt.Sprint(rate), func(t *testing.T) {
task, providers := protocolTask(t, TTSProtocolDashScopeTask)
bound, err := Bind(task, providers)
if err != nil {
t.Fatal(err)
}
bound.ASR.Request.SampleRate = doubaospeech.SampleRate(rate)
pcm := make([]byte, 32000)
for i := range pcm {
pcm[i] = byte(i % 251)
}
var calls atomic.Int32
captured := make(chan []byte, 1)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
var received []byte
defer func() { captured <- received }()
calls.Add(1)
c, err := websocket.Accept(w, r, nil)
if err != nil {
t.Error(err)
return
}
defer c.CloseNow()
ctx, cancel := context.WithTimeout(r.Context(), 2*time.Second)
defer cancel()
cmd := mockRead(t, c, ctx)
p := cmd.Payload["parameters"].(map[string]any)
if cmd.Payload["model"] != "unlisted-asr-model-2029" || p["sample_rate"] != float64(rate) || p["language_hints"].([]any)[0] != "en" {
t.Error("ASR task model/rate/language changed")
}
mockEvent(c, ctx, cmd.Header.TaskID, "task-started", nil)
for {
kind, raw, err := c.Read(ctx)
if err != nil {
t.Error(err)
return
}
if kind == websocket.MessageBinary {
received = append(received, raw...)
continue
}
var finish mockCommand
json.Unmarshal(raw, &finish)
if finish.Header.Action != "finish-task" {
t.Error("ASR not finished explicitly")
}
break
}
mockEvent(c, ctx, cmd.Header.TaskID, "result-generated", map[string]any{"output": map[string]any{"sentence": map[string]any{"text": "Final user text.", "sentence_end": true, "begin_time": 0, "end_time": 1000}}})
mockEvent(c, ctx, cmd.Header.TaskID, "task-finished", nil)
}))
defer server.Close()
bound.ASR.Provider.Endpoint = "ws" + strings.TrimPrefix(server.URL, "http")
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
text, err := bound.Recognize(ctx, pcm)
received := <-captured
if err != nil || text != "Final user text." || len(received) != rate*2 || calls.Load() != 1 {
t.Fatalf("ASR: text=%q err=%v sent=%d calls=%d", text, err, len(received), calls.Load())
}
if rate == 16000 && !bytes.Equal(received, pcm) {
t.Fatal("16k PCM was transformed despite the selected 16k protocol rate")
}
})
}
}
func TestHTTPProtocolForwardsOpaqueModelVoiceAndLanguage(t *testing.T) {
wav, _, err := media.EncodeMonoWAV([]byte{0, 0, 2, 0, 3, 0, 4, 0}, 1024)
if err != nil {
t.Fatal(err)
}
var generation, downloads atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path == "/audio" {
downloads.Add(1)
w.Header().Set("Content-Type", "audio/wav")
w.Write(wav)
return
}
generation.Add(1)
var req struct {
Model string `json:"model"`
Input struct {
Text, Voice string
LanguageType string `json:"language_type"`
} `json:"input"`
}
if json.NewDecoder(r.Body).Decode(&req) != nil || req.Model != "unlisted-tts-model-2029" || req.Input.Voice != "unlisted-voice" || req.Input.LanguageType != "English" || req.Input.Text != "Protocol fixture." {
t.Error("HTTP task parameters changed")
}
w.Header().Set("Content-Type", "application/json")
fmt.Fprintf(w, `{"output":{"audio":{"url":%q}}}`, "http://"+r.Host+"/audio")
}))
defer server.Close()
task, providers := protocolTask(t, TTSProtocolDashScopeHTTP)
bound, err := Bind(task, providers)
if err != nil {
t.Fatal(err)
}
bound.TTS.Provider.Endpoint = server.URL + "/api/v1/services/aigc/multimodal-generation/generation"
ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
defer cancel()
audio, err := bound.Synthesize(ctx, "Protocol fixture.")
if err != nil || len(audio) == 0 || generation.Load() != 1 || downloads.Load() != 1 {
t.Fatalf("HTTP TTS: audio=%d err=%v requests=%d downloads=%d", len(audio), err, generation.Load(), downloads.Load())
}
}
func TestExplicitHTTPProtocolNeverGuessesWebSocketFromFamiliarModel(t *testing.T) {
var httpCalls, wsCalls atomic.Int32
ws := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { wsCalls.Add(1); w.WriteHeader(500) }))
defer ws.Close()
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { httpCalls.Add(1); w.WriteHeader(400) }))
defer server.Close()
task, providers := protocolTask(t, TTSProtocolDashScopeHTTP)
bound, err := Bind(task, providers)
if err != nil {
t.Fatal(err)
}
bound.TTS.Model = "cosyvoice-v3-flash"
bound.TTS.Provider.Endpoint = server.URL + "/api/v1/services/aigc/multimodal-generation/generation"
bound.TTS.Provider.WSEndpoint = "ws" + strings.TrimPrefix(ws.URL, "http")
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
audio, err := bound.Synthesize(ctx, "No fallback.")
if err == nil || len(audio) != 0 || httpCalls.Load() != 1 || wsCalls.Load() != 0 {
t.Fatalf("model/protocol fallback: audio=%d err=%v HTTP=%d WS=%d", len(audio), err, httpCalls.Load(), wsCalls.Load())
}
}