Refactor state caching for SctpDataChannel observers

Replace the internal CachedGetters logic in ObserverAdapter with a
simplified CachedState mechanism in SctpDataChannel.

CacheStateAndCallBackOnSignalingThread captures the channel's state and
error on the network thread before posting a task to the signaling
thread. These values are temporarily cached and returned by the state()
and error() getters during the callback's execution.

This ensures that observers receive the correct state values captured at
the time of the event and avoids thread hops to read the state/error and
also removes some boilerplate ObserverAdapter implementation.

Bug: webrtc:510487699
Change-Id: I62a303efa26c8e84a56feb678d715aac5c8ac433
Reviewed-on: https://webrtc-review.googlesource.com/c/src/+/369142
Reviewed-by: Johannes Kron <kron@webrtc.org>
Commit-Queue: Tomas Gunnarsson <tommi@webrtc.org>
Cr-Commit-Position: refs/heads/main@{#47859}
diff --git a/pc/BUILD.gn b/pc/BUILD.gn
index 3622e7e..1f08a0f 100644
--- a/pc/BUILD.gn
+++ b/pc/BUILD.gn
@@ -891,6 +891,7 @@
     "../rtc_base/containers:flat_set",
     "../rtc_base/system:no_unique_address",
     "../rtc_base/system:unused",
+    "//third_party/abseil-cpp/absl/base:nullability",
     "//third_party/abseil-cpp/absl/functional:any_invocable",
     "//third_party/abseil-cpp/absl/strings:string_view",
   ]
diff --git a/pc/data_channel_controller.cc b/pc/data_channel_controller.cc
index b6d6639..6477e7b 100644
--- a/pc/data_channel_controller.cc
+++ b/pc/data_channel_controller.cc
@@ -320,21 +320,22 @@
   if (!ParseDataChannelOpenMessage(buffer, &label, &config)) {
     RTC_LOG(LS_WARNING) << "Failed to parse the OPEN message for sid "
                         << channel_id;
+    // Return `true` since the open message must be consumed and discarded.
+    return true;
+  }
+  config.open_handshake_role = InternalDataChannelInit::kAcker;
+  auto channel_or_error = CreateDataChannel(label, config);
+  if (channel_or_error.ok()) {
+    signaling_thread()->PostTask(
+        SafeTask(signaling_safety_.flag(),
+                 [this, channel = channel_or_error.MoveValue(),
+                  ready_to_send = data_channel_transport_->IsReadyToSend()] {
+                   RTC_DCHECK_RUN_ON(signaling_thread());
+                   OnDataChannelOpenMessage(std::move(channel), ready_to_send);
+                 }));
   } else {
-    config.open_handshake_role = InternalDataChannelInit::kAcker;
-    auto channel_or_error = CreateDataChannel(label, config);
-    if (channel_or_error.ok()) {
-      signaling_thread()->PostTask(SafeTask(
-          signaling_safety_.flag(),
-          [this, channel = channel_or_error.MoveValue(),
-           ready_to_send = data_channel_transport_->IsReadyToSend()] {
-            RTC_DCHECK_RUN_ON(signaling_thread());
-            OnDataChannelOpenMessage(std::move(channel), ready_to_send);
-          }));
-    } else {
-      RTC_LOG(LS_ERROR) << "Failed to create DataChannel from the OPEN message."
-                        << ToString(channel_or_error.error().type());
-    }
+    RTC_LOG(LS_ERROR) << "Failed to create DataChannel from the OPEN message. "
+                      << channel_or_error;
   }
   return true;
 }
@@ -406,23 +407,22 @@
     config.id = sid->stream_id_int();
   }
 
-  scoped_refptr<SctpDataChannel> channel = SctpDataChannel::Create(
-      weak_factory_.GetWeakPtr(), label, data_channel_transport_ != nullptr,
-      config, signaling_safety_.flag(), signaling_thread(), network_thread());
-  RTC_DCHECK(channel);
-
   // If we have an id already, notify the transport.
   if (sid.has_value()) {
-    RTCError error = AddSctpDataStream(
+    err = AddSctpDataStream(
         *sid, config.priority.value_or(PriorityValue(Priority::kLow)));
-    if (!error.ok()) {
-      return error;
+    if (!err.ok()) {
+      sid_allocator_.ReleaseSid(*sid);
+      return err;
     }
   }
+
+  scoped_refptr<SctpDataChannel> channel = SctpDataChannel::Create(
+      weak_factory_.GetWeakPtr(), label, data_channel_transport_ != nullptr,
+      config, max_message_size_, signaling_safety_.flag(), signaling_thread(),
+      network_thread());
+
   sctp_data_channels_n_.push_back(channel);
-  if (max_message_size_.has_value()) {
-    channel->OnMaxMessageSize(*max_message_size_);
-  }
   return channel;
 }
 
diff --git a/pc/rtc_stats_collector_unittest.cc b/pc/rtc_stats_collector_unittest.cc
index 9883293..75ec9c8 100644
--- a/pc/rtc_stats_collector_unittest.cc
+++ b/pc/rtc_stats_collector_unittest.cc
@@ -2206,10 +2206,12 @@
   FakeDataChannelController controller(pc_->network_thread());
   scoped_refptr<SctpDataChannel> dummy_channel_a = SctpDataChannel::Create(
       controller.weak_ptr(), "DummyChannelA", false, InternalDataChannelInit(),
-      signaling_safety.flag(), Thread::Current(), Thread::Current());
+      std::nullopt, signaling_safety.flag(), Thread::Current(),
+      Thread::Current());
   scoped_refptr<SctpDataChannel> dummy_channel_b = SctpDataChannel::Create(
       controller.weak_ptr(), "DummyChannelB", false, InternalDataChannelInit(),
-      signaling_safety.flag(), Thread::Current(), Thread::Current());
+      std::nullopt, signaling_safety.flag(), Thread::Current(),
+      Thread::Current());
 
   stats_->stats_collector().OnSctpDataChannelStateChanged(
       dummy_channel_a->internal_id(), DataChannelInterface::DataState::kOpen);
diff --git a/pc/sctp_data_channel.cc b/pc/sctp_data_channel.cc
index 12dcbf1..3332134 100644
--- a/pc/sctp_data_channel.cc
+++ b/pc/sctp_data_channel.cc
@@ -19,6 +19,7 @@
 #include <string>
 #include <utility>
 
+#include "absl/base/nullability.h"
 #include "absl/functional/any_invocable.h"
 #include "absl/strings/string_view.h"
 #include "api/data_channel_interface.h"
@@ -36,7 +37,6 @@
 #include "rtc_base/copy_on_write_buffer.h"
 #include "rtc_base/logging.h"
 #include "rtc_base/ssl_stream_adapter.h"
-#include "rtc_base/system/unused.h"
 #include "rtc_base/thread.h"
 #include "rtc_base/thread_annotations.h"
 #include "rtc_base/weak_ptr.h"
@@ -161,6 +161,11 @@
   used_sids_.erase(sid);
 }
 
+struct SctpDataChannel::CachedState {
+  DataChannelInterface::DataState state;
+  RTCError error;
+};
+
 // A DataChannelObserver implementation that offers backwards compatibility with
 // implementations that aren't yet ready to be called back on the network
 // thread. This implementation posts events to the signaling thread where
@@ -182,109 +187,51 @@
   explicit ObserverAdapter(SctpDataChannel* channel)
       : channel_(channel), controller_safety_(channel->controller_safety_) {}
 
-  bool IsInsideCallback() const {
-    RTC_DCHECK_RUN_ON(signaling_thread());
-    return cached_getters_ != nullptr;
-  }
-
-  DataChannelInterface::DataState cached_state() const {
-    RTC_DCHECK_RUN_ON(signaling_thread());
-    RTC_DCHECK(IsInsideCallback());
-    return cached_getters_->state();
-  }
-
-  RTCError cached_error() const {
-    RTC_DCHECK_RUN_ON(signaling_thread());
-    RTC_DCHECK(IsInsideCallback());
-    return cached_getters_->error();
-  }
-
   void SetDelegate(DataChannelObserver* delegate) {
     RTC_DCHECK_RUN_ON(signaling_thread());
     delegate_ = delegate;
     safety_.reset(PendingTaskSafetyFlag::CreateDetached());
   }
 
-  static void DeleteOnSignalingThread(
-      std::unique_ptr<ObserverAdapter> observer) {
-    auto* signaling_thread = observer->signaling_thread();
-    if (!signaling_thread->IsCurrent())
-      signaling_thread->PostTask([observer = std::move(observer)]() {});
-  }
-
  private:
-  class CachedGetters {
-   public:
-    explicit CachedGetters(ObserverAdapter* adapter)
-        : adapter_(adapter),
-          cached_state_(adapter_->channel_->state()),
-          cached_error_(adapter_->channel_->error()) {
-      RTC_DCHECK_RUN_ON(adapter->network_thread());
-    }
-
-    ~CachedGetters() {
-      if (!was_dropped_) {
-        RTC_DCHECK_RUN_ON(adapter_->signaling_thread());
-        RTC_DCHECK_EQ(adapter_->cached_getters_, this);
-        adapter_->cached_getters_ = nullptr;
-      }
-    }
-
-    bool PrepareForCallback() {
-      RTC_DCHECK_RUN_ON(adapter_->signaling_thread());
-      RTC_DCHECK(was_dropped_);
-      was_dropped_ = false;
-      adapter_->cached_getters_ = this;
-      return adapter_->delegate_ && adapter_->controller_safety_->alive();
-    }
-
-    RTCError error() { return cached_error_; }
-    DataChannelInterface::DataState state() { return cached_state_; }
-
-   private:
-    ObserverAdapter* const adapter_;
-    bool was_dropped_ = true;
-    const DataChannelInterface::DataState cached_state_;
-    const RTCError cached_error_;
-  };
-
   void OnStateChange() override {
     RTC_DCHECK_RUN_ON(network_thread());
-    signaling_thread()->PostTask(
-        SafeTask(safety_.flag(),
-                 [this, cached_state = std::make_unique<CachedGetters>(this)] {
-                   RTC_DCHECK_RUN_ON(signaling_thread());
-                   if (cached_state->PrepareForCallback())
-                     delegate_->OnStateChange();
-                 }));
+    RTC_DCHECK_EQ(signaling_thread(), channel_->signaling_thread_);
+    channel_->CacheStateAndCallBackOnSignalingThread(
+        SafeTask(safety_.flag(), [this] {
+          RTC_DCHECK_RUN_ON(signaling_thread());
+          if (delegate_ && controller_safety_->alive()) {
+            delegate_->OnStateChange();
+          }
+        }));
   }
 
   void OnMessage(const DataBuffer& buffer) override {
     RTC_DCHECK_RUN_ON(network_thread());
-    signaling_thread()->PostTask(SafeTask(
-        safety_.flag(), [this, buffer = buffer,
-                         cached_state = std::make_unique<CachedGetters>(this)] {
+    channel_->CacheStateAndCallBackOnSignalingThread(
+        SafeTask(safety_.flag(), [this, buffer = buffer] {
           RTC_DCHECK_RUN_ON(signaling_thread());
-          if (cached_state->PrepareForCallback())
+          if (delegate_ && controller_safety_->alive()) {
             delegate_->OnMessage(buffer);
+          }
         }));
   }
 
   void OnBufferedAmountChange(uint64_t sent_data_size) override {
     RTC_DCHECK_RUN_ON(network_thread());
-    signaling_thread()->PostTask(SafeTask(
-        safety_.flag(), [this, sent_data_size,
-                         cached_state = std::make_unique<CachedGetters>(this)] {
+    channel_->CacheStateAndCallBackOnSignalingThread(
+        SafeTask(safety_.flag(), [this, sent_data_size] {
           RTC_DCHECK_RUN_ON(signaling_thread());
-          if (cached_state->PrepareForCallback())
+          if (delegate_ && controller_safety_->alive()) {
             delegate_->OnBufferedAmountChange(sent_data_size);
+          }
         }));
   }
 
   bool IsOkToCallOnTheNetworkThread() override { return true; }
 
   Thread* signaling_thread() const { return signaling_thread_; }
-  Thread* network_thread() const { return channel_->network_thread_; }
+  Thread* network_thread() const { return network_thread_; }
 
   DataChannelObserver* delegate_ RTC_GUARDED_BY(signaling_thread()) = nullptr;
   SctpDataChannel* const channel_;
@@ -294,27 +241,29 @@
   // Make sure to keep our own signaling_thread_ pointer to avoid dereferencing
   // `channel_` in the `RTC_DCHECK_RUN_ON` checks on the signaling thread.
   Thread* const signaling_thread_{channel_->signaling_thread_};
+  Thread* const network_thread_{channel_->network_thread_};
   ScopedTaskSafety safety_;
-  CachedGetters* cached_getters_ RTC_GUARDED_BY(signaling_thread()) = nullptr;
 };
 
 // static
-scoped_refptr<SctpDataChannel> SctpDataChannel::Create(
+absl_nonnull scoped_refptr<SctpDataChannel> SctpDataChannel::Create(
     WeakPtr<SctpDataChannelControllerInterface> controller,
     absl::string_view label,
     bool connected_to_transport,
     const InternalDataChannelInit& config,
+    std::optional<int> max_message_size,
     scoped_refptr<PendingTaskSafetyFlag> controller_safety,
     Thread* signaling_thread,
     Thread* network_thread) {
   RTC_DCHECK(config.IsValid());
   return make_ref_counted<SctpDataChannel>(
       config, std::move(controller), label, connected_to_transport,
-      std::move(controller_safety), signaling_thread, network_thread);
+      max_message_size, std::move(controller_safety), signaling_thread,
+      network_thread);
 }
 
 // static
-scoped_refptr<DataChannelInterface> SctpDataChannel::CreateProxy(
+absl_nonnull scoped_refptr<DataChannelInterface> SctpDataChannel::CreateProxy(
     scoped_refptr<SctpDataChannel> channel) {
   // Copy thread params to local variables before `std::move()`.
   auto* signaling_thread = channel->signaling_thread_;
@@ -329,6 +278,7 @@
     WeakPtr<SctpDataChannelControllerInterface> controller,
     absl::string_view label,
     bool connected_to_transport,
+    std::optional<int> max_message_size,
     scoped_refptr<PendingTaskSafetyFlag> controller_safety,
     Thread* signaling_thread,
     Thread* network_thread)
@@ -344,13 +294,13 @@
       negotiated_(config.negotiated),
       ordered_(config.ordered),
       observer_(nullptr),
+      max_message_size_(max_message_size),
       controller_(std::move(controller)),
       connected_to_transport_(connected_to_transport),
       controller_safety_(std::move(controller_safety)) {
   RTC_DCHECK_RUN_ON(network_thread_);
   // Since we constructed on the network thread we can't (yet) check the
   // `controller_` pointer since doing so will trigger a thread check.
-  RTC_UNUSED(network_thread_);
   RTC_DCHECK(config.IsValid());
 
   switch (config.open_handshake_role) {
@@ -367,8 +317,14 @@
 }
 
 SctpDataChannel::~SctpDataChannel() {
-  if (observer_adapter_)
-    ObserverAdapter::DeleteOnSignalingThread(std::move(observer_adapter_));
+  if (signaling_thread_->IsCurrent()) {
+    observer_adapter_.reset();
+  } else {
+    signaling_thread_->PostTask(
+        [observer_adapter = std::move(observer_adapter_)]() mutable {
+          observer_adapter.reset();
+        });
+  }
 }
 
 void SctpDataChannel::RegisterObserver(DataChannelObserver* observer) {
@@ -524,9 +480,11 @@
   // fetch a different state value (since pending messages might cause the
   // state to change in the meantime).
   const auto* current_thread = Thread::Current();
-  if (current_thread == signaling_thread_ && observer_adapter_ &&
-      observer_adapter_->IsInsideCallback()) {
-    return observer_adapter_->cached_state();
+  if (current_thread == signaling_thread_) {
+    RTC_DCHECK_RUN_ON(signaling_thread_);
+    if (cached_state_) {
+      return cached_state_->state;
+    }
   }
 
   auto return_state = [&] {
@@ -541,9 +499,11 @@
 
 RTCError SctpDataChannel::error() const {
   const auto* current_thread = Thread::Current();
-  if (current_thread == signaling_thread_ && observer_adapter_ &&
-      observer_adapter_->IsInsideCallback()) {
-    return observer_adapter_->cached_error();
+  if (current_thread == signaling_thread_) {
+    RTC_DCHECK_RUN_ON(signaling_thread_);
+    if (cached_state_) {
+      return cached_state_->error;
+    }
   }
 
   auto return_error = [&] {
@@ -709,6 +669,21 @@
   return stats;
 }
 
+void SctpDataChannel::CacheStateAndCallBackOnSignalingThread(
+    absl::AnyInvocable<void() &&> callback) {
+  RTC_DCHECK_RUN_ON(network_thread_);
+  scoped_refptr<SctpDataChannel> me(this);
+  signaling_thread_->PostTask([me = std::move(me),
+                               cache = CachedState{state_, error_},
+                               callback = std::move(callback)]() mutable {
+    RTC_DCHECK_RUN_ON(me->signaling_thread_);
+    RTC_DCHECK(!me->cached_state_);
+    me->cached_state_ = &cache;
+    std::move(callback)();
+    me->cached_state_ = nullptr;
+  });
+}
+
 void SctpDataChannel::OnDataReceived(DataMessageType type,
                                      const CopyOnWriteBuffer& payload) {
   RTC_DCHECK_RUN_ON(network_thread_);
diff --git a/pc/sctp_data_channel.h b/pc/sctp_data_channel.h
index 7028022..dfefb43 100644
--- a/pc/sctp_data_channel.h
+++ b/pc/sctp_data_channel.h
@@ -17,6 +17,7 @@
 #include <optional>
 #include <string>
 
+#include "absl/base/nullability.h"
 #include "absl/functional/any_invocable.h"
 #include "absl/strings/string_view.h"
 #include "api/data_channel_interface.h"
@@ -145,18 +146,19 @@
   // such callbacks after the peerconnection has been closed. The data
   // controller will update the flag when closed, which will cancel any pending
   // event notifications.
-  static scoped_refptr<SctpDataChannel> Create(
+  static absl_nonnull scoped_refptr<SctpDataChannel> Create(
       WeakPtr<SctpDataChannelControllerInterface> controller,
       absl::string_view label,
       bool connected_to_transport,
       const InternalDataChannelInit& config,
+      std::optional<int> max_message_size,
       scoped_refptr<PendingTaskSafetyFlag> controller_safety,
       Thread* signaling_thread,
       Thread* network_thread);
 
   // Instantiates an API proxy for a SctpDataChannel instance that will be
   // handed out to external callers.
-  static scoped_refptr<DataChannelInterface> CreateProxy(
+  static absl_nonnull scoped_refptr<DataChannelInterface> CreateProxy(
       scoped_refptr<SctpDataChannel> channel);
 
   void RegisterObserver(DataChannelObserver* observer) override;
@@ -244,12 +246,25 @@
                   WeakPtr<SctpDataChannelControllerInterface> controller,
                   absl::string_view label,
                   bool connected_to_transport,
+                  std::optional<int> max_message_size,
                   scoped_refptr<PendingTaskSafetyFlag> controller_safety,
                   Thread* signaling_thread,
                   Thread* network_thread);
   ~SctpDataChannel() override;
 
  private:
+  // Caches the current state on the network thread and makes a call back to the
+  // `callback` object on the signaling thread while applying the cached state
+  // to specific getter functions.
+  // This is useful when a callback to the application is needed and during that
+  // callback, it's expected that this state will be queried (e.g. the
+  // `state()`), but a thread hop should not be required for querying that
+  // state.
+  // Must be called on the network thread.
+  void CacheStateAndCallBackOnSignalingThread(
+      absl::AnyInvocable<void() &&> callback);
+
+  struct CachedState;
   class ObserverAdapter;
 
   // The OPEN(_ACK) signaling state.
@@ -294,6 +309,7 @@
 
   DataChannelObserver* observer_ RTC_GUARDED_BY(network_thread_) = nullptr;
   std::unique_ptr<ObserverAdapter> observer_adapter_;
+  CachedState* cached_state_ RTC_GUARDED_BY(signaling_thread_) = nullptr;
   DataState state_ RTC_GUARDED_BY(network_thread_) = kConnecting;
   RTCError error_ RTC_GUARDED_BY(network_thread_);
   uint32_t messages_sent_ RTC_GUARDED_BY(network_thread_) = 0;
diff --git a/pc/test/fake_data_channel_controller.h b/pc/test/fake_data_channel_controller.h
index c003f4a..ef71133 100644
--- a/pc/test/fake_data_channel_controller.h
+++ b/pc/test/fake_data_channel_controller.h
@@ -12,6 +12,7 @@
 #define PC_TEST_FAKE_DATA_CHANNEL_CONTROLLER_H_
 
 #include <cstddef>
+#include <optional>
 #include <set>
 #include <string>
 #include <utility>
@@ -71,7 +72,7 @@
 
           scoped_refptr<SctpDataChannel> channel = SctpDataChannel::Create(
               std::move(my_weak_ptr), std::string(label), transport_available_,
-              init, signaling_safety_.flag(), signaling_thread_,
+              init, std::nullopt, signaling_safety_.flag(), signaling_thread_,
               network_thread_);
           if (transport_available_ && channel->sid_n().has_value()) {
             AddSctpDataStream(*channel->sid_n(), channel->priority());
diff --git a/pc/test/fake_peer_connection_for_stats.h b/pc/test/fake_peer_connection_for_stats.h
index 18fce85..71e3229 100644
--- a/pc/test/fake_peer_connection_for_stats.h
+++ b/pc/test/fake_peer_connection_for_stats.h
@@ -506,7 +506,7 @@
   void AddSctpDataChannel(const std::string& label,
                           const InternalDataChannelInit& init) {
     AddSctpDataChannel(SctpDataChannel::Create(
-        data_channel_controller_.weak_ptr(), label, false, init,
+        data_channel_controller_.weak_ptr(), label, false, init, std::nullopt,
         signaling_safety_.flag(), signaling_thread_, network_thread_));
   }
 
diff --git a/pc/test/mock_data_channel.h b/pc/test/mock_data_channel.h
index fdc554c..197357b 100644
--- a/pc/test/mock_data_channel.h
+++ b/pc/test/mock_data_channel.h
@@ -12,6 +12,7 @@
 #define PC_TEST_MOCK_DATA_CHANNEL_H_
 
 #include <cstdint>
+#include <optional>
 #include <string>
 #include <utility>
 
@@ -54,6 +55,7 @@
                         std::move(controller),
                         label,
                         false,
+                        std::nullopt,
                         PendingTaskSafetyFlag::Create(),
                         signaling_thread,
                         network_thread) {