Files
magic-sonar/cmd/sip-demo/main_test.go
T
2026-08-18 15:53:53 +08:00

366 lines
11 KiB
Go

package main
import (
"bytes"
"context"
"errors"
"io"
"net"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strconv"
"strings"
"testing"
"time"
"github.com/emiago/sipgo"
"github.com/emiago/sipgo/sip"
)
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"
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/compatible-mode/v1/audio/speech" {
t.Errorf("path = %s", r.URL.Path)
}
if r.Header.Get("Authorization") != "Bearer "+key {
t.Error("missing bearer authorization")
}
w.Header().Set("Content-Type", "application/octet-stream")
_, _ = w.Write(pcmToBytes([]int16{1, 2, 3}))
}))
defer server.Close()
pcm, rate, err := synthesize(context.Background(), server.Client(), server.URL+"/compatible-mode/v1", 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)
}
}
func TestBailianTTSCredentialsRejectPlaintextAndUnsafeRedirects(t *testing.T) {
const key = "redirect-key-that-must-not-leak"
plainRequests := 0
plainClient := &http.Client{Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
plainRequests++
return &http.Response{StatusCode: http.StatusOK, Body: io.NopCloser(strings.NewReader(""))}, nil
})}
if _, _, err := synthesize(context.Background(), plainClient, "http://192.0.2.1", key, "voice", "text"); err == nil || plainRequests != 0 {
t.Fatal("accepted non-loopback HTTP")
}
targetHit := make(chan struct{}, 1)
target := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {
targetHit <- struct{}{}
}))
defer target.Close()
source := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
http.Redirect(w, r, target.URL+"/audio/speech", http.StatusTemporaryRedirect)
}))
defer source.Close()
_, _, err := synthesize(context.Background(), source.Client(), source.URL, key, "voice", "text")
if err == nil || strings.Contains(err.Error(), key) {
t.Fatalf("cross-origin redirect error = %v", err)
}
select {
case <-targetHit:
t.Fatal("credential request reached redirect target")
default:
}
downgradeRequests := 0
downgradeClient := &http.Client{Transport: roundTripFunc(func(req *http.Request) (*http.Response, error) {
downgradeRequests++
if downgradeRequests == 1 {
return &http.Response{
StatusCode: http.StatusTemporaryRedirect,
Header: http.Header{"Location": []string{"http://secure.example/audio/speech"}},
Body: io.NopCloser(strings.NewReader("")),
Request: req,
}, nil
}
return &http.Response{StatusCode: http.StatusOK, Body: io.NopCloser(strings.NewReader(""))}, nil
})}
_, _, err = synthesize(context.Background(), downgradeClient, "https://secure.example", key, "voice", "text")
if err == nil || downgradeRequests != 1 || strings.Contains(err.Error(), key) {
t.Fatalf("HTTPS downgrade error = %v", err)
}
}
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()
}
type roundTripFunc func(*http.Request) (*http.Response, error)
func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
return f(req)
}