From 419cd7e8a9667776aa088876cb8f9cf5ae42a135 Mon Sep 17 00:00:00 2001 From: Haffi Mazhar Date: Wed, 19 Aug 2026 11:20:38 +0100 Subject: [PATCH] add sample observer interface --- localtrack.go | 61 +++++++++++++++++++++++++- readersampleprovider.go | 8 ++++ sample_observer_test.go | 94 +++++++++++++++++++++++++++++++++++++++++ 3 files changed, 161 insertions(+), 2 deletions(-) create mode 100644 sample_observer_test.go diff --git a/localtrack.go b/localtrack.go index c034aee8..2686c05b 100644 --- a/localtrack.go +++ b/localtrack.go @@ -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 @@ -81,6 +100,7 @@ type LocalTrack struct { simulcastID string videoLayer *livekit.VideoLayer onRTCP func(rtcp.Packet) + sampleObserver SampleObserver muted atomic.Bool disabled atomic.Bool @@ -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 { @@ -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() @@ -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 @@ -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 } diff --git a/readersampleprovider.go b/readersampleprovider.go index 7f1bde1c..9ec10d36 100644 --- a/readersampleprovider.go +++ b/readersampleprovider.go @@ -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...) diff --git a/sample_observer_test.go b/sample_observer_test.go new file mode 100644 index 00000000..7bd622df --- /dev/null +++ b/sample_observer_test.go @@ -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") + } +}