diff --git a/pkg/rtc/room_test.go b/pkg/rtc/room_test.go index 4f76d0705..8ccd9574f 100644 --- a/pkg/rtc/room_test.go +++ b/pkg/rtc/room_test.go @@ -51,7 +51,11 @@ const ( ) func init() { - config.InitLoggerFromConfig(&config.DefaultConfig.Logging) + // The logger retains this config and locks it on every component level + // resolution, so it must not be DefaultConfig: NewConfig marshals that + // global from the goroutines these tests spin up. + loggingConf := config.LoggingConfig{PionLevel: config.DefaultConfig.Logging.PionLevel} + config.InitLoggerFromConfig(&loggingConf) roomUpdateInterval = defaultDelay } diff --git a/test/client/client.go b/test/client/client.go index 7a74c4841..14fd108a2 100644 --- a/test/client/client.go +++ b/test/client/client.go @@ -94,10 +94,12 @@ type RTCClient struct { // remote tracks waiting to be processed pendingRemoteTracks []*webrtc.TrackRemote - pendingTrackWriters []TrackWriter - OnConnected func() - OnDataReceived func(data []byte, sid string) - OnDataUnlabeledReceived func(data []byte) + pendingTrackWriters []TrackWriter + // callbacks are set by tests while the client is already running, and read + // from the signal and data channel goroutines + onConnected atomic.Pointer[func()] + onDataReceived atomic.Pointer[func(data []byte, sid string)] + onDataUnlabeledReceived atomic.Pointer[func(data []byte)] refreshToken string // map of livekit.ParticipantID and last packet @@ -346,8 +348,8 @@ func (c *RTCClient) createTransport(rtcconf webrtc.Configuration) error { } c.pendingDataTrackWriters = nil - if c.OnConnected != nil { - go c.OnConnected() + if f := c.onConnected.Load(); f != nil { + go (*f)() } }) publisherHandler.OnOfferCalls(c.onOffer) @@ -440,8 +442,8 @@ func (c *RTCClient) createTransport(rtcconf webrtc.Configuration) error { } c.pendingDataTrackWriters = nil - if c.OnConnected != nil { - go c.OnConnected() + if f := c.onConnected.Load(); f != nil { + go (*f)() } }) subscriberHandler.OnFullyEstablishedCalls(func() { @@ -472,6 +474,21 @@ func (c *RTCClient) ID() livekit.ParticipantID { return c.id } +// SetOnConnected is safe to call after the client is running. +func (c *RTCClient) SetOnConnected(f func()) { + c.onConnected.Store(&f) +} + +// SetOnDataReceived is safe to call after the client is running. +func (c *RTCClient) SetOnDataReceived(f func(data []byte, sid string)) { + c.onDataReceived.Store(&f) +} + +// SetOnDataUnlabeledReceived is safe to call after the client is running. +func (c *RTCClient) SetOnDataUnlabeledReceived(f func(data []byte)) { + c.onDataUnlabeledReceived.Store(&f) +} + // create an offer for the server func (c *RTCClient) Run() error { c.conn.SetCloseHandler(func(code int, text string) error { @@ -1143,15 +1160,15 @@ func (c *RTCClient) handleDataMessage(kind livekit.DataPacket_Kind, data []byte) } dp.Kind = kind if val, ok := dp.Value.(*livekit.DataPacket_User); ok { - if c.OnDataReceived != nil { - c.OnDataReceived(val.User.Payload, val.User.ParticipantSid) + if f := c.onDataReceived.Load(); f != nil { + (*f)(val.User.Payload, val.User.ParticipantSid) } } } func (c *RTCClient) handleDataMessageUnlabeled(data []byte) { - if c.OnDataUnlabeledReceived != nil { - c.OnDataUnlabeledReceived(data) + if f := c.onDataUnlabeledReceived.Load(); f != nil { + (*f)(data) } } diff --git a/test/integration_helpers.go b/test/integration_helpers.go index ec3469620..508d6ee55 100644 --- a/test/integration_helpers.go +++ b/test/integration_helpers.go @@ -63,7 +63,11 @@ const ( var roomClient livekit.RoomService func init() { - config.InitLoggerFromConfig(&config.DefaultConfig.Logging) + // The logger retains this config and locks it on every component level + // resolution, so it must not be DefaultConfig: NewConfig marshals that + // global from the goroutines these tests spin up. + loggingConf := config.LoggingConfig{PionLevel: config.DefaultConfig.Logging.PionLevel} + config.InitLoggerFromConfig(&loggingConf) prometheus.Init("test", livekit.NodeType_SERVER) } diff --git a/test/scenarios.go b/test/scenarios.go index aee71feb1..809c735b7 100644 --- a/test/scenarios.go +++ b/test/scenarios.go @@ -157,11 +157,11 @@ func scenarioDataPublish(t *testing.T) { payload := "test bytes" received := atomic.NewBool(false) - c2.OnDataReceived = func(data []byte, sid string) { + c2.SetOnDataReceived(func(data []byte, sid string) { if string(data) == payload && livekit.ParticipantID(sid) == c1.ID() { received.Store(true) } - } + }) require.NoError(t, c1.PublishData([]byte(payload), livekit.DataPacket_RELIABLE)) @@ -187,11 +187,11 @@ func scenarioDataUnlabeledPublish(t *testing.T) { payload := "test unlabeled bytes" received := atomic.NewBool(false) - c2.OnDataReceived = func(data []byte, _sid string) { + c2.SetOnDataReceived(func(data []byte, _sid string) { if string(data) == payload { received.Store(true) } - } + }) require.NoError(t, c1.PublishDataUnlabeled([]byte(payload))) diff --git a/test/singlenode_test.go b/test/singlenode_test.go index 08f67fd00..14964d12d 100644 --- a/test/singlenode_test.go +++ b/test/singlenode_test.go @@ -1100,30 +1100,30 @@ func TestDataPublishSlowSubscriber(t *testing.T) { // no data should be dropped for fast subscriber var fastDataIndex atomic.Uint64 - fastSub.OnDataReceived = func(data []byte, sid string) { + fastSub.SetOnDataReceived(func(data []byte, sid string) { idx := binary.BigEndian.Uint64(data[len(data)-8:]) require.Equal(t, fastDataIndex.Load()+1, idx) fastDataIndex.Store(idx) - } + }) // no data should be dropped for slow subscriber that is above threshold var slowNoDropDataIndex atomic.Uint64 var drainSlowSubNotDrop atomic.Bool slowNoDropReader := testclient.NewDataChannelReader(dataChannelSlowThreshold * 2) - slowSubNotDrop.OnDataReceived = func(data []byte, sid string) { + slowSubNotDrop.SetOnDataReceived(func(data []byte, sid string) { idx := binary.BigEndian.Uint64(data[len(data)-8:]) require.Equal(t, slowNoDropDataIndex.Load()+1, idx) slowNoDropDataIndex.Store(idx) if !drainSlowSubNotDrop.Load() { slowNoDropReader.Read(data, sid) } - } + }) // data should be dropped for slow subscriber that is below threshold var slowDropDataIndex atomic.Uint64 dropped := make(chan struct{}) slowDropReader := testclient.NewDataChannelReader(dataChannelSlowThreshold / 2) - slowSubDrop.OnDataReceived = func(data []byte, sid string) { + slowSubDrop.SetOnDataReceived(func(data []byte, sid string) { select { case <-dropped: return @@ -1135,7 +1135,7 @@ func TestDataPublishSlowSubscriber(t *testing.T) { } slowDropDataIndex.Store(idx) slowDropReader.Read(data, sid) - } + }) // publisher sends data as fast as possible, it will block by the slowest subscriber above the slow threshold var (