Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
61 changes: 59 additions & 2 deletions localtrack.go
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,25 @@ type SampleWriteOptions struct {
AudioLevel *uint8
}

// SampleWriteResult describes one attempt to write a sample. Timestamps are
// captured around WriteSample. Skipped is true when the track was muted or
// disabled and no write was attempted.
type SampleWriteResult struct {
ReadCompletedAt time.Time
WriteStartedAt time.Time
WriteCompletedAt time.Time
Err error
Skipped bool
}

// SampleObserver observes the synchronous sample read/write/pacing lifecycle.
// Implementations must return quickly and must not block the track write loop.
type SampleObserver interface {
OnSampleRead(sample media.Sample, readCompletedAt time.Time)
OnSampleWriteComplete(sample media.Sample, result SampleWriteResult)
OnSamplePacingLag(sample media.Sample, lag time.Duration)
}

// LocalTrack is a local track that simplifies writing samples.
// It handles timing and publishing of things, so as long as a SampleProvider is provided, the class takes care of
// publishing tracks at the right frequency
Expand All @@ -81,6 +100,7 @@ type LocalTrack struct {
simulcastID string
videoLayer *livekit.VideoLayer
onRTCP func(rtcp.Packet)
sampleObserver SampleObserver

muted atomic.Bool
disabled atomic.Bool
Expand Down Expand Up @@ -125,6 +145,13 @@ func WithRTCPHandler(cb func(rtcp.Packet)) LocalTrackOptions {
}
}

// WithSampleObserver attaches a non-blocking observer to the sample write loop.
func WithSampleObserver(observer SampleObserver) LocalTrackOptions {
return func(s *LocalTrack) {
s.sampleObserver = observer
}
}

func NewLocalTrack(c webrtc.RTPCodecCapability, opts ...LocalTrackOptions) (*LocalTrack, error) {
s := &LocalTrack{log: logger}
for _, o := range opts {
Expand Down Expand Up @@ -659,6 +686,9 @@ func (s *LocalTrack) writeWorker(provider SampleProvider, onComplete func()) {
defer close(writeClosed)

audioProvider, isAudioProvider := provider.(AudioSampleProvider)
s.lock.RLock()
observer := s.sampleObserver
s.lock.RUnlock()

nextSampleTime := time.Now()

Expand All @@ -675,6 +705,10 @@ func (s *LocalTrack) writeWorker(provider SampleProvider, onComplete func()) {
s.log.Errorw("could not get sample from provider", err)
return
}
readCompletedAt := time.Now()
if observer != nil {
observer.OnSampleRead(sample, readCompletedAt)
}

if !s.muted.Load() && !s.disabled.Load() {
var opts *SampleWriteOptions
Expand All @@ -686,15 +720,38 @@ func (s *LocalTrack) writeWorker(provider SampleProvider, onComplete func()) {
}

sample.Timestamp = nextSampleTime
if err := s.WriteSample(sample, opts); err != nil {
s.log.Errorw("could not write sample", err)
writeStartedAt := time.Now()
writeErr := s.WriteSample(sample, opts)
writeCompletedAt := time.Now()
if observer != nil {
observer.OnSampleWriteComplete(sample, SampleWriteResult{
ReadCompletedAt: readCompletedAt,
WriteStartedAt: writeStartedAt,
WriteCompletedAt: writeCompletedAt,
Err: writeErr,
})
}
if writeErr != nil {
s.log.Errorw("could not write sample", writeErr)
return
}
} else if observer != nil {
observer.OnSampleWriteComplete(sample, SampleWriteResult{
ReadCompletedAt: readCompletedAt,
Skipped: true,
})
}

// account for clock drift
nextSampleTime = nextSampleTime.Add(sample.Duration)
sleepDuration := time.Until(nextSampleTime)
if observer != nil {
lag := -sleepDuration
if lag < 0 {
lag = 0
}
observer.OnSamplePacingLag(sample, lag)
}
if sleepDuration <= 0 {
continue
}
Expand Down
8 changes: 8 additions & 0 deletions readersampleprovider.go
Original file line number Diff line number Diff line change
Expand Up @@ -132,6 +132,14 @@ func ReaderTrackWithRTCPHandler(f func(rtcp.Packet)) func(provider *ReaderSample
}
}

// ReaderTrackWithSampleObserver observes sample reads, writes, and pacing lag.
// The observer is invoked synchronously and must not block.
func ReaderTrackWithSampleObserver(observer SampleObserver) func(provider *ReaderSampleProvider) {
return func(provider *ReaderSampleProvider) {
provider.trackOpts = append(provider.trackOpts, WithSampleObserver(observer))
}
}

func ReaderTrackWithSampleOptions(opts ...LocalTrackOptions) func(provider *ReaderSampleProvider) {
return func(provider *ReaderSampleProvider) {
provider.trackOpts = append(provider.trackOpts, opts...)
Expand Down
94 changes: 94 additions & 0 deletions sample_observer_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,94 @@
package lksdk

import (
"context"
"io"
"sync"
"testing"
"time"

"github.com/pion/webrtc/v4/pkg/media"
)

type observerTestProvider struct {
sample media.Sample
done bool
}

func (p *observerTestProvider) OnBind() error { return nil }
func (p *observerTestProvider) OnUnbind() error { return nil }
func (p *observerTestProvider) Close() error { return nil }
func (p *observerTestProvider) NextSample(context.Context) (media.Sample, error) {
if p.done {
return media.Sample{}, io.EOF
}
p.done = true
return p.sample, nil
}

type observerTestRecorder struct {
mu sync.Mutex
reads int
writes []SampleWriteResult
lags []time.Duration
}

func (o *observerTestRecorder) OnSampleRead(media.Sample, time.Time) {
o.mu.Lock()
defer o.mu.Unlock()
o.reads++
}
func (o *observerTestRecorder) OnSampleWriteComplete(_ media.Sample, result SampleWriteResult) {
o.mu.Lock()
defer o.mu.Unlock()
o.writes = append(o.writes, result)
}
func (o *observerTestRecorder) OnSamplePacingLag(_ media.Sample, lag time.Duration) {
o.mu.Lock()
defer o.mu.Unlock()
o.lags = append(o.lags, lag)
}

func TestWriteWorkerNotifiesSampleObserver(t *testing.T) {
recorder := &observerTestRecorder{}
track := &LocalTrack{sampleObserver: recorder}
done := make(chan struct{})
track.writeWorker(&observerTestProvider{sample: media.Sample{Data: []byte{1}, Duration: 0}}, func() { close(done) })
select {
case <-done:
case <-time.After(time.Second):
t.Fatal("write worker did not complete")
}
if recorder.reads != 1 {
t.Fatalf("reads = %d, want 1", recorder.reads)
}
if len(recorder.writes) != 1 || recorder.writes[0].Err != nil || recorder.writes[0].Skipped {
t.Fatalf("unexpected write result: %+v", recorder.writes)
}
if len(recorder.lags) != 1 {
t.Fatalf("lags = %d, want 1", len(recorder.lags))
}
}

func TestWriteWorkerReportsMutedSampleAsSkipped(t *testing.T) {
recorder := &observerTestRecorder{}
track := &LocalTrack{sampleObserver: recorder}
track.muted.Store(true)
track.writeWorker(&observerTestProvider{sample: media.Sample{Data: []byte{1}}}, nil)
if len(recorder.writes) != 1 || !recorder.writes[0].Skipped {
t.Fatalf("unexpected write result: %+v", recorder.writes)
}
}

func TestReaderTrackWithSampleObserverAddsTrackOption(t *testing.T) {
recorder := &observerTestRecorder{}
provider := &ReaderSampleProvider{}
ReaderTrackWithSampleObserver(recorder)(provider)
track := &LocalTrack{}
for _, option := range provider.trackOpts {
option(track)
}
if track.sampleObserver != recorder {
t.Fatal("sample observer was not applied to track")
}
}
Loading