111 lines
3.7 KiB
Go
111 lines
3.7 KiB
Go
package ai
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"encoding/json"
|
|
"github.com/coder/websocket"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"strings"
|
|
"testing"
|
|
"time"
|
|
)
|
|
|
|
func TestTaskTTSProtocolForwardsOpaqueModelVoiceAndParams(t *testing.T) {
|
|
task, ps := protocolTask(t, TTSProtocolDashScopeTask)
|
|
b, err := Bind(task, ps)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
b.TTS.Params = map[string]json.RawMessage{"opaque": json.RawMessage(`"vendor-only"`), "rate": json.RawMessage("9.5"), "unknown_flag": json.RawMessage("false"), "nullable": json.RawMessage("null"), "large_number": json.RawMessage("9007199254740993")}
|
|
expected := []byte{1, 0, 2, 0}
|
|
s := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
c, e := websocket.Accept(w, r, nil)
|
|
if e != nil {
|
|
t.Error(e)
|
|
return
|
|
}
|
|
defer c.CloseNow()
|
|
ctx, cancel := context.WithTimeout(r.Context(), time.Second)
|
|
defer cancel()
|
|
kind, raw, err := c.Read(ctx)
|
|
if err != nil || kind != websocket.MessageText {
|
|
t.Error("missing TTS start payload")
|
|
return
|
|
}
|
|
if !bytes.Contains(raw, []byte(`"large_number":9007199254740993`)) {
|
|
t.Error("large TTS integer changed")
|
|
}
|
|
var cmd mockCommand
|
|
if err := json.Unmarshal(raw, &cmd); err != nil {
|
|
t.Error(err)
|
|
return
|
|
}
|
|
p := cmd.Payload["parameters"].(map[string]any)
|
|
if cmd.Payload["model"] != "unlisted-tts-model-2029" || p["voice"] != "unlisted-voice" || p["opaque"] != "vendor-only" || p["rate"] != 9.5 || p["unknown_flag"] != false || p["nullable"] != nil || len(p) != 6 {
|
|
t.Error("opaque parameters changed")
|
|
}
|
|
mockEvent(c, ctx, cmd.Header.TaskID, "task-started", nil)
|
|
mockRead(t, c, ctx)
|
|
mockRead(t, c, ctx)
|
|
c.Write(ctx, websocket.MessageBinary, expected)
|
|
mockEvent(c, ctx, cmd.Header.TaskID, "task-finished", nil)
|
|
}))
|
|
defer s.Close()
|
|
b.TTS.Provider.WSEndpoint = "ws" + strings.TrimPrefix(s.URL, "http")
|
|
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
|
defer cancel()
|
|
audio, e := b.Synthesize(ctx, "Protocol fixture.")
|
|
if e != nil || !bytes.Equal(audio, expected) {
|
|
t.Fatal("TTS protocol failure", e)
|
|
}
|
|
}
|
|
|
|
func TestTaskASRForwardsParamsWithoutRewritingAudio(t *testing.T) {
|
|
task, ps := protocolTask(t, TTSProtocolDashScopeTask)
|
|
b, e := Bind(task, ps)
|
|
if e != nil {
|
|
t.Fatal(e)
|
|
}
|
|
b.ASR.Params = map[string]json.RawMessage{"sample_rate": json.RawMessage("24000"), "opaque": json.RawMessage(`"unchanged"`), "nullable": json.RawMessage("null")}
|
|
pcm := []byte{1, 0, 2, 0}
|
|
captured := make(chan []byte, 1)
|
|
s := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
c, e := websocket.Accept(w, r, nil)
|
|
if e != nil {
|
|
t.Error(e)
|
|
return
|
|
}
|
|
defer c.CloseNow()
|
|
ctx, cancel := context.WithTimeout(r.Context(), time.Second)
|
|
defer cancel()
|
|
cmd := mockRead(t, c, ctx)
|
|
p := cmd.Payload["parameters"].(map[string]any)
|
|
if len(p) != 3 || p["sample_rate"] != float64(24000) || p["opaque"] != "unchanged" || p["nullable"] != nil {
|
|
t.Error("ASR params rewritten")
|
|
}
|
|
mockEvent(c, ctx, cmd.Header.TaskID, "task-started", nil)
|
|
_, body, e := c.Read(ctx)
|
|
if e != nil {
|
|
t.Error(e)
|
|
return
|
|
}
|
|
captured <- body
|
|
mockRead(t, c, ctx)
|
|
mockEvent(c, ctx, cmd.Header.TaskID, "result-generated", map[string]any{"output": map[string]any{"sentence": map[string]any{"text": "Final.", "sentence_end": true, "begin_time": 0, "end_time": 1}}})
|
|
mockEvent(c, ctx, cmd.Header.TaskID, "task-finished", nil)
|
|
}))
|
|
defer s.Close()
|
|
b.ASR.Provider.Endpoint = "ws" + strings.TrimPrefix(s.URL, "http")
|
|
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
|
defer cancel()
|
|
text, e := b.Recognize(ctx, pcm)
|
|
if e != nil || text != "Final." {
|
|
t.Fatal(e)
|
|
}
|
|
if !bytes.Equal(<-captured, pcm) {
|
|
t.Fatal("input media rewritten from opaque params")
|
|
}
|
|
}
|