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