diff --git a/internal/protocols/webrtc/from_stream.go b/internal/protocols/webrtc/from_stream.go index 518fb5e1..b3ee3dfa 100644 --- a/internal/protocols/webrtc/from_stream.go +++ b/internal/protocols/webrtc/from_stream.go @@ -67,8 +67,7 @@ func timestampToDuration(t int64, clockRate int) time.Duration { func setupVideoTrack( desc *description.Session, r *stream.Reader, - pc *PeerConnection, -) (format.Format, error) { +) (*OutgoingTrack, error) { var av1Format *format.AV1 media := desc.FindFormat(&av1Format) @@ -79,7 +78,6 @@ func setupVideoTrack( ClockRate: 90000, }, } - pc.OutgoingTracks = append(pc.OutgoingTracks, track) encoder := &rtpav1.Encoder{ PayloadType: 105, @@ -112,7 +110,7 @@ func setupVideoTrack( return nil }) - return av1Format, nil + return track, nil } var vp9Format *format.VP9 @@ -126,7 +124,6 @@ func setupVideoTrack( SDPFmtpLine: "profile-id=0", }, } - pc.OutgoingTracks = append(pc.OutgoingTracks, track) encoder := &rtpvp9.Encoder{ PayloadType: 96, @@ -160,7 +157,7 @@ func setupVideoTrack( return nil }) - return vp9Format, nil + return track, nil } var vp8Format *format.VP8 @@ -173,7 +170,6 @@ func setupVideoTrack( ClockRate: 90000, }, } - pc.OutgoingTracks = append(pc.OutgoingTracks, track) encoder := &rtpvp8.Encoder{ PayloadType: 96, @@ -206,7 +202,7 @@ func setupVideoTrack( return nil }) - return vp8Format, nil + return track, nil } var h265Format *format.H265 @@ -220,7 +216,6 @@ func setupVideoTrack( SDPFmtpLine: "level-id=93;profile-id=1;tier-flag=0;tx-mode=SRST", }, } - pc.OutgoingTracks = append(pc.OutgoingTracks, track) encoder := &rtph265.Encoder{ PayloadType: 96, @@ -263,7 +258,7 @@ func setupVideoTrack( return nil }) - return h265Format, nil + return track, nil } var h264Format *format.H264 @@ -277,7 +272,6 @@ func setupVideoTrack( SDPFmtpLine: "level-asymmetry-allowed=1;packetization-mode=1;profile-level-id=42e01f", }, } - pc.OutgoingTracks = append(pc.OutgoingTracks, track) encoder := &rtph264.Encoder{ PayloadType: 96, @@ -320,7 +314,7 @@ func setupVideoTrack( return nil }) - return h264Format, nil + return track, nil } return nil, nil @@ -329,8 +323,7 @@ func setupVideoTrack( func setupAudioTrack( desc *description.Session, r *stream.Reader, - pc *PeerConnection, -) (format.Format, error) { +) (*OutgoingTrack, error) { var opusFormat *format.Opus media := desc.FindFormat(&opusFormat) @@ -367,7 +360,6 @@ func setupAudioTrack( track := &OutgoingTrack{ Caps: caps, } - pc.OutgoingTracks = append(pc.OutgoingTracks, track) curTimestamp, err := randUint32() if err != nil { @@ -396,7 +388,7 @@ func setupAudioTrack( return nil }) - return opusFormat, nil + return track, nil } var g722Format *format.G722 @@ -409,7 +401,6 @@ func setupAudioTrack( ClockRate: 8000, }, } - pc.OutgoingTracks = append(pc.OutgoingTracks, track) r.OnData( media, @@ -423,7 +414,7 @@ func setupAudioTrack( return nil }) - return g722Format, nil + return track, nil } var g711Format *format.G711 @@ -482,7 +473,6 @@ func setupAudioTrack( track := &OutgoingTrack{ Caps: caps, } - pc.OutgoingTracks = append(pc.OutgoingTracks, track) if g711Format.ClockRate() == 8000 { curTimestamp, err := randUint32() @@ -567,7 +557,7 @@ func setupAudioTrack( }) } - return g711Format, nil + return track, nil } var lpcmFormat *format.LPCM @@ -596,7 +586,6 @@ func setupAudioTrack( Channels: uint16(lpcmFormat.ChannelCount), }, } - pc.OutgoingTracks = append(pc.OutgoingTracks, track) encoder := &rtplpcm.Encoder{ PayloadType: 96, @@ -641,7 +630,37 @@ func setupAudioTrack( return nil }) - return lpcmFormat, nil + return track, nil + } + + return nil, nil +} + +func setupKLVDataChannel( + desc *description.Session, + r *stream.Reader, +) (*OutgoingDataChannel, error) { + var klvFormat *format.KLV + media := desc.FindFormat(&klvFormat) + + if klvFormat != nil { + dataChan := &OutgoingDataChannel{ + Label: "KLV", + } + + r.OnData( + media, + klvFormat, + func(u *unit.Unit) error { + if u.NilPayload() { + return nil + } + + dataChan.Write(u.Payload.(unit.PayloadKLV)) + return nil + }) + + return dataChan, nil } return nil, nil @@ -653,17 +672,34 @@ func FromStream( r *stream.Reader, pc *PeerConnection, ) error { - videoFormat, err := setupVideoTrack(desc, r, pc) + videoTrack, err := setupVideoTrack(desc, r) if err != nil { return err } - audioFormat, err := setupAudioTrack(desc, r, pc) + if videoTrack != nil { + pc.OutgoingTracks = append(pc.OutgoingTracks, videoTrack) + } + + audioTrack, err := setupAudioTrack(desc, r) if err != nil { return err } - if videoFormat == nil && audioFormat == nil { + if audioTrack != nil { + pc.OutgoingTracks = append(pc.OutgoingTracks, audioTrack) + } + + klvDataChan, err := setupKLVDataChannel(desc, r) + if err != nil { + return err + } + + if klvDataChan != nil { + pc.OutgoingDataChannels = append(pc.OutgoingDataChannels, klvDataChan) + } + + if len(pc.OutgoingTracks) == 0 && len(pc.OutgoingDataChannels) == 0 { return errNoSupportedCodecsFrom } diff --git a/internal/protocols/webrtc/from_stream_test.go b/internal/protocols/webrtc/from_stream_test.go index 5ee02e08..b03a3d2a 100644 --- a/internal/protocols/webrtc/from_stream_test.go +++ b/internal/protocols/webrtc/from_stream_test.go @@ -28,7 +28,9 @@ func TestFromStreamNoSupportedCodecs(t *testing.T) { }), } - err := FromStream(desc, r, nil) + pc := &PeerConnection{} + + err := FromStream(desc, r, pc) require.Equal(t, errNoSupportedCodecsFrom, err) } diff --git a/internal/protocols/webrtc/outgoing_data_channel.go b/internal/protocols/webrtc/outgoing_data_channel.go new file mode 100644 index 00000000..9b4cd380 --- /dev/null +++ b/internal/protocols/webrtc/outgoing_data_channel.go @@ -0,0 +1,29 @@ +package webrtc + +import ( + "github.com/pion/webrtc/v4" +) + +// OutgoingDataChannel is an outgoing data channel. +type OutgoingDataChannel struct { + Label string + + dataChan *webrtc.DataChannel +} + +func (c *OutgoingDataChannel) setup(p *PeerConnection) error { + var err error + c.dataChan, err = p.wr.CreateDataChannel(c.Label, &webrtc.DataChannelInit{ + Ordered: ptrOf(false), + }) + if err != nil { + return err + } + + return nil +} + +// Write writes data to the channel. +func (c *OutgoingDataChannel) Write(data []byte) { + c.dataChan.Send(data) //nolint:errcheck +} diff --git a/internal/protocols/webrtc/outgoing_track.go b/internal/protocols/webrtc/outgoing_track.go index fad9c1dd..1caf289f 100644 --- a/internal/protocols/webrtc/outgoing_track.go +++ b/internal/protocols/webrtc/outgoing_track.go @@ -10,7 +10,7 @@ import ( "github.com/pion/webrtc/v4" ) -// OutgoingTrack is a WebRTC outgoing track +// OutgoingTrack is an outgoing track. type OutgoingTrack struct { Caps webrtc.RTPCodecCapability diff --git a/internal/protocols/webrtc/peer_connection.go b/internal/protocols/webrtc/peer_connection.go index b43f26be..4f2a1885 100644 --- a/internal/protocols/webrtc/peer_connection.go +++ b/internal/protocols/webrtc/peer_connection.go @@ -145,6 +145,7 @@ type PeerConnection struct { STUNGatherTimeout conf.Duration Publish bool OutgoingTracks []*OutgoingTrack + OutgoingDataChannels []*OutgoingDataChannel Log logger.Writer wr *webrtc.PeerConnection @@ -318,6 +319,14 @@ func (co *PeerConnection) Start() error { return err } } + + for _, dc := range co.OutgoingDataChannels { + err = dc.setup(co) + if err != nil { + co.wr.GracefulClose() //nolint:errcheck + return err + } + } } else { _, err = co.wr.AddTransceiverFromKind(webrtc.RTPCodecTypeVideo, webrtc.RTPTransceiverInit{ Direction: webrtc.RTPTransceiverDirectionRecvonly, diff --git a/internal/protocols/webrtc/peer_connection_test.go b/internal/protocols/webrtc/peer_connection_test.go index 445c6ad9..dcf56d4f 100644 --- a/internal/protocols/webrtc/peer_connection_test.go +++ b/internal/protocols/webrtc/peer_connection_test.go @@ -593,3 +593,63 @@ func TestPeerConnectionFallbackCodecs(t *testing.T) { }, }, s.MediaDescriptions) } + +func TestPeerConnectionPublishDataChannel(t *testing.T) { + pc1, err := webrtc.NewPeerConnection(webrtc.Configuration{}) + require.NoError(t, err) + defer pc1.Close() //nolint:errcheck + + _, err = pc1.CreateDataChannel("", nil) + require.NoError(t, err) + + dataChanCreated := make(chan struct{}) + dataReceived := make(chan struct{}) + + pc1.OnDataChannel(func(dc *webrtc.DataChannel) { + close(dataChanCreated) + + dc.OnMessage(func(msg webrtc.DataChannelMessage) { + require.Equal(t, []byte("test data"), msg.Data) + close(dataReceived) + }) + }) + + offer, err := pc1.CreateOffer(nil) + require.NoError(t, err) + + err = pc1.SetLocalDescription(offer) + require.NoError(t, err) + + pc2 := &PeerConnection{ + LocalRandomUDP: true, + IPsFromInterfaces: true, + HandshakeTimeout: conf.Duration(10 * time.Second), + TrackGatherTimeout: conf.Duration(2 * time.Second), + STUNGatherTimeout: conf.Duration(5 * time.Second), + Publish: true, + OutgoingDataChannels: []*OutgoingDataChannel{ + { + Label: "test-channel", + }, + }, + Log: test.NilLogger, + } + err = pc2.Start() + require.NoError(t, err) + defer pc2.Close() + + answer, err := pc2.CreateFullAnswer(&offer) + require.NoError(t, err) + + err = pc1.SetRemoteDescription(*answer) + require.NoError(t, err) + + err = pc2.WaitUntilConnected() + require.NoError(t, err) + + <-dataChanCreated + + pc2.OutgoingDataChannels[0].Write([]byte("test data")) + + <-dataReceived +} diff --git a/internal/servers/webrtc/reader.js b/internal/servers/webrtc/reader.js index 3048bb1f..d6b4348b 100644 --- a/internal/servers/webrtc/reader.js +++ b/internal/servers/webrtc/reader.js @@ -10,6 +10,11 @@ * @param {RTCTrackEvent} evt - track event. */ +/** + * @callback OnDataChannel + * @param {RTCDataChannelEvent} evt - data channel event. + */ + /** * @typedef Conf * @type {object} @@ -19,6 +24,7 @@ * @property {string} token - token. * @property {OnError} onError - called when there's an error. * @property {OnTrack} onTrack - called when there's a track available. + * @property {OnDataChannel} onDataChannel - called when there's a data channel available. */ /** WebRTC/WHEP reader. */ @@ -442,9 +448,13 @@ class MediaMTXWebRTCReader { this.pc.addTransceiver('video', { direction }); this.pc.addTransceiver('audio', { direction }); + // using data channels requires creating a data channel locally + this.pc.createDataChannel(''); + this.pc.onicecandidate = (evt) => this.#onLocalCandidate(evt); this.pc.onconnectionstatechange = () => this.#onConnectionState(); this.pc.ontrack = (evt) => this.#onTrack(evt); + this.pc.ondatachannel = (evt) => this.#onDataChannel(evt); return this.pc.createOffer() .then((offer) => { @@ -567,6 +577,12 @@ class MediaMTXWebRTCReader { this.conf.onTrack(evt); } } + + #onDataChannel(evt) { + if (this.conf.onDataChannel !== undefined) { + this.conf.onDataChannel(evt); + } + } } window.MediaMTXWebRTCReader = MediaMTXWebRTCReader;