webrtc: rewrite WHIP client (#4299)

This commit is contained in:
Alessandro Ros
2025-03-01 17:01:57 +01:00
committed by GitHub
parent aa101c680c
commit c692f3b78c
10 changed files with 85 additions and 140 deletions
+10 -8
View File
@@ -531,7 +531,7 @@ func TestAPIProtocolListGet(t *testing.T) {
Log: test.NilLogger,
}
_, err = c.Read(context.Background())
err = c.Initialize(context.Background())
require.NoError(t, err)
defer checkClose(t, c.Close)
@@ -1019,12 +1019,6 @@ func TestAPIProtocolKick(t *testing.T) {
u, err := url.Parse("http://localhost:8889/mypath/whip")
require.NoError(t, err)
c := &whip.Client{
HTTPClient: hc,
URL: u,
Log: test.NilLogger,
}
track := &webrtc.OutgoingTrack{
Caps: pwebrtc.RTPCodecCapability{
MimeType: pwebrtc.MimeTypeH264,
@@ -1033,7 +1027,15 @@ func TestAPIProtocolKick(t *testing.T) {
},
}
err = c.Publish(context.Background(), []*webrtc.OutgoingTrack{track})
c := &whip.Client{
HTTPClient: hc,
URL: u,
Log: test.NilLogger,
Publish: true,
OutgoingTracks: []*webrtc.OutgoingTrack{track},
}
err = c.Initialize(context.Background())
require.NoError(t, err)
defer func() {
require.Error(t, c.Close())
+9 -7
View File
@@ -246,12 +246,6 @@ webrtc_sessions_bytes_sent 0
defer tr.CloseIdleConnections()
hc2 := &http.Client{Transport: tr}
s := &whip.Client{
HTTPClient: hc2,
URL: su,
Log: test.NilLogger,
}
track := &webrtc.OutgoingTrack{
Caps: pwebrtc.RTPCodecCapability{
MimeType: pwebrtc.MimeTypeH264,
@@ -260,7 +254,15 @@ webrtc_sessions_bytes_sent 0
},
}
err = s.Publish(context.Background(), []*webrtc.OutgoingTrack{track})
s := &whip.Client{
HTTPClient: hc2,
URL: su,
Log: test.NilLogger,
Publish: true,
OutgoingTracks: []*webrtc.OutgoingTrack{track},
}
err = s.Initialize(context.Background())
require.NoError(t, err)
defer checkClose(t, s.Close)
+1 -1
View File
@@ -503,7 +503,7 @@ func TestPathRunOnRead(t *testing.T) {
Log: test.NilLogger,
}
_, err = c.Read(context.Background())
err = c.Initialize(context.Background())
require.NoError(t, err)
defer checkClose(t, c.Close)
}
+11 -6
View File
@@ -423,7 +423,7 @@ outer:
}
// GatherIncomingTracks gathers incoming tracks.
func (co *PeerConnection) GatherIncomingTracks(ctx context.Context) ([]*IncomingTrack, error) {
func (co *PeerConnection) GatherIncomingTracks(ctx context.Context) error {
var sdp sdp.SessionDescription
sdp.Unmarshal([]byte(co.wr.RemoteDescription().SDP)) //nolint:errcheck
@@ -436,9 +436,9 @@ func (co *PeerConnection) GatherIncomingTracks(ctx context.Context) ([]*Incoming
select {
case <-t.C:
if len(co.incomingTracks) != 0 {
return co.incomingTracks, nil
return nil
}
return nil, fmt.Errorf("deadline exceeded while waiting tracks")
return fmt.Errorf("deadline exceeded while waiting tracks")
case pair := <-co.incomingTrack:
t := &IncomingTrack{
@@ -451,14 +451,14 @@ func (co *PeerConnection) GatherIncomingTracks(ctx context.Context) ([]*Incoming
co.incomingTracks = append(co.incomingTracks, t)
if len(co.incomingTracks) >= maxTrackCount {
return co.incomingTracks, nil
return nil
}
case <-co.Failed():
return nil, fmt.Errorf("peer connection closed")
return fmt.Errorf("peer connection closed")
case <-ctx.Done():
return nil, fmt.Errorf("terminated")
return fmt.Errorf("terminated")
}
}
}
@@ -505,6 +505,11 @@ func (co *PeerConnection) LocalCandidate() string {
return ""
}
// IncomingTracks returns incoming tracks.
func (co *PeerConnection) IncomingTracks() []*IncomingTrack {
return co.incomingTracks
}
// StartReading starts reading all incoming tracks.
func (co *PeerConnection) StartReading() {
for _, track := range co.incomingTracks {
+1 -1
View File
@@ -403,7 +403,7 @@ func TestToStream(t *testing.T) {
})
require.NoError(t, err)
_, err = pc2.GatherIncomingTracks(context.Background())
err = pc2.GatherIncomingTracks(context.Background())
require.NoError(t, err)
var stream *stream.Stream
+38 -104
View File
@@ -26,19 +26,18 @@ const (
// Client is a WHIP client.
type Client struct {
HTTPClient *http.Client
URL *url.URL
Log logger.Writer
URL *url.URL
Publish bool
OutgoingTracks []*webrtc.OutgoingTrack
HTTPClient *http.Client
Log logger.Writer
pc *webrtc.PeerConnection
patchIsSupported bool
}
// Publish publishes tracks.
func (c *Client) Publish(
ctx context.Context,
outgoingTracks []*webrtc.OutgoingTrack,
) error {
// Initialize initializes the Client.
func (c *Client) Initialize(ctx context.Context) error {
iceServers, err := c.optionsICEServers(ctx)
if err != nil {
return err
@@ -50,10 +49,11 @@ func (c *Client) Publish(
IPsFromInterfaces: true,
HandshakeTimeout: conf.Duration(10 * time.Second),
TrackGatherTimeout: conf.Duration(2 * time.Second),
Publish: true,
OutgoingTracks: outgoingTracks,
Publish: c.Publish,
OutgoingTracks: c.OutgoingTracks,
Log: c.Log,
}
err = c.pc.Start()
if err != nil {
return err
@@ -77,6 +77,23 @@ func (c *Client) Publish(
return err
}
if !c.Publish {
var sdp sdp.SessionDescription
err = sdp.Unmarshal([]byte(res.Answer.SDP))
if err != nil {
c.deleteSession(context.Background()) //nolint:errcheck
c.pc.Close()
return err
}
err = webrtc.TracksAreValid(sdp.MediaDescriptions)
if err != nil {
c.deleteSession(context.Background()) //nolint:errcheck
c.pc.Close()
return err
}
}
err = c.pc.SetAnswer(res.Answer)
if err != nil {
c.deleteSession(context.Background()) //nolint:errcheck
@@ -91,7 +108,7 @@ outer:
for {
select {
case ca := <-c.pc.NewLocalCandidate():
err := c.patchCandidate(ctx, offer, res.ETag, ca)
err = c.patchCandidate(ctx, offer, res.ETag, ca)
if err != nil {
c.deleteSession(context.Background()) //nolint:errcheck
c.pc.Close()
@@ -110,104 +127,16 @@ outer:
}
}
return nil
}
// Read reads tracks.
func (c *Client) Read(ctx context.Context) ([]*webrtc.IncomingTrack, error) {
iceServers, err := c.optionsICEServers(ctx)
if err != nil {
return nil, err
}
c.pc = &webrtc.PeerConnection{
LocalRandomUDP: true,
ICEServers: iceServers,
IPsFromInterfaces: true,
HandshakeTimeout: conf.Duration(10 * time.Second),
TrackGatherTimeout: conf.Duration(2 * time.Second),
Publish: false,
Log: c.Log,
}
err = c.pc.Start()
if err != nil {
return nil, err
}
offer, err := c.pc.CreatePartialOffer()
if err != nil {
c.pc.Close()
return nil, err
}
res, err := c.postOffer(ctx, offer)
if err != nil {
c.pc.Close()
return nil, err
}
c.URL, err = c.URL.Parse(res.Location)
if err != nil {
c.pc.Close()
return nil, err
}
var sdp sdp.SessionDescription
err = sdp.Unmarshal([]byte(res.Answer.SDP))
if err != nil {
c.deleteSession(context.Background()) //nolint:errcheck
c.pc.Close()
return nil, err
}
err = webrtc.TracksAreValid(sdp.MediaDescriptions)
if err != nil {
c.deleteSession(context.Background()) //nolint:errcheck
c.pc.Close()
return nil, err
}
err = c.pc.SetAnswer(res.Answer)
if err != nil {
c.deleteSession(context.Background()) //nolint:errcheck
c.pc.Close()
return nil, err
}
t := time.NewTimer(handshakeTimeout)
defer t.Stop()
outer:
for {
select {
case ca := <-c.pc.NewLocalCandidate():
err = c.patchCandidate(ctx, offer, res.ETag, ca)
if err != nil {
c.deleteSession(context.Background()) //nolint:errcheck
c.pc.Close()
return nil, err
}
case <-c.pc.GatheringDone():
case <-c.pc.Connected():
break outer
case <-t.C:
if !c.Publish {
err = c.pc.GatherIncomingTracks(ctx)
if err != nil {
c.deleteSession(context.Background()) //nolint:errcheck
c.pc.Close()
return nil, fmt.Errorf("deadline exceeded while waiting connection")
return err
}
}
tracks, err := c.pc.GatherIncomingTracks(ctx)
if err != nil {
c.deleteSession(context.Background()) //nolint:errcheck
c.pc.Close()
return nil, err
}
return tracks, nil
return nil
}
// PeerConnection returns the underlying peer connection.
@@ -215,6 +144,11 @@ func (c *Client) PeerConnection() *webrtc.PeerConnection {
return c.pc
}
// IncomingTracks returns incoming tracks.
func (c *Client) IncomingTracks() []*webrtc.IncomingTrack {
return c.pc.IncomingTracks()
}
// StartReading starts reading all incoming tracks.
func (c *Client) StartReading() {
c.pc.StartReading()
+2 -2
View File
@@ -126,12 +126,12 @@ func TestClientRead(t *testing.T) {
require.NoError(t, err)
cl := &Client{
HTTPClient: &http.Client{},
URL: u,
HTTPClient: &http.Client{},
Log: test.NilLogger,
}
_, err = cl.Read(context.Background())
err = cl.Initialize(context.Background())
require.NoError(t, err)
defer cl.Close() //nolint:errcheck
}
+11 -9
View File
@@ -275,12 +275,6 @@ func TestServerPublish(t *testing.T) {
su, err := url.Parse("http://myuser:mypass@localhost:8886/teststream/whip?param=value")
require.NoError(t, err)
wc := &whip.Client{
HTTPClient: hc,
URL: su,
Log: test.NilLogger,
}
track := &webrtc.OutgoingTrack{
Caps: pwebrtc.RTPCodecCapability{
MimeType: pwebrtc.MimeTypeH264,
@@ -289,7 +283,15 @@ func TestServerPublish(t *testing.T) {
},
}
err = wc.Publish(context.Background(), []*webrtc.OutgoingTrack{track})
wc := &whip.Client{
HTTPClient: hc,
URL: su,
Publish: true,
OutgoingTracks: []*webrtc.OutgoingTrack{track},
Log: test.NilLogger,
}
err = wc.Initialize(context.Background())
require.NoError(t, err)
defer checkClose(t, wc.Close)
@@ -587,13 +589,13 @@ func TestServerRead(t *testing.T) {
}
}()
tracks, err := wc.Read(context.Background())
err = wc.Initialize(context.Background())
require.NoError(t, err)
defer checkClose(t, wc.Close)
done := make(chan struct{})
tracks[0].OnPacketRTP = func(pkt *rtp.Packet) {
wc.IncomingTracks()[0].OnPacketRTP = func(pkt *rtp.Packet) {
select {
case <-done:
default:
+1 -1
View File
@@ -205,7 +205,7 @@ func (s *session) runPublish() (int, error) {
s.pc = pc
s.mutex.Unlock()
_, err = pc.GatherIncomingTracks(s.ctx)
err = pc.GatherIncomingTracks(s.ctx)
if err != nil {
return 0, err
}
+1 -1
View File
@@ -54,7 +54,7 @@ func (s *Source) Run(params defs.StaticSourceRunParams) error {
Log: s,
}
_, err = client.Read(params.Context)
err = client.Initialize(params.Context)
if err != nil {
return err
}