owpengram-server/internal/sfu/pion_e2e_test.go

445 lines
12 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

package sfu
import (
"context"
"crypto/tls"
"errors"
"fmt"
"net"
"testing"
"time"
"github.com/pion/dtls/v3"
"github.com/pion/dtls/v3/pkg/crypto/selfsign"
dtlsnet "github.com/pion/dtls/v3/pkg/net"
"github.com/pion/ice/v4"
"github.com/pion/rtp"
"github.com/pion/srtp/v3"
"go.uber.org/zap/zaptest"
)
var errClientDTLSStateMissing = errors.New("client dtls state missing")
// fakeTgcallsClient 按 tgcalls GroupNetworkManager 的角色契约模拟客户端:
// ICE CONTROLLED等待 SFU 发起连通性检查、DTLS setup="passive"server 角色,
// 等 SFU 主动握手、SRTP isClient=false。
type fakeTgcallsClient struct {
t *testing.T
ufrag, pwd string
cert tls.Certificate
fingerprint string
ssrc uint32
groups []SsrcGroup
session *srtp.SessionSRTP
write *srtp.WriteStreamSRTP
closeFn []func()
}
func newFakeTgcallsClient(t *testing.T, ssrc uint32) *fakeTgcallsClient {
t.Helper()
cert, err := selfsign.GenerateSelfSigned()
if err != nil {
t.Fatalf("client cert: %v", err)
}
fp, err := certificateFingerprint(cert.Certificate[0])
if err != nil {
t.Fatalf("client fingerprint: %v", err)
}
ufrag, _ := randomICEString(8)
pwd, _ := randomICEString(24)
return &fakeTgcallsClient{t: t, ufrag: ufrag, pwd: pwd, cert: cert, fingerprint: fp, ssrc: ssrc}
}
func (c *fakeTgcallsClient) offer() ClientOffer {
return ClientOffer{
AudioSSRC: c.ssrc,
Ufrag: c.ufrag,
Pwd: c.pwd,
FingerprintSHA256: c.fingerprint,
SsrcGroups: c.groups,
}
}
// connect 完成 ICE(controlled)+DTLS(server)+SRTP 建链。
func (c *fakeTgcallsClient) connect(ctx context.Context, answer ServerAnswer) error {
agent, err := ice.NewAgent(&ice.AgentConfig{
NetworkTypes: []ice.NetworkType{ice.NetworkTypeUDP4},
CandidateTypes: []ice.CandidateType{ice.CandidateTypeHost},
LocalUfrag: c.ufrag,
LocalPwd: c.pwd,
IncludeLoopback: true,
})
if err != nil {
return err
}
c.closeFn = append(c.closeFn, func() { _ = agent.Close() })
if err := agent.OnCandidate(func(ice.Candidate) {}); err != nil {
return err
}
if err := agent.GatherCandidates(); err != nil {
return err
}
for _, cand := range answer.Candidates {
remote, err := ice.NewCandidateHost(&ice.CandidateHostConfig{
Network: "udp",
Address: cand.IP,
Port: cand.Port,
Component: 1,
})
if err != nil {
return err
}
if err := agent.AddRemoteCandidate(remote); err != nil {
return err
}
}
conn, err := agent.Accept(ctx, answer.Ufrag, answer.Pwd) // CONTROLLED
if err != nil {
return err
}
demux := newDemuxer(conn)
c.closeFn = append(c.closeFn, demux.Close)
dtlsRaw := demux.dtlsConn()
dtlsConn, err := dtls.Server(dtlsnet.PacketConnFromConn(dtlsRaw), dtlsRaw.RemoteAddr(), &dtls.Config{
Certificates: []tls.Certificate{c.cert},
ClientAuth: dtls.RequireAnyClientCert,
InsecureSkipVerify: true,
ExtendedMasterSecret: dtls.RequireExtendedMasterSecret,
SRTPProtectionProfiles: []dtls.SRTPProtectionProfile{
dtls.SRTP_AEAD_AES_128_GCM,
dtls.SRTP_AES128_CM_HMAC_SHA1_80,
},
})
if err != nil {
return err
}
hsCtx, hsCancel := context.WithTimeout(ctx, 20*time.Second)
err = dtlsConn.HandshakeContext(hsCtx)
hsCancel()
if err != nil {
return err
}
state, ok := dtlsConn.ConnectionState()
if !ok {
return errClientDTLSStateMissing
}
// 服务端证书指纹应等于信令面下发的 answer.FingerprintSHA256。
gotFP, err := certificateFingerprint(state.PeerCertificates[0])
if err != nil || normalizeFingerprint(gotFP) != normalizeFingerprint(answer.FingerprintSHA256) {
return fmt.Errorf("sfu fingerprint mismatch: %v / %s vs %s", err, gotFP, answer.FingerprintSHA256)
}
profile, _ := dtlsConn.SelectedSRTPProtectionProfile()
srtpConfig := &srtp.Config{}
switch profile {
case dtls.SRTP_AEAD_AES_128_GCM:
srtpConfig.Profile = srtp.ProtectionProfileAeadAes128Gcm
case dtls.SRTP_AES128_CM_HMAC_SHA1_80:
srtpConfig.Profile = srtp.ProtectionProfileAes128CmHmacSha1_80
default:
return fmt.Errorf("unexpected srtp profile %v", profile)
}
if err := srtpConfig.ExtractSessionKeysFromDTLS(&state, false); err != nil {
return err
}
c.session, err = srtp.NewSessionSRTP(demux.srtpConn(), srtpConfig)
if err != nil {
return err
}
c.write, err = c.session.OpenWriteStream()
return err
}
func (c *fakeTgcallsClient) close() {
for i := len(c.closeFn) - 1; i >= 0; i-- {
c.closeFn[i]()
}
}
// sendOpusPacket 构造带 audio-level 扩展one-byte headerid=1的 RTP 包。
func (c *fakeTgcallsClient) sendOpusPacket(seq uint16, payload []byte, audioLevel byte) error {
pkt := &rtp.Packet{
Header: rtp.Header{
Version: 2,
PayloadType: 111, // opus
SequenceNumber: seq,
Timestamp: uint32(seq) * 960,
SSRC: c.ssrc,
},
Payload: payload,
}
pkt.Header.Extension = true
pkt.Header.ExtensionProfile = 0xBEDE
if err := pkt.Header.SetExtension(1, []byte{audioLevel}); err != nil {
return err
}
raw, err := pkt.Marshal()
if err != nil {
return err
}
_, err = c.write.Write(raw)
return err
}
// sendVideoPacket 构造指定 ssrc 的视频 RTP 包VP8 PT=100
func (c *fakeTgcallsClient) sendVideoPacket(ssrc uint32, seq uint16) error {
pkt := &rtp.Packet{
Header: rtp.Header{
Version: 2,
PayloadType: 100, // VP8
SequenceNumber: seq,
Timestamp: uint32(seq) * 3000,
SSRC: ssrc,
},
Payload: []byte{0x90, 0x00, byte(seq)},
}
raw, err := pkt.Marshal()
if err != nil {
return err
}
_, err = c.write.Write(raw)
return err
}
// 视频选层:发布端三层 simulcastSFU 只转发 SIM[0] 层(订阅端只为该 ssrc 建
// 解码 sink高层包客户端会丢弃SFU 侧直接终结省下行)。
func TestPionSFUForwardsOnlyBaseSimulcastLayer(t *testing.T) {
if testing.Short() {
t.Skip("e2e sfu test")
}
port := pickUDPPort(t)
svc, err := NewPion(PionConfig{UDPPort: port, AdvertiseIP: "127.0.0.1", Logger: zaptest.NewLogger(t)})
if err != nil {
t.Fatalf("new pion sfu: %v", err)
}
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
const callID = int64(901)
alice := newFakeTgcallsClient(t, 0x500)
base := alice.ssrc + 1
alice.groups = []SsrcGroup{
{Semantics: "SIM", Sources: []uint32{base, base + 2, base + 4}},
{Semantics: "FID", Sources: []uint32{base, base + 1}},
{Semantics: "FID", Sources: []uint32{base + 2, base + 3}},
{Semantics: "FID", Sources: []uint32{base + 4, base + 5}},
}
bob := newFakeTgcallsClient(t, 0x600)
defer alice.close()
defer bob.close()
answerA, err := svc.Join(ctx, callID, 1, EndpointMain, alice.offer())
if err != nil {
t.Fatalf("join alice: %v", err)
}
answerB, err := svc.Join(ctx, callID, 2, EndpointMain, bob.offer())
if err != nil {
t.Fatalf("join bob: %v", err)
}
errCh := make(chan error, 2)
go func() { errCh <- alice.connect(ctx, answerA) }()
go func() { errCh <- bob.connect(ctx, answerB) }()
for i := 0; i < 2; i++ {
if err := <-errCh; err != nil {
t.Fatalf("client connect: %v", err)
}
}
// alice 同时在 layer0应转发与 layer1应被 SFU 终结)发包。
stop := make(chan struct{})
defer close(stop)
go func() {
seq := uint16(1)
ticker := time.NewTicker(20 * time.Millisecond)
defer ticker.Stop()
for {
select {
case <-stop:
return
case <-ticker.C:
_ = alice.sendVideoPacket(base, seq)
_ = alice.sendVideoPacket(base+2, seq)
seq++
}
}
}()
// bob 第一条转发流必须是 layer0观察窗口内绝不能出现 layer1。
type accepted struct{ ssrc uint32 }
got := make(chan accepted, 4)
go func() {
for {
stream, ssrc, err := bob.session.AcceptStream()
if err != nil {
return
}
got <- accepted{ssrc}
go func() {
buf := make([]byte, 1500)
for {
if _, err := stream.Read(buf); err != nil {
return
}
}
}()
}
}()
deadline := time.After(10 * time.Second)
sawBase := false
for !sawBase {
select {
case in := <-got:
if in.ssrc == base+2 {
t.Fatalf("higher simulcast layer %#x leaked through SFU", in.ssrc)
}
if in.ssrc == base {
sawBase = true
}
case <-deadline:
t.Fatalf("bob never received base layer stream")
}
}
// 再观察一段时间确认高层不漏。
quiet := time.After(2 * time.Second)
for {
select {
case in := <-got:
if in.ssrc == base+2 {
t.Fatalf("higher simulcast layer %#x leaked through SFU", in.ssrc)
}
case <-quiet:
_ = svc.CloseRoom(ctx, callID)
return
}
}
}
func TestPionSFUForwardsOpusBetweenClients(t *testing.T) {
if testing.Short() {
t.Skip("e2e sfu test")
}
port := pickUDPPort(t)
svc, err := NewPion(PionConfig{
UDPPort: port,
AdvertiseIP: "127.0.0.1",
Logger: zaptest.NewLogger(t),
})
if err != nil {
t.Fatalf("new pion sfu: %v", err)
}
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
const callID = int64(900)
alice := newFakeTgcallsClient(t, 0xA11CE)
bob := newFakeTgcallsClient(t, 0xB0B)
defer alice.close()
defer bob.close()
answerA, err := svc.Join(ctx, callID, 1, EndpointMain, alice.offer())
if err != nil {
t.Fatalf("join alice: %v", err)
}
if len(answerA.Candidates) != 1 || answerA.Candidates[0].Port != port {
t.Fatalf("answer candidates = %+v", answerA.Candidates)
}
answerB, err := svc.Join(ctx, callID, 2, EndpointMain, bob.offer())
if err != nil {
t.Fatalf("join bob: %v", err)
}
// 两端建链ICE prflx → DTLS 由 SFU 主动握手 → SRTP
errCh := make(chan error, 2)
go func() { errCh <- alice.connect(ctx, answerA) }()
go func() { errCh <- bob.connect(ctx, answerB) }()
for i := 0; i < 2; i++ {
if err := <-errCh; err != nil {
t.Fatalf("client connect: %v", err)
}
}
// alice 持续发包SFU AcceptStream 在首包后建立读流)。
stop := make(chan struct{})
defer close(stop)
go func() {
seq := uint16(1)
ticker := time.NewTicker(20 * time.Millisecond)
defer ticker.Stop()
for {
select {
case <-stop:
return
case <-ticker.C:
_ = alice.sendOpusPacket(seq, []byte{0xDE, 0xAD, byte(seq)}, 0x7F)
seq++
}
}
}()
// bob 收到来自 alice 的转发流SSRC 不重写、扩展保留。
acceptCtx, acceptCancel := context.WithTimeout(ctx, 15*time.Second)
defer acceptCancel()
type accepted struct {
stream *srtp.ReadStreamSRTP
ssrc uint32
}
got := make(chan accepted, 1)
go func() {
stream, ssrc, err := bob.session.AcceptStream()
if err == nil {
got <- accepted{stream, ssrc}
}
}()
var inbound accepted
select {
case inbound = <-got:
case <-acceptCtx.Done():
t.Fatalf("bob never received forwarded stream")
}
if inbound.ssrc != alice.ssrc {
t.Fatalf("forwarded ssrc = %#x, want alice %#xSFU 不得重写 SSRC", inbound.ssrc, alice.ssrc)
}
buf := make([]byte, 1500)
n, err := inbound.stream.Read(buf)
if err != nil {
t.Fatalf("bob read: %v", err)
}
var pkt rtp.Packet
if err := pkt.Unmarshal(buf[:n]); err != nil {
t.Fatalf("unmarshal forwarded packet: %v", err)
}
if pkt.PayloadType != 111 || pkt.SSRC != alice.ssrc {
t.Fatalf("forwarded packet = %+v", pkt.Header)
}
if ext := pkt.Header.GetExtension(1); len(ext) != 1 || ext[0] != 0x7F {
t.Fatalf("audio-level extension lost: %vspeaking 指示依赖逐包透传)", ext)
}
// 媒体面活性:双端都应在 alive 集合alice 发包、bob 收包均计活性…bob 仅收不发,
// 其活性来自 SFU 写出写出不计bob 至少 connected+握手活性在窗口内)。
alive := svc.AliveUserIDs(callID)
aliveSet := map[int64]bool{}
for _, id := range alive {
aliveSet[id] = true
}
if !aliveSet[1] {
t.Fatalf("alice must be media-alive, got %v", alive)
}
if err := svc.Leave(ctx, callID, 1, EndpointMain); err != nil {
t.Fatalf("leave: %v", err)
}
if err := svc.CloseRoom(ctx, callID); err != nil {
t.Fatalf("close room: %v", err)
}
}
func pickUDPPort(t *testing.T) int {
t.Helper()
conn, err := net.ListenUDP("udp4", &net.UDPAddr{IP: net.IPv4(127, 0, 0, 1)})
if err != nil {
t.Fatalf("pick port: %v", err)
}
port := conn.LocalAddr().(*net.UDPAddr).Port
_ = conn.Close()
return port
}