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