213 lines
7.7 KiB
Go
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())
|
|
}
|
|
}
|