Reorder remote stream updates in BaseChannel::UpdateRemoteStreams_w

Reorder the steps in UpdateRemoteStreams_w so that
RegisterRtpDemuxerSink_w is called before receive stream
removal/addition. Failures in demuxer registration now trigger early
exit before modifying media channel streams.

This prepares transport-level routing (demuxer) before media stream
creation/destruction. Stops demuxing removed streams before deletion in
media channel.

Bug: webrtc:42224170
Change-Id: Ib17d56db13fdb47fa5f565a424df192a9c1724d0
Reviewed-on: https://webrtc-review.googlesource.com/c/src/+/473780
Commit-Queue: Tomas Gunnarsson <tommi@webrtc.org>
Reviewed-by: Harald Alvestrand <hta@webrtc.org>
Cr-Commit-Position: refs/heads/main@{#47752}
diff --git a/pc/channel.cc b/pc/channel.cc
index 26af6a7..1231adc 100644
--- a/pc/channel.cc
+++ b/pc/channel.cc
@@ -536,12 +536,21 @@
 }
 
 bool BaseChannel::RegisterRtpDemuxerSink_w(
-    bool clear_payload_types,
-    std::optional<flat_set<uint32_t>> ssrcs) {
-  media_receive_channel()->OnDemuxerCriteriaUpdatePending();
-  absl::Cleanup cleanup = [this] {
-    media_receive_channel()->OnDemuxerCriteriaUpdateComplete();
-  };
+    const MediaContentDescription* content) {
+  bool clear_payload_types = false;
+  if (!RtpTransceiverDirectionHasSend(content->direction())) {
+    RTC_DLOG(LS_VERBOSE)
+        << "RegisterRtpDemuxerSink_w: remote side will not send "
+           "- disable payload type demuxing for "
+        << ToString();
+    clear_payload_types = true;
+  }
+
+  flat_set<uint32_t> ssrcs;
+  for (const StreamParams& new_stream : content->streams()) {
+    ssrcs.insert(new_stream.ssrcs.begin(), new_stream.ssrcs.end());
+  }
+
   bool ret = network_thread_->BlockingCall([&] {
     RTC_DCHECK_RUN_ON(network_thread());
     if (!rtp_transport_) {
@@ -558,11 +567,9 @@
       needs_re_registration = true;
     }
 
-    if (ssrcs) {
-      if (ssrcs_ != *ssrcs) {
-        ssrcs_ = std::move(*ssrcs);
-        needs_re_registration = true;
-      }
+    if (ssrcs_ != ssrcs) {
+      ssrcs_ = std::move(ssrcs);
+      needs_re_registration = true;
     }
 
     if (!needs_re_registration) {
@@ -957,18 +964,25 @@
     const MediaContentDescription* content,
     SdpType type) {
   RTC_LOG_THREAD_BLOCK_COUNT();
-  bool clear_payload_types = false;
-  if (!RtpTransceiverDirectionHasSend(content->direction())) {
-    RTC_DLOG(LS_VERBOSE) << "UpdateRemoteStreams_w: remote side will not send "
-                            "- disable payload type demuxing for "
-                         << ToString();
-    clear_payload_types = true;
-  }
+  media_receive_channel()->OnDemuxerCriteriaUpdatePending();
+  absl::Cleanup cleanup = [this] {
+    media_receive_channel()->OnDemuxerCriteriaUpdateComplete();
+  };
 
   const std::vector<StreamParams>& streams = content->streams();
   const bool new_has_unsignaled_ssrcs = HasStreamWithNoSsrcs(streams);
   const bool old_has_unsignaled_ssrcs = HasStreamWithNoSsrcs(remote_streams_);
 
+  RTC_DCHECK_BLOCK_COUNT_NO_MORE_THAN(0);
+
+  // Re-register the sink to update after changing the demuxer criteria first.
+  if (!RegisterRtpDemuxerSink_w(content)) {
+    return RTCError::InvalidParameter()
+           << "Failed to set up audio demuxing for mid='" << mid() << "'.";
+  }
+
+  RTC_DCHECK_BLOCK_COUNT_NO_MORE_THAN(1);
+
   // Check for streams that have been removed.
   for (const StreamParams& old_stream : remote_streams_) {
     // If we no longer have an unsignaled stream, we would like to remove
@@ -992,7 +1006,6 @@
   }
 
   // Check for new streams.
-  flat_set<uint32_t> ssrcs;
   for (const StreamParams& new_stream : streams) {
     // We allow a StreamParams with an empty list of SSRCs, in which case the
     // MediaChannel will cache the parameters and use them for any unsignaled
@@ -1014,20 +1027,8 @@
                << " to " << ToString();
       }
     }
-    // Update the receiving SSRCs.
-    ssrcs.insert(new_stream.ssrcs.begin(), new_stream.ssrcs.end());
   }
 
-  RTC_DCHECK_BLOCK_COUNT_NO_MORE_THAN(0);
-
-  // Re-register the sink to update after changing the demuxer criteria.
-  if (!RegisterRtpDemuxerSink_w(clear_payload_types, std::move(ssrcs))) {
-    return RTCError::InvalidParameter()
-           << "Failed to set up audio demuxing for mid='" << mid() << "'.";
-  }
-
-  RTC_DCHECK_BLOCK_COUNT_NO_MORE_THAN(1);
-
   remote_streams_ = streams;
 
   set_remote_content_direction(content->direction());
diff --git a/pc/channel.h b/pc/channel.h
index 2516dcf..f08e7c7 100644
--- a/pc/channel.h
+++ b/pc/channel.h
@@ -265,8 +265,7 @@
   // Registers a demuxer criteria with the transport, on the network thread.
   // This function will fail if there's no transport of if a sink is already
   // registered for this channel's demuxer_critera().
-  bool RegisterRtpDemuxerSink_w(bool clear_payload_types,
-                                std::optional<flat_set<uint32_t>> ssrcs)
+  bool RegisterRtpDemuxerSink_w(const MediaContentDescription* content)
       RTC_RUN_ON(worker_thread());
 
   // Return description of media channel to facilitate logging
diff --git a/pc/channel_unittest.cc b/pc/channel_unittest.cc
index 5e37ee5..15959db 100644
--- a/pc/channel_unittest.cc
+++ b/pc/channel_unittest.cc
@@ -82,11 +82,10 @@
 using ::webrtc::RtpTransceiverDirection;
 using ::webrtc::SdpType;
 using ::webrtc::StreamParams;
+using ::webrtc::Thread;
 
-const webrtc::Codec kPcmuCodec = webrtc::CreateAudioCodec(0, "PCMU", 64000, 1);
-const webrtc::Codec kPcmaCodec = webrtc::CreateAudioCodec(8, "PCMA", 64000, 1);
-const webrtc::Codec kIsacCodec =
-    webrtc::CreateAudioCodec(103, "ISAC", 40000, 1);
+const webrtc::Codec kPcmuCodec = webrtc::CreateAudioCodec(0, "PCMU", 8000, 1);
+const webrtc::Codec kPcmaCodec = webrtc::CreateAudioCodec(8, "PCMA", 8000, 1);
 const webrtc::Codec kH264Codec = webrtc::CreateVideoCodec(97, "H264");
 const webrtc::Codec kH264SvcCodec = webrtc::CreateVideoCodec(99, "H264-SVC");
 constexpr uint32_t kSsrc1 = 0x1111;
@@ -948,6 +947,93 @@
     EXPECT_TRUE(CheckCustomRtp1(kSsrc2, 0));
   }
 
+  class MockReceiveChannel : public T::MediaReceiveChannel {
+   public:
+    MockReceiveChannel(const typename T::Options& options,
+                       Thread* network_thread)
+        : T::MediaReceiveChannel(options, network_thread) {}
+
+    void OnDemuxerCriteriaUpdatePending() override {
+      ++pending_count_;
+      T::MediaReceiveChannel::OnDemuxerCriteriaUpdatePending();
+    }
+
+    void OnDemuxerCriteriaUpdateComplete() override {
+      --pending_count_;
+      T::MediaReceiveChannel::OnDemuxerCriteriaUpdateComplete();
+    }
+
+    bool AddRecvStream(const StreamParams& sp) override {
+      add_stream_called_ = true;
+      if (pending_count_ <= 0) {
+        criteria_not_pending_during_add_stream_ = true;
+      }
+      return T::MediaReceiveChannel::AddRecvStream(sp);
+    }
+
+    bool add_stream_called() const { return add_stream_called_; }
+    bool criteria_not_pending_during_add_stream() const {
+      return criteria_not_pending_during_add_stream_;
+    }
+
+   private:
+    int pending_count_ = 0;
+    bool add_stream_called_ = false;
+    bool criteria_not_pending_during_add_stream_ = false;
+  };
+
+  void TestUpdateRemoteStreamsRaceWithRtpPacket() {
+    auto ch1r = std::make_unique<MockReceiveChannel>(typename T::Options(),
+                                                     network_thread_);
+    MockReceiveChannel* mock_ch1r = ch1r.get();
+
+    CreateChannels(std::make_unique<typename T::MediaSendChannel>(
+                       typename T::Options(), network_thread_),
+                   std::move(ch1r),
+                   std::make_unique<typename T::MediaSendChannel>(
+                       typename T::Options(), network_thread_),
+                   std::make_unique<typename T::MediaReceiveChannel>(
+                       typename T::Options(), network_thread_),
+                   0, 0);
+
+    // Configure a new stream `stream1` to be added
+    StreamParams stream1;
+    stream1.id = "stream1";
+    stream1.ssrcs.push_back(kSsrc1);
+    stream1.cname = "stream1_cname";
+
+    typename T::Content content1;
+    CreateContent(0, kPcmuCodec, kH264Codec, &content1);
+    content1.AddStream(stream1);
+
+    ASSERT_TRUE(channel1_->SetLocalContent(&content1, SdpType::kOffer).ok());
+    channel1_->Enable(true);
+
+    typename T::Content content2;
+    CreateContent(0, kPcmuCodec, kH264Codec, &content2);
+    ASSERT_TRUE(channel2_->SetRemoteContent(&content1, SdpType::kOffer).ok());
+    ConnectFakeTransports();
+
+    // Configure answer adding stream2.
+    StreamParams stream2;
+    stream2.id = "stream2";
+    stream2.ssrcs.push_back(kSsrc2);
+    stream2.cname = "stream2_cname";
+
+    typename T::Content content3;
+    CreateContent(0, kPcmuCodec, kH264Codec, &content3);
+    content3.AddStream(stream2);
+
+    // Call SetRemoteContent.
+    RTCError error = channel1_->SetRemoteContent(&content3, SdpType::kAnswer);
+
+    EXPECT_TRUE(error.ok()) << "SetRemoteContent failed: " << error.message();
+    EXPECT_TRUE(mock_ch1r->add_stream_called());
+    EXPECT_FALSE(mock_ch1r->criteria_not_pending_during_add_stream())
+        << "Race condition detected: demuxer criteria was not pending during "
+           "AddRecvStream!";
+  }
+
   // Test that we only start playout and sending at the right times.
   void TestPlayoutAndSendingStates() {
     CreateChannels(0, 0);
@@ -2027,6 +2113,10 @@
   Base::TestChangeStreamParamsInContent();
 }
 
+TEST_F(VoiceChannelDoubleThreadTest, TestUpdateRemoteStreamsRaceWithRtpPacket) {
+  Base::TestUpdateRemoteStreamsRaceWithRtpPacket();
+}
+
 TEST_F(VoiceChannelDoubleThreadTest, TestPlayoutAndSendingStates) {
   Base::TestPlayoutAndSendingStates();
 }
@@ -2670,6 +2760,10 @@
   Base::TestSendTwoOffers();
 }
 
+TEST_F(VideoChannelDoubleThreadTest, TestUpdateRemoteStreamsRaceWithRtpPacket) {
+  Base::TestUpdateRemoteStreamsRaceWithRtpPacket();
+}
+
 TEST_F(VideoChannelDoubleThreadTest, TestReceiveTwoOffers) {
   Base::TestReceiveTwoOffers();
 }