452 lines
14 KiB
Go
452 lines
14 KiB
Go
package main
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"net"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"os"
|
|
"path/filepath"
|
|
"strconv"
|
|
"strings"
|
|
"sync/atomic"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/emiago/sipgo"
|
|
"github.com/emiago/sipgo/sip"
|
|
"github.com/gobwas/ws"
|
|
"github.com/gobwas/ws/wsutil"
|
|
)
|
|
|
|
func TestOfflineMediaAndSafety(t *testing.T) {
|
|
const testNumber = "138" + "0000" + "1234"
|
|
if err := validateNumber(testNumber); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
for _, number := range []string{"", "1380000123", "23800001234", "1380000123x"} {
|
|
if validateNumber(number) == nil {
|
|
t.Fatalf("accepted invalid number %q", number)
|
|
}
|
|
}
|
|
|
|
addr, pt, err := answerMedia([]byte("v=0\r\nc=IN IP4 192.0.2.10\r\nm=audio 40000 RTP/AVP 0 8\r\n"))
|
|
if err != nil || addr.String() != "192.0.2.10:40000" || pt != 8 {
|
|
t.Fatalf("answerMedia = %v, %d, %v", addr, pt, err)
|
|
}
|
|
payload := []byte{1, 2, 3}
|
|
packet := marshalRTP(8, 7, 160, 42, payload)
|
|
got, gotPT, ok := rtpPayload(packet)
|
|
if !ok || gotPT != 8 || !bytes.Equal(got, payload) {
|
|
t.Fatalf("rtpPayload = %v, %d, %v", got, gotPT, ok)
|
|
}
|
|
|
|
pcm := []int16{-2000, 0, 2000, 1000}
|
|
if got := bytesToPCM(pcmToBytes(pcm)); !equalPCM(got, pcm) {
|
|
t.Fatalf("PCM round trip = %v", got)
|
|
}
|
|
if got := resample(pcm, 8000, 16000); len(got) != 8 {
|
|
t.Fatalf("resampled length = %d", len(got))
|
|
}
|
|
|
|
dir := t.TempDir()
|
|
wav := filepath.Join(dir, "roundtrip.wav")
|
|
if err := writeWAV(wav, pcm, 8000); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
f, err := os.Open(wav)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
gotPCM, rate, err := readWAV(f)
|
|
f.Close()
|
|
if err != nil || rate != 8000 || !equalPCM(gotPCM, pcm) {
|
|
t.Fatalf("WAV round trip = %v, %d, %v", gotPCM, rate, err)
|
|
}
|
|
|
|
logPath := filepath.Join(dir, "sip.log")
|
|
log, err := newSignalLog(logPath, testNumber, caller)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
log.message(">>>", "INVITE sip:"+calleePrefix+testNumber+"@example.invalid From: "+caller)
|
|
log.Close()
|
|
content, err := os.ReadFile(logPath)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
text := string(content)
|
|
if strings.Contains(text, testNumber) || strings.Contains(text, caller) ||
|
|
!strings.Contains(text, "7089*******1234") || !strings.Contains(text, "BD****5882") {
|
|
t.Fatalf("unsafe redaction: %s", text)
|
|
}
|
|
}
|
|
|
|
func TestBailianTTSContract(t *testing.T) {
|
|
const key = "test-key-that-must-not-leak"
|
|
type command struct {
|
|
Raw []byte
|
|
Header struct {
|
|
TaskID string `json:"task_id"`
|
|
} `json:"header"`
|
|
}
|
|
serverErrors := make(chan error, 1)
|
|
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
|
fail := func(format string, args ...any) {
|
|
select {
|
|
case serverErrors <- fmt.Errorf(format, args...):
|
|
default:
|
|
}
|
|
}
|
|
if r.URL.Path != "/api-ws/v1/inference" {
|
|
fail("path = %s", r.URL.Path)
|
|
}
|
|
if r.Header.Get("Authorization") != "Bearer "+key {
|
|
fail("missing bearer authorization")
|
|
}
|
|
conn, _, _, err := ws.UpgradeHTTP(r, w)
|
|
if err != nil {
|
|
fail("upgrade: %v", err)
|
|
return
|
|
}
|
|
defer conn.Close()
|
|
read := func() command {
|
|
var got command
|
|
data, op, err := wsutil.ReadClientData(conn)
|
|
if err != nil {
|
|
fail("read command: %v", err)
|
|
return got
|
|
}
|
|
if op != ws.OpText {
|
|
fail("command opcode = %v", op)
|
|
}
|
|
if err := json.Unmarshal(data, &got); err != nil {
|
|
fail("decode command: %v", err)
|
|
}
|
|
got.Raw = append([]byte(nil), data...)
|
|
return got
|
|
}
|
|
writeEvent := func(event, taskID string) {
|
|
data, _ := json.Marshal(map[string]any{"header": map[string]any{"event": event, "task_id": taskID}, "payload": map[string]any{}})
|
|
if err := wsutil.WriteServerText(conn, data); err != nil {
|
|
fail("write event: %v", err)
|
|
}
|
|
}
|
|
|
|
run := read()
|
|
if !isUUID(run.Header.TaskID) {
|
|
fail("task_id is not a UUID")
|
|
}
|
|
wantRun, _ := json.Marshal(map[string]any{
|
|
"header": map[string]any{"action": "run-task", "task_id": run.Header.TaskID, "streaming": "duplex"},
|
|
"payload": map[string]any{
|
|
"task_group": "audio", "task": "tts", "function": "SpeechSynthesizer", "model": "cosyvoice-v3.5-plus",
|
|
"parameters": map[string]any{"text_type": "PlainText", "voice": "test-voice", "format": "pcm", "sample_rate": 24000},
|
|
"input": map[string]any{},
|
|
},
|
|
})
|
|
if !bytes.Equal(run.Raw, wantRun) {
|
|
fail("invalid run-task payload")
|
|
}
|
|
writeEvent("task-started", run.Header.TaskID)
|
|
|
|
continued := read()
|
|
wantContinue, _ := json.Marshal(map[string]any{
|
|
"header": map[string]any{"action": "continue-task", "task_id": run.Header.TaskID, "streaming": "duplex"},
|
|
"payload": map[string]any{"input": map[string]any{"text": "测试"}},
|
|
})
|
|
if !bytes.Equal(continued.Raw, wantContinue) {
|
|
fail("invalid continue-task payload")
|
|
}
|
|
finished := read()
|
|
wantFinish, _ := json.Marshal(map[string]any{
|
|
"header": map[string]any{"action": "finish-task", "task_id": run.Header.TaskID, "streaming": "duplex"},
|
|
"payload": map[string]any{"input": map[string]any{}},
|
|
})
|
|
if !bytes.Equal(finished.Raw, wantFinish) {
|
|
fail("invalid finish-task payload")
|
|
}
|
|
if err := wsutil.WriteServerBinary(conn, pcmToBytes([]int16{1, 2, 3})); err != nil {
|
|
fail("write audio: %v", err)
|
|
}
|
|
writeEvent("task-finished", run.Header.TaskID)
|
|
}))
|
|
defer server.Close()
|
|
|
|
endpoint := "ws" + strings.TrimPrefix(server.URL, "http") + "/api-ws/v1/inference"
|
|
pcm, rate, err := synthesize(context.Background(), endpoint, key, "test-voice", "测试")
|
|
if err != nil || rate != 24000 || !equalPCM(pcm, []int16{1, 2, 3}) || strings.Contains(errString(err), key) {
|
|
t.Fatalf("synthesize = %v, %d, %v", pcm, rate, err)
|
|
}
|
|
select {
|
|
case err := <-serverErrors:
|
|
t.Fatal(err)
|
|
default:
|
|
}
|
|
}
|
|
|
|
func TestBailianTTSRequiresVoiceAndRedactsFailureDetail(t *testing.T) {
|
|
if _, err := mediaProgram(context.Background(), config{ttsText: "测试"}); err == nil || !strings.Contains(err.Error(), "BAILIAN_TTS_VOICE") {
|
|
t.Fatalf("missing voice error = %v", err)
|
|
}
|
|
|
|
const taskID = "2bf83b9a-baeb-4fda-8d9a-000000000000"
|
|
const secret = "test-key-that-must-not-leak"
|
|
event, err := parseTTSEvent([]byte(`{"header":{"task_id":"`+taskID+`","event":"task-failed","error_code":"InvalidParameter","error_message":"raw server detail `+secret+`"},"payload":{}}`), taskID)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
got := ttsFailure(event).Error()
|
|
if got != "TTS task failed: InvalidParameter" || strings.Contains(got, secret) || strings.Contains(got, "raw server detail") {
|
|
t.Fatalf("failure error = %q", got)
|
|
}
|
|
}
|
|
|
|
func TestBailianTTSEndpointMustBeExactAndSecure(t *testing.T) {
|
|
var hits atomic.Int32
|
|
server := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) { hits.Add(1) }))
|
|
defer server.Close()
|
|
loopback := "ws" + strings.TrimPrefix(server.URL, "http")
|
|
for _, endpoint := range []string{
|
|
loopback + "/compatible-mode/v1",
|
|
loopback + "/api-ws/v1/inference/",
|
|
loopback + "/api-ws/v1/inference?path=other",
|
|
"ws://192.0.2.1/api-ws/v1/inference",
|
|
"https://example.invalid/compatible-mode/v1",
|
|
} {
|
|
_, _, err := synthesize(context.Background(), endpoint, "secret", "voice", "text")
|
|
if err == nil || strings.Contains(err.Error(), "secret") {
|
|
t.Fatalf("accepted endpoint %q: %v", endpoint, err)
|
|
}
|
|
}
|
|
if hits.Load() != 0 {
|
|
t.Fatalf("invalid endpoint received %d requests", hits.Load())
|
|
}
|
|
}
|
|
|
|
func isUUID(value string) bool {
|
|
if len(value) != 36 || value[8] != '-' || value[13] != '-' || value[18] != '-' || value[23] != '-' {
|
|
return false
|
|
}
|
|
for i, c := range value {
|
|
if i == 8 || i == 13 || i == 18 || i == 23 {
|
|
continue
|
|
}
|
|
if !strings.ContainsRune("0123456789abcdef", c) {
|
|
return false
|
|
}
|
|
}
|
|
return true
|
|
}
|
|
|
|
func TestInboundVoiceInterruptsPlayback(t *testing.T) {
|
|
conn, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4zero, Port: 0})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer conn.Close()
|
|
sender, err := net.DialUDP("udp4", nil, conn.LocalAddr().(*net.UDPAddr))
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer sender.Close()
|
|
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
interrupted := make(chan struct{})
|
|
done := make(chan []int16, 1)
|
|
go receiveRTP(ctx, conn, sender.LocalAddr().(*net.UDPAddr), 1000, 8000, interrupted, done)
|
|
frame := make([]int16, 160)
|
|
for i := range frame {
|
|
frame[i] = 4000
|
|
}
|
|
payload := encodeG711(pcmToBytes(frame), 8)
|
|
for seq := uint16(0); seq < 3; seq++ {
|
|
_, _ = sender.Write(marshalRTP(8, seq, uint32(seq)*160, 1, payload))
|
|
}
|
|
select {
|
|
case <-interrupted:
|
|
case <-time.After(time.Second):
|
|
t.Fatal("playback was not interrupted")
|
|
}
|
|
cancel()
|
|
select {
|
|
case pcm := <-done:
|
|
if len(pcm) != 480 {
|
|
t.Fatalf("recorded samples = %d", len(pcm))
|
|
}
|
|
case <-time.After(time.Second):
|
|
t.Fatal("RTP receiver did not stop")
|
|
}
|
|
}
|
|
|
|
func TestSingleCallAgainstLocalSIPAndRTP(t *testing.T) {
|
|
ackFailure := errors.New("injected ACK failure")
|
|
tests := []struct {
|
|
name string
|
|
ack func(context.Context, *sipgo.DialogClientSession) error
|
|
validSDP bool
|
|
byeStatus int
|
|
wantMethods string
|
|
wantErrors []string
|
|
}{
|
|
{name: "normal", validSDP: true, byeStatus: sip.StatusOK, wantMethods: "INVITE,ACK,BYE"},
|
|
{
|
|
name: "ACK failure still sends BYE", validSDP: true, byeStatus: sip.StatusOK, wantMethods: "INVITE,BYE",
|
|
ack: func(context.Context, *sipgo.DialogClientSession) error { return ackFailure },
|
|
wantErrors: []string{"send ACK: injected ACK failure"},
|
|
},
|
|
{
|
|
name: "media and cleanup errors are preserved", validSDP: false, byeStatus: sip.StatusInternalServerError, wantMethods: "INVITE,ACK,BYE",
|
|
wantErrors: []string{"answer SDP lacks usable", "SIP/2.0 500"},
|
|
},
|
|
}
|
|
for _, test := range tests {
|
|
t.Run(test.name, func(t *testing.T) {
|
|
methods, logData, err := localCall(t, test.ack, test.validSDP, test.byeStatus)
|
|
if len(test.wantErrors) == 0 && err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
for _, want := range test.wantErrors {
|
|
if err == nil || !strings.Contains(err.Error(), want) {
|
|
t.Fatalf("error %v does not contain %q", err, want)
|
|
}
|
|
}
|
|
if got := strings.Join(methods, ","); got != test.wantMethods {
|
|
t.Fatalf("SIP methods = %s, want %s", got, test.wantMethods)
|
|
}
|
|
if strings.Contains(string(logData), "138"+"0000"+"1234") || !strings.Contains(string(logData), "200 OK") {
|
|
t.Fatalf("unexpected SIP log: %s", logData)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
func localCall(t *testing.T, ack func(context.Context, *sipgo.DialogClientSession) error, validSDP bool, byeStatus int) ([]string, []byte, error) {
|
|
t.Helper()
|
|
const testNumber = "138" + "0000" + "1234"
|
|
rtpConn, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4zero, Port: 0})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer rtpConn.Close()
|
|
sipConn, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.ParseIP("127.0.0.1"), Port: 0})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer sipConn.Close()
|
|
|
|
type serverResult struct {
|
|
methods []string
|
|
err error
|
|
}
|
|
serverDone := make(chan serverResult, 1)
|
|
go func() {
|
|
parser := sip.NewParser()
|
|
buf := make([]byte, 4096)
|
|
var methods []string
|
|
_ = sipConn.SetReadDeadline(time.Now().Add(3 * time.Second))
|
|
for {
|
|
n, source, err := sipConn.ReadFromUDP(buf)
|
|
if err != nil {
|
|
serverDone <- serverResult{methods, err}
|
|
return
|
|
}
|
|
message, err := parser.ParseSIP(buf[:n])
|
|
if err != nil {
|
|
serverDone <- serverResult{methods, err}
|
|
return
|
|
}
|
|
request, ok := message.(*sip.Request)
|
|
if !ok {
|
|
continue
|
|
}
|
|
methods = append(methods, request.Method.String())
|
|
switch request.Method {
|
|
case sip.INVITE:
|
|
for _, status := range []int{sip.StatusRinging, sip.StatusOK} {
|
|
response := sip.NewResponseFromRequest(request, status, map[int]string{sip.StatusRinging: "Ringing", sip.StatusOK: "OK"}[status], nil)
|
|
response.To().Params.Add("tag", "offline-test")
|
|
if status == sip.StatusOK {
|
|
response.AppendHeader(&sip.ContactHeader{Address: sip.Uri{Scheme: "sip", User: "mock", Host: "127.0.0.1", Port: sipConn.LocalAddr().(*net.UDPAddr).Port}})
|
|
response.AppendHeader(sip.NewHeader("Content-Type", "application/sdp"))
|
|
if validSDP {
|
|
response.SetBody([]byte("v=0\r\nc=IN IP4 127.0.0.1\r\nm=audio " + strconv.Itoa(rtpConn.LocalAddr().(*net.UDPAddr).Port) + " RTP/AVP 8\r\n"))
|
|
} else {
|
|
response.SetBody([]byte("v=0\r\n"))
|
|
}
|
|
}
|
|
if _, err := sipConn.WriteToUDP([]byte(response.String()), source); err != nil {
|
|
serverDone <- serverResult{methods, err}
|
|
return
|
|
}
|
|
if status == sip.StatusRinging {
|
|
time.Sleep(5 * time.Millisecond)
|
|
}
|
|
}
|
|
case sip.BYE:
|
|
reason := "OK"
|
|
if byeStatus != sip.StatusOK {
|
|
reason = "Server Error"
|
|
}
|
|
response := sip.NewResponseFromRequest(request, byeStatus, reason, nil)
|
|
_, err := sipConn.WriteToUDP([]byte(response.String()), source)
|
|
serverDone <- serverResult{methods, err}
|
|
return
|
|
}
|
|
}
|
|
}()
|
|
|
|
dir := t.TempDir()
|
|
cfg := config{
|
|
server: sipConn.LocalAddr().String(),
|
|
number: testNumber,
|
|
ttsText: "",
|
|
ttsVoice: "unused",
|
|
logPath: filepath.Join(dir, "sip.log"),
|
|
recordDir: filepath.Join(dir, "recordings"),
|
|
localAddr: "127.0.0.1:0",
|
|
advertiseIP: "127.0.0.1",
|
|
ringTimeout: time.Second,
|
|
maxDuration: 120 * time.Millisecond,
|
|
vadThreshold: 1000,
|
|
}
|
|
var runErr error
|
|
if ack == nil {
|
|
runErr = run(context.Background(), cfg)
|
|
} else {
|
|
runErr = runWithACK(context.Background(), cfg, ack)
|
|
}
|
|
result := <-serverDone
|
|
if result.err != nil {
|
|
t.Fatal(result.err)
|
|
}
|
|
logData, err := os.ReadFile(cfg.logPath)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
return result.methods, logData, runErr
|
|
}
|
|
|
|
func equalPCM(a, b []int16) bool {
|
|
if len(a) != len(b) {
|
|
return false
|
|
}
|
|
for i := range a {
|
|
if a[i] != b[i] {
|
|
return false
|
|
}
|
|
}
|
|
return true
|
|
}
|
|
|
|
func errString(err error) string {
|
|
if err == nil {
|
|
return ""
|
|
}
|
|
return err.Error()
|
|
}
|