diff --git a/internal/core/api_test.go b/internal/core/api_test.go index a3396f64..7172f4f3 100644 --- a/internal/core/api_test.go +++ b/internal/core/api_test.go @@ -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()) diff --git a/internal/core/metrics_test.go b/internal/core/metrics_test.go index dece1965..1f7b1fbb 100644 --- a/internal/core/metrics_test.go +++ b/internal/core/metrics_test.go @@ -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) diff --git a/internal/core/path_test.go b/internal/core/path_test.go index 6d15e463..39b7f067 100644 --- a/internal/core/path_test.go +++ b/internal/core/path_test.go @@ -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) } diff --git a/internal/protocols/webrtc/peer_connection.go b/internal/protocols/webrtc/peer_connection.go index d2e29abc..ecc86986 100644 --- a/internal/protocols/webrtc/peer_connection.go +++ b/internal/protocols/webrtc/peer_connection.go @@ -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 { diff --git a/internal/protocols/webrtc/to_stream_test.go b/internal/protocols/webrtc/to_stream_test.go index 48049a98..1464a8b6 100644 --- a/internal/protocols/webrtc/to_stream_test.go +++ b/internal/protocols/webrtc/to_stream_test.go @@ -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 diff --git a/internal/protocols/whip/client.go b/internal/protocols/whip/client.go index 8b2a59ba..6563159b 100644 --- a/internal/protocols/whip/client.go +++ b/internal/protocols/whip/client.go @@ -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() diff --git a/internal/protocols/whip/client_test.go b/internal/protocols/whip/client_test.go index 0a79735c..10984ccc 100644 --- a/internal/protocols/whip/client_test.go +++ b/internal/protocols/whip/client_test.go @@ -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 } diff --git a/internal/servers/webrtc/server_test.go b/internal/servers/webrtc/server_test.go index 22fee958..9ba7c06a 100644 --- a/internal/servers/webrtc/server_test.go +++ b/internal/servers/webrtc/server_test.go @@ -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: diff --git a/internal/servers/webrtc/session.go b/internal/servers/webrtc/session.go index 2acc7f26..918ff40b 100644 --- a/internal/servers/webrtc/session.go +++ b/internal/servers/webrtc/session.go @@ -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 } diff --git a/internal/staticsources/webrtc/source.go b/internal/staticsources/webrtc/source.go index aa6817ce..1274eb62 100644 --- a/internal/staticsources/webrtc/source.go +++ b/internal/staticsources/webrtc/source.go @@ -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 }