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

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