107 lines
3.7 KiB
Go
107 lines
3.7 KiB
Go
package ai
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"encoding/json"
|
|
"fmt"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"strings"
|
|
"sync/atomic"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/coder/websocket"
|
|
)
|
|
|
|
func TestProtocolModelsCompleteTurnWithoutPerModelCode(t *testing.T) {
|
|
task, providers := protocolTask(t, TTSProtocolDashScopeTask)
|
|
task = changeCurrentAgent(t, task, func(a map[string]any) {
|
|
a["conversation"].(map[string]any)["opening"] = "Hello."
|
|
a["conversation"].(map[string]any)["hangup_keywords"] = []any{}
|
|
})
|
|
bound, err := Bind(task, providers)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
var asrCalls, ttsCalls, llmCalls atomic.Int32
|
|
expected := []byte{1, 0, 2, 0}
|
|
speech := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
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)
|
|
mockEvent(c, ctx, cmd.Header.TaskID, "task-started", nil)
|
|
switch cmd.Payload["task"] {
|
|
case "asr":
|
|
asrCalls.Add(1)
|
|
if cmd.Payload["model"] != "unlisted-asr-model-2029" {
|
|
t.Error("ASR model replaced")
|
|
}
|
|
for {
|
|
kind, _, err := c.Read(ctx)
|
|
if err != nil {
|
|
t.Error(err)
|
|
return
|
|
}
|
|
if kind == websocket.MessageText {
|
|
break
|
|
}
|
|
}
|
|
mockEvent(c, ctx, cmd.Header.TaskID, "result-generated", map[string]any{"output": map[string]any{"sentence": map[string]any{"text": "Tell me more.", "sentence_end": true, "begin_time": 0, "end_time": 100}}})
|
|
case "tts":
|
|
ttsCalls.Add(1)
|
|
if cmd.Payload["model"] != "unlisted-tts-model-2029" || cmd.Payload["parameters"].(map[string]any)["voice"] != "unlisted-voice" {
|
|
t.Error("TTS model/voice replaced")
|
|
}
|
|
mockRead(t, c, ctx)
|
|
finish := mockRead(t, c, ctx)
|
|
if finish.Header.Action != "finish-task" {
|
|
t.Error("TTS lifecycle changed")
|
|
}
|
|
c.Write(ctx, websocket.MessageBinary, expected)
|
|
default:
|
|
t.Error("unknown role")
|
|
return
|
|
}
|
|
mockEvent(c, ctx, cmd.Header.TaskID, "task-finished", nil)
|
|
}))
|
|
defer speech.Close()
|
|
llm := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
llmCalls.Add(1)
|
|
var request map[string]any
|
|
if r.Method != "POST" || r.URL.Path != "/v1/chat/completions" || json.NewDecoder(r.Body).Decode(&request) != nil || request["model"] != "unlisted-chat-model-2029" {
|
|
t.Error("LLM model/protocol changed")
|
|
}
|
|
w.Header().Set("Content-Type", "application/json")
|
|
fmt.Fprint(w, `{"id":"mock","object":"chat.completion","created":0,"model":"unlisted-chat-model-2029","choices":[{"index":0,"message":{"role":"assistant","content":"Generic response."},"finish_reason":"stop"}]}`)
|
|
}))
|
|
defer llm.Close()
|
|
bound.ASR.Provider.Endpoint = "ws" + strings.TrimPrefix(speech.URL, "http")
|
|
bound.TTS.Provider.WSEndpoint = bound.ASR.Provider.Endpoint
|
|
bound.LLM.Provider.Endpoint = llm.URL + "/v1"
|
|
call, err := NewCall(bound, func(context.Context) error { t.Error("ordinary turn attempted hangup"); return nil })
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
|
defer cancel()
|
|
opening, err := call.Open(ctx)
|
|
if err != nil || !bytes.Equal(opening, expected) {
|
|
t.Fatalf("opening: %v", err)
|
|
}
|
|
result, err := call.RunTurn(ctx, make([]byte, 3200))
|
|
if err != nil || result.Transcript != "Tell me more." || result.Reply != "Generic response." || !bytes.Equal(result.AudioPCM16, expected) || result.EndedByKeyword {
|
|
t.Fatalf("generic model turn: %+v %v", result, err)
|
|
}
|
|
if asrCalls.Load() != 1 || llmCalls.Load() != 1 || ttsCalls.Load() != 2 {
|
|
t.Fatalf("unexpected/retried requests: ASR=%d LLM=%d TTS=%d", asrCalls.Load(), llmCalls.Load(), ttsCalls.Load())
|
|
}
|
|
}
|