webrtc: rewrite WHIP client (#4299)
This commit is contained in:
@@ -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())
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user