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

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")
}
}