From 9b571875ce2b43f4205ec78886037ea57d79b060 Mon Sep 17 00:00:00 2001 From: David Chen Date: Tue, 16 Jun 2026 12:43:40 -0700 Subject: [PATCH] reliable data track --- go.mod | 2 + go.sum | 2 - pkg/rtc/clientinfo.go | 4 ++ pkg/rtc/datadowntrack.go | 10 ++- pkg/rtc/datatrack.go | 14 +++++ pkg/rtc/datatrack_stats.go | 56 ++++++++++++----- pkg/rtc/datatrack_test.go | 33 ++++++++++ pkg/rtc/errors.go | 11 ++-- pkg/rtc/participant_data_track.go | 21 +++++-- pkg/rtc/subscriptionmanager.go | 4 ++ pkg/rtc/transport.go | 51 +++++++++++---- pkg/rtc/transportmanager.go | 43 +++++++++++-- pkg/rtc/types/interfaces.go | 3 +- pkg/rtc/types/typesfakes/fake_data_track.go | 63 +++++++++++++++++++ .../typesfakes/fake_data_track_transport.go | 19 +++--- 15 files changed, 283 insertions(+), 53 deletions(-) diff --git a/go.mod b/go.mod index cd376e81e..8cebb8a57 100644 --- a/go.mod +++ b/go.mod @@ -153,3 +153,5 @@ require ( google.golang.org/genproto/googleapis/rpc v0.0.0-20260526163538-3dc84a4a5aaa // indirect google.golang.org/grpc v1.81.1 // indirect ) + +replace github.com/livekit/protocol => ../protocol-local diff --git a/go.sum b/go.sum index 454ef37c0..fd3b05834 100644 --- a/go.sum +++ b/go.sum @@ -160,8 +160,6 @@ github.com/livekit/mageutil v0.0.0-20250511045019-0f1ff63f7731 h1:9x+U2HGLrSw5AT github.com/livekit/mageutil v0.0.0-20250511045019-0f1ff63f7731/go.mod h1:Rs3MhFwutWhGwmY1VQsygw28z5bWcnEYmS1OG9OxjOQ= github.com/livekit/mediatransportutil v0.0.0-20260608063931-a3417d38cda0 h1:XHNNzebIKZRkLimla/hFGrAIX5EMWHctrgt3hLw7s+I= github.com/livekit/mediatransportutil v0.0.0-20260608063931-a3417d38cda0/go.mod h1:o8CFmAdrVwzJNOCsQCLUzXRjokkufNshnQHOe4fRaqU= -github.com/livekit/protocol v1.46.7-0.20260611165352-04a0fe5b5051 h1:IYqiW7z5pblZBn6o0OHNz8MHd3wJ/TLJG4gh6lCI0/s= -github.com/livekit/protocol v1.46.7-0.20260611165352-04a0fe5b5051/go.mod h1:jO+y05AU9Ec4JswDyuzKCZ4bhziOS0CzMqgnbj60Dzs= github.com/livekit/psrpc v0.7.2 h1:6oZ+NODJ2pLyaT6VqDq1F4Qc/3TpDUSpyphj/P9MhQc= github.com/livekit/psrpc v0.7.2/go.mod h1:rAI+m2+/cb4x9RXhLRtUx5ZwdfjjXOl4zi46IjEetaw= github.com/mackerelio/go-osstat v0.2.7 h1:TCavZi10wF49bT6iQZ9eT2keGZQpC69MTDfdJej5e94= diff --git a/pkg/rtc/clientinfo.go b/pkg/rtc/clientinfo.go index 899d60a70..9cf36c35e 100644 --- a/pkg/rtc/clientinfo.go +++ b/pkg/rtc/clientinfo.go @@ -130,6 +130,10 @@ func (c ClientInfo) SupportsPacketTrailer() bool { return c.HasCapability(livekit.ClientInfo_CAP_PACKET_TRAILER) } +func (c ClientInfo) SupportsReliableDataTrack() bool { + return c.HasCapability(livekit.ClientInfo_CAP_RELIABLE_DATA_TRACK) +} + // compareVersion compares a semver against the current client SDK version // returning 1 if current version is greater than version // 0 if they are the same, and -1 if it's an earlier version diff --git a/pkg/rtc/datadowntrack.go b/pkg/rtc/datadowntrack.go index 45f5e19a8..612978cdc 100644 --- a/pkg/rtc/datadowntrack.go +++ b/pkg/rtc/datadowntrack.go @@ -123,7 +123,7 @@ func (d *DataDownTrack) WritePacket(data []byte, packet *datatrack.Packet, _arri d.logger.Warnw("could not marshal data track message", err) return } - if err := d.params.Transport.SendDataTrackMessage(buf); err != nil { + if err := d.params.Transport.SendDataTrackMessage(buf, d.params.PublishDataTrack.Reliability()); err != nil { d.logger.Warnw("could not send data track message", err) return } @@ -133,5 +133,11 @@ func (d *DataDownTrack) WritePacket(data []byte, packet *datatrack.Packet, _arri } func (d *DataDownTrack) UpdateSubscriptionOptions(subscriptionOptions *livekit.DataTrackSubscriptionOptions) { - // DT-TODO + if subscriptionOptions == nil { + return + } + if subscriptionOptions.GetReliability() == livekit.DataTrackReliability_DTR_RELIABLE && + d.params.PublishDataTrack.Reliability() == livekit.DataTrackReliability_DTR_LOSSY { + d.logger.Warnw("subscriber requested reliable delivery for lossy data track", nil) + } } diff --git a/pkg/rtc/datatrack.go b/pkg/rtc/datatrack.go index fd22f6258..993df6a5c 100644 --- a/pkg/rtc/datatrack.go +++ b/pkg/rtc/datatrack.go @@ -113,6 +113,10 @@ func (d *DataTrack) Name() string { return d.dti.Name } +func (d *DataTrack) Reliability() livekit.DataTrackReliability { + return d.dti.GetReliability() +} + func (d *DataTrack) AddSubscriber(sub types.LocalParticipant) (types.DataDownTrack, error) { d.lock.Lock() defer d.lock.Unlock() @@ -120,6 +124,16 @@ func (d *DataTrack) AddSubscriber(sub types.LocalParticipant) (types.DataDownTra if _, ok := d.subscribedTracks[sub.ID()]; ok { return nil, errAlreadySubscribed } + subscriberInfo := ClientInfo{sub.GetClientInfo()} + if d.Reliability() == livekit.DataTrackReliability_DTR_RELIABLE && !subscriberInfo.SupportsReliableDataTrack() { + d.logger.Warnw( + "subscriber does not support reliable data tracks", + nil, + "subscriberID", sub.ID(), + "subscriberIdentity", sub.Identity(), + ) + return nil, ErrReliableDataTrackUnsupported + } bytesStats := NewBytesTrackStats( sub.GetCountry(), diff --git a/pkg/rtc/datatrack_stats.go b/pkg/rtc/datatrack_stats.go index 5bde6f57f..f3f31b3b8 100644 --- a/pkg/rtc/datatrack_stats.go +++ b/pkg/rtc/datatrack_stats.go @@ -15,6 +15,8 @@ package rtc import ( + "os" + "strconv" "sync" "time" @@ -33,6 +35,8 @@ type dataTrackStats struct { lock sync.Mutex startTime int64 endTime int64 + reportInterval time.Duration + lastReportTime int64 highestSequenceNumber uint16 numPackets int numPacketsLost int @@ -42,8 +46,16 @@ type dataTrackStats struct { } func newDataTrackStats(params dataTrackStatsParams) *dataTrackStats { + reportInterval := time.Duration(0) + if raw := os.Getenv("LIVEKIT_DATA_TRACK_STATS_INTERVAL_MS"); raw != "" { + if intervalMs, err := strconv.Atoi(raw); err == nil && intervalMs > 0 { + reportInterval = time.Duration(intervalMs) * time.Millisecond + } + } + return &dataTrackStats{ - params: params, + params: params, + reportInterval: reportInterval, } } @@ -84,6 +96,11 @@ func (d *dataTrackStats) Update(packet *datatrack.Packet, arrivalTime int64, pay if packet.IsFinalOfFrame { d.numFrames++ } + + if d.reportInterval > 0 && arrivalTime-d.lastReportTime >= int64(d.reportInterval) { + d.lastReportTime = arrivalTime + d.logLocked("data track stats sample", arrivalTime) + } } func (d *dataTrackStats) Close() { @@ -93,18 +110,29 @@ func (d *dataTrackStats) Close() { d.endTime = mono.UnixNano() if d.startTime != 0 { - duration := time.Duration(d.endTime - d.startTime).Seconds() - fps := float64(d.numFrames) / duration - - d.params.Logger.Infow( - "data track stats", - "duration", duration, - "numPackets", d.numPackets, - "numPacketsLost", d.numPacketsLost, - "numPacketsOutOfOrder", d.numPacketsOutOfOrder, - "numFrames", d.numFrames, - "fps", fps, - "numBytes", d.numBytes, - ) + d.logLocked("data track stats", d.endTime) } } + +func (d *dataTrackStats) logLocked(message string, now int64) { + if d.startTime == 0 { + return + } + + duration := time.Duration(now - d.startTime).Seconds() + fps := 0.0 + if duration > 0 { + fps = float64(d.numFrames) / duration + } + + d.params.Logger.Infow( + message, + "duration", duration, + "numPackets", d.numPackets, + "numPacketsLost", d.numPacketsLost, + "numPacketsOutOfOrder", d.numPacketsOutOfOrder, + "numFrames", d.numFrames, + "fps", fps, + "numBytes", d.numBytes, + ) +} diff --git a/pkg/rtc/datatrack_test.go b/pkg/rtc/datatrack_test.go index f2938a1f2..c9bba1939 100644 --- a/pkg/rtc/datatrack_test.go +++ b/pkg/rtc/datatrack_test.go @@ -69,3 +69,36 @@ func TestDataTrackRevokeDisallowedSubscribers(t *testing.T) { require.False(t, dt.IsSubscriber(disallowed.ID())) require.True(t, dt.IsSubscriber(recorder.ID())) } + +func TestDataTrackReliableRequiresSubscriberCapability(t *testing.T) { + dt := NewDataTrack( + DataTrackParams{ + Logger: logger.GetLogger(), + ParticipantID: func() livekit.ParticipantID { return "pubID" }, + ParticipantIdentity: "pub", + }, + &livekit.DataTrackInfo{ + PubHandle: 1, + Sid: "DTR_test", + Name: "test", + Reliability: livekit.DataTrackReliability_DTR_RELIABLE, + }, + ) + defer dt.Close() + + unsupported := newTestDataTrackSubscriber("oldID", "old", false) + unsupported.GetClientInfoReturns(&livekit.ClientInfo{}) + _, err := dt.AddSubscriber(unsupported) + require.ErrorIs(t, err, ErrReliableDataTrackUnsupported) + require.False(t, dt.IsSubscriber(unsupported.ID())) + + supported := newTestDataTrackSubscriber("newID", "new", false) + supported.GetClientInfoReturns(&livekit.ClientInfo{ + Capabilities: []livekit.ClientInfo_Capability{ + livekit.ClientInfo_CAP_RELIABLE_DATA_TRACK, + }, + }) + _, err = dt.AddSubscriber(supported) + require.NoError(t, err) + require.True(t, dt.IsSubscriber(supported.ID())) +} diff --git a/pkg/rtc/errors.go b/pkg/rtc/errors.go index bbe390683..24898180d 100644 --- a/pkg/rtc/errors.go +++ b/pkg/rtc/errors.go @@ -34,11 +34,12 @@ var ( ErrInternalError = errors.New("internal error") // Track subscription related - ErrNoTrackPermission = errors.New("participant is not allowed to subscribe to this track") - ErrNoSubscribePermission = errors.New("participant is not given permission to subscribe to tracks") - ErrTrackNotFound = errors.New("track cannot be found") - ErrTrackNotBound = errors.New("track not bound") - ErrSubscriptionLimitExceeded = errors.New("participant has exceeded its subscription limit") + ErrNoTrackPermission = errors.New("participant is not allowed to subscribe to this track") + ErrNoSubscribePermission = errors.New("participant is not given permission to subscribe to tracks") + ErrTrackNotFound = errors.New("track cannot be found") + ErrTrackNotBound = errors.New("track not bound") + ErrSubscriptionLimitExceeded = errors.New("participant has exceeded its subscription limit") + ErrReliableDataTrackUnsupported = errors.New("participant does not support reliable data tracks") ErrNoSubscribeMetricsPermission = errors.New("participant is not given permission to subscribe to metrics") ) diff --git a/pkg/rtc/participant_data_track.go b/pkg/rtc/participant_data_track.go index dbda1ede0..464e95016 100644 --- a/pkg/rtc/participant_data_track.go +++ b/pkg/rtc/participant_data_track.go @@ -60,6 +60,18 @@ func (p *ParticipantImpl) HandlePublishDataTrackRequest(req *livekit.PublishData return } + if req.GetReliability() == livekit.DataTrackReliability_DTR_RELIABLE && !p.params.ClientInfo.SupportsReliableDataTrack() { + p.pubLogger.Warnw("client does not support reliable data tracks", nil, "req", logger.Proto(req)) + p.sendRequestResponse(&livekit.RequestResponse{ + Reason: livekit.RequestResponse_NOT_ALLOWED, + Message: "client does not support reliable data tracks", + Request: &livekit.RequestResponse_PublishDataTrack{ + PublishDataTrack: utils.CloneProto(req), + }, + }) + return + } + publishedDataTracks := p.UpDataTrackManager.GetPublishedDataTracks() for _, dt := range publishedDataTracks { message := "" @@ -90,10 +102,11 @@ func (p *ParticipantImpl) HandlePublishDataTrackRequest(req *livekit.PublishData } dti := &livekit.DataTrackInfo{ - PubHandle: req.PubHandle, - Sid: guid.New(utils.DataTrackPrefix), - Name: req.Name, - Encryption: req.Encryption, + PubHandle: req.PubHandle, + Sid: guid.New(utils.DataTrackPrefix), + Name: req.Name, + Encryption: req.Encryption, + Reliability: req.GetReliability(), } dt := NewDataTrack( DataTrackParams{ diff --git a/pkg/rtc/subscriptionmanager.go b/pkg/rtc/subscriptionmanager.go index d421d0805..2a3f6196e 100644 --- a/pkg/rtc/subscriptionmanager.go +++ b/pkg/rtc/subscriptionmanager.go @@ -616,6 +616,10 @@ func (m *SubscriptionManager) reconcileDataTrackSubscription(s *dataTrackSubscri // these are errors that are outside of our control, so we'll keep trying // - ErrNoTrackPermission: publisher did not grant subscriber permission, may change any moment // - ErrNoSubscribePermission: participant was not granted canSubscribe, may change any moment + case ErrReliableDataTrackUnsupported: + s.logger.Warnw("unsubscribing from unsupported reliable data track", err) + s.setDesired(false) + m.queueReconcile(s.trackID) case ErrTrackNotFound: // source track was never published or closed // if after timeout we'd unsubscribe from it. diff --git a/pkg/rtc/transport.go b/pkg/rtc/transport.go index 0848ae137..9dfba9cb0 100644 --- a/pkg/rtc/transport.go +++ b/pkg/rtc/transport.go @@ -68,9 +68,10 @@ import ( ) const ( - LossyDataChannel = "_lossy" - ReliableDataChannel = "_reliable" - DataTrackDataChannel = "_data_track" + LossyDataChannel = "_lossy" + ReliableDataChannel = "_reliable" + DataTrackDataChannel = "_data_track" + ReliableDataTrackDataChannel = "_reliable_data_track" fastNegotiationFrequency = 10 * time.Millisecond negotiationFrequency = 150 * time.Millisecond @@ -111,6 +112,10 @@ var ( ErrICECredentialMismatch = errors.New("ice credential mismatch") ) +func isDataTrackDataChannel(label string) bool { + return label == DataTrackDataChannel || label == ReliableDataTrackDataChannel +} + // ------------------------------------------------------------------------- type signal int @@ -221,6 +226,7 @@ type PCTransport struct { lossyDC *datachannel.DataChannelWriter[*webrtc.DataChannel] lossyDCOpened bool dataTrackDC *datachannel.DataChannelWriter[*webrtc.DataChannel] + reliableDataTrackDC *datachannel.DataChannelWriter[*webrtc.DataChannel] unlabeledDataChannels []*datachannel.DataChannelWriter[*webrtc.DataChannel] iceStartedAt time.Time @@ -882,7 +888,7 @@ func (t *PCTransport) onDataChannel(dc *webrtc.DataChannel) { case LossyDataChannel: kind = livekit.DataPacket_LOSSY - case DataTrackDataChannel: + case DataTrackDataChannel, ReliableDataTrackDataChannel: isDataTrack = true default: @@ -910,10 +916,17 @@ func (t *PCTransport) onDataChannel(dc *webrtc.DataChannel) { t.params.Logger.Debugw("data tracks not enabled") isHandled = false } else { - if t.dataTrackDC != nil { - t.dataTrackDC.Close() + if dc.Label() == ReliableDataTrackDataChannel { + if t.reliableDataTrackDC != nil { + t.reliableDataTrackDC.Close() + } + t.reliableDataTrackDC = datachannel.NewDataChannelWriterReliable(dc, rawDC, t.params.DatachannelSlowThreshold) + } else { + if t.dataTrackDC != nil { + t.dataTrackDC.Close() + } + t.dataTrackDC = datachannel.NewDataChannelWriterUnreliable(dc, rawDC, 0, 0) } - t.dataTrackDC = datachannel.NewDataChannelWriterUnreliable(dc, rawDC, 0, 0) } case kind == livekit.DataPacket_RELIABLE: @@ -1228,7 +1241,7 @@ func (t *PCTransport) getNumUnmatchedTransceivers() (uint32, uint32) { } func (t *PCTransport) CreateDataChannel(label string, dci *webrtc.DataChannelInit) error { - if label == DataTrackDataChannel && !t.params.EnableDataTracks { + if isDataTrackDataChannel(label) && !t.params.EnableDataTracks { t.params.Logger.Debugw("data tracks not enabled") return nil } @@ -1262,6 +1275,10 @@ func (t *PCTransport) CreateDataChannel(label string, dci *webrtc.DataChannelIni case DataTrackDataChannel: dcPtr = &t.dataTrackDC isDataTrack = true + + case ReliableDataTrackDataChannel: + dcPtr = &t.reliableDataTrackDC + isDataTrack = true } dc.OnOpen(func() { @@ -1272,7 +1289,7 @@ func (t *PCTransport) CreateDataChannel(label string, dci *webrtc.DataChannelIni } var slowThreshold int - if dc.Label() == ReliableDataChannel || isUnlabeled { + if dc.Label() == ReliableDataChannel || dc.Label() == ReliableDataTrackDataChannel || isUnlabeled { slowThreshold = t.params.DatachannelSlowThreshold } @@ -1293,6 +1310,8 @@ func (t *PCTransport) CreateDataChannel(label string, dci *webrtc.DataChannelIni *dcPtr = datachannel.NewDataChannelWriterUnreliable(dc, rawDC, t.params.DatachannelLossyTargetLatency, uint64(lossyDataChannelMinBufferedAmount)) case dcPtr == &t.dataTrackDC: *dcPtr = datachannel.NewDataChannelWriterUnreliable(dc, rawDC, 0, 0) + case dcPtr == &t.reliableDataTrackDC: + *dcPtr = datachannel.NewDataChannelWriterReliable(dc, rawDC, slowThreshold) } if dcReady != nil { *dcReady = true @@ -1374,7 +1393,7 @@ func (t *PCTransport) CreateReadableDataChannel(label string, dci *webrtc.DataCh } func (t *PCTransport) CreateDataChannelIfEmpty(dcLabel string, dci *webrtc.DataChannelInit) (label string, id uint16, existing bool, err error) { - if dcLabel == DataTrackDataChannel && !t.params.EnableDataTracks { + if isDataTrackDataChannel(dcLabel) && !t.params.EnableDataTracks { t.params.Logger.Debugw("data tracks not enabled") err = errors.New("data tracks not enabled") return @@ -1389,6 +1408,8 @@ func (t *PCTransport) CreateDataChannelIfEmpty(dcLabel string, dci *webrtc.DataC dcw = t.lossyDC case DataTrackDataChannel: dcw = t.dataTrackDC + case ReliableDataTrackDataChannel: + dcw = t.reliableDataTrackDC default: t.params.Logger.Warnw("unknown data channel label", nil, "label", label) err = errors.New("unknown data channel label") @@ -1516,9 +1537,12 @@ func (t *PCTransport) SendDataMessageUnlabeled(data []byte, useRaw bool, sender return t.sendDataMessage(dc, data) } -func (t *PCTransport) SendDataTrackMessage(data []byte) error { +func (t *PCTransport) SendDataTrackMessage(data []byte, reliability livekit.DataTrackReliability) error { t.lock.RLock() dc := t.dataTrackDC + if reliability == livekit.DataTrackReliability_DTR_RELIABLE { + dc = t.reliableDataTrackDC + } t.lock.RUnlock() return t.sendDataMessage(dc, data) @@ -1583,6 +1607,11 @@ func (t *PCTransport) Close() { t.dataTrackDC = nil } + if t.reliableDataTrackDC != nil { + t.reliableDataTrackDC.Close() + t.reliableDataTrackDC = nil + } + for _, dc := range t.unlabeledDataChannels { dc.Close() } diff --git a/pkg/rtc/transportmanager.go b/pkg/rtc/transportmanager.go index 528cd72c7..21b1ef0af 100644 --- a/pkg/rtc/transportmanager.go +++ b/pkg/rtc/transportmanager.go @@ -396,9 +396,13 @@ func (t *TransportManager) handleSendDataResult(err error, kind string, size int func (t *TransportManager) createDataChannelsForSubscriber(pendingDataChannels []*livekit.DataChannelInfo) error { var ( - reliableID, lossyID, dataTrackID uint16 - reliableIDPtr, lossyIDPtr, dataTrackIDPtr *uint16 + reliableID, lossyID, dataTrackID, reliableDataTrackID uint16 + reliableIDPtr, lossyIDPtr, dataTrackIDPtr, reliableDataTrackIDPtr *uint16 ) + dataChannelIDOffset := uint16(6) + if t.params.ClientInfo.SupportsReliableDataTrack() { + dataChannelIDOffset = 8 + } // // For old version migration clients, they don't send subscriber data channel info @@ -410,15 +414,19 @@ func (t *TransportManager) createDataChannelsForSubscriber(pendingDataChannels [ for _, dc := range pendingDataChannels { switch dc.Label { case ReliableDataChannel: - // pion use step 2 for auto generated ID, so we need to add 6 to avoid conflict - reliableID = uint16(dc.Id) + 6 + // pion uses step 2 for auto generated IDs, so offset by the + // number of publisher data channels to avoid conflicts. + reliableID = uint16(dc.Id) + dataChannelIDOffset reliableIDPtr = &reliableID case LossyDataChannel: - lossyID = uint16(dc.Id) + 6 + lossyID = uint16(dc.Id) + dataChannelIDOffset lossyIDPtr = &lossyID case DataTrackDataChannel: - dataTrackID = uint16(dc.Id) + 6 + dataTrackID = uint16(dc.Id) + dataChannelIDOffset dataTrackIDPtr = &dataTrackID + case ReliableDataTrackDataChannel: + reliableDataTrackID = uint16(dc.Id) + dataChannelIDOffset + reliableDataTrackIDPtr = &reliableDataTrackID } } @@ -454,6 +462,18 @@ func (t *TransportManager) createDataChannelsForSubscriber(pendingDataChannels [ return err } + if t.params.ClientInfo.SupportsReliableDataTrack() { + ordered = true + negotiated = t.params.Migration && reliableDataTrackIDPtr == nil + if err := t.subscriber.CreateDataChannel(ReliableDataTrackDataChannel, &webrtc.DataChannelInit{ + Ordered: &ordered, + ID: reliableDataTrackIDPtr, + Negotiated: &negotiated, + }); err != nil { + return err + } + } + return nil } @@ -925,6 +945,17 @@ func (t *TransportManager) ProcessPendingPublisherDataChannels() { Negotiated: &negotiated, ID: &id, }) + case ReliableDataTrackDataChannel: + if !t.params.ClientInfo.SupportsReliableDataTrack() { + continue + } + ordered = true + id := uint16(ci.GetId()) + dcLabel, dcID, dcExisting, err = t.publisher.CreateDataChannelIfEmpty(ReliableDataTrackDataChannel, &webrtc.DataChannelInit{ + Ordered: &ordered, + Negotiated: &negotiated, + ID: &id, + }) } if err != nil { t.params.Logger.Errorw("create migrated data channel failed", err, "label", ci.Label) diff --git a/pkg/rtc/types/interfaces.go b/pkg/rtc/types/interfaces.go index 194f529dd..ccebd936f 100644 --- a/pkg/rtc/types/interfaces.go +++ b/pkg/rtc/types/interfaces.go @@ -818,6 +818,7 @@ type DataTrack interface { ID() livekit.TrackID PubHandle() uint16 Name() string + Reliability() livekit.DataTrackReliability ToProto() *livekit.DataTrackInfo PublisherID() livekit.ParticipantID @@ -856,7 +857,7 @@ type DataTrackSender interface { //counterfeiter:generate . DataTrackTransport type DataTrackTransport interface { - SendDataTrackMessage(data []byte) error + SendDataTrackMessage(data []byte, reliability livekit.DataTrackReliability) error } //counterfeiter:generate . SubscribedTrack diff --git a/pkg/rtc/types/typesfakes/fake_data_track.go b/pkg/rtc/types/typesfakes/fake_data_track.go index 07a49f9b1..4d286ec5f 100644 --- a/pkg/rtc/types/typesfakes/fake_data_track.go +++ b/pkg/rtc/types/typesfakes/fake_data_track.go @@ -91,6 +91,16 @@ type FakeDataTrack struct { pubHandleReturnsOnCall map[int]struct { result1 uint16 } + ReliabilityStub func() livekit.DataTrackReliability + reliabilityMutex sync.RWMutex + reliabilityArgsForCall []struct { + } + reliabilityReturns struct { + result1 livekit.DataTrackReliability + } + reliabilityReturnsOnCall map[int]struct { + result1 livekit.DataTrackReliability + } PublisherIDStub func() livekit.ParticipantID publisherIDMutex sync.RWMutex publisherIDArgsForCall []struct { @@ -581,6 +591,59 @@ func (fake *FakeDataTrack) PubHandleReturnsOnCall(i int, result1 uint16) { }{result1} } +func (fake *FakeDataTrack) Reliability() livekit.DataTrackReliability { + fake.reliabilityMutex.Lock() + ret, specificReturn := fake.reliabilityReturnsOnCall[len(fake.reliabilityArgsForCall)] + fake.reliabilityArgsForCall = append(fake.reliabilityArgsForCall, struct { + }{}) + stub := fake.ReliabilityStub + fakeReturns := fake.reliabilityReturns + fake.recordInvocation("Reliability", []interface{}{}) + fake.reliabilityMutex.Unlock() + if stub != nil { + return stub() + } + if specificReturn { + return ret.result1 + } + return fakeReturns.result1 +} + +func (fake *FakeDataTrack) ReliabilityCallCount() int { + fake.reliabilityMutex.RLock() + defer fake.reliabilityMutex.RUnlock() + return len(fake.reliabilityArgsForCall) +} + +func (fake *FakeDataTrack) ReliabilityCalls(stub func() livekit.DataTrackReliability) { + fake.reliabilityMutex.Lock() + defer fake.reliabilityMutex.Unlock() + fake.ReliabilityStub = stub +} + +func (fake *FakeDataTrack) ReliabilityReturns(result1 livekit.DataTrackReliability) { + fake.reliabilityMutex.Lock() + defer fake.reliabilityMutex.Unlock() + fake.ReliabilityStub = nil + fake.reliabilityReturns = struct { + result1 livekit.DataTrackReliability + }{result1} +} + +func (fake *FakeDataTrack) ReliabilityReturnsOnCall(i int, result1 livekit.DataTrackReliability) { + fake.reliabilityMutex.Lock() + defer fake.reliabilityMutex.Unlock() + fake.ReliabilityStub = nil + if fake.reliabilityReturnsOnCall == nil { + fake.reliabilityReturnsOnCall = make(map[int]struct { + result1 livekit.DataTrackReliability + }) + } + fake.reliabilityReturnsOnCall[i] = struct { + result1 livekit.DataTrackReliability + }{result1} +} + func (fake *FakeDataTrack) PublisherID() livekit.ParticipantID { fake.publisherIDMutex.Lock() ret, specificReturn := fake.publisherIDReturnsOnCall[len(fake.publisherIDArgsForCall)] diff --git a/pkg/rtc/types/typesfakes/fake_data_track_transport.go b/pkg/rtc/types/typesfakes/fake_data_track_transport.go index d93c84f75..257430930 100644 --- a/pkg/rtc/types/typesfakes/fake_data_track_transport.go +++ b/pkg/rtc/types/typesfakes/fake_data_track_transport.go @@ -5,13 +5,15 @@ import ( "sync" "github.com/livekit/livekit-server/pkg/rtc/types" + "github.com/livekit/protocol/livekit" ) type FakeDataTrackTransport struct { - SendDataTrackMessageStub func([]byte) error + SendDataTrackMessageStub func([]byte, livekit.DataTrackReliability) error sendDataTrackMessageMutex sync.RWMutex sendDataTrackMessageArgsForCall []struct { arg1 []byte + arg2 livekit.DataTrackReliability } sendDataTrackMessageReturns struct { result1 error @@ -23,7 +25,7 @@ type FakeDataTrackTransport struct { invocationsMutex sync.RWMutex } -func (fake *FakeDataTrackTransport) SendDataTrackMessage(arg1 []byte) error { +func (fake *FakeDataTrackTransport) SendDataTrackMessage(arg1 []byte, arg2 livekit.DataTrackReliability) error { var arg1Copy []byte if arg1 != nil { arg1Copy = make([]byte, len(arg1)) @@ -33,13 +35,14 @@ func (fake *FakeDataTrackTransport) SendDataTrackMessage(arg1 []byte) error { ret, specificReturn := fake.sendDataTrackMessageReturnsOnCall[len(fake.sendDataTrackMessageArgsForCall)] fake.sendDataTrackMessageArgsForCall = append(fake.sendDataTrackMessageArgsForCall, struct { arg1 []byte - }{arg1Copy}) + arg2 livekit.DataTrackReliability + }{arg1Copy, arg2}) stub := fake.SendDataTrackMessageStub fakeReturns := fake.sendDataTrackMessageReturns - fake.recordInvocation("SendDataTrackMessage", []interface{}{arg1Copy}) + fake.recordInvocation("SendDataTrackMessage", []interface{}{arg1Copy, arg2}) fake.sendDataTrackMessageMutex.Unlock() if stub != nil { - return stub(arg1) + return stub(arg1, arg2) } if specificReturn { return ret.result1 @@ -53,17 +56,17 @@ func (fake *FakeDataTrackTransport) SendDataTrackMessageCallCount() int { return len(fake.sendDataTrackMessageArgsForCall) } -func (fake *FakeDataTrackTransport) SendDataTrackMessageCalls(stub func([]byte) error) { +func (fake *FakeDataTrackTransport) SendDataTrackMessageCalls(stub func([]byte, livekit.DataTrackReliability) error) { fake.sendDataTrackMessageMutex.Lock() defer fake.sendDataTrackMessageMutex.Unlock() fake.SendDataTrackMessageStub = stub } -func (fake *FakeDataTrackTransport) SendDataTrackMessageArgsForCall(i int) []byte { +func (fake *FakeDataTrackTransport) SendDataTrackMessageArgsForCall(i int) ([]byte, livekit.DataTrackReliability) { fake.sendDataTrackMessageMutex.RLock() defer fake.sendDataTrackMessageMutex.RUnlock() argsForCall := fake.sendDataTrackMessageArgsForCall[i] - return argsForCall.arg1 + return argsForCall.arg1, argsForCall.arg2 } func (fake *FakeDataTrackTransport) SendDataTrackMessageReturns(result1 error) {