Move SUBSCRIBE to a dedicated bidirectional stream.

This change refactors MoQT to handle SUBSCRIBE messages on their own bidirectional streams, rather than on the control stream.

Key changes include:
-   New `MoqtSubscribeRequestStream` and `MoqtSubscribeResponseStream` classes.
-   `MoqtSession::Subscribe` now opens a bidi stream and sends the SUBSCRIBE message on it.
-   `MoqtSession::Unsubscribe` now resets the bidi stream associated with the subscription.
-   `MoqtSession::RequestUpdate` sends the update on the existing subscribe bidi stream.
-   The control stream no longer processes SUBSCRIBE, SUBSCRIBE_OK, or UNSUBSCRIBE messages.
-   `MoqtUnsubscribe` message type and serialization/parsing are removed.
-   MoqtSession no longer owns any subscriptions
-   All external state is now destroyed in Detach(). On teardown, streams are destroyed in arbitrary order, making clearing such state in the destructor dangerous.

I wrote better helper functions in MoqtSessionTest to reduce the toil in sending/receiving messages, setting up bidi streams, and not hard-coding message formats.

PiperOrigin-RevId: 949593252
diff --git a/build/source_list.bzl b/build/source_list.bzl
index 30e8046..b1fa635 100644
--- a/build/source_list.bzl
+++ b/build/source_list.bzl
@@ -1607,6 +1607,7 @@
     "quic/moqt/moqt_session_callbacks.h",
     "quic/moqt/moqt_session_interface.h",
     "quic/moqt/moqt_stream_map.h",
+    "quic/moqt/moqt_subscribe_stream.h",
     "quic/moqt/moqt_subscription.h",
     "quic/moqt/moqt_trace_recorder.h",
     "quic/moqt/moqt_track.h",
@@ -1643,6 +1644,7 @@
     "quic/moqt/moqt_relay_track_publisher.cc",
     "quic/moqt/moqt_session.cc",
     "quic/moqt/moqt_stream_map.cc",
+    "quic/moqt/moqt_subscribe_stream.cc",
     "quic/moqt/moqt_subscription.cc",
     "quic/moqt/moqt_trace_recorder.cc",
     "quic/moqt/moqt_track.cc",
@@ -1678,6 +1680,7 @@
     "quic/moqt/moqt_relay_track_publisher_test.cc",
     "quic/moqt/moqt_session_test.cc",
     "quic/moqt/moqt_stream_map_test.cc",
+    "quic/moqt/moqt_subscribe_stream_test.cc",
     "quic/moqt/moqt_subscription_test.cc",
     "quic/moqt/moqt_track_test.cc",
     "quic/moqt/moqt_uni_stream_test.cc",
diff --git a/build/source_list.gni b/build/source_list.gni
index 96933eb..43addc5 100644
--- a/build/source_list.gni
+++ b/build/source_list.gni
@@ -1611,6 +1611,7 @@
     "src/quiche/quic/moqt/moqt_session_callbacks.h",
     "src/quiche/quic/moqt/moqt_session_interface.h",
     "src/quiche/quic/moqt/moqt_stream_map.h",
+    "src/quiche/quic/moqt/moqt_subscribe_stream.h",
     "src/quiche/quic/moqt/moqt_subscription.h",
     "src/quiche/quic/moqt/moqt_trace_recorder.h",
     "src/quiche/quic/moqt/moqt_track.h",
@@ -1647,6 +1648,7 @@
     "src/quiche/quic/moqt/moqt_relay_track_publisher.cc",
     "src/quiche/quic/moqt/moqt_session.cc",
     "src/quiche/quic/moqt/moqt_stream_map.cc",
+    "src/quiche/quic/moqt/moqt_subscribe_stream.cc",
     "src/quiche/quic/moqt/moqt_subscription.cc",
     "src/quiche/quic/moqt/moqt_trace_recorder.cc",
     "src/quiche/quic/moqt/moqt_track.cc",
@@ -1683,6 +1685,7 @@
     "src/quiche/quic/moqt/moqt_relay_track_publisher_test.cc",
     "src/quiche/quic/moqt/moqt_session_test.cc",
     "src/quiche/quic/moqt/moqt_stream_map_test.cc",
+    "src/quiche/quic/moqt/moqt_subscribe_stream_test.cc",
     "src/quiche/quic/moqt/moqt_subscription_test.cc",
     "src/quiche/quic/moqt/moqt_track_test.cc",
     "src/quiche/quic/moqt/moqt_uni_stream_test.cc",
diff --git a/build/source_list.json b/build/source_list.json
index f8008d3..ea8c649 100644
--- a/build/source_list.json
+++ b/build/source_list.json
@@ -1610,6 +1610,7 @@
     "quiche/quic/moqt/moqt_session_callbacks.h",
     "quiche/quic/moqt/moqt_session_interface.h",
     "quiche/quic/moqt/moqt_stream_map.h",
+    "quiche/quic/moqt/moqt_subscribe_stream.h",
     "quiche/quic/moqt/moqt_subscription.h",
     "quiche/quic/moqt/moqt_trace_recorder.h",
     "quiche/quic/moqt/moqt_track.h",
@@ -1646,6 +1647,7 @@
     "quiche/quic/moqt/moqt_relay_track_publisher.cc",
     "quiche/quic/moqt/moqt_session.cc",
     "quiche/quic/moqt/moqt_stream_map.cc",
+    "quiche/quic/moqt/moqt_subscribe_stream.cc",
     "quiche/quic/moqt/moqt_subscription.cc",
     "quiche/quic/moqt/moqt_trace_recorder.cc",
     "quiche/quic/moqt/moqt_track.cc",
@@ -1682,6 +1684,7 @@
     "quiche/quic/moqt/moqt_relay_track_publisher_test.cc",
     "quiche/quic/moqt/moqt_session_test.cc",
     "quiche/quic/moqt/moqt_stream_map_test.cc",
+    "quiche/quic/moqt/moqt_subscribe_stream_test.cc",
     "quiche/quic/moqt/moqt_subscription_test.cc",
     "quiche/quic/moqt/moqt_track_test.cc",
     "quiche/quic/moqt/moqt_uni_stream_test.cc",
diff --git a/quiche/quic/moqt/moqt_bidi_stream.cc b/quiche/quic/moqt/moqt_bidi_stream.cc
index df12569..58bafea 100644
--- a/quiche/quic/moqt/moqt_bidi_stream.cc
+++ b/quiche/quic/moqt/moqt_bidi_stream.cc
@@ -13,6 +13,7 @@
 #include "absl/strings/string_view.h"
 #include "quiche/quic/core/quic_time.h"
 #include "quiche/quic/moqt/moqt_error.h"
+#include "quiche/quic/moqt/moqt_fetch_task.h"
 #include "quiche/quic/moqt/moqt_key_value_pair.h"
 #include "quiche/quic/moqt/moqt_messages.h"
 #include "quiche/quic/moqt/moqt_parser.h"
@@ -83,6 +84,42 @@
                           info.reason_phrase, fin);
 }
 
+absl::Status MoqtBidiStreamBase::SendRequestUpdate(
+    uint64_t request_id, uint64_t existing_request_id,
+    const MessageParameters& parameters, MoqtResponseCallback callback) {
+  MoqtRequestUpdate request_update;
+  request_update.request_id = request_id;
+  request_update.existing_request_id = existing_request_id;
+  request_update.parameters = parameters;
+  pending_responses_.push_back(std::move(callback));
+  pending_updates_.push_back(parameters);
+  return SendOrBufferMessage(framer_->SerializeRequestUpdate(request_update),
+                             /*fin=*/false);
+}
+
+absl::Status MoqtBidiStreamBase::OnControlMessage(
+    const MoqtRequestOk& message) {
+  if (pending_responses_.empty()) {
+    return absl::OkStatus();
+  }
+  std::move(pending_responses_.front())(message.parameters);
+  pending_responses_.pop_front();
+  return absl::OkStatus();
+}
+
+absl::Status MoqtBidiStreamBase::OnControlMessage(
+    const MoqtRequestError& message) {
+  if (pending_responses_.empty()) {
+    return absl::OkStatus();
+  }
+  std::move(pending_responses_.front())(MoqtRequestErrorInfo{
+      message.error_code, message.retry_interval, message.reason_phrase});
+  pending_responses_.clear();
+  pending_updates_.clear();
+  Fin();
+  return absl::OkStatus();
+}
+
 void MoqtBidiStreamBase::OnFatalError(absl::Status status) {
   QUICHE_DCHECK(!status.ok());
   if (session_error_callback_ == nullptr) {
diff --git a/quiche/quic/moqt/moqt_bidi_stream.h b/quiche/quic/moqt/moqt_bidi_stream.h
index 0ebb74c..22af6b9 100644
--- a/quiche/quic/moqt/moqt_bidi_stream.h
+++ b/quiche/quic/moqt/moqt_bidi_stream.h
@@ -13,11 +13,13 @@
 
 #include "absl/base/nullability.h"
 #include "absl/status/status.h"
+#include "absl/status/statusor.h"
 #include "absl/strings/str_cat.h"
 #include "absl/strings/string_view.h"
 #include "quiche/quic/core/quic_time.h"
 #include "quiche/quic/moqt/moqt_control_message_queue.h"
 #include "quiche/quic/moqt/moqt_error.h"
+#include "quiche/quic/moqt/moqt_fetch_task.h"
 #include "quiche/quic/moqt/moqt_framer.h"
 #include "quiche/quic/moqt/moqt_key_value_pair.h"
 #include "quiche/quic/moqt/moqt_messages.h"
@@ -25,6 +27,7 @@
 #include "quiche/common/platform/api/quiche_logging.h"
 #include "quiche/common/quiche_buffer_allocator.h"
 #include "quiche/common/quiche_callbacks.h"
+#include "quiche/common/quiche_circular_deque.h"
 #include "quiche/web_transport/web_transport.h"
 
 namespace moqt {
@@ -102,7 +105,11 @@
       absl::string_view reason_phrase, bool fin = false);
   absl::Status SendRequestError(uint64_t request_id, MoqtRequestErrorInfo info,
                                 bool fin = false);
-
+  // Can be overridden for message-specific constraints.
+  virtual absl::Status SendRequestUpdate(uint64_t request_id,
+                                         uint64_t existing_request_id,
+                                         const MessageParameters& parameters,
+                                         MoqtResponseCallback callback);
   void Fin() {
     CheckStatus(outgoing_message_queue_.Fin());
     Detach();
@@ -122,14 +129,13 @@
     }
   }
 
-  // Removes any state in MoqtSession related to the stream. Overrides of this
-  // method must be robust to multiple invocations.
-  virtual void Detach() = 0;
+  MoqtFramer* framer() const { return framer_; }
 
-  // TODO(martinduke): Remove once SUBSCRIBE moves to a bidi stream. This is
-  // only needed to check whether or not to FIN the bidi stream.
-  bool is_control_stream() const { return control_stream_; }
-  void set_control_stream() { control_stream_ = true; }
+  // Removes any state in MoqtSession related to the stream. Overrides of this
+  // method must be robust to multiple invocations. Only called on sending a FIN
+  // or RESET. If otherwise destroyed, it's due to a larger cleanup where the
+  // state no longer matters.
+  virtual void Detach() = 0;
 
  protected:
   // Called when a WebTransport stream has been associated with the object.
@@ -141,6 +147,18 @@
   virtual absl::Status OnRawControlMessage(
       const MoqtRawControlMessage& message) = 0;
 
+  virtual absl::Status OnControlMessage(const MoqtRequestOk& message);
+  virtual absl::Status OnControlMessage(const MoqtRequestError& message);
+
+  absl::StatusOr<MessageParameters> PopParameters() {
+    if (pending_updates_.empty()) {
+      return absl::NotFoundError("Too many REQUEST_OK received");
+    }
+    MessageParameters parameters = pending_updates_.front();
+    pending_updates_.pop_front();
+    return parameters;
+  }
+
   // Terminates the MoQT session due to a fatal error encountered.
   void OnFatalError(absl::Status status);
 
@@ -148,7 +166,6 @@
   const MoqtControlMessageParser& message_parser() const {
     return message_parser_;
   }
-  MoqtFramer* framer() const { return framer_; }
   webtransport::Stream* stream() const {
     return stream_parser_ != nullptr ? stream_parser_->stream() : nullptr;
   }
@@ -156,12 +173,12 @@
  private:
   friend class test::MoqtBidiStreamTestWrapper;
 
-  // TODO(martinduke): Remove once SUBSCRIBE moves to a bidi stream.
-  bool control_stream_ = false;
   MoqtFramer* absl_nonnull framer_;
   std::unique_ptr<MoqtControlStreamParser> absl_nullable stream_parser_;
   MoqtControlMessageParser message_parser_;
   MoqtControlMessageQueue outgoing_message_queue_;
+  quiche::QuicheCircularDeque<MoqtResponseCallback> pending_responses_;
+  quiche::QuicheCircularDeque<MessageParameters> pending_updates_;
   SessionErrorCallback session_error_callback_;
 };
 
diff --git a/quiche/quic/moqt/moqt_bidi_stream_test.cc b/quiche/quic/moqt/moqt_bidi_stream_test.cc
index 8b5f011..9da7253 100644
--- a/quiche/quic/moqt/moqt_bidi_stream_test.cc
+++ b/quiche/quic/moqt/moqt_bidi_stream_test.cc
@@ -6,15 +6,20 @@
 
 #include <memory>
 #include <optional>
+#include <utility>
+#include <variant>
 
 #include "absl/status/status.h"
 #include "absl/strings/string_view.h"
 #include "quiche/quic/core/quic_types.h"
 #include "quiche/quic/moqt/moqt_error.h"
+#include "quiche/quic/moqt/moqt_fetch_task.h"
 #include "quiche/quic/moqt/moqt_framer.h"
+#include "quiche/quic/moqt/moqt_key_value_pair.h"
 #include "quiche/quic/moqt/moqt_messages.h"
 #include "quiche/quic/moqt/moqt_parser.h"
 #include "quiche/quic/moqt/moqt_session_interface.h"
+#include "quiche/quic/moqt/test_tools/mock_moqt_session.h"
 #include "quiche/quic/moqt/test_tools/moqt_framer_utils.h"
 #include "quiche/common/platform/api/quiche_test.h"
 #include "quiche/common/test_tools/quiche_test_utils.h"
@@ -27,11 +32,6 @@
  public:
   using MoqtBidiStreamBase::MoqtBidiStreamBase;
 
-  absl::Status OnControlMessage(const MoqtRequestOk& message) {
-    ++ok_received_;
-    return absl::OkStatus();
-  }
-
   void OnStreamBound() override {}
   absl::Status OnRawControlMessage(
       const MoqtRawControlMessage& message) override {
@@ -39,9 +39,14 @@
         *this, message_parser(), message, "test");
   }
   void Detach() override { detached_ = true; }
-
-  int ok_received_ = 0;
   bool detached_ = false;
+
+  absl::Status OnControlMessage(const MoqtRequestOk& message) {
+    return MoqtBidiStreamBase::OnControlMessage(message);
+  }
+  absl::Status OnControlMessage(const MoqtRequestError& message) {
+    return MoqtBidiStreamBase::OnControlMessage(message);
+  }
 };
 
 class MoqtBidiStreamTest : public quiche::test::QuicheTest {
@@ -118,7 +123,6 @@
   MoqtFramer framer(/*using_webtrans=*/true, quic::Perspective::IS_SERVER);
   stream.Receive(framer.SerializeRequestOk(MoqtRequestOk()).AsStringView());
   stream_->OnCanRead();
-  EXPECT_EQ(stream_->ok_received_, 1u);
 
   stream.Receive(framer.SerializeGoAway(MoqtGoAway()).AsStringView());
   EXPECT_CALL(error_callback_, Call)
@@ -131,4 +135,98 @@
   stream_->OnCanRead();
 }
 
+TEST_F(MoqtBidiStreamTest, SendRequestOk) {
+  stream_->BindStream(&mock_stream_);
+  EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(testing::Return(true));
+  EXPECT_CALL(
+      mock_stream_,
+      Writev(ControlMessageOfType(MoqtMessageType::kRequestOk), testing::_));
+  MessageParameters parameters;
+  parameters.subscriber_priority = 20;
+  QUICHE_EXPECT_OK(stream_->SendRequestOk(1, parameters, /*fin=*/false));
+  EXPECT_FALSE(stream_->detached_);
+}
+
+TEST_F(MoqtBidiStreamTest, SendRequestOkFin) {
+  stream_->BindStream(&mock_stream_);
+  EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(testing::Return(true));
+  EXPECT_CALL(
+      mock_stream_,
+      Writev(ControlMessageOfType(MoqtMessageType::kRequestOk), testing::_));
+  MessageParameters parameters;
+  QUICHE_EXPECT_OK(stream_->SendRequestOk(1, parameters, /*fin=*/true));
+  EXPECT_TRUE(stream_->detached_);
+}
+
+TEST_F(MoqtBidiStreamTest, SendRequestErrorOverload) {
+  stream_->BindStream(&mock_stream_);
+  EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(testing::Return(true));
+  EXPECT_CALL(
+      mock_stream_,
+      Writev(ControlMessageOfType(MoqtMessageType::kRequestError), testing::_));
+  QUICHE_EXPECT_OK(stream_->SendRequestError(1, RequestErrorCode::kUnauthorized,
+                                             std::nullopt, "reason",
+                                             /*fin=*/true));
+  EXPECT_TRUE(stream_->detached_);
+}
+
+TEST_F(MoqtBidiStreamTest, SendRequestUpdateAndReceiveOk) {
+  stream_->BindStream(&mock_stream_);
+  EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(testing::Return(true));
+  EXPECT_CALL(mock_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kRequestUpdate),
+                     testing::_));
+  MessageParameters parameters;
+  parameters.subscriber_priority = 20;
+  bool callback_called = false;
+  MoqtResponseCallback callback =
+      [&](std::variant<MessageParameters, MoqtRequestErrorInfo> res) {
+        callback_called = true;
+        ASSERT_TRUE(std::holds_alternative<MessageParameters>(res));
+        EXPECT_EQ(std::get<MessageParameters>(res).subscriber_priority, 30);
+      };
+  QUICHE_EXPECT_OK(
+      stream_->SendRequestUpdate(1, 0, parameters, std::move(callback)));
+  // Simulate receiving RequestOk
+  MoqtRequestOk request_ok;
+  request_ok.request_id = 1;
+  request_ok.parameters.subscriber_priority = 30;
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(request_ok));
+  EXPECT_TRUE(callback_called);
+  EXPECT_FALSE(stream_->detached_);
+}
+
+TEST_F(MoqtBidiStreamTest, SendRequestUpdateAndReceiveError) {
+  stream_->BindStream(&mock_stream_);
+  EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(testing::Return(true));
+  EXPECT_CALL(mock_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kRequestUpdate),
+                     testing::_));
+  MessageParameters parameters;
+  bool callback_called = false;
+  MoqtResponseCallback callback =
+      [&](std::variant<MessageParameters, MoqtRequestErrorInfo> res) {
+        callback_called = true;
+        ASSERT_TRUE(std::holds_alternative<MoqtRequestErrorInfo>(res));
+        EXPECT_EQ(std::get<MoqtRequestErrorInfo>(res).error_code,
+                  RequestErrorCode::kUnauthorized);
+      };
+  QUICHE_EXPECT_OK(
+      stream_->SendRequestUpdate(1, 0, parameters, std::move(callback)));
+  // Simulate receiving RequestError
+  MoqtRequestError request_error;
+  request_error.request_id = 1;
+  request_error.error_code = RequestErrorCode::kUnauthorized;
+  request_error.reason_phrase = "unauthorized";
+  ExpectFin(mock_stream_);
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(request_error));
+  EXPECT_TRUE(callback_called);
+  EXPECT_TRUE(stream_->detached_);
+}
+
+TEST_F(MoqtBidiStreamTest, QueueIsFull) {
+  stream_->BindStream(&mock_stream_);
+  EXPECT_FALSE(stream_->QueueIsFull());
+}
+
 }  // namespace moqt::test
diff --git a/quiche/quic/moqt/moqt_framer.cc b/quiche/quic/moqt/moqt_framer.cc
index 17fd0e0..d446967 100644
--- a/quiche/quic/moqt/moqt_framer.cc
+++ b/quiche/quic/moqt/moqt_framer.cc
@@ -544,12 +544,6 @@
       WireStringWithMoqVarIntLength(message.reason_phrase));
 }
 
-quiche::QuicheBuffer MoqtFramer::SerializeUnsubscribe(
-    const MoqtUnsubscribe& message) {
-  return SerializeControlMessage(MoqtMessageType::kUnsubscribe,
-                                 WireMoqVarInt(message.request_id));
-}
-
 quiche::QuicheBuffer MoqtFramer::SerializePublishDone(
     const MoqtPublishDone& message) {
   return SerializeControlMessage(
@@ -707,8 +701,8 @@
 quiche::QuicheBuffer MoqtFramer::SerializeObjectAck(
     const MoqtObjectAck& message) {
   return SerializeControlMessage(
-      MoqtMessageType::kObjectAck, WireMoqVarInt(message.subscribe_id),
-      WireMoqVarInt(message.group_id), WireMoqVarInt(message.object_id),
+      MoqtMessageType::kObjectAck, WireMoqVarInt(message.group_id),
+      WireMoqVarInt(message.object_id),
       WireMoqVarInt(SignedVarintSerializedForm(
           message.delta_from_deadline.ToMicroseconds())));
 }
diff --git a/quiche/quic/moqt/moqt_framer.h b/quiche/quic/moqt/moqt_framer.h
index 96a7251..0d8df94 100644
--- a/quiche/quic/moqt/moqt_framer.h
+++ b/quiche/quic/moqt/moqt_framer.h
@@ -54,7 +54,6 @@
   quiche::QuicheBuffer SerializeSubscribeOk(
       const MoqtSubscribeOk& message,
       MoqtMessageType message_type = MoqtMessageType::kSubscribeOk);
-  quiche::QuicheBuffer SerializeUnsubscribe(const MoqtUnsubscribe& message);
   quiche::QuicheBuffer SerializePublishDone(const MoqtPublishDone& message);
   quiche::QuicheBuffer SerializeRequestUpdate(const MoqtRequestUpdate& message);
   quiche::QuicheBuffer SerializePublishNamespace(
diff --git a/quiche/quic/moqt/moqt_framer_test.cc b/quiche/quic/moqt/moqt_framer_test.cc
index 482e5c4..2000e69 100644
--- a/quiche/quic/moqt/moqt_framer_test.cc
+++ b/quiche/quic/moqt/moqt_framer_test.cc
@@ -46,7 +46,6 @@
       MoqtMessageType::kRequestError,
       MoqtMessageType::kSubscribe,
       MoqtMessageType::kSubscribeOk,
-      MoqtMessageType::kUnsubscribe,
       MoqtMessageType::kPublishDone,
       MoqtMessageType::kPublishNamespace,
       MoqtMessageType::kPublishNamespaceDone,
@@ -152,10 +151,6 @@
         auto data = std::get<MoqtSubscribeOk>(structured_data);
         return framer_.SerializeSubscribeOk(data);
       }
-      case MoqtMessageType::kUnsubscribe: {
-        auto data = std::get<MoqtUnsubscribe>(structured_data);
-        return framer_.SerializeUnsubscribe(data);
-      }
       case MoqtMessageType::kPublishDone: {
         auto data = std::get<MoqtPublishDone>(structured_data);
         return framer_.SerializePublishDone(data);
diff --git a/quiche/quic/moqt/moqt_integration_test.cc b/quiche/quic/moqt/moqt_integration_test.cc
index c00cd60..89a2d22 100644
--- a/quiche/quic/moqt/moqt_integration_test.cc
+++ b/quiche/quic/moqt/moqt_integration_test.cc
@@ -849,6 +849,10 @@
       test_harness_.RunUntilWithDefaultTimeout([&]() { return stream_reset; });
   EXPECT_TRUE(success);
   EXPECT_EQ(bytes_received, 2000);
+  // On teardown, streams are destroyed in arbitrary order. If the uni stream
+  // is destroyed before the bidi stream, there will be a second stream reset
+  // notification to the visitor. If not, there won't be.
+  EXPECT_CALL(subscribe_visitor_, OnStreamReset).Times(testing::AnyNumber());
 }
 
 TEST_F(MoqtIntegrationTest, BandwidthProbe) {
diff --git a/quiche/quic/moqt/moqt_messages.h b/quiche/quic/moqt/moqt_messages.h
index edfe38b..b05a392 100644
--- a/quiche/quic/moqt/moqt_messages.h
+++ b/quiche/quic/moqt/moqt_messages.h
@@ -386,10 +386,6 @@
   TrackExtensions extensions;
 };
 
-struct QUICHE_EXPORT MoqtUnsubscribe {
-  uint64_t request_id;
-};
-
 struct QUICHE_EXPORT MoqtPublishDone {
   uint64_t request_id;
   PublishDoneCode status_code;
@@ -526,11 +522,10 @@
   TrackExtensions extensions;
 };
 
-// All of the four values in this message are encoded as varints.
+// All of the three values in this message are encoded as varints.
 // `delta_from_deadline` is encoded as an absolute value, with the lowest bit
 // indicating the sign (0 if positive).
 struct QUICHE_EXPORT MoqtObjectAck {
-  uint64_t subscribe_id;
   uint64_t group_id;
   uint64_t object_id;
   // Positive if the object has been received before the deadline.
diff --git a/quiche/quic/moqt/moqt_namespace_stream.h b/quiche/quic/moqt/moqt_namespace_stream.h
index a80c4a1..34e64d1 100644
--- a/quiche/quic/moqt/moqt_namespace_stream.h
+++ b/quiche/quic/moqt/moqt_namespace_stream.h
@@ -50,7 +50,7 @@
         request_id_(request_id),
         remove_callback_(std::move(remove_callback)),
         response_callback_(std::move(response_callback)) {}
-  ~MoqtNamespaceSubscriberStream() override;
+  ~MoqtNamespaceSubscriberStream();
 
   // MoqtBidiStreamBase overrides.
   void OnStreamBound() override;
@@ -154,7 +154,7 @@
       AddPrefixCallback add_callback, RemovePrefixCallback remove_callback,
       SessionErrorCallback session_error_callback,
       MoqtIncomingSubscribeNamespaceCallback& application);
-  ~MoqtNamespacePublisherStream() override { Detach(); }
+  ~MoqtNamespacePublisherStream() { Detach(); }
 
   void OnStreamBound() override {
     // TODO(martinduke): Set the priority for this stream.
diff --git a/quiche/quic/moqt/moqt_parser.cc b/quiche/quic/moqt/moqt_parser.cc
index 13bfe26..5a297a6 100644
--- a/quiche/quic/moqt/moqt_parser.cc
+++ b/quiche/quic/moqt/moqt_parser.cc
@@ -685,17 +685,6 @@
   return request_error;
 }
 
-absl::StatusOr<MoqtUnsubscribe> MoqtControlMessageParser::ProcessUnsubscribe(
-    absl::string_view data) const {
-  quic::QuicDataReader reader(data);
-  MoqtUnsubscribe unsubscribe;
-  if (!reader.ReadMoqVarInt(&unsubscribe.request_id)) {
-    return absl::InvalidArgumentError("Message missing fields");
-  }
-  QUICHE_RETURN_IF_ERROR(CheckForTrailingData(reader));
-  return unsubscribe;
-}
-
 absl::StatusOr<MoqtPublishDone> MoqtControlMessageParser::ProcessPublishDone(
     absl::string_view data) const {
   quic::QuicDataReader reader(data);
@@ -1002,8 +991,7 @@
   quic::QuicDataReader reader(data);
   MoqtObjectAck object_ack;
   uint64_t raw_delta;
-  if (!reader.ReadMoqVarInt(&object_ack.subscribe_id) ||
-      !reader.ReadMoqVarInt(&object_ack.group_id) ||
+  if (!reader.ReadMoqVarInt(&object_ack.group_id) ||
       !reader.ReadMoqVarInt(&object_ack.object_id) ||
       !reader.ReadMoqVarInt(&raw_delta)) {
     return absl::InvalidArgumentError("Message missing fields");
diff --git a/quiche/quic/moqt/moqt_parser.h b/quiche/quic/moqt/moqt_parser.h
index f36415c..eca79f4 100644
--- a/quiche/quic/moqt/moqt_parser.h
+++ b/quiche/quic/moqt/moqt_parser.h
@@ -127,8 +127,6 @@
   absl::StatusOr<MoqtSubscribe> ProcessSubscribe(absl::string_view data) const;
   absl::StatusOr<MoqtSubscribeOk> ProcessSubscribeOk(
       absl::string_view data) const;
-  absl::StatusOr<MoqtUnsubscribe> ProcessUnsubscribe(
-      absl::string_view data) const;
   absl::StatusOr<MoqtPublishDone> ProcessPublishDone(
       absl::string_view data) const;
   absl::StatusOr<MoqtRequestUpdate> ProcessRequestUpdate(
@@ -184,8 +182,6 @@
         return parse(&MoqtControlMessageParser::ProcessSubscribe);
       case MoqtMessageType::kSubscribeOk:
         return parse(&MoqtControlMessageParser::ProcessSubscribeOk);
-      case MoqtMessageType::kUnsubscribe:
-        return parse(&MoqtControlMessageParser::ProcessUnsubscribe);
       case MoqtMessageType::kPublishDone:
         return parse(&MoqtControlMessageParser::ProcessPublishDone);
       case MoqtMessageType::kRequestUpdate:
diff --git a/quiche/quic/moqt/moqt_parser_test.cc b/quiche/quic/moqt/moqt_parser_test.cc
index 7b9d69c..74fcd62 100644
--- a/quiche/quic/moqt/moqt_parser_test.cc
+++ b/quiche/quic/moqt/moqt_parser_test.cc
@@ -52,7 +52,6 @@
     MoqtMessageType::kSubscribe,
     MoqtMessageType::kSubscribeOk,
     MoqtMessageType::kRequestUpdate,
-    MoqtMessageType::kUnsubscribe,
     MoqtMessageType::kPublishDone,
     MoqtMessageType::kTrackStatus,
     MoqtMessageType::kPublishNamespace,
@@ -1312,8 +1311,8 @@
 
 TEST_F(MoqtMessageSpecificTest, ObjectAckNegativeDelta) {
   char object_ack[] = {
-      0xb1, 0x84, 0x00, 0x05,  // type
-      0x01, 0x10, 0x20,        // subscribe ID, group, object
+      0xb1, 0x84, 0x00, 0x04,  // type
+      0x10, 0x20,              // group, object
       0x80, 0x81,              // -0x40 time delta
   };
   absl::StatusOr<std::vector<AnyMoqtControlMessage>> parsed =
@@ -1322,7 +1321,6 @@
   ASSERT_TRUE(parsed.ok());
   ASSERT_EQ(parsed->size(), 1);
   MoqtObjectAck message = std::get<MoqtObjectAck>((*parsed)[0]);
-  EXPECT_EQ(message.subscribe_id, 0x01);
   EXPECT_EQ(message.group_id, 0x10);
   EXPECT_EQ(message.object_id, 0x20);
   EXPECT_EQ(message.delta_from_deadline,
diff --git a/quiche/quic/moqt/moqt_publish_stream.cc b/quiche/quic/moqt/moqt_publish_stream.cc
index dd7c36e..1ffe56c 100644
--- a/quiche/quic/moqt/moqt_publish_stream.cc
+++ b/quiche/quic/moqt/moqt_publish_stream.cc
@@ -38,7 +38,12 @@
       response_callback_(std::move(response_callback)),
       stream_deleted_callback_(std::move(stream_deleted_callback)) {}
 
-MoqtPublishPublisherStream::~MoqtPublishPublisherStream() { Detach(); }
+MoqtPublishPublisherStream::~MoqtPublishPublisherStream() {
+  if (publisher_ != nullptr) {
+    publisher_->IgnoreResetAllStreams();
+  }
+  Detach();
+}
 
 void MoqtPublishPublisherStream::OnStreamBound() {
   stream_parser()->set_allow_fin(true);
@@ -165,6 +170,11 @@
                              }},
               response);
         }));
+  } else {
+    // Since the application already called SUBSCRIBE, there will be no
+    // invocation of the request callback. Send REQUEST_OK immediately.
+    CheckStatus(
+        SendRequestOk(message.request_id, subscriber_->const_parameters()));
   }
   incoming_publish_callback_ = nullptr;
   if (subscriber_->visitor() == nullptr) {
@@ -183,6 +193,10 @@
 
 absl::Status MoqtPublishSubscriberStream::OnControlMessage(
     const MoqtRequestUpdate& message) {
+  if (subscriber_ == nullptr) {
+    // Stream is already closing.
+    return absl::OkStatus();
+  }
   subscriber_->Update(message.parameters);
   CheckStatus(SendRequestOk(message.request_id, MessageParameters()));
   return absl::OkStatus();
diff --git a/quiche/quic/moqt/moqt_publish_stream.h b/quiche/quic/moqt/moqt_publish_stream.h
index a3cef79..43e6dbe 100644
--- a/quiche/quic/moqt/moqt_publish_stream.h
+++ b/quiche/quic/moqt/moqt_publish_stream.h
@@ -48,6 +48,10 @@
   absl::Status OnControlMessage(const MoqtRequestOk& message);
   absl::Status OnControlMessage(const MoqtRequestError& message);
   absl::Status OnControlMessage(const MoqtRequestUpdate& message);
+  absl::Status OnControlMessage(const MoqtObjectAck& message) {
+    publisher_->ProcessObjectAck(message);
+    return absl::OkStatus();
+  }
 
   void SetPublisher(std::unique_ptr<SubscriptionPublisher> publisher) {
     publisher_ = std::move(publisher);
@@ -61,6 +65,8 @@
         std::move(stream_deleted_callback_);
     stream_deleted_callback_ = nullptr;
     std::move(callback)(publisher_.get());
+    publisher_->ResetAllStreams();
+    publisher_ = nullptr;
   }
 
  private:
@@ -96,6 +102,8 @@
   absl::Status OnControlMessage(const MoqtRequestError& message);
   absl::Status OnControlMessage(const MoqtPublishDone& message);
 
+  SubscribeRemoteTrack* track() { return subscriber_.get(); }
+
   void Detach() override {
     if (remove_callback_ != nullptr) {
       SubscribeRemoteTrack::RemoveCallback callback =
@@ -103,6 +111,7 @@
       remove_callback_ = nullptr;
       std::move(callback)(subscriber_.get());
     }
+    subscriber_ = nullptr;
   }
 
  private:
diff --git a/quiche/quic/moqt/moqt_publish_stream_test.cc b/quiche/quic/moqt/moqt_publish_stream_test.cc
index e59cb95..9de992f 100644
--- a/quiche/quic/moqt/moqt_publish_stream_test.cc
+++ b/quiche/quic/moqt/moqt_publish_stream_test.cc
@@ -14,7 +14,6 @@
 #include "absl/status/status.h"
 #include "absl/strings/string_view.h"
 #include "absl/types/span.h"
-#include "quiche/quic/core/quic_alarm_factory.h"
 #include "quiche/quic/core/quic_time.h"
 #include "quiche/quic/core/quic_types.h"
 #include "quiche/quic/moqt/moqt_error.h"
@@ -31,6 +30,7 @@
 #include "quiche/quic/moqt/moqt_trace_recorder.h"
 #include "quiche/quic/moqt/moqt_track.h"
 #include "quiche/quic/moqt/moqt_types.h"
+#include "quiche/quic/moqt/test_tools/mock_moqt_session.h"
 #include "quiche/quic/moqt/test_tools/moqt_framer_utils.h"
 #include "quiche/quic/moqt/test_tools/moqt_mock_visitor.h"
 #include "quiche/quic/test_tools/mock_clock.h"
@@ -66,18 +66,6 @@
 using ::testing::Return;
 using ::testing::StrictMock;
 
-class MockSessionToPublisherInterface : public SessionToPublisherInterface {
- public:
-  ~MockSessionToPublisherInterface() override = default;
-  MOCK_METHOD(bool, alternate_delivery_timeout, (), (const, override));
-  MOCK_METHOD(void, UpdateTrackPriority,
-              (uint64_t, std::optional<MoqtTrackPriority>, MoqtTrackPriority),
-              (override));
-  MOCK_METHOD(quic::QuicAlarmFactory*, alarm_factory, (), (override));
-  MOCK_METHOD(void, PublishIsDone, (uint64_t), (override));
-  MOCK_METHOD(webtransport::Session*, session, (), (override));
-};
-
 constexpr uint64_t kRequestId = 1;
 constexpr uint64_t kTrackAlias = 10;
 const FullTrackName kTrackName("foo", "bar");
@@ -105,8 +93,7 @@
     EXPECT_CALL(visitor_, session).WillRepeatedly(Return(&webtrans_));
     auto publisher = std::make_unique<SubscriptionPublisher>(
         framer_, track_publisher_, stream_.get(), kRequestId, kTrackAlias,
-        parameters_, &visitor_, /*monitoring_interface=*/nullptr, &mock_clock_,
-        trace_recorder_, /*is_publish=*/true);
+        parameters_, visitor_.weak_ptr_factory_.Create(), /*is_publish=*/true);
 
     publisher_ = publisher.get();  // Keep raw pointer for testing
     stream_->SetPublisher(std::move(publisher));
@@ -229,6 +216,30 @@
   EXPECT_EQ(pub_params.subscription_filter->start(), Location(1, 3));
 }
 
+TEST_F(MoqtPublishPublisherStreamTest, ReceiveObjectAck) {
+  EXPECT_CALL(mock_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kPublish), _))
+      .WillOnce(Return(absl::OkStatus()));
+  stream_->BindStream(&mock_stream_);
+
+  MoqtObjectAck ack;
+  ack.group_id = 1;
+  ack.object_id = 2;
+  ack.delta_from_deadline = quic::QuicTimeDelta::FromMilliseconds(100);
+
+  EXPECT_CALL(visitor_, trace_recorder())
+      .WillOnce(testing::ReturnRef(trace_recorder_));
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(ack));
+}
+
+TEST_F(MoqtPublishPublisherStreamTest, Detach) {
+  EXPECT_CALL(deleted_callback_, Call(publisher_));
+  stream_->Detach();
+  // Verifying second detach is a no-op
+  EXPECT_CALL(deleted_callback_, Call).Times(0);
+  stream_->Detach();
+}
+
 class MoqtPublishSubscriberStreamTest : public quiche::test::QuicheTest {
  public:
   MoqtPublishSubscriberStreamTest()
@@ -561,5 +572,34 @@
   EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(kTrackName));
 }
 
+TEST_F(MoqtPublishSubscriberStreamTest, ReceiveRequestOkAndErrorTodo) {
+  MoqtRequestOk request_ok;
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(request_ok));
+
+  MoqtRequestError request_error;
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(request_error));
+}
+
+TEST_F(MoqtPublishSubscriberStreamTest, TrackAndDetach) {
+  MoqtPublish publish = DefaultPublish();
+  EXPECT_CALL(incoming_publish_callback_mock_, Call(kTrackName, _, _, _))
+      .WillOnce(Return(&mock_subscribe_visitor_));
+  EXPECT_CALL(mock_add_callback_, Call(NotNull())).WillOnce(Return(true));
+  EXPECT_CALL(mock_subscribe_visitor_, OnReply(kTrackName, _))
+      .WillOnce(
+          [](const FullTrackName&,
+             const std::variant<SubscribeOkData, MoqtRequestErrorInfo>& reply) {
+            EXPECT_TRUE(std::holds_alternative<SubscribeOkData>(reply));
+          });
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(publish));
+  EXPECT_NE(stream_->track(), nullptr);
+  EXPECT_EQ(stream_->track()->track_alias(), kTrackAlias);
+
+  EXPECT_CALL(mock_remove_callback_, Call(stream_->track()));
+  EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(kTrackName));
+  stream_->Detach();
+  EXPECT_EQ(stream_->track(), nullptr);
+}
+
 }  // namespace
 }  // namespace moqt::test
diff --git a/quiche/quic/moqt/moqt_session.cc b/quiche/quic/moqt/moqt_session.cc
index f7f8333..9dd7d62 100644
--- a/quiche/quic/moqt/moqt_session.cc
+++ b/quiche/quic/moqt/moqt_session.cc
@@ -4,13 +4,13 @@
 
 #include "quiche/quic/moqt/moqt_session.h"
 
+#include <array>
 #include <cstdint>
 #include <memory>
 #include <optional>
 #include <string>
 #include <utility>
 #include <variant>
-#include <vector>
 
 #include "absl/base/casts.h"
 #include "absl/base/nullability.h"
@@ -24,6 +24,7 @@
 #include "absl/status/status.h"
 #include "absl/status/statusor.h"
 #include "absl/strings/string_view.h"
+#include "absl/types/span.h"
 #include "quiche/quic/core/quic_alarm_factory.h"
 #include "quiche/quic/core/quic_time.h"
 #include "quiche/quic/core/quic_types.h"
@@ -42,6 +43,7 @@
 #include "quiche/quic/moqt/moqt_publisher.h"
 #include "quiche/quic/moqt/moqt_session_callbacks.h"
 #include "quiche/quic/moqt/moqt_session_interface.h"
+#include "quiche/quic/moqt/moqt_subscribe_stream.h"
 #include "quiche/quic/moqt/moqt_subscription.h"
 #include "quiche/quic/moqt/moqt_track.h"
 #include "quiche/quic/moqt/moqt_types.h"
@@ -50,7 +52,7 @@
 #include "quiche/common/platform/api/quiche_bug_tracker.h"
 #include "quiche/common/platform/api/quiche_logging.h"
 #include "quiche/common/quiche_buffer_allocator.h"
-#include "quiche/common/quiche_status_utils.h"
+#include "quiche/common/quiche_mem_slice.h"
 #include "quiche/common/quiche_weak_ptr.h"
 #include "quiche/web_transport/web_transport.h"
 
@@ -91,6 +93,7 @@
       local_max_request_id_(parameters.max_request_id),
       alarm_factory_(std::move(alarm_factory)),
       weak_ptr_factory_(this),
+      weak_ptr_factory_for_publishers_(this),
       liveness_token_(std::make_shared<Empty>()) {
   if (parameters_.using_webtrans) {
     session_->SetOnDraining([this]() {
@@ -165,6 +168,23 @@
 void MoqtSession::OnIncomingBidirectionalStreamAvailable() {
   while (webtransport::Stream* stream =
              session_->AcceptIncomingBidirectionalStream()) {
+    if (sent_goaway_) {
+      // Immediately reject new requests with REQUEST_ERROR. If the stream
+      // cannot be written, just reset it.
+      if (!stream->CanWrite()) {
+        stream->ResetWithUserCode(kResetCodeSessionClosed);
+        continue;
+      }
+      webtransport::StreamWriteOptions options;
+      options.set_send_fin(true);
+      std::array write_vector = {
+          quiche::QuicheMemSlice(framer_.SerializeRequestError(MoqtRequestError{
+              0, RequestErrorCode::kGoingAway, std::nullopt, ""}))};
+      if (!stream->Writev(absl::MakeSpan(write_vector), options).ok()) {
+        stream->ResetWithUserCode(kResetCodeSessionClosed);
+      };
+      continue;
+    }
     auto bidi_stream = std::make_unique<UnknownBidiStream>(this, stream);
     stream->SetVisitor(std::move(bidi_stream));
     stream->visitor()->OnCanRead();
@@ -196,7 +216,7 @@
                     << message.object_id << " priority "
                     << message.publisher_priority << " length "
                     << payload->size();
-  SubscribeRemoteTrack* track = RemoteTrackByAlias(message.track_alias);
+  SubscribeRemoteTrack* track = SubscribeByAlias(message.track_alias);
   if (track == nullptr) {
     return;
   }
@@ -444,10 +464,9 @@
 }
 
 bool MoqtSession::Subscribe(const FullTrackName& name,
-                            SubscribeVisitor* visitor,
+                            SubscribeVisitor* absl_nonnull visitor,
                             const MessageParameters& parameters) {
   QUICHE_DCHECK(name.IsValid());
-
   if (next_request_id_ >= peer_max_request_id_) {
     if (!last_requests_blocked_sent_.has_value() ||
         peer_max_request_id_ > *last_requests_blocked_sent_) {
@@ -471,24 +490,19 @@
     QUIC_DLOG(INFO) << ENDPOINT << "Tried to send SUBSCRIBE after GOAWAY";
     return false;
   }
-  MoqtSubscribe message(next_request_id_, name, parameters);
-  next_request_id_ += 2;
-  if (SupportsObjectAck() && visitor != nullptr) {
-    // Since we do not expose subscribe IDs directly in the API, instead wrap
-    // the session and subscribe ID in a callback.
-    visitor->OnCanAckObjects(absl::bind_front(&MoqtSession::SendObjectAck, this,
-                                              message.request_id));
-  } else {
-    QUICHE_DLOG_IF(WARNING, message.parameters.oack_window_size.has_value())
-        << "Attempting to set object_ack_window on a connection that does not "
-           "support it.";
-    message.parameters.oack_window_size = std::nullopt;
+  if (!session_->CanOpenNextOutgoingBidirectionalStream()) {
+    return false;  // Do not retry opening a SUBSCRIBE stream.
   }
-  SendControlMessage(framer_.SerializeSubscribe(message));
-  QUIC_DLOG(INFO) << ENDPOINT << "Sent SUBSCRIBE message for "
-                  << message.full_track_name;
-  auto track = std::make_unique<SubscribeRemoteTrack>(
-      message, visitor,
+  auto stream_visitor = std::make_unique<MoqtSubscribeRequestStream>(
+      &framer_, ControlMessageParser(), NextRequestId(),
+      [weak_session = GetWeakPtr()](MoqtError code, absl::string_view reason) {
+        MoqtSessionInterface* session = weak_session.GetIfAvailable();
+        if (session == nullptr) {
+          return;
+        }
+        session->Error(code, reason);
+      },
+      name, visitor, parameters,
       [weakptr = GetWeakPtr()](SubscribeRemoteTrack* track) {
         MoqtSession* session = MoqtSessionFromWeakPtr(weakptr);
         if (session == nullptr || !track->track_alias().has_value()) {
@@ -507,10 +521,18 @@
         if (track->track_alias().has_value()) {
           session->subscribe_by_alias_.erase(*track->track_alias());
         }
-        session->upstream_by_id_.erase(track->request_id());
-      });
-  subscribe_by_name_.emplace(message.full_track_name, track.get());
-  upstream_by_id_.emplace(message.request_id, std::move(track));
+      },
+      callbacks_.clock, alarm_factory_.get());
+  webtransport::Stream* stream = session_->OpenOutgoingBidirectionalStream();
+  QUICHE_CHECK(stream != nullptr);
+  MoqtSubscribeRequestStream* stream_visitor_ptr = stream_visitor.get();
+  stream->SetVisitor(std::move(stream_visitor));
+  stream_visitor_ptr->BindStream(stream);
+  subscribe_by_name_[name] = stream_visitor_ptr->track();
+  if (SupportsObjectAck()) {
+    visitor->OnCanAckObjects(
+        absl::bind_front(&MoqtSession::SendObjectAck, this, name));
+  }
   return true;
 }
 
@@ -518,32 +540,22 @@
                                   const MessageParameters& parameters,
                                   MoqtResponseCallback response_callback) {
   QUICHE_DCHECK(name.IsValid());
-  if (next_request_id_ >= peer_max_request_id_) {
-    if (!last_requests_blocked_sent_.has_value() ||
-        peer_max_request_id_ > *last_requests_blocked_sent_) {
-      MoqtRequestsBlocked requests_blocked;
-      requests_blocked.max_request_id = peer_max_request_id_;
-      SendControlMessage(framer_.SerializeRequestsBlocked(requests_blocked));
-      last_requests_blocked_sent_ = peer_max_request_id_;
-    }
-    QUIC_DLOG(INFO) << ENDPOINT << "Tried to send SUBSCRIBE with ID "
-                    << next_request_id_
-                    << " which is greater than the maximum ID "
-                    << peer_max_request_id_;
-    return false;
-  }
   auto it = subscribe_by_name_.find(name);
   if (it == subscribe_by_name_.end()) {
     return false;
   }
-  // TODO(martinduke): Support Update on PUBLISH streams.
-  pending_subscribe_updates_[next_request_id_] = {name, parameters,
-                                                  std::move(response_callback)};
-  MoqtRequestUpdate update{next_request_id_, it->second->request_id(),
-                           parameters};
-  next_request_id_ += 2;
-  SendControlMessage(framer_.SerializeRequestUpdate(update));
-  return true;
+  // sending zero because related request ID is ignored for SUBSCRIBE.
+  return it->second->request_stream()
+      ->SendRequestUpdate(NextRequestId(), 0, parameters,
+                          std::move(response_callback))
+      .ok();
+}
+
+bool MoqtSession::PublishUpdate(const FullTrackName& name,
+                                const MessageParameters& parameters,
+                                MoqtResponseCallback response_callback) {
+  // TODO(martinduke): Implement this.
+  return false;
 }
 
 void MoqtSession::Unsubscribe(const FullTrackName& name) {
@@ -551,16 +563,11 @@
     return;
   }
   QUICHE_DCHECK(name.IsValid());
-  SubscribeRemoteTrack* track = RemoteTrackByName(name);
-  if (track == nullptr) {
+  auto it = subscribe_by_name_.find(name);
+  if (it == subscribe_by_name_.end()) {
     return;
   }
-  QUICHE_DCHECK(name.IsValid());
-  QUIC_DLOG(INFO) << ENDPOINT << "Sent UNSUBSCRIBE message for " << name;
-  MoqtUnsubscribe message;
-  message.request_id = track->request_id();
-  SendControlMessage(framer_.SerializeUnsubscribe(message));
-  track->Destroy();
+  it->second->request_stream()->Reset(kResetCodeCancelled);
 }
 
 bool MoqtSession::Publish(
@@ -576,10 +583,16 @@
   if (!session_->CanOpenNextOutgoingBidirectionalStream()) {
     return false;  // Do not retry opening a PUBLISH stream.
   }
-  if (!subscribed_track_names_.insert(name).second) {
-    QUICHE_DLOG(INFO) << ENDPOINT << "Tried to send PUBLISH for track " << name
-                      << " which is already published";
-    return false;
+  auto it = subscribed_track_names_.find(name);
+  if (it != subscribed_track_names_.end()) {
+    if (it->second->established()) {
+      QUICHE_DLOG(INFO) << ENDPOINT << "Tried to send PUBLISH for track "
+                        << name << " which is already published";
+      return false;
+    }
+    it->second->OnSubscribeRejected(
+        MoqtRequestErrorInfo{RequestErrorCode::kDuplicateSubscription,
+                             std::nullopt, "PUBLISH is coming"});
   }
   auto stream_visitor = std::make_unique<MoqtPublishPublisherStream>(
       &framer_, ControlMessageParser(),
@@ -602,8 +615,8 @@
       std::move(response_callback));
   auto publish_state = std::make_unique<SubscriptionPublisher>(
       framer_, publisher, stream_visitor.get(), next_request_id_,
-      next_local_track_alias_, parameters, this, nullptr, callbacks_.clock,
-      trace_recorder_, true);
+      next_local_track_alias_, parameters,
+      weak_ptr_factory_for_publishers_.Create(), true);
   SubscriptionPublisher* publisher_ptr = publish_state.get();
   stream_visitor->SetPublisher(std::move(publish_state));
   webtransport::Stream* stream = session_->OpenOutgoingBidirectionalStream();
@@ -644,10 +657,10 @@
   QUIC_DLOG(INFO) << ENDPOINT << "Sent FETCH message for " << name;
   auto fetch = std::make_unique<UpstreamFetch>(
       message, std::get<StandaloneFetch>(message.fetch), std::move(callback),
-      [this, id = message.request_id]() {  // Deletion callback
-        upstream_by_id_.erase(id);
+      [this, id = message.request_id]() {
+        fetch_by_id_.erase(id);  // Deletion callback
       });
-  upstream_by_id_.emplace(message.request_id, std::move(fetch));
+  fetch_by_id_.emplace(message.request_id, std::move(fetch));
   return true;
 }
 
@@ -658,15 +671,13 @@
   QUICHE_DCHECK(name.IsValid());
   return RelativeJoiningFetch(
       name, visitor,
-      [this, id = next_request_id_](std::unique_ptr<MoqtFetchTask> fetch_task) {
+      [this, track_name = name](std::unique_ptr<MoqtFetchTask> fetch_task) {
         // Move the fetch_task to the subscribe to plumb into its visitor.
-        RemoteTrack* track = RemoteTrackById(id);
-        if (track == nullptr || track->is_fetch()) {
+        SubscribeRemoteTrack* subscribe = SubscribeByName(track_name);
+        if (subscribe == nullptr || subscribe->is_fetch()) {
           fetch_task.release();
           return;
         }
-        auto* subscribe = absl::down_cast<SubscribeRemoteTrack*>(track);
-        RemoteTrackByName(track->full_track_name());
         subscribe->OnJoiningFetchReady(std::move(fetch_task));
       },
       num_previous_groups, parameters);
@@ -702,9 +713,9 @@
   auto upstream_fetch = std::make_unique<UpstreamFetch>(
       fetch, name, std::move(callback),
       /*Deletion callback=*/[this, id = fetch.request_id]() {
-        upstream_by_id_.erase(id);
+        fetch_by_id_.erase(id);
       });
-  upstream_by_id_.emplace(fetch.request_id, std::move(upstream_fetch));
+  fetch_by_id_.emplace(fetch.request_id, std::move(upstream_fetch));
   return true;
 }
 
@@ -733,19 +744,6 @@
                   "Peer did not close session after GOAWAY");
 }
 
-void MoqtSession::PublishIsDone(uint64_t request_id) {
-  if (is_closing_) {
-    return;
-  }
-  auto it = published_subscriptions_.find(request_id);
-  if (it == published_subscriptions_.end()) {
-    // If a PUBLISH, we will end up here.
-    return;
-  }
-  subscribed_track_names_.erase(it->second->publisher().GetTrackName());
-  published_subscriptions_.erase(it);
-}
-
 void MoqtSession::UpdateTrackPriority(
     uint64_t request_id, std::optional<MoqtTrackPriority> old_priority,
     MoqtTrackPriority new_priority) {
@@ -762,6 +760,25 @@
   subscriptions_with_queued_streams_.emplace(new_priority, request_id);
 }
 
+std::shared_ptr<MoqtTrackPublisher> MoqtSession::GetTrackPublisher(
+    const FullTrackName& name) {
+  if (publisher_ == nullptr) {
+    return nullptr;
+  }
+  return publisher_->GetTrack(name);
+}
+
+MoqtPublishingMonitorInterface* MoqtSession::ReleaseMonitoringInterface(
+    const FullTrackName& name) {
+  auto it = monitoring_interfaces_for_published_tracks_.find(name);
+  if (it == monitoring_interfaces_for_published_tracks_.end()) {
+    return nullptr;
+  }
+  MoqtPublishingMonitorInterface* interface = it->second;
+  monitoring_interfaces_for_published_tracks_.erase(it);
+  return interface;
+}
+
 bool MoqtSession::OpenDataStream(PublishedFetch* fetch,
                                  webtransport::SendOrder send_order) {
   webtransport::Stream* new_stream =
@@ -793,7 +810,7 @@
   return true;
 }
 
-SubscribeRemoteTrack* MoqtSession::RemoteTrackByAlias(uint64_t track_alias) {
+SubscribeRemoteTrack* MoqtSession::SubscribeByAlias(uint64_t track_alias) {
   auto it = subscribe_by_alias_.find(track_alias);
   if (it == subscribe_by_alias_.end()) {
     return nullptr;
@@ -801,24 +818,23 @@
   return it->second;
 }
 
-RemoteTrack* MoqtSession::RemoteTrackById(uint64_t request_id) {
-  auto it = upstream_by_id_.find(request_id);
-  if (it == upstream_by_id_.end()) {
-    return nullptr;
-  }
-  return it->second.get();
-}
-
-SubscribeRemoteTrack* MoqtSession::RemoteTrackByName(
-    const FullTrackName& name) {
-  QUICHE_DCHECK(name.IsValid());
-  auto it = subscribe_by_name_.find(name);
+SubscribeRemoteTrack* MoqtSession::SubscribeByName(
+    const FullTrackName& track_name) {
+  auto it = subscribe_by_name_.find(track_name);
   if (it == subscribe_by_name_.end()) {
     return nullptr;
   }
   return it->second;
 }
 
+UpstreamFetch* MoqtSession::FetchById(uint64_t request_id) {
+  auto it = fetch_by_id_.find(request_id);
+  if (it == fetch_by_id_.end()) {
+    return nullptr;
+  }
+  return it->second.get();
+}
+
 void MoqtSession::OnCanCreateNewOutgoingUnidirectionalStream() {
   while (!subscriptions_with_queued_streams_.empty() &&
          session_->CanOpenNextOutgoingUnidirectionalStream()) {
@@ -863,8 +879,9 @@
     Error(MoqtError::kInvalidRequestId, "Request ID evenness incorrect");
     return false;
   }
-  if (published_subscriptions_.contains(request_id) ||
-      incoming_fetches_.contains(request_id) ||
+  // TODO(martinduke): Write new checks for duplicate request IDs. It's
+  // probably best to track the largest observed plus a set of holes.
+  if (incoming_fetches_.contains(request_id) ||
       incoming_track_status_.contains(request_id) ||
       incoming_publish_namespaces_by_id_.contains(request_id)) {
     QUICHE_DLOG(INFO) << ENDPOINT << "Duplicate request ID";
@@ -978,7 +995,9 @@
             auto it =
                 session->subscribe_by_name_.find(track->full_track_name());
             if (it != session->subscribe_by_name_.end()) {
-              // It's a pending SUBSCRIBE; kill it.
+              // It's a pending SUBSCRIBE; kill it, but use the parameters and
+              // visitor from the SUBSCRIBE.
+              track->Update(it->second->const_parameters());
               track->set_visitor(it->second->ReleaseVisitor());
               session->Unsubscribe(it->second->full_track_name());
             }
@@ -1005,6 +1024,51 @@
       temp_stream->OnCanRead();
       break;
     }
+    case MoqtMessageType::kSubscribe: {
+      auto subscribe_stream = std::make_unique<MoqtSubscribeResponseStream>(
+          &session_->framer_, session_->ControlMessageParser(),
+          session_->next_local_track_alias_++,
+          [weakptr =
+               session_->GetWeakPtr()](SubscriptionPublisher* subscription) {
+            MoqtSession* session = MoqtSessionFromWeakPtr(weakptr);
+            if (session == nullptr) {
+              return true;
+            }
+            auto [it, success] = session->published_subscriptions_.try_emplace(
+                subscription->request_id(), subscription);
+            if (!success) {
+              return false;
+            }
+            auto [it2, success2] = session->subscribed_track_names_.try_emplace(
+                subscription->publisher().GetTrackName(), subscription);
+            return success2;
+          },
+          [weakptr =
+               session_->GetWeakPtr()](SubscriptionPublisher* subscription) {
+            MoqtSession* session = MoqtSessionFromWeakPtr(weakptr);
+            if (session == nullptr) {
+              return;
+            }
+            session->published_subscriptions_.erase(subscription->request_id());
+            session->subscribed_track_names_.erase(
+                subscription->publisher().GetTrackName());
+          },
+          [weakptr = session_->GetWeakPtr()](MoqtError code,
+                                             absl::string_view reason) {
+            MoqtSession* session = MoqtSessionFromWeakPtr(weakptr);
+            if (session != nullptr) {
+              session->Error(code, reason);
+            }
+          },
+          session_->weak_ptr_factory_for_publishers_.Create());
+      subscribe_stream->BindStream(std::move(parser_));
+      MoqtSubscribeResponseStream* temp_stream = subscribe_stream.get();
+      stream_->SetVisitor(std::move(subscribe_stream));
+      // The UnknownBidiStream object is deleted; no class access after this
+      // point.
+      temp_stream->OnCanRead();
+      break;
+    }
     default:
       session_->Error(MoqtError::kProtocolViolation,
                       "Unexpected message type received to start bidi stream");
@@ -1050,113 +1114,9 @@
   }
 }
 
-absl::Status MoqtSession::OnControlMessage(const MoqtSubscribe& message) {
-  if (!ValidateRequestId(message.request_id)) {
-    return absl::OkStatus();
-  }
-  QUIC_DLOG(INFO) << ENDPOINT << "Received a SUBSCRIBE for "
-                  << message.full_track_name;
-  if (sent_goaway_) {
-    QUIC_DLOG(INFO) << ENDPOINT << "Received a SUBSCRIBE after GOAWAY";
-    SendRequestErrorOnControlStream(message.request_id,
-                                    RequestErrorCode::kUnauthorized,
-                                    std::nullopt, "SUBSCRIBE after GOAWAY");
-    return absl::OkStatus();
-  }
-  if (subscribed_track_names_.contains(message.full_track_name)) {
-    SendRequestErrorOnControlStream(message.request_id,
-                                    RequestErrorCode::kDuplicateSubscription,
-                                    std::nullopt, "");
-    return absl::OkStatus();
-  }
-  const FullTrackName& track_name = message.full_track_name;
-  std::shared_ptr<MoqtTrackPublisher> track_publisher =
-      publisher_->GetTrack(track_name);
-  if (track_publisher == nullptr) {
-    QUIC_DLOG(INFO) << ENDPOINT << "SUBSCRIBE for " << track_name
-                    << " rejected by the application: does not exist";
-    SendRequestErrorOnControlStream(message.request_id,
-                                    RequestErrorCode::kDoesNotExist,
-                                    std::nullopt, "not found");
-    return absl::OkStatus();
-  }
-
-  MoqtPublishingMonitorInterface* monitoring = nullptr;
-  auto monitoring_it =
-      monitoring_interfaces_for_published_tracks_.find(track_name);
-  if (monitoring_it != monitoring_interfaces_for_published_tracks_.end()) {
-    monitoring = monitoring_it->second;
-    monitoring_interfaces_for_published_tracks_.erase(monitoring_it);
-  }
-
-  MoqtTrackPublisher* track_publisher_ptr = track_publisher.get();
-  auto subscription = std::make_unique<SubscriptionPublisher>(
-      framer_, track_publisher, GetControlStream(), message.request_id,
-      next_local_track_alias_++, message.parameters, this, monitoring,
-      callbacks_.clock, trace_recorder_, false);
-  SubscriptionPublisher* subscription_ptr = subscription.get();
-  auto [it, success] = published_subscriptions_.emplace(
-      message.request_id, std::move(subscription));
-  if (!success) {
-    QUICHE_NOTREACHED();  // ValidateRequestId() should have caught this.
-  }
-  subscribed_track_names_.insert(message.full_track_name);
-  track_publisher_ptr->AddObjectListener(subscription_ptr);
-  return absl::OkStatus();
-}
-
-absl::Status MoqtSession::OnControlMessage(const MoqtSubscribeOk& message) {
-  RemoteTrack* track = RemoteTrackById(message.request_id);
-  if (track == nullptr) {
-    QUIC_DLOG(INFO) << ENDPOINT << "Received the SUBSCRIBE_OK for "
-                    << "request_id = " << message.request_id
-                    << " but no track exists";
-    // Subscription state might have been destroyed for internal reasons.
-    return absl::OkStatus();
-  }
-  if (track->is_fetch()) {
-    return absl::InvalidArgumentError("Received SUBSCRIBE_OK for a FETCH");
-  }
-  if (message.parameters.largest_object.has_value()) {
-    QUIC_DLOG(INFO) << ENDPOINT << "Received the SUBSCRIBE_OK for "
-                    << "request_id = " << message.request_id << " "
-                    << track->full_track_name()
-                    << " largest_id = " << *message.parameters.largest_object;
-  } else {
-    QUIC_DLOG(INFO) << ENDPOINT << "Received the SUBSCRIBE_OK for "
-                    << "request_id = " << message.request_id << " "
-                    << track->full_track_name();
-  }
-  SubscribeRemoteTrack* subscribe =
-      absl::down_cast<SubscribeRemoteTrack*>(track);
-  if (!subscribe->set_track_alias(message.track_alias)) {
-    return absl::AlreadyExistsError("Duplicate track alias");
-  }
-  subscribe->OnObjectOrOk(
-      SubscribeOkData(message.parameters, message.extensions));
-  return absl::OkStatus();
-}
-
 absl::Status MoqtSession::OnControlMessage(const MoqtRequestOk& message) {
-  if (upstream_by_id_.contains(message.request_id)) {
-    return absl::InvalidArgumentError(
-        "Received REQUEST_OK for SUBSCRIBE, FETCH, or PUBLISH");
-  }
-  // Response to REQUEST_UPDATE for a subscribe.
-  auto ru_it = pending_subscribe_updates_.find(message.request_id);
-  if (ru_it != pending_subscribe_updates_.end()) {
-    auto sub_it = subscribe_by_name_.find(ru_it->second.name);
-    if (sub_it == subscribe_by_name_.end()) {
-      std::move(ru_it->second.response_callback)(
-          MoqtRequestErrorInfo{RequestErrorCode::kDoesNotExist, std::nullopt,
-                               "subscription does not exist anymore"});
-      pending_subscribe_updates_.erase(ru_it);
-      return absl::OkStatus();
-    }
-    sub_it->second->Update(ru_it->second.parameters);
-    std::move(ru_it->second.response_callback)(MessageParameters());
-    pending_subscribe_updates_.erase(ru_it);
-    return absl::OkStatus();
+  if (fetch_by_id_.contains(message.request_id)) {
+    return absl::InvalidArgumentError("Received REQUEST_OK for FETCH");
   }
   // Response to PUBLISH_NAMESPACE.
   auto pn_it = publish_namespace_by_id_.find(message.request_id);
@@ -1181,43 +1141,27 @@
   MoqtRequestErrorInfo error_info{message.error_code, message.retry_interval,
                                   message.reason_phrase};
   // TODO(martinduke): Do something with retry_interval.
-  RemoteTrack* track = RemoteTrackById(message.request_id);
-  if (track != nullptr) {
-    // It's in response to SUBSCRIBE or FETCH.
-    if (!track->ErrorIsAllowed()) {
+  UpstreamFetch* fetch = FetchById(message.request_id);
+  if (fetch != nullptr) {
+    // It's in response to FETCH.
+    if (!fetch->ErrorIsAllowed()) {
       return absl::InvalidArgumentError(
           "Received REQUEST_ERROR after REQUEST_OK or objects");
     }
     QUIC_DLOG(INFO) << ENDPOINT << "Received the REQUEST_ERROR for "
                     << "request_id = " << message.request_id << " ("
-                    << track->full_track_name() << ")"
+                    << fetch->full_track_name() << ")"
                     << ", error = " << static_cast<uint64_t>(message.error_code)
                     << " (" << message.reason_phrase << ")";
-    if (track->is_fetch()) {
-      UpstreamFetch* fetch = absl::down_cast<UpstreamFetch*>(track);
-      absl::Status status =
-          RequestErrorCodeToStatus(message.error_code, message.reason_phrase);
-      fetch->OnFetchResult(Location(0, 0), status, nullptr);
-    } else {
-      SubscribeRemoteTrack* subscribe =
-          absl::down_cast<SubscribeRemoteTrack*>(track);
-      if (subscribe->visitor() != nullptr) {
-        subscribe->visitor()->OnReply(subscribe->full_track_name(), error_info);
-      }
-    }
+    absl::Status status =
+        RequestErrorCodeToStatus(message.error_code, message.reason_phrase);
+    fetch->OnFetchResult(Location(0, 0), status, nullptr);
     if (!is_closing_) {
       // The visitor might have closed the session.
-      track->Destroy();
+      fetch->Destroy();
     }
     return absl::OkStatus();
   }
-  // Response to REQUEST_UPDATE for a subscribe.
-  auto ru_it = pending_subscribe_updates_.find(message.request_id);
-  if (ru_it != pending_subscribe_updates_.end()) {
-    std::move(ru_it->second.response_callback)(error_info);
-    pending_subscribe_updates_.erase(ru_it);
-    return absl::OkStatus();
-  }
   // Response to PUBLISH_NAMESPACE.
   auto pn_it = publish_namespace_by_id_.find(message.request_id);
   if (pn_it != publish_namespace_by_id_.end()) {
@@ -1238,41 +1182,7 @@
   return absl::OkStatus();
 }
 
-absl::Status MoqtSession::OnControlMessage(const MoqtUnsubscribe& message) {
-  auto it = published_subscriptions_.find(message.request_id);
-  if (it == published_subscriptions_.end()) {
-    return absl::OkStatus();
-  }
-  QUIC_DLOG(INFO) << ENDPOINT << "Received an UNSUBSCRIBE for "
-                  << it->second->publisher().GetTrackName();
-  PublishIsDone(message.request_id);
-  return absl::OkStatus();
-}
-
-absl::Status MoqtSession::OnControlMessage(const MoqtPublishDone& message) {
-  auto it = upstream_by_id_.find(message.request_id);
-  if (it == upstream_by_id_.end()) {
-    return absl::OkStatus();
-  }
-  auto* subscribe = absl::down_cast<SubscribeRemoteTrack*>(it->second.get());
-  QUIC_DLOG(INFO) << ENDPOINT << "Received a PUBLISH_DONE for "
-                  << it->second->full_track_name();
-  subscribe->OnPublishDone(message.stream_count, callbacks_.clock,
-                           alarm_factory_.get());
-  return absl::OkStatus();
-}
-
 absl::Status MoqtSession::OnControlMessage(const MoqtRequestUpdate& message) {
-  auto it = published_subscriptions_.find(message.existing_request_id);
-  if (it != published_subscriptions_.end()) {
-    // It's updating SUBSCRIBE.
-    it->second->Update(message.parameters);
-    // TODO(martinduke): There should be an MoqtResponseCallback sent to the
-    // application, rather than automatic OK.
-    SendControlMessage(framer_.SerializeRequestOk(
-        MoqtRequestOk{.request_id = message.request_id}));
-    return absl::OkStatus();
-  }
   auto pn_it = publish_namespace_by_id_.find(message.existing_request_id);
   if (pn_it != publish_namespace_by_id_.end()) {
     // It's updating PUBLISH_NAMESPACE.
@@ -1622,7 +1532,7 @@
 }
 
 absl::Status MoqtSession::OnControlMessage(const MoqtFetchOk& message) {
-  RemoteTrack* track = RemoteTrackById(message.request_id);
+  UpstreamFetch* track = FetchById(message.request_id);
   if (track == nullptr) {
     QUIC_DLOG(INFO) << ENDPOINT << "Received the FETCH_OK for "
                     << "request_id = " << message.request_id
@@ -1630,9 +1540,6 @@
     // Subscription state might have been destroyed for internal reasons.
     return absl::OkStatus();
   }
-  if (!track->is_fetch()) {
-    return absl::InvalidArgumentError("Received FETCH_OK for a SUBSCRIBE");
-  }
   QUIC_DLOG(INFO) << ENDPOINT << "Received the FETCH_OK for request_id = "
                   << message.request_id << " " << track->full_track_name();
   UpstreamFetch* fetch = absl::down_cast<UpstreamFetch*>(track);
@@ -1646,24 +1553,12 @@
   return absl::OkStatus();
 }
 
-absl::Status MoqtSession::OnControlMessage(const MoqtPublish& message) {
-  if (!ValidateRequestId(message.request_id)) {
-    return absl::OkStatus();
-  }
-  RequestErrorCode error_code = sent_goaway_ ? RequestErrorCode::kUnauthorized
-                                             : RequestErrorCode::kNotSupported;
-  absl::string_view error_reason = sent_goaway_
-                                       ? "Received a PUBLISH after GOAWAY"
-                                       : "PUBLISH is not supported";
-  SendRequestErrorOnControlStream(message.request_id, error_code, std::nullopt,
-                                  error_reason);
-  return absl::OkStatus();
-}
-
 void MoqtSession::OnMalformedTrack(RemoteTrack* track) {
   if (!track->is_fetch()) {
-    absl::down_cast<SubscribeRemoteTrack*>(track)->visitor()->OnMalformedTrack(
-        track->full_track_name());
+    auto* subscribe = absl::down_cast<SubscribeRemoteTrack*>(track);
+    if (subscribe->visitor() != nullptr) {
+      subscribe->visitor()->OnMalformedTrack(track->full_track_name());
+    }
     Unsubscribe(track->full_track_name());
     return;
   }
@@ -1687,7 +1582,6 @@
   // Incoming SUBSCRIBE_NAMESPACE is automatically cleaned up; the destroyed
   // session owns the webtransport stream, which owns the StreamVisitor, which
   // owns the task. Destroying the task notifies the application.
-  published_subscriptions_.clear();
   for (auto& it : incoming_publish_namespaces_by_namespace_) {
     callbacks_.incoming_publish_namespace_callback(it.first, std::nullopt,
                                                    nullptr);
@@ -1696,8 +1590,17 @@
     std::move(it.second.cancel_callback)(MoqtRequestErrorInfo{
         RequestErrorCode::kUninterested, std::nullopt, "Session closed"});
   }
-  while (!upstream_by_id_.empty()) {
-    upstream_by_id_.begin()->second->Destroy();
+  while (!fetch_by_id_.empty()) {
+    fetch_by_id_.begin()->second->Destroy();
+  }
+  for (auto& [track_name, subscriber] : subscribe_by_name_) {
+    // It's possible the application is going to destroy its visitor as early
+    // as session_deleted_callback is called. So call OnPublishDone() now and
+    // clear the visitor.
+    if (subscriber->visitor() != nullptr) {
+      subscriber->visitor()->OnPublishDone(subscriber->full_track_name());
+      subscriber->ReleaseVisitor();
+    }
   }
 }
 
@@ -1705,8 +1608,8 @@
   if (is_closing_) {
     return;
   }
-  auto it = upstream_by_id_.find(request_id);
-  if (it == upstream_by_id_.end()) {
+  auto it = fetch_by_id_.find(request_id);
+  if (it == fetch_by_id_.end()) {
     return;
   }
   it->second->Destroy();
diff --git a/quiche/quic/moqt/moqt_session.h b/quiche/quic/moqt/moqt_session.h
index d0bcb21..fd45f9d 100644
--- a/quiche/quic/moqt/moqt_session.h
+++ b/quiche/quic/moqt/moqt_session.h
@@ -87,11 +87,15 @@
   MoqtSessionCallbacks& callbacks() override { return callbacks_; }
   void Error(MoqtError code, absl::string_view error) override;
   // Returns false if the SUBSCRIBE isn't sent.
-  bool Subscribe(const FullTrackName& name, SubscribeVisitor* visitor,
+  bool Subscribe(const FullTrackName& name,
+                 SubscribeVisitor* absl_nonnull visitor,
                  const MessageParameters& parameters) override;
   bool SubscribeUpdate(const FullTrackName& name,
                        const MessageParameters& parameters,
                        MoqtResponseCallback response_callback) override;
+  bool PublishUpdate(const FullTrackName& name,
+                     const MessageParameters& parameters,
+                     MoqtResponseCallback response_callback) override;
   void Unsubscribe(const FullTrackName& name) override;
   bool Publish(std::shared_ptr<MoqtTrackPublisher> absl_nonnull publisher,
                const MessageParameters& parameters,
@@ -148,7 +152,12 @@
   quic::QuicAlarmFactory* alarm_factory() override {
     return alarm_factory_.get();
   }
-  void PublishIsDone(uint64_t request_id) override;
+  std::shared_ptr<MoqtTrackPublisher> GetTrackPublisher(
+      const FullTrackName& name) override;
+  MoqtPublishingMonitorInterface* ReleaseMonitoringInterface(
+      const FullTrackName& name) override;
+  const quic::QuicClock* clock() override { return callbacks_.clock; }
+  MoqtTraceRecorder& trace_recorder() override { return trace_recorder_; }
   webtransport::Session* session() override {
     return is_closing_ ? nullptr : session_;
   }
@@ -161,16 +170,17 @@
   // draft-ietf-moqt-moq-transport-12. Unsubscribe and notify the application so
   // the error can be propagated downstream, if necessary.
   void OnMalformedTrack(RemoteTrack* track);
-  quiche::QuicheWeakPtr<RemoteTrack> GetSubscribe(uint64_t track_alias) {
-    auto it = subscribe_by_alias_.find(track_alias);
-    if (it == subscribe_by_alias_.end()) {
+  quiche::QuicheWeakPtr<RemoteTrack> GetSubscribe(
+      uint64_t track_alias) override {
+    RemoteTrack* track = SubscribeByAlias(track_alias);
+    if (track == nullptr) {
       return quiche::QuicheWeakPtr<RemoteTrack>();
     }
-    return it->second->weak_ptr();
+    return track->weak_ptr();
   }
   quiche::QuicheWeakPtr<RemoteTrack> GetFetch(uint64_t request_id) {
-    auto it = upstream_by_id_.find(request_id);
-    if (it == upstream_by_id_.end()) {
+    auto it = fetch_by_id_.find(request_id);
+    if (it == fetch_by_id_.end()) {
       return quiche::QuicheWeakPtr<RemoteTrack>();
     }
     return it->second->weak_ptr();
@@ -208,8 +218,6 @@
 
   void UseAlternateDeliveryTimeout() { alternate_delivery_timeout_ = true; }
 
-  MoqtTraceRecorder& trace_recorder() { return trace_recorder_; }
-
  private:
   friend class ControlMessageDispatcher;
   friend class test::MoqtSessionPeer;
@@ -254,18 +262,12 @@
                 }
               }),
           session_(session),
-          weak_ptr_factory_(this) {
-      this->set_control_stream();
-    }
-    // TODO(martinduke): Remove constructor body once SUBSCRIBE moves to a bidi
-    // stream.
+          weak_ptr_factory_(this) {}
 
     void OnStreamBound() override;
     absl::Status OnRawControlMessage(
         const MoqtRawControlMessage& message) override;
 
-    // MoqtControlParserVisitor implementation.
-
     // webtransport::StreamVisitor overrides
     void OnResetStreamReceived(webtransport::StreamErrorCode error) override {
       session_->Error(MoqtError::kProtocolViolation,
@@ -338,7 +340,7 @@
       MessageParameters parameters;
       parameters.expires = publisher_->expiration();
       parameters.largest_object = publisher_->largest_location();
-      MoqtBidiStreamBase* control_stream = session_->GetControlStream();
+      ControlStream* control_stream = session_->GetControlStream();
       if (control_stream != nullptr) {
         control_stream->CheckStatus(
             control_stream->SendRequestOk(request_id_, parameters));
@@ -348,7 +350,7 @@
     }
 
     void OnSubscribeRejected(MoqtRequestErrorInfo info) override {
-      MoqtBidiStreamBase* control_stream = session_->GetControlStream();
+      ControlStream* control_stream = session_->GetControlStream();
       if (control_stream != nullptr) {
         control_stream->CheckStatus(control_stream->SendRequestError(
             request_id_, info.error_code, info.retry_interval,
@@ -397,10 +399,9 @@
   // Returns false if creation failed.
   [[nodiscard]] bool OpenDataStream(PublishedFetch* fetch,
                                     webtransport::SendOrder send_order);
-
-  SubscribeRemoteTrack* RemoteTrackByAlias(uint64_t track_alias);
-  RemoteTrack* RemoteTrackById(uint64_t request_id);
-  SubscribeRemoteTrack* RemoteTrackByName(const FullTrackName& name);
+  SubscribeRemoteTrack* SubscribeByAlias(uint64_t track_alias);
+  SubscribeRemoteTrack* SubscribeByName(const FullTrackName& track_name);
+  UpstreamFetch* FetchById(uint64_t request_id);
 
   // Checks that a subscribe ID from a SUBSCRIBE or FETCH is valid, and throws
   // a session error if is not.
@@ -409,20 +410,17 @@
   void CancelFetch(uint64_t request_id);
 
   // Sends an OBJECT_ACK message for a specific subscribe ID.
-  void SendObjectAck(uint64_t subscribe_id, uint64_t group_id,
+  void SendObjectAck(FullTrackName track_name, uint64_t group_id,
                      uint64_t object_id,
                      quic::QuicTimeDelta delta_from_deadline) {
     if (!SupportsObjectAck()) {
       return;
     }
-    MoqtObjectAck ack;
-    ack.subscribe_id = subscribe_id;
-    ack.group_id = group_id;
-    ack.object_id = object_id;
-    ack.delta_from_deadline = delta_from_deadline;
-    SendControlMessage(framer_.SerializeObjectAck(ack));
+    SubscribeRemoteTrack* track = SubscribeByName(track_name);
+    if (track != nullptr) {
+      track->SendObjectAck(group_id, object_id, delta_from_deadline);
+    }
   }
-
   // Indicates if OBJECT_ACK is supported by both sides.
   bool SupportsObjectAck() const {
     return parameters_.support_object_acks && peer_supports_object_ack_;
@@ -440,12 +438,11 @@
 
   // Handlers for the control messages on the main control stream.
   absl::Status OnControlMessage(const MoqtSetup& message);
+
+  // TODO(martinduke): All of these should be moved to bidi streams or
+  // deleted.
   absl::Status OnControlMessage(const MoqtRequestOk& message);
   absl::Status OnControlMessage(const MoqtRequestError& message);
-  absl::Status OnControlMessage(const MoqtSubscribe& message);
-  absl::Status OnControlMessage(const MoqtSubscribeOk& message);
-  absl::Status OnControlMessage(const MoqtUnsubscribe& message);
-  absl::Status OnControlMessage(const MoqtPublishDone& /*message*/);
   absl::Status OnControlMessage(const MoqtRequestUpdate& message);
   absl::Status OnControlMessage(const MoqtPublishNamespace& message);
   absl::Status OnControlMessage(const MoqtPublishNamespaceDone& /*message*/);
@@ -459,15 +456,6 @@
   }
   absl::Status OnControlMessage(const MoqtFetchOk& message);
   absl::Status OnControlMessage(const MoqtRequestsBlocked& message);
-  absl::Status OnControlMessage(const MoqtPublish& message);
-  absl::Status OnControlMessage(const MoqtObjectAck& message) {
-    auto subscription_it = published_subscriptions_.find(message.subscribe_id);
-    if (subscription_it == published_subscriptions_.end()) {
-      return absl::OkStatus();
-    }
-    subscription_it->second->ProcessObjectAck(message);
-    return absl::OkStatus();
-  }
 
   // TODO(vasilvv): remove this once all requests are moved into individual
   // streams.
@@ -483,6 +471,12 @@
     SendControlMessage(framer_.SerializeRequestError(request_error));
   }
 
+  uint64_t NextRequestId() {
+    uint64_t id = next_request_id_;
+    next_request_id_ += 2;
+    return id;
+  }
+
   bool is_closing_ = false;
   webtransport::Session* session_;
   MoqtSessionParameters parameters_;
@@ -501,24 +495,12 @@
 
   MoqtTraceRecorder trace_recorder_;
 
-  // Upstream SUBSCRIBE state.
-  // Upstream SUBSCRIBEs and FETCHes, indexed by subscribe_id. Do not erase
-  // directly, call RemoteTrack::Destroy(), except in deletion callbacks passed
-  // to RemoteTrack.
-  absl::flat_hash_map<uint64_t, std::unique_ptr<RemoteTrack>> upstream_by_id_;
-  // All SUBSCRIBEs, indexed by track_alias.
+  // Upstream FETCHes, indexed by request_id. Do not erase.
+  absl::flat_hash_map<uint64_t, std::unique_ptr<UpstreamFetch>> fetch_by_id_;
+  // All outgoing SUBSCRIBE and incoming PUBLISH, indexed by track_alias.
   absl::flat_hash_map<uint64_t, SubscribeRemoteTrack*> subscribe_by_alias_;
-  // All SUBSCRIBEs, indexed by track name.
+  // All outgoing SUBSCRIBE and incoming PUBLISH, indexed by track name.
   absl::flat_hash_map<FullTrackName, SubscribeRemoteTrack*> subscribe_by_name_;
-  struct SubscribeUpdateStatus {
-    FullTrackName name;
-    MessageParameters parameters;
-    MoqtResponseCallback response_callback;
-  };
-  // Outgoing Subscribe Updates. We should not update parameters until a
-  // REQUEST_OK arrives.
-  absl::flat_hash_map<uint64_t, SubscribeUpdateStatus>
-      pending_subscribe_updates_;
 
   // The next subscribe ID that the local endpoint can send.
   uint64_t next_request_id_ = 0;
@@ -528,15 +510,16 @@
 
   // All open incoming subscriptions, indexed by track name, used to check for
   // duplicates.
-  absl::flat_hash_set<FullTrackName> subscribed_track_names_;
+  absl::flat_hash_map<FullTrackName, SubscriptionPublisher*>
+      subscribed_track_names_;
   // Application object representing the publisher for all of the tracks that
   // can be subscribed to via this connection.  Must outlive this object.
   MoqtPublisher* publisher_;
-  // Subscriptions for local tracks by the remote peer, indexed by subscribe ID.
-  absl::flat_hash_map<uint64_t, std::unique_ptr<SubscriptionPublisher>>
+  // Subscriptions for local tracks by the remote peer, indexed by request ID.
+  absl::flat_hash_map<uint64_t, SubscriptionPublisher*>
       published_subscriptions_;
-  // Keeps track of all request IDs that have queued outgoing data streams. The
-  // first element is the highest priority (lowest integer).
+  // Keeps track of all request IDs that have queued outgoing data streams.
+  // The first element is the highest priority (lowest integer).
   absl::btree_multimap<MoqtTrackPriority, uint64_t>
       subscriptions_with_queued_streams_;
   // This is only used to check for track_alias collisions.
@@ -587,6 +570,8 @@
   bool alternate_delivery_timeout_ = false;
 
   quiche::QuicheWeakPtrFactory<MoqtSessionInterface> weak_ptr_factory_;
+  quiche::QuicheWeakPtrFactory<SessionToPublisherInterface>
+      weak_ptr_factory_for_publishers_;
 
   // Must be last.  Token used to make sure that the streams do not call into
   // the session when the session has already been destroyed.
diff --git a/quiche/quic/moqt/moqt_session_interface.h b/quiche/quic/moqt/moqt_session_interface.h
index e1455a6..72c5164 100644
--- a/quiche/quic/moqt/moqt_session_interface.h
+++ b/quiche/quic/moqt/moqt_session_interface.h
@@ -97,13 +97,19 @@
   virtual void Error(MoqtError code, absl::string_view error) = 0;
 
   // Return true if SUBSCRIBE was actually sent.
-  virtual bool Subscribe(const FullTrackName& name, SubscribeVisitor* visitor,
+  virtual bool Subscribe(const FullTrackName& name,
+                         SubscribeVisitor* absl_nonnull visitor,
                          const MessageParameters& parameters) = 0;
   // If a parameter is nullopt, there is no change to the current value.
-  // Returns false if the subscription is not found.
+  // Returns false if the subscription is not found. Used by the subscriber for
+  // a SUBSCRIBE or PUBLISH.
   virtual bool SubscribeUpdate(const FullTrackName& name,
                                const MessageParameters& parameters,
                                MoqtResponseCallback response_callback) = 0;
+  // Used by the publisher of a PUBLISH message.
+  virtual bool PublishUpdate(const FullTrackName& name,
+                             const MessageParameters& parameters,
+                             MoqtResponseCallback response_callback) = 0;
 
   // Sends an UNSUBSCRIBE message and removes all of the state related to the
   // subscription.  Returns false if the subscription is not found.
diff --git a/quiche/quic/moqt/moqt_session_test.cc b/quiche/quic/moqt/moqt_session_test.cc
index c327f5e..afff2a7 100644
--- a/quiche/quic/moqt/moqt_session_test.cc
+++ b/quiche/quic/moqt/moqt_session_test.cc
@@ -7,7 +7,6 @@
 #include <algorithm>
 #include <cstdint>
 #include <cstring>
-#include <functional>
 #include <memory>
 #include <optional>
 #include <queue>
@@ -17,7 +16,6 @@
 #include <vector>
 
 #include "absl/base/casts.h"
-#include "absl/memory/memory.h"
 #include "absl/status/status.h"
 #include "absl/strings/string_view.h"
 #include "absl/types/span.h"
@@ -32,16 +30,13 @@
 #include "quiche/quic/moqt/moqt_known_track_publisher.h"
 #include "quiche/quic/moqt/moqt_messages.h"
 #include "quiche/quic/moqt/moqt_names.h"
-#include "quiche/quic/moqt/moqt_namespace_stream.h"
 #include "quiche/quic/moqt/moqt_object.h"
-#include "quiche/quic/moqt/moqt_parser.h"
 #include "quiche/quic/moqt/moqt_priority.h"
 #include "quiche/quic/moqt/moqt_publisher.h"
 #include "quiche/quic/moqt/moqt_session_callbacks.h"
 #include "quiche/quic/moqt/moqt_session_interface.h"
 #include "quiche/quic/moqt/moqt_track.h"
 #include "quiche/quic/moqt/moqt_types.h"
-#include "quiche/quic/moqt/session_namespace_tree.h"
 #include "quiche/quic/moqt/test_tools/moqt_framer_utils.h"
 #include "quiche/quic/moqt/test_tools/moqt_mock_visitor.h"
 #include "quiche/quic/moqt/test_tools/moqt_session_peer.h"
@@ -51,7 +46,6 @@
 #include "quiche/common/quiche_data_reader.h"
 #include "quiche/common/quiche_mem_slice.h"
 #include "quiche/common/quiche_weak_ptr.h"
-#include "quiche/common/test_tools/quiche_test_utils.h"
 #include "quiche/web_transport/test_tools/in_memory_stream.h"
 #include "quiche/web_transport/test_tools/mock_web_transport.h"
 #include "quiche/web_transport/web_transport.h"
@@ -146,7 +140,8 @@
     session_.set_publisher(&publisher_);
     MoqtSessionPeer::set_peer_max_request_id(&session_,
                                              kDefaultInitialMaxRequestId);
-    ON_CALL(mock_session_, GetStreamById).WillByDefault(Return(&mock_stream_));
+    ON_CALL(mock_session_, GetStreamById)
+        .WillByDefault(Return(&mock_bidi_stream_));
     EXPECT_EQ(MoqtSessionPeer::GetImplementationString(&session_),
               kImplementationName);
   }
@@ -168,6 +163,71 @@
     ON_CALL(*publisher, largest_location()).WillByDefault(Return(largest_id));
   }
 
+  // Opens an incoming request stream, determines the type based on first_byte
+  // and returns MoqtBidiStreamBase that can be downcast to the correct type.
+  // |wt_stream| is the underlying mock WebTransport stream; if nullptr, use
+  // mock_bidi_stream_.
+  static constexpr absl::string_view kSubscribeByte = "\x03";
+  static constexpr absl::string_view kSubscribeNamespaceByte = "\x11";
+  static constexpr absl::string_view kPublishByte = "\x1d";
+  std::unique_ptr<MoqtBidiStreamBase> ResponseStream(
+      absl::string_view first_byte,
+      webtransport::test::MockStream* wt_stream = nullptr) {
+    webtransport::test::MockStream* stream =
+        wt_stream != nullptr ? wt_stream : &mock_bidi_stream_;
+    EXPECT_CALL(mock_session_, AcceptIncomingBidirectionalStream())
+        .WillOnce(Return(stream))
+        .WillOnce(Return(nullptr));
+    std::unique_ptr<webtransport::StreamVisitor> unknown_bidi_stream;
+    std::unique_ptr<MoqtBidiStreamBase> final_stream;
+    EXPECT_CALL(*stream, SetVisitor)
+        .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) {
+          unknown_bidi_stream = std::move(visitor);
+        })
+        .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) {
+          final_stream = std::unique_ptr<MoqtBidiStreamBase>(
+              absl::down_cast<MoqtBidiStreamBase*>(visitor.release()));
+        });
+    EXPECT_CALL(*stream, visitor())
+        .WillOnce([&]() { return unknown_bidi_stream.get(); })
+        .WillRepeatedly([&]() { return final_stream.get(); });
+    EXPECT_CALL(*stream, PeekNextReadableRegion)
+        .WillOnce(
+            Return(webtransport::Stream::PeekResult(first_byte, false, false)))
+        .WillRepeatedly(
+            Return(webtransport::Stream::PeekResult("", false, false)));
+    EXPECT_CALL(*stream, ReadableBytes())
+        .WillOnce(Return(first_byte.length()))
+        .WillRepeatedly(Return(0));
+    EXPECT_CALL(*stream, Read(::testing::An<absl::Span<char>>()))
+        .WillOnce([&](absl::Span<char> bytes_to_read) {
+          memcpy(bytes_to_read.data(), first_byte.data(), first_byte.length());
+          return webtransport::Stream::ReadResult(first_byte.length(), false);
+        });
+    session_.OnIncomingBidirectionalStreamAvailable();
+    EXPECT_NE(final_stream, nullptr);
+    EXPECT_CALL(*stream, CanWrite).WillRepeatedly(Return(true));
+    return final_stream;
+  }
+
+  void PrepareRequestStream(
+      std::unique_ptr<MoqtBidiStreamTestWrapper>& stream_wrapper,
+      webtransport::test::MockStream* wt_stream = nullptr) {
+    webtransport::test::MockStream* stream =
+        wt_stream != nullptr ? wt_stream : &mock_bidi_stream_;
+    EXPECT_CALL(mock_session_, CanOpenNextOutgoingBidirectionalStream())
+        .WillOnce(Return(true));
+    EXPECT_CALL(mock_session_, OpenOutgoingBidirectionalStream())
+        .WillOnce(Return(stream));
+    EXPECT_CALL(*stream, SetVisitor)
+        .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) {
+          stream_wrapper = std::make_unique<MoqtBidiStreamTestWrapper>(
+              std::unique_ptr<MoqtBidiStreamBase>(
+                  absl::down_cast<MoqtBidiStreamBase*>(visitor.release())));
+        });
+    EXPECT_CALL(*stream, CanWrite).WillRepeatedly(Return(true));
+  }
+
   // The publisher receives SUBSCRIBE and synchronously publishes namespaces it
   // supports.
   MoqtObjectListener* ReceiveSubscribeSynchronousOk(
@@ -189,7 +249,8 @@
         parameters,
         extensions,
     };
-    EXPECT_CALL(mock_stream_, Writev(SerializedControlMessage(expected_ok), _));
+    EXPECT_CALL(mock_bidi_stream_,
+                Writev(SerializedControlMessage(expected_ok), _));
     control_parser->ReceiveMessage(subscribe);
     return listener_ptr;
   }
@@ -258,12 +319,14 @@
     }
   }
 
-  webtransport::test::MockStream mock_stream_, control_stream_;
-  MockSessionCallbacks session_callbacks_;
-  webtransport::test::MockSession mock_session_;
   MockSubscribeRemoteTrackVisitor remote_track_visitor_;
-  MoqtSession session_;
   MoqtKnownTrackPublisher publisher_;
+  webtransport::test::MockSession mock_session_;
+  MockSessionCallbacks session_callbacks_;
+  std::unique_ptr<MoqtBidiStreamTestWrapper> bidi_wrapper_;
+  MoqtSession session_;
+  webtransport::test::MockStream mock_bidi_stream_, mock_uni_stream_;
+  //  std::shared_ptr<IncomingSubscribeInfo> last_incoming_subscribe_;
 };
 
 TEST_F(MoqtSessionTest, Queries) {
@@ -277,28 +340,28 @@
   EXPECT_CALL(mock_session_, CanOpenNextOutgoingBidirectionalStream())
       .WillOnce(Return(true));
   EXPECT_CALL(mock_session_, OpenOutgoingBidirectionalStream())
-      .WillOnce(Return(&mock_stream_));
-  EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true));
+      .WillOnce(Return(&mock_bidi_stream_));
+  EXPECT_CALL(mock_bidi_stream_, CanWrite).WillRepeatedly(Return(true));
   std::unique_ptr<webtransport::StreamVisitor> visitor;
   // Save a reference to MoqtSession::Stream
-  EXPECT_CALL(mock_stream_, SetVisitor(_))
+  EXPECT_CALL(mock_bidi_stream_, SetVisitor(_))
       .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> new_visitor) {
         visitor = std::move(new_visitor);
       });
-  EXPECT_CALL(mock_stream_, GetStreamId())
+  EXPECT_CALL(mock_bidi_stream_, GetStreamId())
       .WillRepeatedly(Return(webtransport::StreamId(4)));
-  EXPECT_CALL(mock_stream_,
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kSetup), _));
   session_.OnSessionReady();
 
   // Receive SERVER_SETUP
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
+  bidi_wrapper_ =
       MoqtSessionPeer::FetchParserVisitorFromWebtransportStreamVisitor(
           std::move(visitor));
   // Handle the server setup
   MoqtSetup setup;  // No fields are set.
   EXPECT_CALL(session_callbacks_.session_established_callback, Call()).Times(1);
-  stream_input->ReceiveMessage(setup);
+  bidi_wrapper_->ReceiveMessage(setup);
 }
 
 TEST_F(MoqtSessionTest, OnSessionReadyNoControlStream) {
@@ -316,16 +379,16 @@
       std::make_unique<quic::test::TestAlarmFactory>(),
       session_callbacks_.AsSessionCallbacks());
   EXPECT_CALL(mock_session_, AcceptIncomingBidirectionalStream())
-      .WillOnce(Return(&mock_stream_))
+      .WillOnce(Return(&mock_bidi_stream_))
       .WillOnce(Return(nullptr));
   std::unique_ptr<webtransport::StreamVisitor> visitor;
   webtransport::test::MockStreamVisitor mock_stream_visitor;
-  EXPECT_CALL(mock_stream_, SetVisitor)
+  EXPECT_CALL(mock_bidi_stream_, SetVisitor)
       .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> new_visitor) {
         visitor = std::move(new_visitor);
-        EXPECT_CALL(mock_stream_, visitor).WillOnce(Return(visitor.get()));
+        EXPECT_CALL(mock_bidi_stream_, visitor).WillOnce(Return(visitor.get()));
       });
-  EXPECT_CALL(mock_stream_, PeekNextReadableRegion())
+  EXPECT_CALL(mock_bidi_stream_, PeekNextReadableRegion())
       .WillRepeatedly(Return(
           webtransport::Stream::PeekResult(absl::string_view(), false, false)));
   server_session.OnIncomingBidirectionalStreamAvailable();
@@ -371,9 +434,10 @@
   ::testing::InSequence seq;
   StrictMock<webtransport::test::MockStreamVisitor> mock_stream_visitor;
   EXPECT_CALL(mock_session_, AcceptIncomingBidirectionalStream())
-      .WillOnce(Return(&mock_stream_));
-  EXPECT_CALL(mock_stream_, SetVisitor(_)).Times(1);
-  EXPECT_CALL(mock_stream_, visitor()).WillOnce(Return(&mock_stream_visitor));
+      .WillOnce(Return(&mock_bidi_stream_));
+  EXPECT_CALL(mock_bidi_stream_, SetVisitor).Times(1);
+  EXPECT_CALL(mock_bidi_stream_, visitor())
+      .WillOnce(Return(&mock_stream_visitor));
   EXPECT_CALL(mock_stream_visitor, OnCanRead()).Times(1);
   EXPECT_CALL(mock_session_, AcceptIncomingBidirectionalStream())
       .WillOnce(Return(nullptr));
@@ -384,9 +448,10 @@
   ::testing::InSequence seq;
   StrictMock<webtransport::test::MockStreamVisitor> mock_stream_visitor;
   EXPECT_CALL(mock_session_, AcceptIncomingUnidirectionalStream())
-      .WillOnce(Return(&mock_stream_));
-  EXPECT_CALL(mock_stream_, SetVisitor(_)).Times(1);
-  EXPECT_CALL(mock_stream_, visitor()).WillOnce(Return(&mock_stream_visitor));
+      .WillOnce(Return(&mock_uni_stream_));
+  EXPECT_CALL(mock_uni_stream_, SetVisitor(_)).Times(1);
+  EXPECT_CALL(mock_uni_stream_, visitor())
+      .WillOnce(Return(&mock_stream_visitor));
   EXPECT_CALL(mock_stream_visitor, OnCanRead()).Times(1);
   EXPECT_CALL(mock_session_, AcceptIncomingUnidirectionalStream())
       .WillOnce(Return(nullptr));
@@ -409,18 +474,20 @@
 
 TEST_F(MoqtSessionTest, AddLocalTrack) {
   MoqtSubscribe request = DefaultSubscribe();
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeByte));
   // Request for track returns REQUEST_ERROR.
-  EXPECT_CALL(mock_stream_,
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _));
-  stream_input->ReceiveMessage(request);
+  bidi_wrapper_->ReceiveMessage(request);
 
   // Add the track. Now Subscribe should succeed.
   MockTrackPublisher* track = CreateTrackPublisher();
-  std::make_shared<MockTrackPublisher>(request.full_track_name);
   request.request_id += 2;
-  ReceiveSubscribeSynchronousOk(track, request, stream_input.get());
+  bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeByte));
+  ReceiveSubscribeSynchronousOk(track, request, bidi_wrapper_.get(),
+                                MoqtSessionPeer::GetLastTrackAlias(&session_));
 }
 
 TEST_F(MoqtSessionTest, IncomingPublishRejected) {
@@ -431,22 +498,22 @@
       .parameters = MessageParameters(),
   };
   publish.parameters.largest_object = Location(4, 5);
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
-  // Request for track returns REQUEST_ERROR.
-  EXPECT_CALL(mock_stream_,
+  bidi_wrapper_ =
+      std::make_unique<MoqtBidiStreamTestWrapper>(ResponseStream(kPublishByte));
+  // With the default incoming_publish_callbackm, will return REQUEST_ERROR.
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _));
-  stream_input->ReceiveMessage(publish);
+  bidi_wrapper_->ReceiveMessage(publish);
 }
 
 TEST_F(MoqtSessionTest, PublishNamespaceWithOkAndCancel) {
   testing::MockFunction<void(
       std::variant<MessageParameters, MoqtRequestErrorInfo> error_message)>
       publish_namespace_response_callback;
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  std::unique_ptr<MoqtBidiStreamTestWrapper> bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   EXPECT_CALL(
-      mock_stream_,
+      mock_bidi_stream_,
       Writev(ControlMessageOfType(MoqtMessageType::kPublishNamespace), _));
   MoqtRequestErrorInfo cancel_error_info;
   session_.PublishNamespace(
@@ -460,14 +527,14 @@
           [&](std::variant<MessageParameters, MoqtRequestErrorInfo> response) {
             EXPECT_TRUE(std::holds_alternative<MessageParameters>(response));
           });
-  stream_input->ReceiveMessage(ok);
+  bidi_wrapper_->ReceiveMessage(ok);
 
   MoqtPublishNamespaceCancel cancel = {
       /*request_id=*/0,
       RequestErrorCode::kInternalError,
       /*error_reason=*/"Test error",
   };
-  stream_input->ReceiveMessage(cancel);
+  bidi_wrapper_->ReceiveMessage(cancel);
   EXPECT_EQ(cancel_error_info.error_code, RequestErrorCode::kInternalError);
   EXPECT_EQ(cancel_error_info.reason_phrase, "Test error");
   // State is gone.
@@ -478,10 +545,10 @@
   testing::MockFunction<void(
       std::variant<MessageParameters, MoqtRequestErrorInfo>)>
       publish_namespace_resolved_callback;
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  std::unique_ptr<MoqtBidiStreamTestWrapper> bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   EXPECT_CALL(
-      mock_stream_,
+      mock_bidi_stream_,
       Writev(ControlMessageOfType(MoqtMessageType::kPublishNamespace), _));
   session_.PublishNamespace(TrackNamespace{"foo"}, MessageParameters(),
                             publish_namespace_resolved_callback.AsStdFunction(),
@@ -493,10 +560,10 @@
           [&](std::variant<MessageParameters, MoqtRequestErrorInfo> response) {
             EXPECT_TRUE(std::holds_alternative<MessageParameters>(response));
           });
-  stream_input->ReceiveMessage(ok);
+  bidi_wrapper_->ReceiveMessage(ok);
 
   EXPECT_CALL(
-      mock_stream_,
+      mock_bidi_stream_,
       Writev(ControlMessageOfType(MoqtMessageType::kPublishNamespaceDone), _));
   session_.PublishNamespaceDone(TrackNamespace{"foo"});
   // State is gone.
@@ -507,10 +574,10 @@
   testing::MockFunction<void(
       std::variant<MessageParameters, MoqtRequestErrorInfo>)>
       publish_namespace_resolved_callback;
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  std::unique_ptr<MoqtBidiStreamTestWrapper> bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   EXPECT_CALL(
-      mock_stream_,
+      mock_bidi_stream_,
       Writev(ControlMessageOfType(MoqtMessageType::kPublishNamespace), _));
   session_.PublishNamespace(TrackNamespace{"foo"}, MessageParameters(),
                             publish_namespace_resolved_callback.AsStdFunction(),
@@ -527,23 +594,23 @@
             EXPECT_EQ(error.error_code, RequestErrorCode::kInternalError);
             EXPECT_EQ(error.reason_phrase, "Test error");
           });
-  stream_input->ReceiveMessage(error);
+  bidi_wrapper_->ReceiveMessage(error);
   // State is gone.
   EXPECT_FALSE(session_.PublishNamespaceDone(TrackNamespace{"foo"}));
 }
 
 TEST_F(MoqtSessionTest, AsynchronousSubscribeReturnsOk) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeByte));
   MoqtSubscribe request = DefaultSubscribe();
   MockTrackPublisher* track = CreateTrackPublisher();
   MoqtObjectListener* listener;
   EXPECT_CALL(*track, AddObjectListener)
       .WillOnce(
           [&](MoqtObjectListener* listener_ptr) { listener = listener_ptr; });
-  stream_input->ReceiveMessage(request);
+  bidi_wrapper_->ReceiveMessage(request);
 
-  EXPECT_CALL(mock_stream_,
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kSubscribeOk), _));
   listener->OnSubscribeAccepted();
   EXPECT_TRUE(MoqtSessionPeer::RequestIdIsSubscriptionPublisher(
@@ -551,61 +618,66 @@
 }
 
 TEST_F(MoqtSessionTest, AsynchronousSubscribeReturnsError) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeByte));
   MoqtSubscribe request = DefaultSubscribe();
   MockTrackPublisher* track = CreateTrackPublisher();
   MoqtObjectListener* listener;
   EXPECT_CALL(*track, AddObjectListener)
       .WillOnce(
           [&](MoqtObjectListener* listener_ptr) { listener = listener_ptr; });
-  stream_input->ReceiveMessage(request);
-  EXPECT_CALL(mock_stream_,
-              Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _));
+  bidi_wrapper_->ReceiveMessage(request);
+  EXPECT_CALL(mock_bidi_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _))
+      .WillOnce([](absl::Span<quiche::QuicheMemSlice>,
+                   const webtransport::StreamWriteOptions& options) {
+        EXPECT_TRUE(options.send_fin());
+        return absl::OkStatus();
+      });
   listener->OnSubscribeRejected(MoqtRequestErrorInfo(
       RequestErrorCode::kInternalError, std::nullopt, "Test error"));
-  EXPECT_FALSE(MoqtSessionPeer::RequestIdIsSubscriptionPublisher(
-      &session_, kDefaultPeerRequestId));
 }
 
 TEST_F(MoqtSessionTest, SynchronousSubscribeReturnsError) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeByte));
   MoqtSubscribe request = DefaultSubscribe();
   MockTrackPublisher* track = CreateTrackPublisher();
   EXPECT_CALL(*track, AddObjectListener)
       .WillOnce([&](MoqtObjectListener* listener) {
-        EXPECT_CALL(
-            mock_stream_,
-            Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _));
         EXPECT_CALL(*track, RemoveObjectListener);
         listener->OnSubscribeRejected(MoqtRequestErrorInfo(
             RequestErrorCode::kInternalError, std::nullopt, "Test error"));
       });
-  stream_input->ReceiveMessage(request);
-  EXPECT_FALSE(MoqtSessionPeer::RequestIdIsSubscriptionPublisher(
-      &session_, kDefaultPeerRequestId));
+  EXPECT_CALL(mock_bidi_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _))
+      .WillOnce([](absl::Span<quiche::QuicheMemSlice>,
+                   const webtransport::StreamWriteOptions& options) {
+        EXPECT_TRUE(options.send_fin());
+        return absl::OkStatus();
+      });
+  bidi_wrapper_->ReceiveMessage(request);
 }
 
 TEST_F(MoqtSessionTest, SubscribeForPast) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeByte));
   MockTrackPublisher* track = CreateTrackPublisher();
   SetLargestId(track, Location(10, 20));
   MoqtSubscribe request = DefaultSubscribe();
-  ReceiveSubscribeSynchronousOk(track, request, stream_input.get());
+  ReceiveSubscribeSynchronousOk(track, request, bidi_wrapper_.get());
 }
 
 TEST_F(MoqtSessionTest, SubscribeDoNotForward) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeByte));
   MockTrackPublisher* track = CreateTrackPublisher();
   MoqtSubscribe request = DefaultSubscribe();
   request.parameters.set_forward(false);
   request.parameters.subscription_filter.emplace(
       MoqtFilterType::kLargestObject);
   MoqtObjectListener* listener =
-      ReceiveSubscribeSynchronousOk(track, request, stream_input.get());
+      ReceiveSubscribeSynchronousOk(track, request, bidi_wrapper_.get());
   // forward=false, so incoming objects are ignored.
   EXPECT_CALL(mock_session_, CanOpenNextOutgoingUnidirectionalStream())
       .Times(0);
@@ -613,13 +685,13 @@
 }
 
 TEST_F(MoqtSessionTest, SubscribeAbsoluteStartNoDataYet) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeByte));
   MockTrackPublisher* track = CreateTrackPublisher();
   MoqtSubscribe request = DefaultSubscribe();
   request.parameters.subscription_filter.emplace(Location(1, 0));
   MoqtObjectListener* listener =
-      ReceiveSubscribeSynchronousOk(track, request, stream_input.get());
+      ReceiveSubscribeSynchronousOk(track, request, bidi_wrapper_.get());
   // Window was not set to (0, 0) by SUBSCRIBE acceptance.
   EXPECT_CALL(mock_session_, CanOpenNextOutgoingUnidirectionalStream())
       .Times(0);
@@ -627,15 +699,15 @@
 }
 
 TEST_F(MoqtSessionTest, SubscribeNextGroup) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeByte));
   MockTrackPublisher* track = CreateTrackPublisher();
   MoqtSubscribe request = DefaultSubscribe();
   request.parameters.subscription_filter.emplace(
       MoqtFilterType::kNextGroupStart);
   SetLargestId(track, Location(10, 20));
   MoqtObjectListener* listener =
-      ReceiveSubscribeSynchronousOk(track, request, stream_input.get());
+      ReceiveSubscribeSynchronousOk(track, request, bidi_wrapper_.get());
   // Later objects in group 10 ignored.
   EXPECT_CALL(mock_session_, CanOpenNextOutgoingUnidirectionalStream())
       .Times(0);
@@ -648,88 +720,83 @@
 }
 
 TEST_F(MoqtSessionTest, TwoSubscribesForTrack) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeByte));
   MockTrackPublisher* track = CreateTrackPublisher();
   MoqtSubscribe request = DefaultSubscribe();
-  ReceiveSubscribeSynchronousOk(track, request, stream_input.get());
+  ReceiveSubscribeSynchronousOk(track, request, bidi_wrapper_.get());
 
   request.request_id = 3;
   request.parameters.subscription_filter.emplace(Location(12, 0));
-  EXPECT_CALL(mock_stream_,
+  webtransport::test::MockStream bidi_stream_2;
+  auto bidi_wrapper_2 = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeByte, &bidi_stream_2));
+  EXPECT_CALL(bidi_stream_2,
               Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _));
-  stream_input->ReceiveMessage(request);
+  bidi_wrapper_2->ReceiveMessage(request);
 }
 
 TEST_F(MoqtSessionTest, UnsubscribeAllowsSecondSubscribe) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeByte));
   MockTrackPublisher* track = CreateTrackPublisher();
   MoqtSubscribe request = DefaultSubscribe();
-  ReceiveSubscribeSynchronousOk(track, request, stream_input.get());
+  ReceiveSubscribeSynchronousOk(track, request, bidi_wrapper_.get());
 
   // Peer unsubscribes.
-  MoqtUnsubscribe unsubscribe = {
-      kDefaultPeerRequestId,
-  };
-  stream_input->ReceiveMessage(unsubscribe);
+  bidi_wrapper_->stream().Reset(kResetCodeCancelled);
+  bidi_wrapper_ = nullptr;
   EXPECT_FALSE(MoqtSessionPeer::RequestIdIsSubscriptionPublisher(&session_, 1));
 
   // Subscribe again, succeeds.
   request.request_id = 3;
   request.parameters.subscription_filter.emplace(Location(12, 0));
-  ReceiveSubscribeSynchronousOk(track, request, stream_input.get(),
+  bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeByte));
+  ReceiveSubscribeSynchronousOk(track, request, bidi_wrapper_.get(),
                                 /*track_alias=*/1);
 }
 
-TEST_F(MoqtSessionTest, RequestIdTooHigh) {
-  // Peer subscribes to (0, 0)
-  MoqtSubscribe request = DefaultSubscribe();
-  request.request_id = kDefaultInitialMaxRequestId + 1;
-
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
-  EXPECT_CALL(mock_session_,
-              CloseSession(static_cast<uint64_t>(MoqtError::kTooManyRequests),
-                           "Received request with too large ID"));
-  stream_input->ReceiveMessage(request);
-}
-
 TEST_F(MoqtSessionTest, RequestIdWrongLsb) {
   // TODO(martinduke): Implement this test.
 }
 
 TEST_F(MoqtSessionTest, SubscribeIdNotIncreasing) {
   MoqtSubscribe request = DefaultSubscribe();
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeByte));
   MockTrackPublisher* track = CreateTrackPublisher();
   EXPECT_CALL(*track, AddObjectListener);
-  stream_input->ReceiveMessage(request);
+  bidi_wrapper_->ReceiveMessage(request);
 
   // Second request is a protocol violation.
   request.full_track_name = FullTrackName({"dead", "beef"});
-  EXPECT_CALL(mock_session_,
-              CloseSession(static_cast<uint64_t>(MoqtError::kInvalidRequestId),
-                           "Duplicate request ID"));
-  stream_input->ReceiveMessage(request);
+  auto publisher =
+      std::make_shared<MockTrackPublisher>(request.full_track_name);
+  publisher_.Add(publisher);
+  webtransport::test::MockStream bidi_stream_2;
+  auto bidi_wrapper_2 = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeByte, &bidi_stream_2));
+  EXPECT_CALL(bidi_stream_2,
+              Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _));
+  bidi_wrapper_2->ReceiveMessage(request);
 }
 
 TEST_F(MoqtSessionTest, TooManySubscribes) {
   MoqtSessionPeer::set_next_request_id(&session_,
                                        kDefaultInitialMaxRequestId - 1);
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
-  EXPECT_CALL(mock_session_, GetStreamById(_))
-      .WillRepeatedly(Return(&mock_stream_));
-  EXPECT_CALL(mock_stream_,
+  PrepareRequestStream(bidi_wrapper_);
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _));
   MessageParameters parameters(SubscribeForTest());
   parameters.subscription_filter.emplace(MoqtFilterType::kLargestObject);
   EXPECT_TRUE(session_.Subscribe(FullTrackName("foo", "bar"),
                                  &remote_track_visitor_, parameters));
+  webtransport::test::MockStream control_stream;
+  std::unique_ptr<MoqtBidiStreamTestWrapper> control_wrapper =
+      MoqtSessionPeer::CreateControlStream(&session_, &control_stream);
   EXPECT_CALL(
-      mock_stream_,
+      control_stream,
       Writev(ControlMessageOfType(MoqtMessageType::kRequestsBlocked), _))
       .Times(1);
   EXPECT_FALSE(session_.Subscribe(FullTrackName("foo2", "bar2"),
@@ -740,11 +807,8 @@
 }
 
 TEST_F(MoqtSessionTest, SubscribeDuplicateTrackName) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
-  EXPECT_CALL(mock_session_, GetStreamById(_))
-      .WillRepeatedly(Return(&mock_stream_));
-  EXPECT_CALL(mock_stream_,
+  PrepareRequestStream(bidi_wrapper_);
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _));
   MessageParameters parameters(SubscribeForTest());
   EXPECT_TRUE(session_.Subscribe(FullTrackName("foo", "bar"),
@@ -754,9 +818,8 @@
 }
 
 TEST_F(MoqtSessionTest, SubscribeWithOk) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
-  EXPECT_CALL(mock_stream_,
+  PrepareRequestStream(bidi_wrapper_);
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _));
   MessageParameters parameters(SubscribeForTest());
   session_.Subscribe(FullTrackName("foo", "bar"), &remote_track_visitor_,
@@ -775,16 +838,16 @@
             EXPECT_EQ(ftn, FullTrackName("foo", "bar"));
             EXPECT_TRUE(std::holds_alternative<SubscribeOkData>(response));
           });
-  stream_input->ReceiveMessage(ok);
+  bidi_wrapper_->ReceiveMessage(ok);
 }
 
 TEST_F(MoqtSessionTest, SubscribeNextGroupWithOk) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  PrepareRequestStream(bidi_wrapper_);
   MoqtSubscribe subscribe = DefaultLocalSubscribe();
   subscribe.parameters.subscription_filter.emplace(
       MoqtFilterType::kNextGroupStart);
-  EXPECT_CALL(mock_stream_, Writev(SerializedControlMessage(subscribe), _));
+  EXPECT_CALL(mock_bidi_stream_,
+              Writev(SerializedControlMessage(subscribe), _));
   session_.Subscribe(FullTrackName("foo", "bar"), &remote_track_visitor_,
                      subscribe.parameters);
 
@@ -801,15 +864,12 @@
             EXPECT_EQ(ftn, FullTrackName("foo", "bar"));
             EXPECT_TRUE(std::holds_alternative<SubscribeOkData>(response));
           });
-  stream_input->ReceiveMessage(ok);
+  bidi_wrapper_->ReceiveMessage(ok);
 }
 
 TEST_F(MoqtSessionTest, OutgoingSubscribeUpdate) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
-  EXPECT_CALL(mock_session_, GetStreamById)
-      .WillRepeatedly(Return(&mock_stream_));
-  EXPECT_CALL(mock_stream_,
+  PrepareRequestStream(bidi_wrapper_);
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _));
   MessageParameters parameters(SubscribeForTest());
   parameters.subscription_filter.emplace(Location(1, 0), 10);
@@ -822,8 +882,8 @@
       TrackExtensions(),
   };
   EXPECT_CALL(remote_track_visitor_, OnReply);
-  stream_input->ReceiveMessage(ok);
-  EXPECT_CALL(mock_stream_,
+  bidi_wrapper_->ReceiveMessage(ok);
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kRequestUpdate), _));
   MessageParameters update_parameters;
   update_parameters.subscription_filter.emplace(Location(2, 1), 9);
@@ -836,7 +896,7 @@
         ASSERT_TRUE(std::holds_alternative<MessageParameters>(info));
         EXPECT_EQ(std::get<MessageParameters>(info), MessageParameters());
       }));
-  stream_input->ReceiveMessage(MoqtRequestOk{
+  bidi_wrapper_->ReceiveMessage(MoqtRequestOk{
       /*request_id=*/2,
       MessageParameters(),
   });
@@ -867,12 +927,11 @@
 
 TEST_F(MoqtSessionTest, MaxRequestIdChangesResponse) {
   MoqtSessionPeer::set_next_request_id(&session_, kDefaultInitialMaxRequestId);
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
-  EXPECT_CALL(mock_session_, GetStreamById(_))
-      .WillRepeatedly(Return(&mock_stream_));
+  webtransport::test::MockStream control_stream;
+  std::unique_ptr<MoqtBidiStreamTestWrapper> control_wrapper =
+      MoqtSessionPeer::CreateControlStream(&session_, &control_stream);
   EXPECT_CALL(
-      mock_stream_,
+      control_stream,
       Writev(ControlMessageOfType(MoqtMessageType::kRequestsBlocked), _));
   MessageParameters parameters(SubscribeForTest());
   parameters.subscription_filter.emplace(MoqtFilterType::kLargestObject);
@@ -881,9 +940,10 @@
   MoqtMaxRequestId max_request_id = {
       /*max_request_id=*/kDefaultInitialMaxRequestId + 1,
   };
-  stream_input->ReceiveMessage(max_request_id);
+  control_wrapper->ReceiveMessage(max_request_id);
 
-  EXPECT_CALL(mock_stream_,
+  PrepareRequestStream(bidi_wrapper_);
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _));
   EXPECT_TRUE(session_.Subscribe(FullTrackName("foo", "bar"),
                                  &remote_track_visitor_, parameters));
@@ -893,32 +953,34 @@
   MoqtMaxRequestId max_request_id = {
       /*max_request_id=*/kDefaultInitialMaxRequestId - 1,
   };
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   EXPECT_CALL(mock_session_,
               CloseSession(static_cast<uint64_t>(MoqtError::kProtocolViolation),
                            "MAX_REQUEST_ID has lower value than previous"))
       .Times(1);
-  stream_input->ReceiveMessage(max_request_id);
+  bidi_wrapper_->ReceiveMessage(max_request_id);
 }
 
 TEST_F(MoqtSessionTest, GrantMoreRequests) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
-  EXPECT_CALL(mock_stream_,
+  webtransport::test::MockStream control_stream;
+  std::unique_ptr<MoqtBidiStreamTestWrapper> control_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &control_stream);
+  EXPECT_CALL(control_stream,
               Writev(ControlMessageOfType(MoqtMessageType::kMaxRequestId), _));
   session_.GrantMoreRequests(1);
   // Peer subscribes to (0, 0)
+  bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeByte));
   MoqtSubscribe request = DefaultSubscribe();
   request.request_id = kDefaultInitialMaxRequestId + 1;
   MockTrackPublisher* track = CreateTrackPublisher();
-  ReceiveSubscribeSynchronousOk(track, request, stream_input.get());
+  ReceiveSubscribeSynchronousOk(track, request, bidi_wrapper_.get());
 }
 
 TEST_F(MoqtSessionTest, SubscribeWithError) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
-  EXPECT_CALL(mock_stream_,
+  PrepareRequestStream(bidi_wrapper_);
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _));
   MessageParameters parameters(SubscribeForTest());
   parameters.subscription_filter.emplace(MoqtFilterType::kLargestObject);
@@ -941,31 +1003,30 @@
                 std::get<MoqtRequestErrorInfo>(response).reason_phrase ==
                     "deadbeef");
           });
-  stream_input->ReceiveMessage(error);
+  EXPECT_CALL(mock_bidi_stream_, Writev(testing::IsEmpty(), _));  // FIN.
+  bidi_wrapper_->ReceiveMessage(error);
 }
 
 TEST_F(MoqtSessionTest, Unsubscribe) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  PrepareRequestStream(bidi_wrapper_);
   FullTrackName ftn = FullTrackName("foo", "bar");
-  EXPECT_CALL(mock_stream_,
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _));
   EXPECT_TRUE(
       session_.Subscribe(ftn, &remote_track_visitor_, MessageParameters()));
-  EXPECT_CALL(mock_stream_,
-              Writev(ControlMessageOfType(MoqtMessageType::kUnsubscribe), _));
+  EXPECT_CALL(mock_bidi_stream_, ResetWithUserCode);
   EXPECT_CALL(remote_track_visitor_, OnPublishDone);
   session_.Unsubscribe(ftn);
   // Verify it was destroyed.
-  EXPECT_CALL(mock_stream_, Writev).Times(0);
+  EXPECT_CALL(mock_bidi_stream_, ResetWithUserCode).Times(0);
   EXPECT_CALL(remote_track_visitor_, OnPublishDone).Times(0);
   session_.Unsubscribe(ftn);
 }
 
 TEST_F(MoqtSessionTest, ReplyToPublishNamespaceWithOkThenPublishNamespaceDone) {
   TrackNamespace track_namespace{"foo"};
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   MessageParameters parameters;
   parameters.authorization_tokens.emplace_back(AuthTokenType::kOutOfBand,
                                                "foo");
@@ -981,11 +1042,11 @@
                    MoqtResponseCallback callback) {
         std::move(callback)(MessageParameters());
       });
-  EXPECT_CALL(mock_stream_,
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(SerializedControlMessage(MoqtRequestOk{
                          kDefaultPeerRequestId, MessageParameters()}),
                      _));
-  stream_input->ReceiveMessage(publish_namespace);
+  bidi_wrapper_->ReceiveMessage(publish_namespace);
   MoqtPublishNamespaceDone publish_namespace_done = {
       /*request_id=*/0,
   };
@@ -994,15 +1055,15 @@
       .WillOnce(
           [](const TrackNamespace&, const std::optional<MessageParameters>&,
              MoqtResponseCallback callback) { EXPECT_EQ(callback, nullptr); });
-  stream_input->ReceiveMessage(publish_namespace_done);
+  bidi_wrapper_->ReceiveMessage(publish_namespace_done);
 }
 
 TEST_F(MoqtSessionTest,
        ReplyToPublishNamespaceWithOkThenPublishNamespaceCancel) {
   TrackNamespace track_namespace{"foo"};
 
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   MessageParameters parameters;
   parameters.authorization_tokens.emplace_back(AuthTokenType::kOutOfBand,
                                                "foo");
@@ -1018,12 +1079,12 @@
                    MoqtResponseCallback callback) {
         std::move(callback)(MessageParameters());
       });
-  EXPECT_CALL(mock_stream_,
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(SerializedControlMessage(MoqtRequestOk{
                          kDefaultPeerRequestId, MessageParameters()}),
                      _));
-  stream_input->ReceiveMessage(publish_namespace);
-  EXPECT_CALL(mock_stream_,
+  bidi_wrapper_->ReceiveMessage(publish_namespace);
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(SerializedControlMessage(MoqtPublishNamespaceCancel{
                          kDefaultPeerRequestId,
                          RequestErrorCode::kInternalError, "deadbeef"}),
@@ -1035,8 +1096,8 @@
 TEST_F(MoqtSessionTest, ReplyToPublishNamespaceWithError) {
   TrackNamespace track_namespace{"foo"};
 
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   MessageParameters parameters;
   parameters.authorization_tokens.emplace_back(AuthTokenType::kOutOfBand,
                                                "foo");
@@ -1055,31 +1116,20 @@
       .WillOnce(
           [&](const TrackNamespace&, const std::optional<MessageParameters>&,
               MoqtResponseCallback callback) { std::move(callback)(error); });
-  EXPECT_CALL(mock_stream_,
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(SerializedControlMessage(MoqtRequestError{
                          kDefaultPeerRequestId, error.error_code,
                          error.retry_interval, error.reason_phrase}),
                      _));
-  stream_input->ReceiveMessage(publish_namespace);
+  bidi_wrapper_->ReceiveMessage(publish_namespace);
 }
 
 TEST_F(MoqtSessionTest, SubscribeNamespaceLifeCycle) {
   TrackNamespace prefix({"foo"});
   bool got_callback = false;
-  EXPECT_CALL(mock_session_, CanOpenNextOutgoingBidirectionalStream())
-      .WillOnce(Return(true));
-  EXPECT_CALL(mock_session_, OpenOutgoingBidirectionalStream())
-      .WillOnce(Return(&mock_stream_));
-  std::unique_ptr<MoqtNamespaceSubscriberStream> stream_input;
-  EXPECT_CALL(mock_stream_, SetVisitor)
-      .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) {
-        stream_input = absl::WrapUnique(
-            absl::down_cast<MoqtNamespaceSubscriberStream*>(visitor.release()));
-        ASSERT_NE(stream_input, nullptr);
-      });
-  EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true));
+  PrepareRequestStream(bidi_wrapper_);
   EXPECT_CALL(
-      mock_stream_,
+      mock_bidi_stream_,
       Writev(ControlMessageOfType(MoqtMessageType::kSubscribeNamespace), _));
   std::unique_ptr<MoqtNamespaceTask> task = session_.SubscribeNamespace(
       prefix, SubscribeNamespaceOption::kNamespace, MessageParameters(),
@@ -1088,28 +1138,17 @@
         EXPECT_TRUE(std::holds_alternative<MessageParameters>(response));
       });
   MoqtRequestOk ok = {kDefaultLocalRequestId, MessageParameters()};
-  QUICHE_ASSERT_OK(stream_input->OnControlMessage(ok));
+  bidi_wrapper_->ReceiveMessage(ok);
   EXPECT_TRUE(got_callback);
-  EXPECT_CALL(mock_stream_, ResetWithUserCode);
+  EXPECT_CALL(mock_bidi_stream_, ResetWithUserCode);
 }
 
 TEST_F(MoqtSessionTest, SubscribeNamespaceError) {
   TrackNamespace prefix({"foo"});
   bool got_callback = false;
-  EXPECT_CALL(mock_session_, CanOpenNextOutgoingBidirectionalStream())
-      .WillOnce(Return(true));
-  EXPECT_CALL(mock_session_, OpenOutgoingBidirectionalStream())
-      .WillOnce(Return(&mock_stream_));
-  std::unique_ptr<MoqtNamespaceSubscriberStream> stream_input;
-  EXPECT_CALL(mock_stream_, SetVisitor)
-      .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) {
-        stream_input = std::unique_ptr<MoqtNamespaceSubscriberStream>(
-            absl::down_cast<MoqtNamespaceSubscriberStream*>(visitor.release()));
-        ASSERT_NE(stream_input, nullptr);
-      });
-  EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true));
+  PrepareRequestStream(bidi_wrapper_);
   EXPECT_CALL(
-      mock_stream_,
+      mock_bidi_stream_,
       Writev(ControlMessageOfType(MoqtMessageType::kSubscribeNamespace), _));
   std::unique_ptr<MoqtNamespaceTask> task = session_.SubscribeNamespace(
       prefix, SubscribeNamespaceOption::kNamespace, MessageParameters(),
@@ -1124,7 +1163,7 @@
   MoqtRequestError error = {kDefaultLocalRequestId,
                             RequestErrorCode::kInvalidRange, std::nullopt,
                             "deadbeef"};
-  QUICHE_ASSERT_OK(stream_input->OnControlMessage(error));
+  bidi_wrapper_->ReceiveMessage(error);
   EXPECT_TRUE(got_callback);
 }
 
@@ -1136,17 +1175,8 @@
                 [&](std::variant<MessageParameters, MoqtRequestErrorInfo>) {}),
             nullptr);
   // kBoth is treated as kNamespace.
-  EXPECT_CALL(mock_session_, CanOpenNextOutgoingBidirectionalStream())
-      .WillOnce(Return(true));
-  EXPECT_CALL(mock_session_, OpenOutgoingBidirectionalStream())
-      .WillOnce(Return(&mock_stream_));
-  std::unique_ptr<webtransport::StreamVisitor> stream_visitor;
-  EXPECT_CALL(mock_stream_, SetVisitor)
-      .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) {
-        stream_visitor = std::move(visitor);
-      });
-  EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true));
-  EXPECT_CALL(mock_stream_,
+  PrepareRequestStream(bidi_wrapper_);
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(SerializedControlMessage(MoqtSubscribeNamespace{
                          0, prefix, SubscribeNamespaceOption::kNamespace,
                          MessageParameters()}),
@@ -1158,9 +1188,7 @@
 }
 
 TEST_F(MoqtSessionTest, SubscribeOkWithBadTrackAlias) {
-  webtransport::test::MockStream mock_control_stream;
-  std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_control_stream);
+  PrepareRequestStream(bidi_wrapper_);
   session_.Subscribe(FullTrackName("foo", "bar"), &remote_track_visitor_,
                      MessageParameters());
   MoqtSubscribeOk subscribe_ok = {
@@ -1169,42 +1197,44 @@
       MessageParameters(),
       TrackExtensions(),
   };
-  control_stream->ReceiveMessage(subscribe_ok);
+  bidi_wrapper_->ReceiveMessage(subscribe_ok);
   // Second subscribe, but OK has the same track alias.
+  webtransport::test::MockStream bidi_stream_2;
+  std::unique_ptr<MoqtBidiStreamTestWrapper> bidi_wrapper_2 =
+      MoqtSessionPeer::CreateControlStream(&session_, &bidi_stream_2);
+  PrepareRequestStream(bidi_wrapper_2, &bidi_stream_2);
   session_.Subscribe(FullTrackName("foo2", "bar2"), &remote_track_visitor_,
                      MessageParameters());
   subscribe_ok.request_id += 2;
   EXPECT_CALL(
       mock_session_,
       CloseSession(static_cast<uint64_t>(MoqtError::kDuplicateTrackAlias),
-                   "Duplicate track alias"));
-  control_stream->ReceiveMessage(subscribe_ok);
+                   "Track alias already exists"));
+  bidi_wrapper_2->ReceiveMessage(subscribe_ok);
 }
 
 TEST_F(MoqtSessionTest, ReceiveUnsubscribe) {
   MockTrackPublisher* track = CreateTrackPublisher();
   MoqtSubscribe request = DefaultSubscribe();
   const MoqtPriority kLocalDefaultPriority = 0x20;
-  std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeByte));
   TrackExtensions extensions(std::nullopt, std::nullopt, kLocalDefaultPriority,
                              std::nullopt, std::nullopt, std::nullopt);
   EXPECT_CALL(*track, extensions)
       .WillRepeatedly(testing::ReturnRef(extensions));
   MoqtObjectListener* listener = ReceiveSubscribeSynchronousOk(
-      track, request, control_stream.get(), /*track_alias=*/0, extensions);
-  MoqtUnsubscribe unsubscribe = {/*request_id=*/1};
+      track, request, bidi_wrapper_.get(), /*track_alias=*/0, extensions);
   EXPECT_CALL(*track, RemoveObjectListener(listener));
-  control_stream->ReceiveMessage(unsubscribe);
+  bidi_wrapper_->stream().OnResetStreamReceived(kResetCodeCancelled);
 }
 
 TEST_F(MoqtSessionTest, ReceiveDatagram) {
   FullTrackName ftn("foo", "bar");
   const MoqtPriority kPeerDefaultPriority = 0x20;
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  PrepareRequestStream(bidi_wrapper_);
   std::string payload = "deadbeef";
-  EXPECT_CALL(mock_stream_,
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _));
   session_.Subscribe(ftn, &remote_track_visitor_, MessageParameters());
   MoqtSubscribeOk ok;
@@ -1214,7 +1244,7 @@
       TrackExtensions(std::nullopt, std::nullopt, kPeerDefaultPriority,
                       std::nullopt, std::nullopt, std::nullopt);
   EXPECT_CALL(remote_track_visitor_, OnReply);
-  stream_input->ReceiveMessage(ok);
+  bidi_wrapper_->ReceiveMessage(ok);
 
   MoqtObject object = {
       /*track_alias=*/2,
@@ -1249,9 +1279,8 @@
 TEST_F(MoqtSessionTest, UsePeerDefaultPriority) {
   FullTrackName ftn("foo", "bar");
   const MoqtPriority kPeerDefaultPriority = 0x20;
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
-  EXPECT_CALL(mock_stream_,
+  PrepareRequestStream(bidi_wrapper_);
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _));
   session_.Subscribe(ftn, &remote_track_visitor_, MessageParameters());
   MoqtSubscribeOk ok;
@@ -1261,7 +1290,7 @@
       TrackExtensions(std::nullopt, std::nullopt, kPeerDefaultPriority,
                       std::nullopt, std::nullopt, std::nullopt);
   EXPECT_CALL(remote_track_visitor_, OnReply);
-  stream_input->ReceiveMessage(ok);
+  bidi_wrapper_->ReceiveMessage(ok);
   // Omit priority from a datagram.
   char datagram[] = {0x0c, 0x02, 0x05, 0x64, 0x65, 0x61,
                      0x64, 0x62, 0x65, 0x65, 0x66};
@@ -1296,8 +1325,8 @@
 TEST_F(MoqtSessionTest, OmitPublisherPriority) {
   MoqtSubscribe request = DefaultSubscribe();
   const MoqtPriority kLocalDefaultPriority = 0x20;
-  std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeByte));
   // Create the publisher and the SUBSCRIBE with kLocalDefaultPriority.
   MockTrackPublisher* track = CreateTrackPublisher();
   std::make_shared<MockTrackPublisher>(request.full_track_name);
@@ -1306,25 +1335,25 @@
   EXPECT_CALL(*track, extensions)
       .WillRepeatedly(testing::ReturnRef(extensions));
   MoqtObjectListener* listener = ReceiveSubscribeSynchronousOk(
-      track, request, control_stream.get(), /*track_alias=*/0, extensions);
+      track, request, bidi_wrapper_.get(), /*track_alias=*/0, extensions);
 
   // Deliver an object with kLocalDefaultPriority; stream_type will omit
   // the priority.
   EXPECT_CALL(mock_session_, CanOpenNextOutgoingUnidirectionalStream())
       .WillOnce(Return(true));
   EXPECT_CALL(mock_session_, OpenOutgoingUnidirectionalStream())
-      .WillOnce(Return(&mock_stream_));
-  EXPECT_CALL(mock_stream_, GetStreamId()).WillRepeatedly(Return(1));
+      .WillOnce(Return(&mock_uni_stream_));
+  EXPECT_CALL(mock_uni_stream_, GetStreamId()).WillRepeatedly(Return(1));
   std::unique_ptr<webtransport::StreamVisitor> stream_visitor;
-  EXPECT_CALL(mock_stream_, SetVisitor)
+  EXPECT_CALL(mock_uni_stream_, SetVisitor)
       .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) {
         stream_visitor = std::move(visitor);
       });
-  EXPECT_CALL(mock_stream_, SetPriority);
-  EXPECT_CALL(mock_stream_, visitor()).WillRepeatedly([&]() {
+  EXPECT_CALL(mock_uni_stream_, SetPriority);
+  EXPECT_CALL(mock_uni_stream_, visitor()).WillRepeatedly([&]() {
     return stream_visitor.get();
   });
-  EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true));
+  EXPECT_CALL(mock_uni_stream_, CanWrite).WillRepeatedly(Return(true));
   EXPECT_CALL(*track, GetCachedObject(_, _, _, _))
       .WillOnce(Return(PublishedObject{
           PublishedObjectMetadata{
@@ -1332,7 +1361,7 @@
               kLocalDefaultPriority, true, 8, MoqtSessionPeer::Now(&session_)},
           PayloadFromString("deadbeef")}))
       .WillOnce(Return(std::nullopt));
-  EXPECT_CALL(mock_stream_, Writev)
+  EXPECT_CALL(mock_uni_stream_, Writev)
       .WillOnce([&](absl::Span<quiche::QuicheMemSlice> data,
                     const webtransport::StreamWriteOptions& options) {
         // The stream type omits the priority.
@@ -1380,9 +1409,8 @@
 TEST_F(MoqtSessionTest, DatagramOutOfWindow) {
   FullTrackName ftn("foo", "bar");
   const MoqtPriority kPeerDefaultPriority = 0x20;
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
-  EXPECT_CALL(mock_stream_,
+  PrepareRequestStream(bidi_wrapper_);
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _));
   MessageParameters params;
   params.subscription_filter.emplace(Location(1, 0));
@@ -1394,7 +1422,7 @@
       TrackExtensions(std::nullopt, std::nullopt, kPeerDefaultPriority,
                       std::nullopt, std::nullopt, std::nullopt);
   EXPECT_CALL(remote_track_visitor_, OnReply);
-  stream_input->ReceiveMessage(ok);
+  bidi_wrapper_->ReceiveMessage(ok);
   char datagram[] = {0x01, 0x02, 0x00, 0x00, 0x80, 0x00, 0x08, 0x64,
                      0x65, 0x61, 0x64, 0x62, 0x65, 0x65, 0x66};
   EXPECT_CALL(remote_track_visitor_, OnObjectFragment).Times(0);
@@ -1517,8 +1545,8 @@
 
 // All callbacks are called asynchronously.
 TEST_F(MoqtSessionTest, ProcessFetchGetEverythingFromUpstream) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   MoqtFetch fetch = DefaultFetch();
   MockTrackPublisher* track = CreateTrackPublisher();
 
@@ -1527,14 +1555,15 @@
   MockFetchTask* fetch_task = fetch_task_ptr.get();
   EXPECT_CALL(*track, StandaloneFetch)
       .WillOnce(Return(std::move(fetch_task_ptr)));
-  stream_input->ReceiveMessage(fetch);
+  bidi_wrapper_->ReceiveMessage(fetch);
 
   // Compose and send the FETCH_OK.
   MoqtFetchOk expected_ok;
   expected_ok.request_id = fetch.request_id;
   expected_ok.end_of_track = false;
   expected_ok.end_location = Location(1, 4);
-  EXPECT_CALL(mock_stream_, Writev(SerializedControlMessage(expected_ok), _));
+  EXPECT_CALL(mock_bidi_stream_,
+              Writev(SerializedControlMessage(expected_ok), _));
   fetch_task->CallFetchResponseCallback(expected_ok);
   // Data arrives.
   webtransport::test::MockStream data_stream;
@@ -1549,8 +1578,8 @@
 // All callbacks are called synchronously. All relevant data is cached (or this
 // is the original publisher).
 TEST_F(MoqtSessionTest, ProcessFetchWholeRangeIsPresent) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   MoqtFetch fetch = DefaultFetch();
   MockTrackPublisher* track = CreateTrackPublisher();
 
@@ -1563,7 +1592,8 @@
   MockFetchTask* fetch_task = fetch_task_ptr.get();
   EXPECT_CALL(*track, StandaloneFetch)
       .WillOnce(Return(std::move(fetch_task_ptr)));
-  EXPECT_CALL(mock_stream_, Writev(SerializedControlMessage(expected_ok), _));
+  EXPECT_CALL(mock_bidi_stream_,
+              Writev(SerializedControlMessage(expected_ok), _));
   webtransport::test::MockStream data_stream;
   std::unique_ptr<webtransport::StreamVisitor> stream_visitor;
   ExpectStreamOpen(mock_session_, fetch_task, data_stream, stream_visitor);
@@ -1572,13 +1602,13 @@
                    MoqtFetchTask::GetNextObjectResult::kPending);
   // Everything spins upon message receipt. FetchTask is generating the
   // necessary callbacks.
-  stream_input->ReceiveMessage(fetch);
+  bidi_wrapper_->ReceiveMessage(fetch);
 }
 
 TEST_F(MoqtSessionTest, SendFragmentedFetchObject) {
   using ::testing::ByMove;
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   MoqtFetch fetch = DefaultFetch();
   fetch.request_id = 3;  // Use an odd ID for peer request in client session.
   MockTrackPublisher* track = CreateTrackPublisher();
@@ -1591,27 +1621,27 @@
       .WillOnce(Return(ByMove(std::move(fetch_task_ptr))));
 
   // Receive FETCH, send FETCH_OK.
-  stream_input->ReceiveMessage(fetch);
+  bidi_wrapper_->ReceiveMessage(fetch);
   // FETCH_OK responding to the request.
   MoqtFetchOk expected_ok;
   expected_ok.request_id = fetch.request_id;
   expected_ok.end_of_track = false;
   expected_ok.end_location = Location(1, 0);
-  EXPECT_CALL(mock_stream_, Writev(SerializedControlMessage(expected_ok), _));
+  EXPECT_CALL(mock_bidi_stream_,
+              Writev(SerializedControlMessage(expected_ok), _));
   fetch_task->CallFetchResponseCallback(expected_ok);
 
-  webtransport::test::MockStream data_stream;
   std::unique_ptr<webtransport::StreamVisitor> stream_visitor;
   EXPECT_CALL(mock_session_, CanOpenNextOutgoingUnidirectionalStream)
       .WillRepeatedly(Return(true));
   EXPECT_CALL(mock_session_, OpenOutgoingUnidirectionalStream())
-      .WillOnce(Return(&data_stream));
-  EXPECT_CALL(data_stream, SetVisitor)
+      .WillOnce(Return(&mock_uni_stream_));
+  EXPECT_CALL(mock_uni_stream_, SetVisitor)
       .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) {
         stream_visitor = std::move(visitor);
       });
-  EXPECT_CALL(data_stream, SetPriority);
-  EXPECT_CALL(data_stream, CanWrite).WillRepeatedly(Return(true));
+  EXPECT_CALL(mock_uni_stream_, SetPriority);
+  EXPECT_CALL(mock_uni_stream_, CanWrite).WillRepeatedly(Return(true));
   // Trigger stream opening (calls SetObjectAvailableCallback with lambda1).
   // Setting the stream visitor will cause a second call to the callback.
   PublishedObjectMetadata metadata = {
@@ -1623,7 +1653,7 @@
         return MoqtFetchTask::GetNextObjectResult::kSuccess;
       })
       .WillOnce(Return(MoqtFetchTask::GetNextObjectResult::kPending));
-  EXPECT_CALL(data_stream, Writev)
+  EXPECT_CALL(mock_uni_stream_, Writev)
       .WillOnce([&](absl::Span<const quiche::QuicheMemSlice> data,
                     const webtransport::StreamWriteOptions& options) {
         EXPECT_EQ(data.size(), 2);
@@ -1631,7 +1661,7 @@
         return absl::OkStatus();
       });
   fetch_task->CallObjectsAvailableCallback();
-  // lambda1 ran, data_stream captured, stream_visitor set.
+  // lambda1 ran, mock_uni_stream_ captured, stream_visitor set.
   ASSERT_NE(stream_visitor, nullptr);
 
   // The second fragment is available.
@@ -1642,7 +1672,7 @@
         return MoqtFetchTask::GetNextObjectResult::kSuccess;
       })
       .WillRepeatedly(Return(MoqtFetchTask::GetNextObjectResult::kPending));
-  EXPECT_CALL(data_stream, Writev)
+  EXPECT_CALL(mock_uni_stream_, Writev)
       .WillOnce([&](absl::Span<const quiche::QuicheMemSlice> data,
                     const webtransport::StreamWriteOptions& options) {
         EXPECT_EQ(data.size(), 1);  // No header.
@@ -1655,8 +1685,8 @@
 // The publisher has the first object locally, but has to go upstream to get
 // the rest.
 TEST_F(MoqtSessionTest, FetchReturnsObjectBeforeOk) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   MoqtFetch fetch = DefaultFetch();
   MockTrackPublisher* track = CreateTrackPublisher();
 
@@ -1672,19 +1702,20 @@
   ExpectSendObject(fetch_task, data_stream, MoqtObjectStatus::kNormal,
                    Location(0, 0), "foo",
                    MoqtFetchTask::GetNextObjectResult::kPending);
-  stream_input->ReceiveMessage(fetch);
+  bidi_wrapper_->ReceiveMessage(fetch);
 
   MoqtFetchOk expected_ok;
   expected_ok.request_id = fetch.request_id;
   expected_ok.end_of_track = false;
   expected_ok.end_location = Location(1, 4);
-  EXPECT_CALL(mock_stream_, Writev(SerializedControlMessage(expected_ok), _));
+  EXPECT_CALL(mock_bidi_stream_,
+              Writev(SerializedControlMessage(expected_ok), _));
   fetch_task->CallFetchResponseCallback(expected_ok);
 }
 
 TEST_F(MoqtSessionTest, FetchReturnsObjectBeforeError) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   MoqtFetch fetch = DefaultFetch();
   MockTrackPublisher* track = CreateTrackPublisher();
 
@@ -1699,34 +1730,33 @@
   ExpectSendObject(fetch_task, data_stream, MoqtObjectStatus::kNormal,
                    Location(0, 0), "foo",
                    MoqtFetchTask::GetNextObjectResult::kPending);
-  stream_input->ReceiveMessage(fetch);
+  bidi_wrapper_->ReceiveMessage(fetch);
 
   MoqtRequestError expected_error{
       fetch.request_id, RequestErrorCode::kDoesNotExist, std::nullopt, "foo"};
-  EXPECT_CALL(mock_stream_,
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(SerializedControlMessage(expected_error), _));
   fetch_task->CallFetchResponseCallback(expected_error);
 }
 
 TEST_F(MoqtSessionTest, InvalidFetch) {
-  webtransport::test::MockStream control_stream;
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &control_stream);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   MockTrackPublisher* track = CreateTrackPublisher();
   MoqtFetch fetch = DefaultFetch();
   EXPECT_CALL(*track, StandaloneFetch)
       .WillOnce(Return(std::make_unique<MockFetchTask>()));
-  stream_input->ReceiveMessage(fetch);
+  bidi_wrapper_->ReceiveMessage(fetch);
   EXPECT_CALL(mock_session_,
               CloseSession(static_cast<uint64_t>(MoqtError::kInvalidRequestId),
                            "Duplicate request ID"))
       .Times(1);
-  stream_input->ReceiveMessage(fetch);
+  bidi_wrapper_->ReceiveMessage(fetch);
 }
 
 TEST_F(MoqtSessionTest, FetchFails) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   MoqtFetch fetch = DefaultFetch();
   MockTrackPublisher* track = CreateTrackPublisher();
 
@@ -1736,14 +1766,14 @@
       .WillOnce(Return(std::move(fetch_task_ptr)));
   EXPECT_CALL(*fetch_task, GetStatus())
       .WillRepeatedly(Return(absl::Status(absl::StatusCode::kInternal, "foo")));
-  EXPECT_CALL(mock_stream_,
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _));
-  stream_input->ReceiveMessage(fetch);
+  bidi_wrapper_->ReceiveMessage(fetch);
 }
 
 TEST_F(MoqtSessionTest, FullFetchDeliveryWithFlowControl) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   MoqtFetch fetch = DefaultFetch();
   MockTrackPublisher* track = CreateTrackPublisher();
 
@@ -1753,7 +1783,7 @@
   EXPECT_CALL(*track, StandaloneFetch)
       .WillOnce(Return(std::move(fetch_task_ptr)));
 
-  stream_input->ReceiveMessage(fetch);
+  bidi_wrapper_->ReceiveMessage(fetch);
   EXPECT_CALL(mock_session_, CanOpenNextOutgoingUnidirectionalStream())
       .WillOnce(Return(false));
   fetch_task->CallObjectsAvailableCallback();
@@ -1776,12 +1806,15 @@
   // Give it the latest object filter.
   subscribe.parameters.subscription_filter.emplace(
       MoqtFilterType::kLargestObject);
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeByte));
   MockTrackPublisher* track = CreateTrackPublisher();
   SetLargestId(track, Location(4, 10));
-  ReceiveSubscribeSynchronousOk(track, subscribe, stream_input.get());
+  ReceiveSubscribeSynchronousOk(track, subscribe, bidi_wrapper_.get());
 
+  webtransport::test::MockStream control_stream;
+  std::unique_ptr<MoqtBidiStreamTestWrapper> control_wrapper =
+      MoqtSessionPeer::CreateControlStream(&session_, &control_stream);
   ASSERT_TRUE(MoqtSessionPeer::RequestIdIsSubscriptionPublisher(
       &session_, subscribe.request_id));
   MoqtFetch fetch = DefaultFetch();
@@ -1789,7 +1822,7 @@
   fetch.fetch = JoiningFetchRelative(1, 2);
   EXPECT_CALL(*track, StandaloneFetch(Location(2, 0), Location(4, 10), _))
       .WillOnce(Return(std::make_unique<MockFetchTask>()));
-  stream_input->ReceiveMessage(fetch);
+  control_wrapper->ReceiveMessage(fetch);
 }
 
 TEST_F(MoqtSessionTest, IncomingAbsoluteJoiningFetch) {
@@ -1797,25 +1830,28 @@
   // Give it the latest object filter.
   subscribe.parameters.subscription_filter.emplace(
       MoqtFilterType::kLargestObject);
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeByte));
   MockTrackPublisher* track = CreateTrackPublisher();
   SetLargestId(track, Location(4, 10));
-  ReceiveSubscribeSynchronousOk(track, subscribe, stream_input.get());
+  ReceiveSubscribeSynchronousOk(track, subscribe, bidi_wrapper_.get());
 
   ASSERT_TRUE(MoqtSessionPeer::RequestIdIsSubscriptionPublisher(
       &session_, subscribe.request_id));
+  webtransport::test::MockStream control_stream;
+  std::unique_ptr<MoqtBidiStreamTestWrapper> control_wrapper =
+      MoqtSessionPeer::CreateControlStream(&session_, &control_stream);
   MoqtFetch fetch = DefaultFetch();
   fetch.request_id = 3;
   fetch.fetch = JoiningFetchAbsolute(1, 2);
   EXPECT_CALL(*track, StandaloneFetch(Location(2, 0), Location(4, 10), _))
       .WillOnce(Return(std::make_unique<MockFetchTask>()));
-  stream_input->ReceiveMessage(fetch);
+  control_wrapper->ReceiveMessage(fetch);
 }
 
 TEST_F(MoqtSessionTest, IncomingJoiningFetchBadRequestId) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   MoqtFetch fetch = DefaultFetch();
   fetch.fetch = JoiningFetchRelative(1, 2);
   MoqtRequestError expected_error = {
@@ -1824,20 +1860,23 @@
       /*retry_interval=*/std::nullopt,
       "Joining Fetch for non-existent request",
   };
-  EXPECT_CALL(mock_stream_,
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(SerializedControlMessage(expected_error), _));
-  stream_input->ReceiveMessage(fetch);
+  bidi_wrapper_->ReceiveMessage(fetch);
 }
 
 TEST_F(MoqtSessionTest, IncomingJoiningFetchForwardZero) {
   MoqtSubscribe subscribe = DefaultSubscribe();
   subscribe.parameters.set_forward(false);
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeByte));
   MockTrackPublisher* track = CreateTrackPublisher();
   SetLargestId(track, Location(2, 10));
-  ReceiveSubscribeSynchronousOk(track, subscribe, stream_input.get());
+  ReceiveSubscribeSynchronousOk(track, subscribe, bidi_wrapper_.get());
 
+  webtransport::test::MockStream control_stream;
+  std::unique_ptr<MoqtBidiStreamTestWrapper> control_wrapper =
+      MoqtSessionPeer::CreateControlStream(&session_, &control_stream);
   MoqtFetch fetch = DefaultFetch();
   fetch.request_id = 3;
   fetch.fetch = JoiningFetchRelative(1, 2);
@@ -1845,14 +1884,14 @@
               CloseSession(static_cast<uint64_t>(MoqtError::kProtocolViolation),
                            "Joining Fetch for non-forwarding subscribe"))
       .Times(1);
-  stream_input->ReceiveMessage(fetch);
+  control_wrapper->ReceiveMessage(fetch);
 }
 
 TEST_F(MoqtSessionTest, SendJoiningFetch) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
-  EXPECT_CALL(mock_session_, GetStreamById(_))
-      .WillRepeatedly(Return(&mock_stream_));
+  PrepareRequestStream(bidi_wrapper_);
+  webtransport::test::MockStream control_stream;
+  std::unique_ptr<MoqtBidiStreamTestWrapper> control_wrapper =
+      MoqtSessionPeer::CreateControlStream(&session_, &control_stream);
   MoqtSubscribe expected_subscribe(
       0, FullTrackName("foo", "bar"),
       MessageParameters(MoqtFilterType::kLargestObject));
@@ -1861,9 +1900,9 @@
       /*fetch=*/JoiningFetchRelative(0, 1),
       MessageParameters(),
   };
-  EXPECT_CALL(mock_stream_,
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(SerializedControlMessage(expected_subscribe), _));
-  EXPECT_CALL(mock_stream_,
+  EXPECT_CALL(control_stream,
               Writev(SerializedControlMessage(expected_fetch), _));
   EXPECT_TRUE(session_.RelativeJoiningFetch(expected_subscribe.full_track_name,
                                             &remote_track_visitor_, nullptr, 1,
@@ -1871,13 +1910,13 @@
 }
 
 TEST_F(MoqtSessionTest, SendJoiningFetchNoFlowControl) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
-  EXPECT_CALL(mock_session_, GetStreamById(_))
-      .WillRepeatedly(Return(&mock_stream_));
-  EXPECT_CALL(mock_stream_,
+  PrepareRequestStream(bidi_wrapper_);
+  webtransport::test::MockStream control_stream;
+  std::unique_ptr<MoqtBidiStreamTestWrapper> control_wrapper =
+      MoqtSessionPeer::CreateControlStream(&session_, &control_stream);
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _));
-  EXPECT_CALL(mock_stream_,
+  EXPECT_CALL(control_stream,
               Writev(ControlMessageOfType(MoqtMessageType::kFetch), _));
   EXPECT_TRUE(session_.RelativeJoiningFetch(FullTrackName("foo", "bar"),
                                             &remote_track_visitor_, 0,
@@ -1886,9 +1925,9 @@
   EXPECT_CALL(remote_track_visitor_, OnReply).Times(1);
   MessageParameters parameters;
   parameters.largest_object = Location(2, 0);
-  stream_input->ReceiveMessage(
+  bidi_wrapper_->ReceiveMessage(
       MoqtSubscribeOk(0, 2, parameters, TrackExtensions()));
-  stream_input->ReceiveMessage(MoqtFetchOk(
+  control_wrapper->ReceiveMessage(MoqtFetchOk(
       2, false, Location(2, 0), MessageParameters(), TrackExtensions()));
   // Packet arrives on FETCH stream.
   MoqtObject object = {
@@ -1912,7 +1951,7 @@
   data_stream.Receive(header.AsStringView(), false);
   EXPECT_CALL(remote_track_visitor_, OnObjectFragment).Times(1);
   // Last object of the FETCH causes FETCH_CANCEL.
-  EXPECT_CALL(mock_stream_,
+  EXPECT_CALL(control_stream,
               Writev(ControlMessageOfType(MoqtMessageType::kFetchCancel), _));
   data_stream.Receive("foo", false);
 }
@@ -1922,17 +1961,9 @@
   MessageParameters parameters;
   parameters.authorization_tokens.emplace_back(AuthTokenType::kOutOfBand,
                                                "foo");
-  auto bidi_stream =
-      std::make_unique<webtransport::test::InMemoryStreamWithWriteBuffer>(4);
-  MoqtFramer framer(true, quic::Perspective::IS_SERVER);
+
   MoqtSubscribeNamespace subscribe_namespace = {
       /*request_id=*/1, prefix, SubscribeNamespaceOption::kBoth, parameters};
-  bidi_stream->Receive(
-      framer.SerializeSubscribeNamespace(subscribe_namespace).AsStringView(),
-      /*fin=*/false);
-  EXPECT_CALL(mock_session_, AcceptIncomingBidirectionalStream())
-      .WillOnce(Return(bidi_stream.get()))
-      .WillOnce(Return(nullptr));
   quiche::QuicheWeakPtr<MockNamespaceTask> task;
   MoqtRequestOk expected_ok(/*request_id=*/1);
   expected_ok.parameters.expires = quic::QuicTimeDelta::FromSeconds(60);
@@ -1946,10 +1977,12 @@
         task = task_ptr->GetWeakPtr();
         return task_ptr;
       });
-  session_.OnIncomingBidirectionalStreamAvailable();
-  EXPECT_EQ(bidi_stream->write_buffer(),
-            framer.SerializeRequestOk(expected_ok).AsStringView());
-  bidi_stream->write_buffer().clear();
+  bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeNamespaceByte));
+  EXPECT_CALL(mock_bidi_stream_,
+              Writev(SerializedControlMessage(expected_ok), _))
+      .WillOnce(Return(absl::OkStatus()));
+  bidi_wrapper_->ReceiveMessage(subscribe_namespace);
 
   // Deliver a NAMESPACE
   ASSERT_TRUE(task.IsValid());
@@ -1960,14 +1993,16 @@
         return GetNextResult::kSuccess;
       })
       .WillOnce(Return(GetNextResult::kPending));
+  MoqtNamespace expected_namespace = {
+      TrackNamespace({"bar"}),
+  };
+  EXPECT_CALL(mock_bidi_stream_,
+              Writev(SerializedControlMessage(expected_namespace), _))
+      .WillOnce(Return(absl::OkStatus()));
   task.GetIfAvailable()->InvokeCallback();
-  char expected_data[] = {0x08, 0x00, 0x05, 0x01, 0x03, 'b', 'a', 'r'};
-  absl::string_view expected_data_view(expected_data, sizeof(expected_data));
-  EXPECT_EQ(expected_data_view,
-            bidi_stream->write_buffer().substr(0, expected_data_view.length()));
 
   // Unsubscribe
-  bidi_stream.reset();
+  bidi_wrapper_.reset();
   EXPECT_FALSE(task.IsValid());
 }
 
@@ -1976,16 +2011,10 @@
   MessageParameters parameters;
   parameters.authorization_tokens.emplace_back(AuthTokenType::kOutOfBand,
                                                "foo");
-  webtransport::test::InMemoryStreamWithWriteBuffer bidi_stream(4);
-  MoqtFramer framer(true, quic::Perspective::IS_SERVER);
+  bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeNamespaceByte));
   MoqtSubscribeNamespace subscribe_namespace = {
       /*request_id=*/1, prefix, SubscribeNamespaceOption::kBoth, parameters};
-  bidi_stream.Receive(
-      framer.SerializeSubscribeNamespace(subscribe_namespace).AsStringView(),
-      /*fin=*/false);
-  EXPECT_CALL(mock_session_, AcceptIncomingBidirectionalStream())
-      .WillOnce(Return(&bidi_stream))
-      .WillOnce(Return(nullptr));
   EXPECT_CALL(session_callbacks_.incoming_subscribe_namespace_callback,
               Call(prefix, SubscribeNamespaceOption::kBoth, parameters, _))
       .WillOnce([&](const TrackNamespace&, SubscribeNamespaceOption,
@@ -1995,10 +2024,14 @@
             RequestErrorCode::kUnauthorized, std::nullopt, "foo"});
         return nullptr;
       });
-  session_.OnIncomingBidirectionalStreamAvailable();
-  EXPECT_EQ(PeekControlMessageType(bidi_stream.write_buffer()),
-            MoqtMessageType::kRequestError);
-  EXPECT_TRUE(bidi_stream.fin_sent());
+  EXPECT_CALL(mock_bidi_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _))
+      .WillOnce([](absl::Span<quiche::QuicheMemSlice>,
+                   const webtransport::StreamWriteOptions& options) {
+        EXPECT_TRUE(options.send_fin());
+        return absl::OkStatus();
+      });
+  bidi_wrapper_->ReceiveMessage(subscribe_namespace);
 }
 
 TEST_F(MoqtSessionTest, IncomingSubscribeNamespaceWithPrefixOverlap) {
@@ -2006,17 +2039,10 @@
   MessageParameters parameters;
   parameters.authorization_tokens.emplace_back(AuthTokenType::kOutOfBand,
                                                "foo");
-  webtransport::test::InMemoryStreamWithWriteBuffer bidi_stream1(4),
-      bidi_stream2(8);
-  MoqtFramer framer(true, quic::Perspective::IS_SERVER);
+  bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeNamespaceByte));
   MoqtSubscribeNamespace subscribe_namespace = {
       /*request_id=*/1, foo, SubscribeNamespaceOption::kBoth, parameters};
-  bidi_stream1.Receive(
-      framer.SerializeSubscribeNamespace(subscribe_namespace).AsStringView(),
-      /*fin=*/false);
-  EXPECT_CALL(mock_session_, AcceptIncomingBidirectionalStream())
-      .WillOnce(Return(&bidi_stream1))
-      .WillOnce(Return(nullptr));
   EXPECT_CALL(session_callbacks_.incoming_subscribe_namespace_callback,
               Call(foo, SubscribeNamespaceOption::kBoth, parameters, _))
       .WillOnce([&](const TrackNamespace& prefix, SubscribeNamespaceOption,
@@ -2026,27 +2052,29 @@
         auto task_ptr = std::make_unique<MockNamespaceTask>(prefix);
         return task_ptr;
       });
-  session_.OnIncomingBidirectionalStreamAvailable();
-  EXPECT_EQ(PeekControlMessageType(bidi_stream1.write_buffer()),
-            MoqtMessageType::kRequestOk);
+  EXPECT_CALL(mock_bidi_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kRequestOk), _))
+      .WillOnce(Return(absl::OkStatus()));
+  bidi_wrapper_->ReceiveMessage(subscribe_namespace);
 
   subscribe_namespace.request_id += 2;
   subscribe_namespace.track_namespace_prefix = foobar;
-  bidi_stream2.Receive(
-      framer.SerializeSubscribeNamespace(subscribe_namespace).AsStringView(),
-      /*fin=*/false);
-  EXPECT_CALL(mock_session_, AcceptIncomingBidirectionalStream())
-      .WillOnce(Return(&bidi_stream2))
-      .WillOnce(Return(nullptr));
-  session_.OnIncomingBidirectionalStreamAvailable();
-  EXPECT_EQ(PeekControlMessageType(bidi_stream2.write_buffer()),
-            MoqtMessageType::kRequestError);
-  EXPECT_TRUE(bidi_stream2.fin_sent());
+  webtransport::test::MockStream bidi_stream_2;
+  auto bidi_wrapper_2 = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeNamespaceByte, &bidi_stream_2));
+  EXPECT_CALL(bidi_stream_2,
+              Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _))
+      .WillOnce([](absl::Span<quiche::QuicheMemSlice>,
+                   const webtransport::StreamWriteOptions& options) {
+        EXPECT_TRUE(options.send_fin());
+        return absl::OkStatus();
+      });
+  bidi_wrapper_2->ReceiveMessage(subscribe_namespace);
 }
 
 TEST_F(MoqtSessionTest, FetchThenOkThenCancel) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   std::unique_ptr<MoqtFetchTask> fetch_task;
   session_.Fetch(
       FullTrackName("foo", "bar"),
@@ -2059,21 +2087,21 @@
       /*end_of_track=*/false, Location(3, 25),
       MessageParameters(),    TrackExtensions(),
   };
-  stream_input->ReceiveMessage(ok);
+  bidi_wrapper_->ReceiveMessage(ok);
   ASSERT_NE(fetch_task, nullptr);
   EXPECT_TRUE(fetch_task->GetStatus().ok());
   PublishedObject object;
   EXPECT_EQ(fetch_task->GetNextObject(object),
             MoqtFetchTask::GetNextObjectResult::kPending);
   // Cancel the fetch.
-  EXPECT_CALL(mock_stream_,
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kFetchCancel), _));
   fetch_task.reset();
 }
 
 TEST_F(MoqtSessionTest, FetchThenError) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   std::unique_ptr<MoqtFetchTask> fetch_task;
   session_.Fetch(
       FullTrackName("foo", "bar"),
@@ -2087,7 +2115,7 @@
       /*retry_interval=*/std::nullopt,
       "No username provided",
   };
-  stream_input->ReceiveMessage(error);
+  bidi_wrapper_->ReceiveMessage(error);
   ASSERT_NE(fetch_task, nullptr);
   EXPECT_TRUE(absl::IsPermissionDenied(fetch_task->GetStatus()));
   EXPECT_EQ(fetch_task->GetStatus().message(), "No username provided");
@@ -2095,8 +2123,8 @@
 
 // The application takes objects as they arrive.
 TEST_F(MoqtSessionTest, IncomingFetchObjectsGreedyApp) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   std::unique_ptr<MoqtFetchTask> fetch_task;
   uint64_t expected_object_id = 0;
   session_.Fetch(
@@ -2165,7 +2193,7 @@
       MessageParameters(),
       TrackExtensions(),
   };
-  stream_input->ReceiveMessage(ok);
+  bidi_wrapper_->ReceiveMessage(ok);
   ASSERT_NE(fetch_task, nullptr);
   EXPECT_EQ(expected_object_id, 2);
 
@@ -2180,8 +2208,8 @@
 }
 
 TEST_F(MoqtSessionTest, IncomingFetchObjectsSlowApp) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   std::unique_ptr<MoqtFetchTask> fetch_task;
   uint64_t expected_object_id = 0;
   bool objects_available = false;
@@ -2237,7 +2265,7 @@
       /*end_of_track=*/false, Location(3, 25),
       MessageParameters(),    TrackExtensions(),
   };
-  stream_input->ReceiveMessage(ok);
+  bidi_wrapper_->ReceiveMessage(ok);
   ASSERT_NE(fetch_task, nullptr);
   EXPECT_TRUE(objects_available);
 
@@ -2278,10 +2306,10 @@
 TEST_F(MoqtSessionTest, DeliveryTimeoutParameter) {
   MoqtSubscribe request = DefaultSubscribe();
   request.parameters.delivery_timeout = quic::QuicTimeDelta::FromSeconds(1);
-  std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeByte));
   MockTrackPublisher* track = CreateTrackPublisher();
-  ReceiveSubscribeSynchronousOk(track, request, control_stream.get());
+  ReceiveSubscribeSynchronousOk(track, request, bidi_wrapper_.get());
   std::optional<quic::QuicTimeDelta> delivery_timeout =
       MoqtSessionPeer::GetDeliveryTimeout(&session_, request.request_id);
   EXPECT_TRUE(delivery_timeout.has_value() &&
@@ -2289,12 +2317,12 @@
 }
 
 TEST_F(MoqtSessionTest, ReceiveGoAwayEnforcement) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   EXPECT_CALL(session_callbacks_.goaway_received_callback, Call("foo"));
-  stream_input->ReceiveMessage(MoqtGoAway("foo"));
+  bidi_wrapper_->ReceiveMessage(MoqtGoAway("foo"));
   // New requests not allowed.
-  EXPECT_CALL(mock_stream_, Writev).Times(0);
+  EXPECT_CALL(mock_bidi_stream_, Writev).Times(0);
   MessageParameters parameters = SubscribeForTest();
   parameters.subscription_filter.emplace(MoqtFilterType::kLargestObject);
   EXPECT_FALSE(session_.Subscribe(FullTrackName("foo", "bar"),
@@ -2324,50 +2352,52 @@
         reported_error = true;
         EXPECT_EQ(error_message, "Received multiple GOAWAY messages");
       });
-  stream_input->ReceiveMessage(MoqtGoAway("foo"));
+  bidi_wrapper_->ReceiveMessage(MoqtGoAway("foo"));
 }
 
 TEST_F(MoqtSessionTest, SendGoAwayEnforcement) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   CreateTrackPublisher();
-  EXPECT_CALL(mock_stream_,
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kGoAway), _));
   session_.GoAway("");
-  EXPECT_CALL(mock_stream_,
+
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _));
-  stream_input->ReceiveMessage(DefaultSubscribe());
-  EXPECT_CALL(mock_stream_,
-              Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _));
-  stream_input->ReceiveMessage(
+  bidi_wrapper_->ReceiveMessage(
       MoqtPublishNamespace(3, TrackNamespace({"foo"}), MessageParameters()));
-  EXPECT_CALL(mock_stream_,
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _));
   MoqtFetch fetch = DefaultFetch();
   fetch.request_id = 5;
-  stream_input->ReceiveMessage(fetch);
+  bidi_wrapper_->ReceiveMessage(fetch);
 
-  MoqtFramer framer(true, quic::Perspective::IS_CLIENT);
-  SessionNamespaceTree tree;
-  MoqtIncomingSubscribeNamespaceCallback callback =
-      DefaultIncomingSubscribeNamespaceCallback;
-  MoqtNamespacePublisherStream namespace_stream(
-      &framer,
-      MoqtControlMessageParser(kDefaultMoqtVersion, true,
-                               quic::Perspective::IS_CLIENT),
-      [](const TrackNamespace&) { return true; }, nullptr, nullptr, callback);
-  namespace_stream.BindStream(&mock_stream_);
-  EXPECT_CALL(mock_stream_,
-              Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _));
-  QUICHE_ASSERT_OK(
-      namespace_stream.OnControlMessage(MoqtSubscribeNamespace(7)));
-  MoqtTrackStatus track_status = DefaultSubscribe();
-  track_status.request_id = 7;
-  EXPECT_CALL(mock_stream_,
-              Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _));
-  stream_input->ReceiveMessage(track_status);
-  // Block all outgoing SUBSCRIBE, PUBLISH_NAMESPACE, GOAWAY,etc.
-  EXPECT_CALL(mock_stream_, Writev).Times(0);
+  // All new bidi streams types are immediately rejected.
+  webtransport::test::MockStream new_request_stream;
+  EXPECT_CALL(mock_session_, AcceptIncomingBidirectionalStream)
+      .WillOnce(Return(&new_request_stream))
+      .WillOnce(Return(nullptr));
+  EXPECT_CALL(new_request_stream, CanWrite()).WillOnce(Return(true));
+  EXPECT_CALL(new_request_stream,
+              Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _))
+      .WillOnce([](absl::Span<quiche::QuicheMemSlice>,
+                   const webtransport::StreamWriteOptions& options) {
+        EXPECT_TRUE(options.send_fin());
+        return absl::OkStatus();
+      });
+  session_.OnIncomingBidirectionalStreamAvailable();
+
+  // If a new stream can't be written, reset it.
+  webtransport::test::MockStream new_request_stream_2;
+  EXPECT_CALL(mock_session_, AcceptIncomingBidirectionalStream)
+      .WillOnce(Return(&new_request_stream_2))
+      .WillOnce(Return(nullptr));
+  EXPECT_CALL(new_request_stream_2, CanWrite()).WillOnce(Return(false));
+  EXPECT_CALL(new_request_stream_2, ResetWithUserCode);
+  session_.OnIncomingBidirectionalStreamAvailable();
+
+  // Block all outgoing PUBLISH_NAMESPACE, GOAWAY,etc.
   MessageParameters parameters = SubscribeForTest();
   parameters.subscription_filter.emplace(MoqtFilterType::kLargestObject);
   EXPECT_FALSE(session_.Subscribe(FullTrackName({"foo"}, "bar"),
@@ -2400,10 +2430,10 @@
 
 TEST_F(MoqtSessionTest, ClientCannotSendNewSessionUri) {
   // session_ is a client session.
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   // Client GOAWAY not sent.
-  EXPECT_CALL(mock_stream_, Writev).Times(0);
+  EXPECT_CALL(mock_bidi_stream_, Writev).Times(0);
   session_.GoAway("foo");
 }
 
@@ -2413,8 +2443,8 @@
                       MoqtSessionParameters(quic::Perspective::IS_SERVER),
                       std::make_unique<quic::test::TestAlarmFactory>(),
                       session_callbacks_.AsSessionCallbacks());
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session, &mock_stream_);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session, &mock_bidi_stream_);
   EXPECT_CALL(
       mock_session,
       CloseSession(static_cast<uint64_t>(MoqtError::kProtocolViolation),
@@ -2427,14 +2457,13 @@
         EXPECT_EQ(error_message,
                   "Received GOAWAY with new_session_uri on the server");
       });
-  stream_input->ReceiveMessage(MoqtGoAway("foo"));
+  bidi_wrapper_->ReceiveMessage(MoqtGoAway("foo"));
   EXPECT_TRUE(reported_error);
 }
 
 TEST_F(MoqtSessionTest, IncomingTrackStatusThenSynchronousOk) {
-  webtransport::test::MockStream control_stream;
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &control_stream);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   auto* track = CreateTrackPublisher();
 
   MoqtTrackStatus track_status = DefaultSubscribe();
@@ -2450,25 +2479,24 @@
         expected_ok.parameters.expires =
             quic::QuicTimeDelta::FromMilliseconds(10000);
         expected_ok.parameters.largest_object = Location(5, 30);
-        EXPECT_CALL(control_stream,
+        EXPECT_CALL(mock_bidi_stream_,
                     Writev(SerializedControlMessage(expected_ok), _));
         EXPECT_CALL(*track, RemoveObjectListener);
         listener->OnSubscribeAccepted();
       });
-  stream_input->ReceiveMessage(track_status);
+  bidi_wrapper_->ReceiveMessage(track_status);
 }
 
 TEST_F(MoqtSessionTest, IncomingTrackStatusThenAsynchronousOk) {
-  webtransport::test::MockStream control_stream;
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &control_stream);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   auto* track = CreateTrackPublisher();
 
   MoqtTrackStatus track_status = DefaultSubscribe();
   MoqtObjectListener* listener = nullptr;
   EXPECT_CALL(*track, AddObjectListener)
       .WillOnce(testing::SaveArg<0>(&listener));
-  stream_input->ReceiveMessage(track_status);
+  bidi_wrapper_->ReceiveMessage(track_status);
   ASSERT_NE(listener, nullptr);
   EXPECT_CALL(*track, expiration)
       .WillRepeatedly(Return(quic::QuicTimeDelta::FromMilliseconds(10000)));
@@ -2477,15 +2505,15 @@
   expected_ok.request_id = track_status.request_id;
   expected_ok.parameters.expires = quic::QuicTimeDelta::FromMilliseconds(10000);
   expected_ok.parameters.largest_object = Location(5, 30);
-  EXPECT_CALL(control_stream, Writev(SerializedControlMessage(expected_ok), _));
+  EXPECT_CALL(mock_bidi_stream_,
+              Writev(SerializedControlMessage(expected_ok), _));
   EXPECT_CALL(*track, RemoveObjectListener(listener));
   listener->OnSubscribeAccepted();
 }
 
 TEST_F(MoqtSessionTest, IncomingTrackStatusThenSynchronousError) {
-  webtransport::test::MockStream control_stream;
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &control_stream);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   auto* track = CreateTrackPublisher();
 
   MoqtTrackStatus track_status = DefaultSubscribe();
@@ -2493,30 +2521,29 @@
   EXPECT_CALL(*track, AddObjectListener)
       .WillOnce([&](MoqtObjectListener* listener) {
         EXPECT_CALL(
-            control_stream,
+            mock_bidi_stream_,
             Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _));
         EXPECT_CALL(*track, RemoveObjectListener);
         listener->OnSubscribeRejected(MoqtRequestErrorInfo(
             RequestErrorCode::kInternalError, std::nullopt, "Test error"));
         executed_AddObjectListener = true;
       });
-  stream_input->ReceiveMessage(track_status);
+  bidi_wrapper_->ReceiveMessage(track_status);
   EXPECT_TRUE(executed_AddObjectListener);
 }
 
 TEST_F(MoqtSessionTest, IncomingTrackStatusThenAsynchronousError) {
-  webtransport::test::MockStream control_stream;
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &control_stream);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   auto* track = CreateTrackPublisher();
 
   MoqtTrackStatus track_status = DefaultSubscribe();
   MoqtObjectListener* listener;
   EXPECT_CALL(*track, AddObjectListener)
       .WillOnce(testing::SaveArg<0>(&listener));
-  stream_input->ReceiveMessage(track_status);
+  bidi_wrapper_->ReceiveMessage(track_status);
   ASSERT_NE(listener, nullptr);
-  EXPECT_CALL(control_stream,
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _));
   EXPECT_CALL(*track, RemoveObjectListener(listener));
   listener->OnSubscribeRejected(MoqtRequestErrorInfo(
@@ -2524,11 +2551,8 @@
 }
 
 TEST_F(MoqtSessionTest, FinReportedToVisitor) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream =
-      MoqtSessionPeer::CreateControlStream(&session_, &control_stream_);
-  EXPECT_CALL(mock_session_, GetStreamById)
-      .WillRepeatedly(Return(&control_stream_));
-  EXPECT_CALL(control_stream_,
+  PrepareRequestStream(bidi_wrapper_);
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _));
   MessageParameters parameters = SubscribeForTest();
   parameters.subscription_filter.emplace(MoqtFilterType::kLargestObject);
@@ -2543,7 +2567,7 @@
             EXPECT_EQ(ftn, FullTrackName("foo", "bar"));
             EXPECT_TRUE(std::holds_alternative<SubscribeOkData>(response));
           });
-  control_stream->ReceiveMessage(ok);
+  bidi_wrapper_->ReceiveMessage(ok);
   MoqtObject object = {
       /*track_alias=*/2,
       /*group_id=*/0,
@@ -2555,13 +2579,13 @@
       /*first_object_in_subgroup=*/true,
       /*payload_length=*/0,
   };
-  EXPECT_CALL(mock_stream_, GetStreamId())
+  EXPECT_CALL(mock_uni_stream_, GetStreamId())
       .WillRepeatedly(Return(kIncomingUniStreamId));
   EXPECT_CALL(mock_session_, GetStreamById(kIncomingUniStreamId))
-      .WillRepeatedly(Return(&mock_stream_));
+      .WillRepeatedly(Return(&mock_uni_stream_));
   std::unique_ptr<webtransport::StreamVisitor> data_stream;
-  DeliverObject(object, /*fin=*/true, mock_session_, &mock_stream_, data_stream,
-                &remote_track_visitor_);
+  DeliverObject(object, /*fin=*/true, mock_session_, &mock_uni_stream_,
+                data_stream, &remote_track_visitor_);
   // The data stream died and destroyed the visitor (IncomingDataStream).
   EXPECT_CALL(remote_track_visitor_,
               OnStreamFin(FullTrackName("foo", "bar"), DataStreamIndex(0, 0)));
@@ -2569,11 +2593,8 @@
 }
 
 TEST_F(MoqtSessionTest, ResetReportedToVisitor) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream =
-      MoqtSessionPeer::CreateControlStream(&session_, &control_stream_);
-  EXPECT_CALL(mock_session_, GetStreamById)
-      .WillRepeatedly(Return(&control_stream_));
-  EXPECT_CALL(control_stream_,
+  PrepareRequestStream(bidi_wrapper_);
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _));
   MessageParameters parameters = SubscribeForTest();
   parameters.subscription_filter.emplace(MoqtFilterType::kLargestObject);
@@ -2588,7 +2609,7 @@
             EXPECT_EQ(ftn, FullTrackName("foo", "bar"));
             EXPECT_TRUE(std::holds_alternative<SubscribeOkData>(response));
           });
-  control_stream->ReceiveMessage(ok);
+  bidi_wrapper_->ReceiveMessage(ok);
   MoqtObject object = {
       /*track_alias=*/2,
       /*group_id=*/0,
@@ -2600,12 +2621,12 @@
       /*first_object_in_subgroup=*/true,
       /*payload_length=*/0,
   };
-  EXPECT_CALL(mock_stream_, GetStreamId())
+  EXPECT_CALL(mock_uni_stream_, GetStreamId())
       .WillRepeatedly(Return(kIncomingUniStreamId));
   EXPECT_CALL(mock_session_, GetStreamById(kIncomingUniStreamId))
-      .WillRepeatedly(Return(&mock_stream_));
+      .WillRepeatedly(Return(&mock_uni_stream_));
   std::unique_ptr<webtransport::StreamVisitor> data_stream;
-  DeliverObject(object, /*fin=*/false, mock_session_, &mock_stream_,
+  DeliverObject(object, /*fin=*/false, mock_session_, &mock_uni_stream_,
                 data_stream, &remote_track_visitor_);
   // The data stream died and destroyed the visitor (IncomingDataStream).
   data_stream->OnResetStreamReceived(kResetCodeCancelled);
@@ -2615,9 +2636,8 @@
 }
 
 TEST_F(MoqtSessionTest, IncomingPublishNamespaceCleanup) {
-  webtransport::test::MockStream control_stream;
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &control_stream);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   // Register two incoming PUBLISH_NAMESPACE.
   MoqtPublishNamespace publish_namespace{
       /*request_id=*/1, TrackNamespace{"foo"}, MessageParameters()};
@@ -2630,8 +2650,9 @@
                     MoqtResponseCallback callback) {
         std::move(callback)(expected_ok.parameters);
       });
-  EXPECT_CALL(control_stream, Writev(SerializedControlMessage(expected_ok), _));
-  stream_input->ReceiveMessage(publish_namespace);
+  EXPECT_CALL(mock_bidi_stream_,
+              Writev(SerializedControlMessage(expected_ok), _));
+  bidi_wrapper_->ReceiveMessage(publish_namespace);
 
   publish_namespace = MoqtPublishNamespace(
       /*request_id=*/3, TrackNamespace{"bar"}, MessageParameters());
@@ -2642,9 +2663,9 @@
                     MoqtResponseCallback callback) {
         std::move(callback)(MessageParameters());
       });
-  EXPECT_CALL(control_stream,
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kRequestOk), _));
-  stream_input->ReceiveMessage(publish_namespace);
+  bidi_wrapper_->ReceiveMessage(publish_namespace);
 
   // Revoke "bar"
   MoqtPublishNamespaceDone done{/*request_id=*/3};
@@ -2654,7 +2675,7 @@
       .WillOnce(
           [](const TrackNamespace&, const std::optional<MessageParameters>&,
              MoqtResponseCallback callback) { EXPECT_EQ(callback, nullptr); });
-  stream_input->ReceiveMessage(done);
+  bidi_wrapper_->ReceiveMessage(done);
 
   // Destroying the session should revoke "foo".
   EXPECT_CALL(
@@ -2683,85 +2704,77 @@
   session_.OnSessionReady();
 }
 
-TEST_F(MoqtSessionTest, SubscribeThenRequestOk) {
-  webtransport::test::MockStream control_stream;
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &control_stream);
-  MessageParameters parameters = SubscribeForTest();
-  parameters.subscription_filter.emplace(MoqtFilterType::kLargestObject);
-  session_.Subscribe(FullTrackName("foo", "bar"), &remote_track_visitor_,
-                     parameters);
-  EXPECT_CALL(mock_session_, CloseSession);
-  EXPECT_CALL(session_callbacks_.session_terminated_callback, Call);
-  stream_input->ReceiveMessage(MoqtRequestOk{0, MessageParameters()});
-}
 
 TEST_F(MoqtSessionTest, ClientSetupNotAllowedOnControlStream) {
   // While technically on the Control stream, when it arrives, it's an
   // UnknownBidiStream
-  std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   EXPECT_CALL(mock_session_, CloseSession);
   EXPECT_CALL(session_callbacks_.session_terminated_callback, Call);
-  control_stream->ReceiveMessage(
+  bidi_wrapper_->ReceiveMessage(
       MoqtSetup(SetupParameters("/", "example.com", 0)));
 }
 
 TEST_F(MoqtSessionTest, NamespaceNotAllowedOnControlStream) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   EXPECT_CALL(mock_session_, CloseSession);
   EXPECT_CALL(session_callbacks_.session_terminated_callback, Call);
-  control_stream->ReceiveMessage(MoqtNamespace());
+  bidi_wrapper_->ReceiveMessage(MoqtNamespace());
 }
 
 TEST_F(MoqtSessionTest, NamespaceDoneNotAllowedOnControlStream) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   EXPECT_CALL(mock_session_, CloseSession);
   EXPECT_CALL(session_callbacks_.session_terminated_callback, Call);
-  control_stream->ReceiveMessage(MoqtNamespaceDone());
+  bidi_wrapper_->ReceiveMessage(MoqtNamespaceDone());
 }
 
 TEST_F(MoqtSessionTest, IncomingRequestUpdateTriggersRequestOk) {
   MoqtSubscribe subscribe = DefaultSubscribe();
   MockTrackPublisher* track = CreateTrackPublisher();
-  std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
-  ReceiveSubscribeSynchronousOk(track, subscribe, control_stream.get(), 0);
-  EXPECT_CALL(mock_stream_,
+  bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeByte));
+  ReceiveSubscribeSynchronousOk(track, subscribe, bidi_wrapper_.get(), 0);
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kRequestOk), _));
-  control_stream->ReceiveMessage(MoqtRequestUpdate{3, 1, MessageParameters()});
+  bidi_wrapper_->ReceiveMessage(MoqtRequestUpdate{3, 1, MessageParameters()});
 }
 
 TEST_F(MoqtSessionTest, IncomingRequestUpdateTriggersRequestError) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
-  EXPECT_CALL(mock_stream_,
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _));
-  control_stream->ReceiveMessage(MoqtRequestUpdate{3, 1, MessageParameters()});
+  bidi_wrapper_->ReceiveMessage(MoqtRequestUpdate{3, 1, MessageParameters()});
 }
 
 TEST_F(MoqtSessionTest, StopSendingBlocksSubgroup) {
   MoqtSubscribe subscribe = DefaultSubscribe();
   MockTrackPublisher* track = CreateTrackPublisher();
-  std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ = std::make_unique<MoqtBidiStreamTestWrapper>(
+      ResponseStream(kSubscribeByte));
   MoqtObjectListener* listener =
-      ReceiveSubscribeSynchronousOk(track, subscribe, control_stream.get(), 0);
+      ReceiveSubscribeSynchronousOk(track, subscribe, bidi_wrapper_.get(), 0);
   EXPECT_CALL(mock_session_, CanOpenNextOutgoingUnidirectionalStream)
       .WillOnce(Return(true));
   EXPECT_CALL(mock_session_, OpenOutgoingUnidirectionalStream)
-      .WillOnce(Return(&mock_stream_));
+      .WillOnce(Return(&mock_uni_stream_));
   std::unique_ptr<webtransport::StreamVisitor> data_stream_visitor;
-  EXPECT_CALL(mock_stream_, SetVisitor)
+  EXPECT_CALL(mock_uni_stream_, SetVisitor)
       .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) {
         data_stream_visitor = std::move(visitor);
       });
-  EXPECT_CALL(mock_stream_, visitor).WillRepeatedly([&]() {
+  EXPECT_CALL(mock_uni_stream_, GetStreamId())
+      .WillRepeatedly(Return(kOutgoingUniStreamId));
+  EXPECT_CALL(mock_session_, GetStreamById(kOutgoingUniStreamId))
+      .WillRepeatedly(Return(&mock_uni_stream_));
+  EXPECT_CALL(mock_uni_stream_, visitor).WillRepeatedly([&]() {
     return data_stream_visitor.get();
   });
-  EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true));
+  EXPECT_CALL(mock_uni_stream_, CanWrite).WillRepeatedly(Return(true));
   EXPECT_CALL(*track, GetCachedObject(0, Optional(1), 0, 0))
       .WillOnce(Return(PublishedObject{
           PublishedObjectMetadata{Location(0, 0), 1, "",
@@ -2771,14 +2784,14 @@
   EXPECT_CALL(*track, GetCachedObject(0, Optional(1), 1, 0))
       .WillOnce(Return(std::nullopt));
   SetLargestId(track, Location(0, 0));
-  EXPECT_CALL(mock_stream_, Writev).WillOnce(Return(absl::OkStatus()));
+  EXPECT_CALL(mock_uni_stream_, Writev).WillOnce(Return(absl::OkStatus()));
   listener->OnNewObjectAvailable(Location(0, 0), 1, 0x80);
 
-  EXPECT_CALL(mock_stream_, ResetWithUserCode(kResetCodeCancelled));
+  EXPECT_CALL(mock_uni_stream_, ResetWithUserCode(kResetCodeCancelled));
   data_stream_visitor->OnStopSendingReceived(kResetCodeCancelled);
   // New object in the same subgroup should not be sent.
   EXPECT_CALL(*track, GetCachedObject).Times(0);
-  EXPECT_CALL(mock_stream_, Writev).Times(0);
+  EXPECT_CALL(mock_uni_stream_, Writev).Times(0);
   listener->OnNewObjectAvailable(Location(0, 1), 1, 0x80);
 }
 
@@ -2786,21 +2799,10 @@
   CreateTrackPublisher();
   std::shared_ptr<MoqtTrackPublisher> track_publisher =
       publisher_.GetTrack(kDefaultTrackName());
-  webtransport::test::MockStream mock_publish_stream;
-  EXPECT_CALL(mock_session_, CanOpenNextOutgoingBidirectionalStream())
-      .WillOnce(Return(true));
-  EXPECT_CALL(mock_session_, OpenOutgoingBidirectionalStream())
-      .WillOnce(Return(&mock_publish_stream));
 
-  std::unique_ptr<webtransport::StreamVisitor> publish_stream_visitor;
-  EXPECT_CALL(mock_publish_stream, SetVisitor)
-      .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) {
-        publish_stream_visitor = std::move(visitor);
-      });
-  EXPECT_CALL(mock_publish_stream, CanWrite).WillRepeatedly(Return(true));
-
+  PrepareRequestStream(bidi_wrapper_);
   // Verify PUBLISH message is sent on the publish stream.
-  EXPECT_CALL(mock_publish_stream,
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kPublish), _))
       .WillOnce(Return(absl::OkStatus()));
 
@@ -2810,17 +2812,11 @@
       [&](std::variant<MessageParameters, MoqtRequestErrorInfo> resp) {
         response = resp;
       }));
-  ASSERT_NE(publish_stream_visitor, nullptr);
-
-  std::unique_ptr<MoqtBidiStreamBase> bidi_stream(
-      absl::down_cast<MoqtBidiStreamBase*>(publish_stream_visitor.release()));
-  MoqtBidiStreamTestWrapper wrapper(std::move(bidi_stream));
 
   MoqtRequestOk request_ok;
   request_ok.request_id = 0;
   request_ok.parameters.delivery_timeout = quic::QuicTimeDelta::FromSeconds(2);
-
-  wrapper.ReceiveMessage(request_ok);
+  bidi_wrapper_->ReceiveMessage(request_ok);
 
   ASSERT_TRUE(response.has_value());
   EXPECT_TRUE(std::holds_alternative<MessageParameters>(*response));
@@ -2840,11 +2836,11 @@
 }
 
 TEST_F(MoqtSessionTest, PublishAfterGoaway) {
-  std::unique_ptr<MoqtBidiStreamTestWrapper> stream_input =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
+  bidi_wrapper_ =
+      MoqtSessionPeer::CreateControlStream(&session_, &mock_bidi_stream_);
   MoqtGoAway goaway;
   goaway.new_session_uri = "";
-  stream_input->ReceiveMessage(goaway);
+  bidi_wrapper_->ReceiveMessage(goaway);
   CreateTrackPublisher();
   std::shared_ptr<MoqtTrackPublisher> track_publisher =
       publisher_.GetTrack(kDefaultTrackName());
@@ -2854,10 +2850,8 @@
 }
 
 TEST_F(MoqtSessionTest, IncomingPublishAbortsPendingSubscribe) {
-  // 1. Start a pending SUBSCRIBE.
-  std::unique_ptr<MoqtBidiStreamTestWrapper> control_stream =
-      MoqtSessionPeer::CreateControlStream(&session_, &mock_stream_);
-  EXPECT_CALL(mock_stream_,
+  PrepareRequestStream(bidi_wrapper_);
+  EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _));
   MessageParameters parameters(SubscribeForTest());
   session_.Subscribe(kDefaultTrackName(), &remote_track_visitor_, parameters);
@@ -2873,110 +2867,24 @@
       };
 
   // Prepare PUBLISH message.
-  MoqtPublish publish;
-  publish.request_id = 0;  // Matches pending SUBSCRIBE request ID
-  publish.full_track_name = kDefaultTrackName();
-  publish.track_alias = 10;
-  publish.parameters.delivery_timeout = quic::QuicTimeDelta::FromSeconds(5);
-
-  MoqtFramer peer_framer(true, quic::Perspective::IS_SERVER);
-  quiche::QuicheBuffer serialized_buffer =
-      peer_framer.SerializePublish(publish);
-  std::string serialized(serialized_buffer.data(), serialized_buffer.size());
-
-  // Setup mock_publish_stream.
-  webtransport::test::MockStream mock_publish_stream;
-  EXPECT_CALL(mock_session_, AcceptIncomingBidirectionalStream())
-      .WillOnce(Return(&mock_publish_stream))
-      .WillOnce(Return(nullptr));
-
-  // Setup mock_publish_stream to return serialized data.
-  size_t data_read = 0;
-  auto peek_lambda = [&data_read,
-                      &serialized]() -> webtransport::Stream::PeekResult {
-    webtransport::Stream::PeekResult result;
-    result.peeked_data = absl::string_view(serialized.data() + data_read,
-                                           serialized.size() - data_read);
-    result.fin_next = (data_read == serialized.size());
-    result.all_data_received = false;
-    return result;
-  };
-  std::function<webtransport::Stream::PeekResult()> peek_action = peek_lambda;
-
-  auto readable_bytes_lambda = [&data_read, &serialized]() -> size_t {
-    if (data_read >= serialized.size()) {
-      return 0;
-    }
-    return serialized.size() - data_read;
-  };
-  std::function<size_t()> readable_bytes_action = readable_bytes_lambda;
-
-  auto read_lambda =
-      [&data_read, &serialized](
-          absl::Span<char> bytes_to_read) -> webtransport::Stream::ReadResult {
-    size_t read_size =
-        std::min(bytes_to_read.size(), serialized.size() - data_read);
-    memcpy(bytes_to_read.data(), serialized.data() + data_read, read_size);
-    data_read += read_size;
-    webtransport::Stream::ReadResult result;
-    result.bytes_read = read_size;
-    result.fin = (data_read == serialized.size());
-    return result;
-  };
-  std::function<webtransport::Stream::ReadResult(absl::Span<char>)>
-      read_action = read_lambda;
-
-  auto skip_lambda = [&data_read, &serialized](size_t bytes) -> bool {
-    data_read += bytes;
-    return data_read == serialized.size();
-  };
-  std::function<bool(size_t)> skip_action = skip_lambda;
-
-  EXPECT_CALL(mock_publish_stream, PeekNextReadableRegion())
-      .WillRepeatedly(peek_action);
-  EXPECT_CALL(mock_publish_stream, ReadableBytes())
-      .WillRepeatedly(readable_bytes_action);
-  EXPECT_CALL(mock_publish_stream, Read(testing::An<absl::Span<char>>()))
-      .WillRepeatedly(read_action);
-  EXPECT_CALL(mock_publish_stream, SkipBytes).WillRepeatedly(skip_action);
-
-  // Capture SetVisitor calls and mock visitor().
-  std::unique_ptr<webtransport::StreamVisitor> unknown_bidi_stream_visitor;
-  std::unique_ptr<webtransport::StreamVisitor> upgraded_visitor;
-
-  {
-    testing::InSequence seq;
-    EXPECT_CALL(mock_publish_stream, SetVisitor)
-        .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) {
-          unknown_bidi_stream_visitor = std::move(visitor);
-        })
-        .WillOnce([&](std::unique_ptr<webtransport::StreamVisitor> visitor) {
-          upgraded_visitor = std::move(visitor);
-        });
-  }
-
-  EXPECT_CALL(mock_publish_stream, visitor())
-      .WillRepeatedly([&]() -> webtransport::StreamVisitor* {
-        return upgraded_visitor ? upgraded_visitor.get()
-                                : unknown_bidi_stream_visitor.get();
-      });
-
-  // The RemoteTrackVisitor is being reused, not destroyed.
-  EXPECT_CALL(remote_track_visitor_,
-              OnReply(kDefaultTrackName(),
-                      testing::VariantWith<SubscribeOkData>(
-                          testing::Field(&SubscribeOkData::parameters,
-                                         testing::Eq(publish.parameters)))));
-  EXPECT_CALL(mock_stream_,
-              Writev(ControlMessageOfType(MoqtMessageType::kUnsubscribe), _))
-      .WillOnce(Return(absl::OkStatus()));
-  EXPECT_CALL(remote_track_visitor_, OnPublishDone(kDefaultTrackName()))
-      .Times(0);
+  MoqtPublish publish{1, kDefaultTrackName(), 10, MessageParameters(),
+                      TrackExtensions()};
+  webtransport::test::MockStream publish_stream;
+  std::unique_ptr<MoqtBidiStreamTestWrapper> publish_wrapper =
+      std::make_unique<MoqtBidiStreamTestWrapper>(
+          ResponseStream(kPublishByte, &publish_stream));
+  EXPECT_CALL(mock_bidi_stream_, ResetWithUserCode(kResetCodeCancelled));
+  MoqtRequestOk expected_request_ok;
+  expected_request_ok.request_id = publish.request_id;
+  expected_request_ok.parameters = parameters;  // params from the SUBSCRIBE.
+  // group_order can be in SUBSCRIBE but not REQUEST_OK.
+  expected_request_ok.parameters.group_order = std::nullopt;
+  EXPECT_CALL(publish_stream,
+              Writev(SerializedControlMessage(expected_request_ok), _));
+  // remote_track_visitor_ is reused, not destroyed.
+  EXPECT_CALL(remote_track_visitor_, OnReply);
+  publish_wrapper->ReceiveMessage(publish);
   EXPECT_FALSE(incoming_publish_callback_called);
-
-  // Trigger the read by making incoming stream available.
-  session_.OnIncomingBidirectionalStreamAvailable();
-
   // Verify it was aborted immediately (not at teardown).
   EXPECT_TRUE(
       testing::Mock::VerifyAndClearExpectations(&remote_track_visitor_));
@@ -3048,22 +2956,27 @@
   control_stream.write_buffer().clear();
 
   // 4. Subscribe
+  webtransport::test::InMemoryStreamWithWriteBuffer sub_stream(5);
+  EXPECT_CALL(mock_session_, CanOpenNextOutgoingBidirectionalStream)
+      .WillOnce(Return(true));
+  EXPECT_CALL(mock_session_, OpenOutgoingBidirectionalStream)
+      .WillOnce(Return(&sub_stream));
   FullTrackName track_name1("namespace2", "track1");
   bool s1 = session_.Subscribe(track_name1, &remote_track_visitor_,
                                MessageParameters());
   EXPECT_TRUE(s1);
-  EXPECT_EQ(get_request_id(control_stream), next_request_id);
+  EXPECT_EQ(get_request_id(sub_stream), next_request_id);
   next_request_id += 2;
-  control_stream.write_buffer().clear();
+  sub_stream.write_buffer().clear();
 
   // 5. SubscribeUpdate
   bool s_update = session_.SubscribeUpdate(
       track_name1, MessageParameters(),
       [](std::variant<MessageParameters, MoqtRequestErrorInfo>) {});
   EXPECT_TRUE(s_update);
-  EXPECT_EQ(get_request_id(control_stream), next_request_id);
+  EXPECT_EQ(get_request_id(sub_stream), next_request_id);
   next_request_id += 2;
-  control_stream.write_buffer().clear();
+  sub_stream.write_buffer().clear();
 
   // 6. Fetch
   FullTrackName fetch_track("namespace2", "fetch_track");
diff --git a/quiche/quic/moqt/moqt_subscribe_stream.cc b/quiche/quic/moqt/moqt_subscribe_stream.cc
new file mode 100644
index 0000000..473de89
--- /dev/null
+++ b/quiche/quic/moqt/moqt_subscribe_stream.cc
@@ -0,0 +1,224 @@
+// Copyright (c) 2026 The Chromium Authors. All rights reserved.
+// Use of this source code is governed by a BSD-style license that can be
+// found in the LICENSE file.
+
+#include "quiche/quic/moqt/moqt_subscribe_stream.h"
+
+#include <cstdint>
+#include <memory>
+#include <optional>
+#include <utility>
+#include <variant>
+
+#include "absl/base/nullability.h"
+#include "absl/status/status.h"
+#include "quiche/quic/core/quic_alarm_factory.h"
+#include "quiche/quic/core/quic_clock.h"
+#include "quiche/quic/moqt/moqt_bidi_stream.h"
+#include "quiche/quic/moqt/moqt_error.h"
+#include "quiche/quic/moqt/moqt_framer.h"
+#include "quiche/quic/moqt/moqt_key_value_pair.h"
+#include "quiche/quic/moqt/moqt_messages.h"
+#include "quiche/quic/moqt/moqt_names.h"
+#include "quiche/quic/moqt/moqt_parser.h"
+#include "quiche/quic/moqt/moqt_publisher.h"
+#include "quiche/quic/moqt/moqt_session_callbacks.h"
+#include "quiche/quic/moqt/moqt_subscription.h"
+#include "quiche/quic/moqt/moqt_track.h"
+#include "quiche/common/quiche_weak_ptr.h"
+
+namespace moqt {
+
+MoqtSubscribeRequestStream::MoqtSubscribeRequestStream(
+    MoqtFramer* absl_nonnull framer,
+    const MoqtControlMessageParser& message_parser, uint64_t request_id,
+    SessionErrorCallback session_error_callback, const FullTrackName& name,
+    SubscribeVisitor* absl_nonnull visitor, const MessageParameters& parameters,
+    SubscribeRemoteTrack::AddCallback add_callback,
+    SubscribeRemoteTrack::RemoveCallback remove_callback,
+    const quic::QuicClock* absl_nonnull clock,
+    quic::QuicAlarmFactory* absl_nonnull alarm_factory)
+    : MoqtBidiStreamBase(framer, message_parser,
+                         std::move(session_error_callback)),
+      track_(std::make_unique<SubscribeRemoteTrack>(
+          MoqtSubscribe{request_id, name, parameters}, visitor, this)),
+      add_callback_(std::move(add_callback)),
+      remove_callback_(std::move(remove_callback)),
+      clock_(clock),
+      alarm_factory_(alarm_factory) {}
+
+void MoqtSubscribeRequestStream::OnStreamBound() {
+  stream_parser()->set_allow_fin(true);
+  SendOrBufferMessageOrFatal(framer()->SerializeSubscribe(
+      MoqtSubscribe{track_->request_id(), track_->full_track_name(),
+                    track_->const_parameters()}));
+}
+
+absl::Status MoqtSubscribeRequestStream::OnRawControlMessage(
+    const MoqtRawControlMessage& message) {
+  return ControlMessageDispatcher::DispatchControlMessage(
+      *this, message_parser(), message, "subscribe request");
+}
+
+absl::Status MoqtSubscribeRequestStream::OnControlMessage(
+    const MoqtSubscribeOk& message) {
+  if (message.request_id != track_->request_id()) {
+    return absl::InvalidArgumentError("SUBSCRIBE_OK request ID mismatch");
+  }
+  if (add_callback_ == nullptr) {
+    return absl::InvalidArgumentError(
+        "Multiple SUBSCRIBE_OK on the same stream");
+  }
+  track_->set_track_alias(message.track_alias);
+  if (!std::move(add_callback_)(track_.get())) {
+    add_callback_ = nullptr;
+    OnFatalError(absl::AlreadyExistsError("Track alias already exists"));
+    return absl::OkStatus();
+  }
+  add_callback_ = nullptr;
+
+  track_->OnObjectOrOk(SubscribeOkData(message.parameters, message.extensions));
+  return absl::OkStatus();
+}
+
+absl::Status MoqtSubscribeRequestStream::OnControlMessage(
+    const MoqtRequestOk& message) {
+  if (!track_->track_alias().has_value()) {
+    // Not yet established.
+    OnFatalError(
+        absl::InvalidArgumentError("REQUEST_OK received before SUBSCRIBE_OK"));
+    return absl::OkStatus();
+  }
+  auto status_or_params = PopParameters();
+  if (status_or_params.ok()) {
+    MessageParameters parameters = status_or_params.value();
+    // EXPIRES or LARGEST_OBJECT could be present in REQUEST_OK.
+    if (message.parameters.largest_object.has_value()) {
+      parameters.largest_object = message.parameters.largest_object;
+    }
+    if (message.parameters.expires.has_value()) {
+      parameters.expires = message.parameters.expires;
+    }
+    track_->Update(parameters);
+  }
+  return MoqtBidiStreamBase::OnControlMessage(message);
+}
+
+absl::Status MoqtSubscribeRequestStream::OnControlMessage(
+    const MoqtRequestError& message) {
+  MoqtRequestErrorInfo error_info{message.error_code, message.retry_interval,
+                                  message.reason_phrase};
+  if (track_->ErrorIsAllowed()) {
+    if (track_->visitor() != nullptr) {
+      track_->visitor()->OnReply(track_->full_track_name(), error_info);
+    }
+    Fin();
+    return absl::OkStatus();
+  }
+  // In response to REQUEST_UPDATE, utilize the ResponseCallback and do not
+  // update parameters.
+  return MoqtBidiStreamBase::OnControlMessage(message);
+}
+
+absl::Status MoqtSubscribeRequestStream::OnControlMessage(
+    const MoqtPublishDone& message) {
+  if (track_ == nullptr) {
+    // PUBLISH_DONE can be sent before the subscriber rejects the track.
+    return absl::OkStatus();
+  }
+  track_->OnPublishDone(message.stream_count, clock_, alarm_factory_);
+  return absl::OkStatus();
+}
+
+void MoqtSubscribeRequestStream::Detach() {
+  if (remove_callback_ != nullptr) {
+    SubscribeRemoteTrack::RemoveCallback remove_callback =
+        std::move(remove_callback_);
+    remove_callback_ = nullptr;
+    std::move(remove_callback)(track_.get());
+  }
+  track_ = nullptr;
+}
+
+MoqtSubscribeResponseStream::MoqtSubscribeResponseStream(
+    MoqtFramer* absl_nonnull framer,
+    const MoqtControlMessageParser& message_parser, uint64_t track_alias,
+    SubscriptionPublisher::AddCallback add_callback,
+    SubscriptionPublisher::RemoveCallback remove_callback,
+    SessionErrorCallback session_error_callback,
+    quiche::QuicheWeakPtr<SessionToPublisherInterface> session)
+    : MoqtBidiStreamBase(framer, message_parser,
+                         std::move(session_error_callback)),
+      track_alias_(track_alias),
+      add_callback_(std::move(add_callback)),
+      remove_callback_(std::move(remove_callback)),
+      session_(std::move(session)) {}
+
+absl::Status MoqtSubscribeResponseStream::OnRawControlMessage(
+    const MoqtRawControlMessage& message) {
+  return ControlMessageDispatcher::DispatchControlMessage(
+      *this, message_parser(), message, "subscribe response");
+}
+
+absl::Status MoqtSubscribeResponseStream::OnControlMessage(
+    const MoqtSubscribe& message) {
+  if (subscription_ != nullptr) {
+    return absl::InvalidArgumentError(
+        "SUBSCRIBE received on stream that already has a subscription");
+  }
+  QUIC_DLOG(INFO) << "Received a SUBSCRIBE for " << message.full_track_name;
+  if (session() == nullptr) {
+    return absl::OkStatus();
+  }
+  std::shared_ptr<MoqtTrackPublisher> track_publisher =
+      session()->GetTrackPublisher(message.full_track_name);
+  if (track_publisher == nullptr) {
+    QUIC_DLOG(INFO) << "SUBSCRIBE for " << message.full_track_name
+                    << " rejected by the application: does not exist";
+    return SendRequestError(message.request_id, RequestErrorCode::kDoesNotExist,
+                            std::nullopt, "not found", /*fin=*/true);
+  }
+  subscription_ = std::make_unique<SubscriptionPublisher>(
+      *framer(), track_publisher, this, message.request_id, track_alias_,
+      message.parameters, session_, false);
+  if (add_callback_ != nullptr) {
+    bool result = std::move(add_callback_)(subscription_.get());
+    add_callback_ = nullptr;
+    if (!result) {
+      return SendRequestError(message.request_id,
+                              RequestErrorCode::kDuplicateSubscription,
+                              std::nullopt, "duplicate subscription",
+                              /*fin=*/true);
+    }
+  }
+  // Don't add the publisher until we know it's successful.
+  track_publisher->AddObjectListener(subscription_.get());
+  return absl::OkStatus();
+}
+
+absl::Status MoqtSubscribeResponseStream::OnControlMessage(
+    const MoqtRequestUpdate& message) {
+  if (subscription_ == nullptr) {
+    QUICHE_BUG(INFO) << "Received REQUEST_UPDATE, no subscription state";
+    return SendRequestError(message.request_id,
+                            RequestErrorCode::kInternalError, std::nullopt,
+                            "no subscription", /*fin=*/true);
+  }
+  subscription_->Update(message.parameters);
+  return SendRequestOk(message.request_id, MessageParameters());
+}
+
+void MoqtSubscribeResponseStream::Detach() {
+  if (remove_callback_ != nullptr && subscription_ != nullptr) {
+    SubscriptionPublisher::RemoveCallback remove_callback =
+        std::move(remove_callback_);
+    remove_callback_ = nullptr;
+    std::move(remove_callback)(subscription_.get());
+  }
+  if (subscription_ != nullptr) {
+    subscription_->ResetAllStreams();
+    subscription_ = nullptr;
+  }
+}
+
+}  // namespace moqt
diff --git a/quiche/quic/moqt/moqt_subscribe_stream.h b/quiche/quic/moqt/moqt_subscribe_stream.h
new file mode 100644
index 0000000..3b6f5aa
--- /dev/null
+++ b/quiche/quic/moqt/moqt_subscribe_stream.h
@@ -0,0 +1,114 @@
+// Copyright (c) 2026 The Chromium Authors. All rights reserved.
+// Use of this source code is governed by a BSD-style license that can be
+// found in the LICENSE file.
+
+#ifndef QUICHE_QUIC_MOQT_MOQT_SUBSCRIBE_STREAM_H_
+#define QUICHE_QUIC_MOQT_MOQT_SUBSCRIBE_STREAM_H_
+
+#include <cstdint>
+#include <memory>
+
+#include "absl/base/nullability.h"
+#include "absl/status/status.h"
+#include "quiche/quic/core/quic_alarm_factory.h"
+#include "quiche/quic/core/quic_clock.h"
+#include "quiche/quic/moqt/moqt_bidi_stream.h"
+#include "quiche/quic/moqt/moqt_framer.h"
+#include "quiche/quic/moqt/moqt_key_value_pair.h"
+#include "quiche/quic/moqt/moqt_messages.h"
+#include "quiche/quic/moqt/moqt_names.h"
+#include "quiche/quic/moqt/moqt_parser.h"
+#include "quiche/quic/moqt/moqt_session_callbacks.h"
+#include "quiche/quic/moqt/moqt_subscription.h"
+#include "quiche/quic/moqt/moqt_track.h"
+#include "quiche/common/quiche_weak_ptr.h"
+
+namespace moqt {
+
+class MoqtSubscribeRequestStream : public MoqtBidiStreamBase {
+ public:
+  MoqtSubscribeRequestStream(
+      MoqtFramer* absl_nonnull framer,
+      const MoqtControlMessageParser& message_parser, uint64_t request_id,
+      SessionErrorCallback session_error_callback, const FullTrackName& name,
+      SubscribeVisitor* absl_nonnull visitor,
+      const MessageParameters& parameters,
+      SubscribeRemoteTrack::AddCallback add_callback,
+      SubscribeRemoteTrack::RemoveCallback remove_callback,
+      const quic::QuicClock* absl_nonnull clock,
+      quic::QuicAlarmFactory* absl_nonnull alarm_factory);
+  ~MoqtSubscribeRequestStream() { Detach(); }
+
+  // StreamBase overrides.
+  void OnStreamBound() override;
+  absl::Status OnRawControlMessage(
+      const MoqtRawControlMessage& message) override;
+  absl::Status OnControlMessage(const MoqtRequestOk& message) override;
+  absl::Status OnControlMessage(const MoqtRequestError& message) override;
+  absl::Status OnControlMessage(const MoqtSubscribeOk& message);
+  absl::Status OnControlMessage(const MoqtPublishDone& message);
+
+  SubscribeRemoteTrack* track() const { return track_.get(); }
+  void Detach() override;
+
+ private:
+  std::unique_ptr<SubscribeRemoteTrack> track_;
+  SubscribeRemoteTrack::AddCallback add_callback_;
+  SubscribeRemoteTrack::RemoveCallback remove_callback_;
+  const quic::QuicClock* clock_;
+  quic::QuicAlarmFactory* alarm_factory_;
+};
+
+class MoqtSubscribeResponseStream : public MoqtBidiStreamBase {
+ public:
+  MoqtSubscribeResponseStream(
+      MoqtFramer* absl_nonnull framer,
+      const MoqtControlMessageParser& message_parser, uint64_t track_alias,
+      SubscriptionPublisher::AddCallback add_callback,
+      SubscriptionPublisher::RemoveCallback remove_callback,
+      SessionErrorCallback session_error_callback,
+      quiche::QuicheWeakPtr<SessionToPublisherInterface> session);
+  ~MoqtSubscribeResponseStream() {
+    if (subscription_ != nullptr) {
+      subscription_->IgnoreResetAllStreams();
+    }
+    Detach();
+  }
+
+  // MoqtBidiStreamBase overrides.
+  void OnStreamBound() override { stream_parser()->set_allow_fin(true); }
+  absl::Status OnRawControlMessage(
+      const MoqtRawControlMessage& message) override;
+  absl::Status OnControlMessage(const MoqtRequestOk& message) override {
+    return absl::InvalidArgumentError(
+        "REQUEST_OK not allowed from Subscriber on SUBSCRIBE stream");
+  }
+  absl::Status OnControlMessage(const MoqtRequestError& message) override {
+    return absl::InvalidArgumentError(
+        "REQUEST_ERROR not allowed from Subscriber on SUBSCRIBE stream");
+  }
+
+  absl::Status OnControlMessage(const MoqtSubscribe& message);
+  absl::Status OnControlMessage(const MoqtRequestUpdate& message);
+  absl::Status OnControlMessage(const MoqtObjectAck& message) {
+    subscription_->ProcessObjectAck(message);
+    return absl::OkStatus();
+  }
+  void Detach() override;
+
+ private:
+  // Returns nullptr if MoqtSession is gone.
+  SessionToPublisherInterface* absl_nullable session() const {
+    return session_.GetIfAvailable();
+  }
+
+  uint64_t track_alias_;
+  std::unique_ptr<SubscriptionPublisher> subscription_;
+  SubscriptionPublisher::AddCallback add_callback_;
+  SubscriptionPublisher::RemoveCallback remove_callback_;
+  quiche::QuicheWeakPtr<SessionToPublisherInterface> session_;
+};
+
+}  // namespace moqt
+
+#endif  // QUICHE_QUIC_MOQT_MOQT_SUBSCRIBE_STREAM_H_
diff --git a/quiche/quic/moqt/moqt_subscribe_stream_test.cc b/quiche/quic/moqt/moqt_subscribe_stream_test.cc
new file mode 100644
index 0000000..300dafb
--- /dev/null
+++ b/quiche/quic/moqt/moqt_subscribe_stream_test.cc
@@ -0,0 +1,334 @@
+// Copyright (c) 2026 The Chromium Authors. All rights reserved.
+// Use of this source code is governed by a BSD-style license that can be
+// found in the LICENSE file.
+
+#include "quiche/quic/moqt/moqt_subscribe_stream.h"
+
+#include <cstdint>
+#include <memory>
+#include <optional>
+#include <utility>
+#include <variant>
+
+#include "absl/status/status.h"
+#include "absl/strings/string_view.h"
+#include "quiche/quic/core/quic_time.h"
+#include "quiche/quic/core/quic_types.h"
+#include "quiche/quic/moqt/moqt_error.h"
+#include "quiche/quic/moqt/moqt_fetch_task.h"
+#include "quiche/quic/moqt/moqt_framer.h"
+#include "quiche/quic/moqt/moqt_key_value_pair.h"
+#include "quiche/quic/moqt/moqt_messages.h"
+#include "quiche/quic/moqt/moqt_names.h"
+#include "quiche/quic/moqt/moqt_parser.h"
+#include "quiche/quic/moqt/moqt_session_callbacks.h"
+#include "quiche/quic/moqt/moqt_session_interface.h"
+#include "quiche/quic/moqt/moqt_subscription.h"
+#include "quiche/quic/moqt/moqt_trace_recorder.h"
+#include "quiche/quic/moqt/moqt_track.h"
+#include "quiche/quic/moqt/test_tools/mock_moqt_session.h"
+#include "quiche/quic/moqt/test_tools/moqt_framer_utils.h"
+#include "quiche/quic/moqt/test_tools/moqt_mock_visitor.h"
+#include "quiche/quic/test_tools/mock_clock.h"
+#include "quiche/quic/test_tools/quic_test_utils.h"
+#include "quiche/common/platform/api/quiche_test.h"
+#include "quiche/common/test_tools/quiche_test_utils.h"
+#include "quiche/web_transport/test_tools/mock_web_transport.h"
+
+namespace moqt::test {
+namespace {
+
+using ::testing::_;
+using ::testing::Return;
+using ::testing::StrictMock;
+
+class MoqtSubscribeRequestStreamTest : public quiche::test::QuicheTest {
+ public:
+  MoqtSubscribeRequestStreamTest()
+      : framer_(/*using_webtrans=*/true, quic::Perspective::IS_CLIENT),
+        message_parser_(kDefaultMoqtVersion, /*uses_web_transport=*/true,
+                        quic::Perspective::IS_CLIENT),
+        track_name_("foo", "bar") {
+    stream_ = std::make_unique<MoqtSubscribeRequestStream>(
+        &framer_, message_parser_, kRequestId, error_callback_.AsStdFunction(),
+        track_name_, &mock_subscribe_visitor_, parameters_,
+        mock_add_callback_.AsStdFunction(),
+        mock_remove_callback_.AsStdFunction(), &mock_clock_,
+        &mock_alarm_factory_);
+    EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(testing::Return(true));
+    EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(track_name_))
+        .Times(testing::AnyNumber());
+  }
+
+  MoqtFramer framer_;
+  MoqtControlMessageParser message_parser_;
+  uint64_t kRequestId = 1;
+  uint64_t kTrackAlias = 100;
+  FullTrackName track_name_;
+  MessageParameters parameters_;
+  testing::StrictMock<testing::MockFunction<void(MoqtError, absl::string_view)>>
+      error_callback_;
+  testing::MockFunction<bool(SubscribeRemoteTrack*)> mock_add_callback_;
+  testing::MockFunction<void(SubscribeRemoteTrack*)> mock_remove_callback_;
+  quic::MockClock mock_clock_;
+  quic::test::MockAlarmFactory mock_alarm_factory_;
+  StrictMock<MockSubscribeRemoteTrackVisitor> mock_subscribe_visitor_;
+  webtransport::test::MockStream mock_stream_;
+  std::unique_ptr<MoqtSubscribeRequestStream> stream_;
+};
+
+TEST_F(MoqtSubscribeRequestStreamTest, OnStreamBound) {
+  EXPECT_CALL(mock_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _))
+      .WillOnce(Return(absl::OkStatus()));
+  stream_->BindStream(&mock_stream_);
+}
+
+TEST_F(MoqtSubscribeRequestStreamTest, ReceiveSubscribeOk) {
+  EXPECT_CALL(mock_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _))
+      .WillOnce(Return(absl::OkStatus()));
+  stream_->BindStream(&mock_stream_);
+  EXPECT_CALL(mock_add_callback_, Call(stream_->track()))
+      .WillOnce(Return(true));
+  EXPECT_CALL(mock_subscribe_visitor_, OnReply(track_name_, _));
+  MoqtSubscribeOk subscribe_ok;
+  subscribe_ok.request_id = kRequestId;
+  subscribe_ok.track_alias = kTrackAlias;
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(subscribe_ok));
+  EXPECT_EQ(stream_->track()->track_alias(), kTrackAlias);
+  // Test cleanup.
+  EXPECT_CALL(mock_remove_callback_, Call);
+}
+
+TEST_F(MoqtSubscribeRequestStreamTest, ReceiveSubscribeOkAliasDuplicate) {
+  EXPECT_CALL(mock_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _))
+      .WillOnce(Return(absl::OkStatus()));
+  stream_->BindStream(&mock_stream_);
+  EXPECT_CALL(mock_add_callback_, Call(stream_->track()))
+      .WillOnce(Return(false));
+  EXPECT_CALL(error_callback_, Call(MoqtError::kDuplicateTrackAlias, _));
+  MoqtSubscribeOk subscribe_ok;
+  subscribe_ok.request_id = kRequestId;
+  subscribe_ok.track_alias = kTrackAlias;
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(subscribe_ok));
+  // Test cleanup.
+  EXPECT_CALL(mock_remove_callback_, Call);
+}
+
+TEST_F(MoqtSubscribeRequestStreamTest, RequestOkBeforeSubscribeOk) {
+  EXPECT_CALL(mock_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _))
+      .WillOnce(Return(absl::OkStatus()));
+  stream_->BindStream(&mock_stream_);
+  EXPECT_CALL(error_callback_, Call(MoqtError::kProtocolViolation, _));
+  MoqtRequestOk request_ok;
+  request_ok.request_id = kRequestId;
+  request_ok.parameters.expires = quic::QuicTimeDelta::FromSeconds(30);
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(request_ok));
+  // Test cleanup.
+  EXPECT_CALL(mock_remove_callback_, Call);
+}
+
+TEST_F(MoqtSubscribeRequestStreamTest, ReceiveRequestOk) {
+  // SUBSCRIBE handshake.
+  EXPECT_CALL(mock_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _))
+      .WillOnce(Return(absl::OkStatus()));
+  stream_->BindStream(&mock_stream_);
+  EXPECT_CALL(mock_add_callback_, Call(stream_->track()))
+      .WillOnce(Return(true));
+  EXPECT_CALL(mock_subscribe_visitor_, OnReply(track_name_, _));
+  MoqtSubscribeOk subscribe_ok;
+  subscribe_ok.request_id = kRequestId;
+  subscribe_ok.track_alias = kTrackAlias;
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(subscribe_ok));
+  // REQUEST_UPDATE.
+  EXPECT_CALL(mock_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kRequestUpdate), _))
+      .WillOnce(Return(absl::OkStatus()));
+  bool callback_called = false;
+  MoqtResponseCallback callback =
+      [&](std::variant<MessageParameters, MoqtRequestErrorInfo> res) {
+        callback_called = true;
+        ASSERT_TRUE(std::holds_alternative<MessageParameters>(res));
+        EXPECT_EQ(std::get<MessageParameters>(res).expires,
+                  quic::QuicTimeDelta::FromSeconds(30));
+      };
+  parameters_.subscriber_priority = 20;
+  QUICHE_EXPECT_OK(stream_->SendRequestUpdate(
+      kRequestId, kRequestId, parameters_, std::move(callback)));
+  // Params not yet updated.
+  EXPECT_EQ(stream_->track()->const_parameters().subscriber_priority,
+            std::nullopt);
+  MoqtRequestOk request_ok;
+  request_ok.request_id = kRequestId;
+  request_ok.parameters.expires = quic::QuicTimeDelta::FromSeconds(30);
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(request_ok));
+  EXPECT_EQ(stream_->track()->const_parameters().subscriber_priority, 20);
+  EXPECT_EQ(stream_->track()->const_parameters().expires,
+            quic::QuicTimeDelta::FromSeconds(30));
+  EXPECT_TRUE(callback_called);
+  // Test cleanup.
+  EXPECT_CALL(mock_remove_callback_, Call);
+}
+
+TEST_F(MoqtSubscribeRequestStreamTest, ReceiveRequestError) {
+  EXPECT_CALL(mock_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kSubscribe), _))
+      .WillOnce(Return(absl::OkStatus()));
+  stream_->BindStream(&mock_stream_);
+  EXPECT_CALL(mock_subscribe_visitor_, OnReply(track_name_, _))
+      .WillOnce(
+          [](const FullTrackName&,
+             const std::variant<SubscribeOkData, MoqtRequestErrorInfo>& reply) {
+            ASSERT_TRUE(std::holds_alternative<MoqtRequestErrorInfo>(reply));
+            EXPECT_EQ(std::get<MoqtRequestErrorInfo>(reply).error_code,
+                      RequestErrorCode::kUnauthorized);
+          });
+  EXPECT_CALL(mock_stream_, Writev(testing::IsEmpty(), _))
+      .WillOnce(Return(absl::OkStatus()));
+  EXPECT_CALL(mock_remove_callback_, Call);
+  MoqtRequestError request_error;
+  request_error.request_id = kRequestId;
+  request_error.error_code = RequestErrorCode::kUnauthorized;
+  request_error.reason_phrase = "unauthorized";
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(request_error));
+}
+
+TEST_F(MoqtSubscribeRequestStreamTest, ReceivePublishDone) {
+  MoqtPublishDone publish_done;
+  publish_done.request_id = kRequestId;
+  publish_done.stream_count = 5;
+  EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(track_name_));
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(publish_done));
+  EXPECT_CALL(mock_remove_callback_, Call);
+}
+
+class MoqtSubscribeResponseStreamTest : public quiche::test::QuicheTest {
+ public:
+  MoqtSubscribeResponseStreamTest()
+      : framer_(/*using_webtrans=*/true, quic::Perspective::IS_SERVER),
+        message_parser_(kDefaultMoqtVersion, /*uses_web_transport=*/true,
+                        quic::Perspective::IS_SERVER),
+        track_publisher_(std::make_shared<TestTrackPublisher>(kTrackName)) {
+    stream_ = std::make_unique<MoqtSubscribeResponseStream>(
+        &framer_, message_parser_, kTrackAlias,
+        mock_add_callback_.AsStdFunction(),
+        mock_remove_callback_.AsStdFunction(), error_callback_.AsStdFunction(),
+        visitor_.weak_ptr_factory_.Create());
+    stream_->BindStream(&mock_stream_);
+    EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(testing::Return(true));
+  }
+
+  MoqtFramer framer_;
+  MoqtControlMessageParser message_parser_;
+  uint64_t kRequestId = 1;
+  uint64_t kTrackAlias = 100;
+  FullTrackName kTrackName{"foo", "bar"};
+  std::shared_ptr<TestTrackPublisher> track_publisher_;
+  testing::MockFunction<void(MoqtError, absl::string_view)> error_callback_;
+  testing::MockFunction<bool(SubscriptionPublisher*)> mock_add_callback_;
+  testing::MockFunction<void(SubscriptionPublisher*)> mock_remove_callback_;
+  MockSessionToPublisherInterface visitor_;
+  webtransport::test::MockSession webtrans_;
+  webtransport::test::MockStream mock_stream_;
+  MoqtTraceRecorder trace_recorder_;
+  std::unique_ptr<MoqtSubscribeResponseStream> stream_;
+};
+
+TEST_F(MoqtSubscribeResponseStreamTest, ReceiveSubscribeSuccess) {
+  EXPECT_CALL(visitor_, session).WillRepeatedly(Return(&webtrans_));
+  EXPECT_CALL(visitor_, GetTrackPublisher(kTrackName))
+      .WillOnce(Return(track_publisher_));
+  EXPECT_CALL(mock_add_callback_, Call(testing::NotNull()))
+      .WillOnce(Return(true));
+  MoqtSubscribe subscribe;
+  subscribe.request_id = kRequestId;
+  subscribe.full_track_name = kTrackName;
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(subscribe));
+  // Test cleanup.
+  EXPECT_CALL(mock_remove_callback_, Call);
+}
+
+TEST_F(MoqtSubscribeResponseStreamTest, ReceiveSubscribeDoesNotExist) {
+  EXPECT_CALL(visitor_, session).WillRepeatedly(Return(&webtrans_));
+  EXPECT_CALL(visitor_, GetTrackPublisher(kTrackName))
+      .WillOnce(Return(nullptr));
+  EXPECT_CALL(mock_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _))
+      .WillOnce(Return(absl::OkStatus()));
+
+  MoqtSubscribe subscribe;
+  subscribe.request_id = kRequestId;
+  subscribe.full_track_name = kTrackName;
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(subscribe));
+}
+
+TEST_F(MoqtSubscribeResponseStreamTest, ReceiveSubscribeDuplicate) {
+  EXPECT_CALL(visitor_, session).WillRepeatedly(Return(&webtrans_));
+  EXPECT_CALL(visitor_, GetTrackPublisher(kTrackName))
+      .WillOnce(Return(track_publisher_));
+  EXPECT_CALL(mock_add_callback_, Call(testing::NotNull()))
+      .WillOnce(Return(false));
+  EXPECT_CALL(mock_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _))
+      .WillOnce(Return(absl::OkStatus()));
+  MoqtSubscribe subscribe;
+  subscribe.request_id = kRequestId;
+  subscribe.full_track_name = kTrackName;
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(subscribe));
+}
+
+TEST_F(MoqtSubscribeResponseStreamTest, ReceiveRequestUpdate) {
+  EXPECT_CALL(visitor_, session).WillRepeatedly(Return(&webtrans_));
+  EXPECT_CALL(visitor_, GetTrackPublisher(kTrackName))
+      .WillOnce(Return(track_publisher_));
+  EXPECT_CALL(mock_add_callback_, Call(testing::NotNull()))
+      .WillOnce(Return(true));
+  MoqtSubscribe subscribe;
+  subscribe.request_id = kRequestId;
+  subscribe.full_track_name = kTrackName;
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(subscribe));
+  EXPECT_CALL(mock_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kRequestOk), _))
+      .WillOnce(Return(absl::OkStatus()));
+  MoqtRequestUpdate update;
+  update.request_id = kRequestId;
+  update.parameters.subscriber_priority = 10;
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(update));
+  // Test cleanup.
+  EXPECT_CALL(mock_remove_callback_, Call(_));
+}
+
+TEST_F(MoqtSubscribeResponseStreamTest, ReceiveInvalidControlMessages) {
+  MoqtRequestOk request_ok;
+  EXPECT_FALSE(stream_->OnControlMessage(request_ok).ok());
+  MoqtRequestError request_error;
+  EXPECT_FALSE(stream_->OnControlMessage(request_error).ok());
+}
+
+TEST_F(MoqtSubscribeResponseStreamTest, ReceiveObjectAck) {
+  EXPECT_CALL(visitor_, session).WillRepeatedly(Return(&webtrans_));
+  EXPECT_CALL(visitor_, GetTrackPublisher(kTrackName))
+      .WillOnce(Return(track_publisher_));
+  EXPECT_CALL(mock_add_callback_, Call(testing::NotNull()))
+      .WillOnce(Return(true));
+  MoqtSubscribe subscribe;
+  subscribe.request_id = kRequestId;
+  subscribe.full_track_name = kTrackName;
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(subscribe));
+  MoqtObjectAck ack;
+  ack.group_id = 1;
+  ack.object_id = 2;
+  ack.delta_from_deadline = quic::QuicTimeDelta::FromMilliseconds(100);
+  EXPECT_CALL(visitor_, trace_recorder())
+      .WillOnce(testing::ReturnRef(trace_recorder_));
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(ack));
+  // Test cleanup.
+  EXPECT_CALL(mock_remove_callback_, Call(_));
+}
+
+}  // namespace
+}  // namespace moqt::test
diff --git a/quiche/quic/moqt/moqt_subscription.cc b/quiche/quic/moqt/moqt_subscription.cc
index 9bf4513..fa83093 100644
--- a/quiche/quic/moqt/moqt_subscription.cc
+++ b/quiche/quic/moqt/moqt_subscription.cc
@@ -30,6 +30,7 @@
 #include "quiche/quic/moqt/moqt_types.h"
 #include "quiche/quic/moqt/moqt_uni_stream.h"
 #include "quiche/common/quiche_buffer_allocator.h"
+#include "quiche/common/quiche_weak_ptr.h"
 #include "quiche/web_transport/web_transport.h"
 
 namespace moqt {
@@ -38,10 +39,7 @@
     MoqtFramer framer, std::shared_ptr<MoqtTrackPublisher> track_publisher,
     MoqtBidiStreamBase* absl_nonnull bidi_stream, uint64_t request_id,
     uint64_t track_alias, const MessageParameters& parameters,
-    SessionToPublisherInterface* absl_nonnull visitor,
-    MoqtPublishingMonitorInterface* monitoring_interface,
-    const quic::QuicClock* absl_nonnull clock,
-    MoqtTraceRecorder& trace_recorder, bool is_publish)
+    quiche::QuicheWeakPtr<SessionToPublisherInterface> visitor, bool is_publish)
     : track_publisher_(track_publisher),
       bidi_stream_(bidi_stream),
       visitor_(visitor),
@@ -49,11 +47,14 @@
       established_(is_publish),
       track_alias_(track_alias),
       framer_(framer),
-      trace_recorder_(trace_recorder),
       parameters_(parameters),
-      monitoring_interface_(monitoring_interface),
-      clock_(clock),
       weak_ptr_factory_(this) {
+  SessionToPublisherInterface* session_info = visitor_.GetIfAvailable();
+  if (session_info == nullptr) {
+    return;
+  }
+  monitoring_interface_ = session_info->ReleaseMonitoringInterface(
+      track_publisher_->GetTrackName());
   if (monitoring_interface_ != nullptr) {
     monitoring_interface_->OnObjectAckSupportKnown(parameters.oack_window_size);
   }
@@ -66,13 +67,6 @@
   if (track_publisher_ != nullptr) {
     track_publisher_->RemoveObjectListener(this);
   }
-  // Reset all streams.
-  for (const webtransport::StreamId stream_id : stream_map_.GetAllStreams()) {
-    webtransport::Stream* stream = GetStreamById(stream_id);
-    if (stream != nullptr) {
-      stream->ResetWithUserCode(kResetCodeCancelled);
-    }
-  }
 }
 
 void SubscriptionPublisher::Update(const MessageParameters& parameters) {
@@ -103,13 +97,28 @@
         pending_streams_.rbegin()->second.publisher_priority.value_or(
             track_publisher_->extensions().default_publisher_priority());
     MoqtTrackPriority old_track_priority = {old_priority, publisher_priority};
-    visitor_->UpdateTrackPriority(
+    if (visitor() == nullptr) {
+      return;
+    }
+    visitor()->UpdateTrackPriority(
         request_id_, old_track_priority,
         MoqtTrackPriority{new_priority, publisher_priority});
     // Don't bother to update all the pending stream send orders.
   }
 }
 
+void SubscriptionPublisher::ResetAllStreams() {
+  if (ignore_reset_all_streams_) {
+    return;
+  }
+  for (const webtransport::StreamId stream_id : stream_map_.GetAllStreams()) {
+    webtransport::Stream* stream = GetStreamById(stream_id);
+    if (stream != nullptr) {
+      stream->ResetWithUserCode(kResetCodeCancelled);
+    }
+  }
+}
+
 void SubscriptionPublisher::OnSubscribeAccepted() {
   if (established_) {
     return;  // It's a PUBLISH.
@@ -142,11 +151,9 @@
 }
 
 void SubscriptionPublisher::OnSubscribeRejected(MoqtRequestErrorInfo info) {
-  bidi_stream_->CheckStatus(bidi_stream_->SendRequestError(
-      request_id_, info, /*fin=*/!bidi_stream_->is_control_stream()));
-  if (bidi_stream_->is_control_stream()) {
-    visitor_->PublishIsDone(request_id_);
-  }
+  bidi_stream_->CheckStatus(bidi_stream_->SendRequestError(request_id_, info,
+                                                           /*fin=*/true));
+  // Sending FIN will delete the class.
 }
 
 void SubscriptionPublisher::OnNewObjectAvailable(
@@ -174,7 +181,11 @@
 
   // TODO(vasilvv): This currently sends UINT64_MAX for datagram subgroups.
   // Maybe do something more satisfactory?
-  trace_recorder_.RecordNewObjectAvaliable(
+  SessionToPublisherInterface* session_info = visitor();
+  if (session_info == nullptr) {
+    return;  // Session is gone.
+  }
+  session_info->trace_recorder().RecordNewObjectAvaliable(
       track_alias_, *track_publisher_, location, subgroup.value_or(UINT64_MAX),
       publisher_priority);
 
@@ -187,7 +198,7 @@
     }
     stream_id = stream_map_.GetStreamFor(index);
   }
-  if (visitor_->alternate_delivery_timeout() &&
+  if (session_info->alternate_delivery_timeout() &&
       !delivery_timeout().IsInfinite() && largest_sent_.has_value() &&
       location.group >= largest_sent_->group) {
     // Start the delivery timeout timer on all previous groups.
@@ -201,7 +212,7 @@
         }
         OutgoingSubgroupStream* stream =
             absl::down_cast<OutgoingSubgroupStream*>(raw_stream->visitor());
-        stream->CreateAndSetAlarm(clock_->ApproximateNow() +
+        stream->CreateAndSetAlarm(session_info->clock()->ApproximateNow() +
                                   delivery_timeout());
       }
     }
@@ -229,7 +240,7 @@
   if (raw_stream == nullptr) {
     StreamRank rank = StreamRankFor(parameters);
     if (pending_streams_.empty() || rank > pending_streams_.rbegin()->first) {
-      visitor_->UpdateTrackPriority(
+      session_info->UpdateTrackPriority(
           request_id_,
           /*old_priority=*/pending_streams_.empty()
               ? std::optional<MoqtTrackPriority>()
@@ -267,6 +278,7 @@
   OutgoingSubgroupStream* stream =
       absl::down_cast<OutgoingSubgroupStream*>(raw_stream->visitor());
   stream->Fin(location);
+  // Sending FIN will delete the class.
 }
 
 void SubscriptionPublisher::OnSubgroupAbandoned(
@@ -345,34 +357,39 @@
   quiche::QuicheBuffer datagram = framer_.SerializeObjectDatagram(
       header, object->payload[0].AsStringView(),
       default_publisher_priority_.value_or(kDefaultPublisherPriority));
-  if (visitor_->session() == nullptr) {
+  if (visitor() == nullptr) {
     return;
   }
-  visitor_->session()->SendOrQueueDatagram(datagram.AsStringView());
+  visitor()->session()->SendOrQueueDatagram(datagram.AsStringView());
   OnObjectSent(object->metadata.location);
 }
 
 void SubscriptionPublisher::ProcessObjectAck(const MoqtObjectAck& message) {
-  trace_recorder_.RecordObjectAck(track_alias_,
-                                  Location(message.group_id, message.object_id),
-                                  message.delta_from_deadline);
-
-  if (monitoring_interface_ == nullptr) {
+  SessionToPublisherInterface* session_info = visitor();
+  if (session_info == nullptr) {
     return;
   }
-  monitoring_interface_->OnObjectAckReceived(
-      Location(message.group_id, message.object_id),
+  session_info->trace_recorder().RecordObjectAck(
+      track_alias_, Location(message.group_id, message.object_id),
       message.delta_from_deadline);
+  if (monitoring_interface_ != nullptr) {
+    monitoring_interface_->OnObjectAckReceived(
+        Location(message.group_id, message.object_id),
+        message.delta_from_deadline);
+  }
 }
 
 webtransport::Stream* absl_nullable SubscriptionPublisher::OpenDataStream(
     const NewDataStreamParameters& parameters) {
-  if (visitor_->session() == nullptr ||
-      !visitor_->session()->CanOpenNextOutgoingUnidirectionalStream()) {
+  SessionToPublisherInterface* session_info = visitor();
+  if (session_info == nullptr) {
+    return nullptr;
+  }
+  if (!session_info->session()->CanOpenNextOutgoingUnidirectionalStream()) {
     return nullptr;
   }
   webtransport::Stream* new_stream =
-      visitor_->session()->OpenOutgoingUnidirectionalStream();
+      session_info->session()->OpenOutgoingUnidirectionalStream();
   if (new_stream == nullptr) {
     return nullptr;
   }
@@ -380,7 +397,8 @@
   new_stream->SetVisitor(std::make_unique<OutgoingSubgroupStream>(
       framer_, new_stream, parameters.index, parameters.first_object,
       weak_ptr_factory_.Create(), track_publisher_,
-      StreamPriorityFor(parameters), track_alias_, &trace_recorder_));
+      StreamPriorityFor(parameters), track_alias_,
+      &session_info->trace_recorder()));
   ++streams_opened_;
   new_stream->visitor()->OnCanWrite();
   return new_stream;
@@ -398,17 +416,9 @@
   // FIN naturally, where possible.
   QUICHE_DLOG(INFO) << "Sending PUBLISH_DONE message for "
                     << track_publisher_->GetTrackName();
-  // TODO(martinduke): For SUBSCRIBE, no FIN because it's the control stream.
   bidi_stream_->SendOrBufferMessageOrFatal(
-      framer_.SerializePublishDone(publish_done),
-      /*fin=*/!bidi_stream_->is_control_stream());
-  if (bidi_stream_->is_control_stream()) {
-    visitor_->PublishIsDone(request_id_);
-  } else {
-    // Only detach immediately for PUBLISH flow.
-    track_publisher_->RemoveObjectListener(this);
-    track_publisher_ = nullptr;
-  }
+      framer_.SerializePublishDone(publish_done), /*fin=*/true);
+  // sending FIN will delete the class.
 }
 
 void SubscriptionPublisher::OnDataStreamDestroyed(
@@ -417,8 +427,12 @@
 }
 
 void SubscriptionPublisher::OnCanCreateNewUniStream() {
-  while (visitor_->session() != nullptr &&
-         visitor_->session()->CanOpenNextOutgoingUnidirectionalStream()) {
+  SessionToPublisherInterface* session_info = visitor();
+  if (session_info == nullptr) {
+    return;
+  }
+  while (visitor_.IsValid() &&
+         session_info->session()->CanOpenNextOutgoingUnidirectionalStream()) {
     auto it = pending_streams_.rbegin();
     while (it != pending_streams_.rend() &&
            (it->second.index.group < first_active_group_ ||
@@ -434,7 +448,7 @@
     }
     pending_streams_.erase(--(it.base()));
     if (!pending_streams_.empty()) {
-      visitor_->UpdateTrackPriority(
+      session_info->UpdateTrackPriority(
           request_id_, std::nullopt,
           MoqtTrackPriority{
               subscriber_priority(),
diff --git a/quiche/quic/moqt/moqt_subscription.h b/quiche/quic/moqt/moqt_subscription.h
index d32ed4c..d4af902 100644
--- a/quiche/quic/moqt/moqt_subscription.h
+++ b/quiche/quic/moqt/moqt_subscription.h
@@ -21,6 +21,7 @@
 #include "quiche/quic/moqt/moqt_framer.h"
 #include "quiche/quic/moqt/moqt_key_value_pair.h"
 #include "quiche/quic/moqt/moqt_messages.h"
+#include "quiche/quic/moqt/moqt_names.h"
 #include "quiche/quic/moqt/moqt_priority.h"
 #include "quiche/quic/moqt/moqt_publisher.h"
 #include "quiche/quic/moqt/moqt_stream_map.h"
@@ -76,17 +77,18 @@
   virtual ~SessionToPublisherInterface() = default;
   virtual bool alternate_delivery_timeout() const = 0;
   // If |old_priority| is nullopt, the subscription does not have any pending
-  // streams. If it has a value, |old_priority| is the old value to be replaced
-  // by |new_priority|.
+  // streams. If it has a value, |old_priority| is the old value to be
+  // replaced by |new_priority|.
   virtual void UpdateTrackPriority(
       uint64_t request_id, std::optional<MoqtTrackPriority> old_priority,
       MoqtTrackPriority new_priority) = 0;
   virtual quic::QuicAlarmFactory* alarm_factory() = 0;
-  // Destroy any state associated with the subscription. It is OK destroy
-  // SubscriptionPublisher in this method.
-  // TODO(martinduke): Delete once SUBSCRIBE is on the bidi stream.
-  virtual void PublishIsDone(uint64_t request_id) = 0;
-  // Returns nullptr if MoqtSession is closing.
+  virtual std::shared_ptr<MoqtTrackPublisher> GetTrackPublisher(
+      const FullTrackName& name) = 0;
+  virtual MoqtPublishingMonitorInterface* ReleaseMonitoringInterface(
+      const FullTrackName& name) = 0;
+  virtual const quic::QuicClock* clock() = 0;
+  virtual MoqtTraceRecorder& trace_recorder() = 0;
   virtual webtransport::Session* session() = 0;
 };
 
@@ -95,20 +97,19 @@
 class SubscriptionPublisher : public MoqtObjectListener,
                               public SubscriptionPublisherInterface {
  public:
-  // The provider of this callback will delete whatever state it is tracking for
-  // the subscription. This will be used by both PUBLISH and SUBSCRIBE streams.
-  // For control stream SUBSCRIBE, visitor->PublishIsDone() does this instead.
+  // The provider of this callback will add/delete whatever state it is tracking
+  // for the subscription. This will be used by both PUBLISH and SUBSCRIBE
+  // streams. AddCallback returns |false| if the add fails because the key
+  // already exists.
+  using AddCallback = quiche::SingleUseCallback<bool(SubscriptionPublisher*)>;
   using RemoveCallback =
       quiche::SingleUseCallback<void(SubscriptionPublisher*)>;
-  SubscriptionPublisher(MoqtFramer framer,
-                        std::shared_ptr<MoqtTrackPublisher> track_publisher,
-                        MoqtBidiStreamBase* absl_nonnull bidi_stream,
-                        uint64_t request_id, uint64_t track_alias,
-                        const MessageParameters& parameters,
-                        SessionToPublisherInterface* absl_nonnull visitor,
-                        MoqtPublishingMonitorInterface* monitoring_interface,
-                        const quic::QuicClock* absl_nonnull clock,
-                        MoqtTraceRecorder& trace_recorder, bool is_publish);
+  SubscriptionPublisher(
+      MoqtFramer framer, std::shared_ptr<MoqtTrackPublisher> track_publisher,
+      MoqtBidiStreamBase* absl_nonnull bidi_stream, uint64_t request_id,
+      uint64_t track_alias, const MessageParameters& parameters,
+      quiche::QuicheWeakPtr<SessionToPublisherInterface> visitor,
+      bool is_publish);
   ~SubscriptionPublisher();
 
   SubscriptionPublisher(const SubscriptionPublisher&) = delete;
@@ -143,28 +144,39 @@
              parameters_.subscription_filter->InWindow(location)));
   };
   bool alternate_delivery_timeout() override {
-    return visitor_->alternate_delivery_timeout();
+    if (visitor() == nullptr) {
+      return false;
+    }
+    return visitor()->alternate_delivery_timeout();
   }
-  const quic::QuicClock* clock() override { return clock_; }
+  const quic::QuicClock* clock() override {
+    if (visitor() == nullptr) {
+      return nullptr;
+    }
+    return visitor()->clock();
+  }
   quic::QuicTimeDelta delivery_timeout() override {
     return std::min(
         parameters_.delivery_timeout.value_or(kDefaultDeliveryTimeout),
         publisher_delivery_timeout_.value_or(kDefaultDeliveryTimeout));
   }
   quic::QuicAlarmFactory* alarm_factory() override {
-    return visitor_->alarm_factory();
+    if (visitor() == nullptr) {
+      return nullptr;
+    }
+    return visitor()->alarm_factory();
   }
   void OnObjectSent(Location sequence) override;
   void OnStreamTimeout(DataStreamIndex index) override {
     reset_subgroups_.insert(index);
-    if (visitor_->alternate_delivery_timeout()) {
+    if (visitor()->alternate_delivery_timeout()) {
       first_active_group_ = std::max(first_active_group_, index.group + 1);
     }
   }
   // OnSubgroupAbandoned() is declared above with MoqtObjectListener.
   void OnDataStreamDestroyed(DataStreamIndex) override;
 
-  // Called by MoqtSession when this subscription can open a new stream.
+  // MoqtObjectPublisher implementation.
   void OnCanCreateNewUniStream();
 
   // Called when the parameters_ needs an update.
@@ -174,6 +186,17 @@
 
   bool established() const { return established_; }
 
+  quiche::QuicheWeakPtr<SubscriptionPublisherInterface> GetWeakPtr() {
+    return weak_ptr_factory_.Create();
+  }
+
+  // Resets all unidirectional streams.
+  void ResetAllStreams();
+  // Called when the bidi stream is being destroyed. If the result of a session
+  // teardown, it is not safe to access the streams via GetStreamById, and the
+  // uni streams will be destroyed anyway.
+  void IgnoreResetAllStreams() { ignore_reset_all_streams_ = true; }
+
  private:
   friend class test::SubscriptionPublisherPeer;
 
@@ -224,20 +247,25 @@
   }
 
   webtransport::Stream* GetStreamById(webtransport::StreamId stream_id) {
-    return visitor_->session() == nullptr
-               ? nullptr
-               : visitor_->session()->GetStreamById(stream_id);
+    if (visitor() != nullptr) {
+      return visitor()->session()->GetStreamById(stream_id);
+    }
+    return nullptr;
+  }
+
+  // If nullptr, MoqtSession is gone.
+  SessionToPublisherInterface* absl_nullable visitor() const {
+    return visitor_.GetIfAvailable();
   }
 
   std::shared_ptr<MoqtTrackPublisher> track_publisher_;
   MoqtBidiStreamBase* absl_nonnull bidi_stream_;
-  SessionToPublisherInterface* absl_nonnull visitor_;
+  quiche::QuicheWeakPtr<SessionToPublisherInterface> visitor_;
   uint64_t request_id_;
   // Subscription is in the ESTABLISHED state.
   bool established_;
   const uint64_t track_alias_;
   MoqtFramer framer_;
-  MoqtTraceRecorder& trace_recorder_;
   // These are (mostly) the parameters from the SUBSCRIBE message. However,
   // group_order and largest_object may be updated by SUBSCRIBE_OK because
   // have no effect in a future REQUEST_UPDATE message.
@@ -257,11 +285,11 @@
   // Largest sequence number ever sent via this subscription.
   std::optional<Location> largest_sent_;
   SendStreamMap stream_map_;
+  bool ignore_reset_all_streams_ = false;
   // Store the StreamRank of queued outgoing data streams. High StreamRank is
   // highest priority, so use rbegin() to get the highest priority pending
   // stream.
   absl::btree_multimap<StreamRank, NewDataStreamParameters> pending_streams_;
-  const quic::QuicClock* absl_nonnull clock_;
   // Must be last.
   quiche::QuicheWeakPtrFactory<SubscriptionPublisherInterface>
       weak_ptr_factory_;
diff --git a/quiche/quic/moqt/moqt_subscription_test.cc b/quiche/quic/moqt/moqt_subscription_test.cc
index 154725b..b35818d 100644
--- a/quiche/quic/moqt/moqt_subscription_test.cc
+++ b/quiche/quic/moqt/moqt_subscription_test.cc
@@ -18,8 +18,8 @@
 #include "absl/strings/match.h"
 #include "absl/strings/string_view.h"
 #include "absl/types/span.h"
-#include "quiche/quic/core/quic_alarm_factory.h"
 #include "quiche/quic/core/quic_time.h"
+#include "quiche/quic/core/quic_types.h"
 #include "quiche/quic/moqt/moqt_bidi_stream.h"
 #include "quiche/quic/moqt/moqt_error.h"
 #include "quiche/quic/moqt/moqt_framer.h"
@@ -33,8 +33,7 @@
 #include "quiche/quic/moqt/moqt_session_interface.h"
 #include "quiche/quic/moqt/moqt_trace_recorder.h"
 #include "quiche/quic/moqt/moqt_types.h"
-#include "quiche/quic/moqt/moqt_uni_stream.h"
-#include "quiche/quic/moqt/test_tools/moqt_framer_utils.h"
+#include "quiche/quic/moqt/test_tools/mock_moqt_session.h"
 #include "quiche/quic/moqt/test_tools/moqt_mock_visitor.h"
 #include "quiche/quic/moqt/test_tools/moqt_session_peer.h"
 #include "quiche/quic/platform/api/quic_test.h"
@@ -71,27 +70,13 @@
 using ::webtransport::DatagramStatus;
 using ::webtransport::DatagramStatusCode;
 
-class MockSessionToPublisherInterface : public SessionToPublisherInterface {
- public:
-  ~MockSessionToPublisherInterface() override = default;
-  MOCK_METHOD(bool, alternate_delivery_timeout, (), (const, override));
-  MOCK_METHOD(void, UpdateTrackPriority,
-              (uint64_t, std::optional<MoqtTrackPriority>, MoqtTrackPriority),
-              (override));
-  MOCK_METHOD(quic::QuicAlarmFactory*, alarm_factory, (), (override));
-  MOCK_METHOD(void, PublishIsDone, (uint64_t), (override));
-  MOCK_METHOD(webtransport::Session*, session, (), (override));
-};
-
 class TestMoqtBidiStream : public MoqtBidiStreamBase {
  public:
   TestMoqtBidiStream(MoqtFramer* absl_nonnull framer,
                      const MoqtControlMessageParser& message_parser,
                      SessionErrorCallback session_error_callback)
       : MoqtBidiStreamBase(framer, message_parser,
-                           std::move(session_error_callback)) {
-    set_control_stream();  // TODO(martinduke): Delete
-  }
+                           std::move(session_error_callback)) {}
   ~TestMoqtBidiStream() override = default;
   void OnStreamBound() override {};
   absl::Status OnRawControlMessage(
@@ -134,21 +119,25 @@
     EXPECT_CALL(monitoring_interface_, OnObjectAckSupportKnown)
         .Times(AtLeast(0));
     EXPECT_CALL(visitor_, session).WillRepeatedly(Return(&webtrans_));
+    ON_CALL(visitor_, ReleaseMonitoringInterface)
+        .WillByDefault(Return(&monitoring_interface_));
     publisher_ = std::make_unique<SubscriptionPublisher>(
         framer_, track_publisher_, &bidi_stream_, kRequestId, kTrackAlias,
-        parameters_, &visitor_, &monitoring_interface_, &mock_clock_,
-        trace_recorder_, /*is_publish=*/false);
+        parameters_, visitor_.weak_ptr_factory_.Create(),
+        /*is_publish=*/false);
     ON_CALL(visitor_, alternate_delivery_timeout).WillByDefault(Return(false));
     ON_CALL(webtrans_, GetStreamById(kStreamId))
         .WillByDefault(Return(&mock_uni_stream_));
     ON_CALL(visitor_, alarm_factory).WillByDefault(Return(&alarm_factory_));
+    ON_CALL(visitor_, clock).WillByDefault(Return(&mock_clock_));
+    ON_CALL(visitor_, trace_recorder).WillByDefault(ReturnRef(trace_recorder_));
   }
 
   ~SubscriptionPublisherTest() override {
+    if (track_publisher_ == nullptr) {
+      return;
+    }
     EXPECT_CALL(*track_publisher_, RemoveObjectListener(publisher_.get()));
-    size_t num_open_streams =
-        SubscriptionPublisherPeer::num_open_streams(publisher_.get());
-    EXPECT_CALL(mock_uni_stream_, ResetWithUserCode).Times(num_open_streams);
   }
 
   MoqtPriority subscriber_priority() const {
@@ -308,7 +297,6 @@
   EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _))
       .WillOnce(Return(absl::OkStatus()));
-  EXPECT_CALL(visitor_, PublishIsDone(1));
   publisher_->OnSubscribeRejected(MoqtRequestErrorInfo(
       RequestErrorCode::kDoesNotExist, std::nullopt, "reason"));
 }
@@ -471,9 +459,15 @@
   };
   EXPECT_CALL(mock_bidi_stream_, CanWrite()).WillRepeatedly(Return(true));
   EXPECT_CALL(mock_bidi_stream_,
-              Writev(SerializedControlMessage(expected_publish_done), _));
-  EXPECT_CALL(visitor_, PublishIsDone(kRequestId));
+              Writev(SerializedControlMessage(expected_publish_done), _))
+      .WillOnce([&](absl::Span<quiche::QuicheMemSlice> data,
+                    const webtransport::StreamWriteOptions& options) {
+        EXPECT_TRUE(options.send_fin());
+        return absl::OkStatus();
+      });
+  EXPECT_CALL(*track_publisher_, RemoveObjectListener);
   publisher_->OnGroupAbandoned(5);
+  track_publisher_ = nullptr;
 }
 
 TEST_F(SubscriptionPublisherTest, OnCanCreateNewUniStreamPendingCleanup) {
@@ -502,8 +496,9 @@
   EXPECT_CALL(mock_bidi_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kPublishDone), _))
       .WillOnce(Return(absl::OkStatus()));
-  EXPECT_CALL(visitor_, PublishIsDone(1));
+  EXPECT_CALL(*track_publisher_, RemoveObjectListener);
   publisher_->OnTrackPublisherGone();
+  track_publisher_ = nullptr;
 }
 
 TEST_F(SubscriptionPublisherTest, ProcessObjectAck) {
@@ -701,6 +696,34 @@
   publisher_->OnNewObjectAvailable(Location(5, 1), 0, 128);
 }
 
+TEST_F(SubscriptionPublisherTest, OnNewFinAvailable) {
+  CreateStream(
+      Location(1, 0), 0, 127,
+      {0x51, static_cast<uint8_t>(kTrackAlias), 0x01, 0x7f, 0x00, 0x0a});
+  EXPECT_CALL(mock_uni_stream_, Writev(testing::IsEmpty(), _))
+      .WillOnce([](absl::Span<quiche::QuicheMemSlice> data,
+                   const webtransport::StreamWriteOptions& options) {
+        EXPECT_TRUE(options.send_fin());
+        return absl::OkStatus();
+      });
+  publisher_->OnNewFinAvailable(Location(1, 0), 0);
+}
+
+TEST_F(SubscriptionPublisherTest, OnSubgroupAbandoned) {
+  CreateStream(
+      Location(1, 0), 0, 127,
+      {0x51, static_cast<uint8_t>(kTrackAlias), 0x01, 0x7f, 0x00, 0x0a});
+  EXPECT_CALL(mock_uni_stream_, ResetWithUserCode(1234));
+  publisher_->OnSubgroupAbandoned(1, 0, 1234);
+}
+
+TEST_F(SubscriptionPublisherTest, OnSubgroupAbandonedOutsideWindow) {
+  parameters_.subscription_filter = SubscriptionFilter(Location(20, 0));
+  publisher_->Update(parameters_);
+  EXPECT_CALL(mock_uni_stream_, ResetWithUserCode).Times(0);
+  publisher_->OnSubgroupAbandoned(1, 0, 1234);
+}
+
 }  // namespace
 
 }  // namespace moqt::test
diff --git a/quiche/quic/moqt/moqt_track.cc b/quiche/quic/moqt/moqt_track.cc
index 36e06cf..23060f0 100644
--- a/quiche/quic/moqt/moqt_track.cc
+++ b/quiche/quic/moqt/moqt_track.cc
@@ -43,14 +43,9 @@
   if (publish_done_alarm_ != nullptr) {
     publish_done_alarm_->PermanentCancel();
   }
-  if (remove_callback_ != nullptr) {
-    RemoveCallback callback = std::move(remove_callback_);
-    remove_callback_ = nullptr;
-    std::move(callback)(this);
-  }
   if (visitor_ != nullptr) {
+    // The application expects OnPublishDone even if the Session has gone away.
     visitor_->OnPublishDone(full_track_name());
-    visitor_ = nullptr;
   }
 }
 
@@ -62,10 +57,6 @@
   publisher_delivery_timeout_ = data.extensions.delivery_timeout();
   // TODO(martinduke): Is there anything to do with EXPIRES?
   default_publisher_priority_ = data.extensions.default_publisher_priority();
-  if (!parameters().group_order.has_value()) {
-    // Use publisher default because the subscriber didn't care.
-    parameters().group_order = data.extensions.default_publisher_group_order();
-  }
   dynamic_groups_ = data.extensions.dynamic_groups();
   visitor_->OnReply(full_track_name(), data);
   OnObjectOrOk();
@@ -83,7 +74,7 @@
   ++streams_closed_;
   --currently_open_streams_;
   QUICHE_DCHECK_GE(currently_open_streams_, -1);
-  if (index.has_value()) {
+  if (index.has_value() && visitor_ != nullptr) {
     // If index is nullopt, there was not an object received on the stream.
     if (fin_received) {
       visitor_->OnStreamFin(full_track_name(), *index);
@@ -92,7 +83,8 @@
     }
   }
   if (all_streams_closed()) {
-    Destroy();
+    request_stream()->Fin();
+    // Fin will destroy the class.
     return;
   }
   if (publish_done_alarm_ == nullptr) {
@@ -107,11 +99,12 @@
   total_streams_ = stream_count;
   clock_ = clock;
   if (all_streams_closed()) {
-    Destroy();
+    request_stream()->Fin();
+    // Fin will destroy the class.
     return;
   }
   publish_done_alarm_ = std::unique_ptr<quic::QuicAlarm>(
-      alarm_factory->CreateAlarm(new PublishDoneDelegate(this)));
+      alarm_factory->CreateAlarm(new PublishDoneDelegate(weak_ptr())));
   MaybeSetPublishDoneAlarm();
 }
 
@@ -177,6 +170,14 @@
   }
 }
 
+void SubscribeRemoteTrack::SendObjectAck(
+    uint64_t group_id, uint64_t object_id,
+    quic::QuicTimeDelta delta_from_deadline) {
+  request_stream()->SendOrBufferMessageOrFatal(
+      request_stream()->framer()->SerializeObjectAck(
+          {group_id, object_id, delta_from_deadline}));
+}
+
 UpstreamFetch::~UpstreamFetch() {
   UpstreamFetchTask* task = task_.GetIfAvailable();
   if (task != nullptr) {
@@ -184,10 +185,7 @@
     // If this has already been called, UpstreamFetchTask will ignore it.
     task->OnStreamAndFetchClosed(kResetCodeCancelled, "");
   }
-  if (remove_callback_ != nullptr) {
-    std::move(remove_callback_)();
-    remove_callback_ = nullptr;
-  }
+  task = nullptr;
 }
 
 void UpstreamFetch::OnFetchResult(Location largest_location,
diff --git a/quiche/quic/moqt/moqt_track.h b/quiche/quic/moqt/moqt_track.h
index 54716d1..269ef2e 100644
--- a/quiche/quic/moqt/moqt_track.h
+++ b/quiche/quic/moqt/moqt_track.h
@@ -12,12 +12,14 @@
 #include <optional>
 #include <utility>
 
+#include "absl/base/nullability.h"
 #include "absl/status/status.h"
 #include "absl/strings/string_view.h"
 #include "quiche/quic/core/quic_alarm.h"
 #include "quiche/quic/core/quic_alarm_factory.h"
 #include "quiche/quic/core/quic_time.h"
 #include "quiche/quic/moqt/moqt_bidi_stream.h"
+#include "quiche/quic/moqt/moqt_error.h"
 #include "quiche/quic/moqt/moqt_fetch_task.h"
 #include "quiche/quic/moqt/moqt_key_value_pair.h"
 #include "quiche/quic/moqt/moqt_messages.h"
@@ -69,8 +71,6 @@
 
   virtual bool is_fetch() const = 0;
 
-  virtual void Destroy() = 0;
-
   // A REQUEST_UPDATE changes any field that is present in |parameters|.
   void Update(const MessageParameters& parameters) {
     parameters_.Update(parameters);
@@ -78,8 +78,9 @@
 
   MoqtBidiStreamBase* request_stream() { return request_stream_; }
 
- protected:
   const MessageParameters& const_parameters() const { return parameters_; }
+
+ protected:
   MessageParameters& parameters() { return parameters_; }
 
  private:
@@ -102,15 +103,11 @@
   using AddCallback = quiche::SingleUseCallback<bool(SubscribeRemoteTrack*)>;
   using RemoveCallback = quiche::SingleUseCallback<void(SubscribeRemoteTrack*)>;
   SubscribeRemoteTrack(const MoqtSubscribe& subscribe,
-                       SubscribeVisitor* visitor, AddCallback add_callback,
-                       RemoveCallback remove_callback)
+                       SubscribeVisitor* visitor,
+                       MoqtBidiStreamBase* request_stream)
       : RemoteTrack(subscribe.full_track_name, subscribe.request_id,
-                    subscribe.parameters, /*request_stream=*/nullptr),
-        visitor_(visitor),
-        add_callback_(std::move(add_callback)),
-        remove_callback_(std::move(remove_callback)) {}
-  // For incoming PUBLISH, all |callbacks| will be nullptr because it's handled
-  // by the stream.
+                    subscribe.parameters, request_stream),
+        visitor_(visitor) {}
   SubscribeRemoteTrack(const MoqtPublish& publish, SubscribeVisitor* visitor,
                        MoqtBidiStreamBase* request_stream)
       : RemoteTrack(publish.full_track_name, publish.request_id,
@@ -127,12 +124,8 @@
   std::optional<uint64_t> track_alias() const { return track_alias_; }
   // Returns false if the callback returns false, meaning the session has been
   // destroyed.
-  bool set_track_alias(uint64_t track_alias) {
+  void set_track_alias(uint64_t track_alias) {
     track_alias_.emplace(track_alias);
-    if (add_callback_ != nullptr) {
-      return std::move(add_callback_)(this);
-    }
-    return true;
   }
   void OnStreamOpened();
   void OnStreamClosed(bool fin_received, std::optional<DataStreamIndex> index);
@@ -169,19 +162,8 @@
   }
   void set_visitor(SubscribeVisitor* visitor) { visitor_ = visitor; }
 
-  void Destroy() {
-    if (request_stream() != nullptr) {
-      request_stream()->Fin();
-    }
-    if (remove_callback_) {
-      // Null the callback before calling, because the session owns this and
-      // the callback will call the destructor. When SUBSCRIBE moves to a stream
-      // this won't be a problem.
-      RemoveCallback callback = std::move(remove_callback_);
-      remove_callback_ = nullptr;
-      std::move(callback)(this);
-    }
-  }
+  void SendObjectAck(uint64_t group_id, uint64_t object_id,
+                     quic::QuicTimeDelta delta_from_deadline);
 
  private:
   friend class test::MoqtSessionPeer;
@@ -189,13 +171,19 @@
 
   class PublishDoneDelegate : public quic::QuicAlarm::DelegateWithoutContext {
    public:
-    PublishDoneDelegate(SubscribeRemoteTrack* subscribe)
+    PublishDoneDelegate(quiche::QuicheWeakPtr<RemoteTrack> subscribe)
         : subscribe_(subscribe) {}
 
-    void OnAlarm() override { subscribe_->Destroy(); }
+    void OnAlarm() override {
+      RemoteTrack* subscribe = subscribe_.GetIfAvailable();
+      if (subscribe == nullptr) {
+        return;
+      }
+      subscribe->request_stream()->Reset(kResetCodeCancelled);
+    }
 
    private:
-    SubscribeRemoteTrack* subscribe_;
+    quiche::QuicheWeakPtr<RemoteTrack> subscribe_;
   };
 
   void MaybeSetPublishDoneAlarm();
@@ -216,9 +204,6 @@
   int currently_open_streams_ = 0;
   // Every stream that has received FIN or RESET_STREAM.
   uint64_t streams_closed_ = 0;
-  // For PUBLISH (and later SUBSCRIBE), will be handled in the request stream.
-  AddCallback add_callback_ = nullptr;
-  RemoveCallback remove_callback_ = nullptr;
   // Value assigned on PUBLISH_DONE. Can destroy subscription state if
   // streams_closed_ == total_streams_.
   std::optional<uint64_t> total_streams_;
@@ -292,7 +277,7 @@
   // Called when the data stream is destroyed.
   void OnStreamClosed() { Destroy(); }
 
-  void Destroy() override {
+  void Destroy() {
     if (remove_callback_) {
       RemoveFetchCallback callback = std::move(remove_callback_);
       remove_callback_ = nullptr;
@@ -421,7 +406,6 @@
 
   // Initial values from Fetch() call.
   FetchResponseCallback ok_callback_;  // Will be destroyed on FETCH_OK.
-
   RemoveFetchCallback remove_callback_;
 };
 
diff --git a/quiche/quic/moqt/moqt_track_test.cc b/quiche/quic/moqt/moqt_track_test.cc
index 010cda9..bb50891 100644
--- a/quiche/quic/moqt/moqt_track_test.cc
+++ b/quiche/quic/moqt/moqt_track_test.cc
@@ -10,17 +10,20 @@
 
 #include "absl/status/status.h"
 #include "quiche/quic/core/quic_alarm.h"
+#include "quiche/quic/moqt/moqt_error.h"
 #include "quiche/quic/moqt/moqt_fetch_task.h"
 #include "quiche/quic/moqt/moqt_key_value_pair.h"
 #include "quiche/quic/moqt/moqt_messages.h"
 #include "quiche/quic/moqt/moqt_names.h"
 #include "quiche/quic/moqt/moqt_object.h"
 #include "quiche/quic/moqt/moqt_types.h"
+#include "quiche/quic/moqt/test_tools/mock_moqt_session.h"
 #include "quiche/quic/moqt/test_tools/moqt_mock_visitor.h"
 #include "quiche/quic/platform/api/quic_test.h"
 #include "quiche/quic/test_tools/mock_clock.h"
 #include "quiche/quic/test_tools/quic_test_utils.h"
 #include "quiche/common/quiche_mem_slice.h"
+#include "quiche/web_transport/test_tools/mock_web_transport.h"
 #include "quiche/web_transport/web_transport.h"
 
 namespace moqt {
@@ -52,21 +55,16 @@
 
 class SubscribeRemoteTrackTest : public quic::test::QuicTest {
  public:
-  SubscribeRemoteTrackTest()
-      : track_(
-            subscribe_, &visitor_,
-            [&](SubscribeRemoteTrack*) {
-              alias_registered_ = true;
-              return true;
-            },
-            [this](SubscribeRemoteTrack*) { deleted_ = true; }) {}
+  SubscribeRemoteTrackTest() : track_(subscribe_, &visitor_, &stream_) {
+    stream_.BindStream(&wt_stream_);
+  }
 
   MockSubscribeRemoteTrackVisitor visitor_;
   MoqtSubscribe subscribe_ = {/*request_id=*/1, FullTrackName("foo", "bar"),
                               MessageParameters(Location(2, 0))};
+  MockBidiStream stream_;
+  webtransport::test::MockStream wt_stream_;
   SubscribeRemoteTrack track_;
-  bool alias_registered_ = false;
-  bool deleted_ = false;
   quic::MockClock clock_;
   quic::test::MockAlarmFactory alarm_factory_;
 };
@@ -77,9 +75,8 @@
   EXPECT_FALSE(track_.track_alias().has_value());
   EXPECT_EQ(track_.visitor(), &visitor_);
   EXPECT_FALSE(track_.is_fetch());
-  EXPECT_TRUE(track_.set_track_alias(1));
+  track_.set_track_alias(1);
   EXPECT_EQ(track_.track_alias(), 1);
-  EXPECT_TRUE(alias_registered_);
 }
 
 TEST_F(SubscribeRemoteTrackTest, AllowError) {
@@ -97,34 +94,35 @@
   track_.OnStreamOpened();
   track_.OnStreamClosed(true, std::nullopt);
   EXPECT_CALL(visitor_, OnPublishDone);
+  ExpectFin(wt_stream_);
   track_.OnPublishDone(1, &clock_, &alarm_factory_);
-  EXPECT_TRUE(deleted_);
 }
 
 TEST_F(SubscribeRemoteTrackTest, OnPublishDoneAllStreamsCloseLater) {
   track_.OnStreamOpened();
   EXPECT_CALL(visitor_, OnPublishDone).Times(0);
+  EXPECT_CALL(wt_stream_, Writev).Times(0);
   track_.OnPublishDone(2, &clock_, &alarm_factory_);
   track_.OnStreamClosed(true, std::nullopt);
   track_.OnStreamOpened();
+  ExpectFin(wt_stream_);
   EXPECT_CALL(visitor_, OnPublishDone);
   track_.OnStreamClosed(true, std::nullopt);
-  EXPECT_TRUE(deleted_);
 }
 
 TEST_F(SubscribeRemoteTrackTest, OnPublishDoneTimesOut) {
   track_.OnStreamOpened();
   EXPECT_CALL(visitor_, OnPublishDone).Times(0);
+  EXPECT_CALL(wt_stream_, Writev).Times(0);
   track_.OnPublishDone(2, &clock_, &alarm_factory_);
-  EXPECT_FALSE(deleted_);
   track_.OnStreamClosed(true, std::nullopt);  // No streams are open; timer set.
   quic::QuicAlarm* alarm =
       SubscribeRemoteTrackPeer::GetPublishDoneAlarm(&track_);
   EXPECT_NE(alarm, nullptr);
   EXPECT_TRUE(alarm->IsSet());
   EXPECT_CALL(visitor_, OnPublishDone);
+  EXPECT_CALL(wt_stream_, ResetWithUserCode(kResetCodeCancelled));
   alarm_factory_.FireAlarm(alarm);
-  EXPECT_TRUE(deleted_);
 }
 
 TEST_F(SubscribeRemoteTrackTest, JoiningFetchMultiObject) {
diff --git a/quiche/quic/moqt/moqt_uni_stream_test.cc b/quiche/quic/moqt/moqt_uni_stream_test.cc
index 9efcabd..65c8ad2 100644
--- a/quiche/quic/moqt/moqt_uni_stream_test.cc
+++ b/quiche/quic/moqt/moqt_uni_stream_test.cc
@@ -491,15 +491,9 @@
         subscribe_message_(1, ftn_, MessageParameters()) {
     EXPECT_CALL(session_, deliver_partial_objects())
         .WillRepeatedly(Return(false));
-    track_ = std::make_unique<SubscribeRemoteTrack>(
-        subscribe_message_, &visitor_,
-        [this](SubscribeRemoteTrack* track) {
-          alias_track_ = track;
-          alias_ = track->track_alias().value();
-          return true;
-        },
-        nullptr);
-    EXPECT_TRUE(track_->set_track_alias(2));
+    track_ = std::make_unique<SubscribeRemoteTrack>(subscribe_message_,
+                                                    &visitor_, nullptr);
+    track_->set_track_alias(2);
     CreateStream();
   }
 
@@ -519,8 +513,6 @@
     EXPECT_CALL(session_, GetSubscribe(alias))
         .WillOnce(Return(track_->weak_ptr()));
     stream_->OnCanRead();
-    EXPECT_EQ(alias_, alias);
-    EXPECT_EQ(alias_track_, track_.get());
   }
 
   webtransport::test::InMemoryStream mock_stream_;
@@ -531,8 +523,6 @@
   testing::NiceMock<MockSubscribeRemoteTrackVisitor> visitor_;
   std::unique_ptr<SubscribeRemoteTrack> track_;
   std::unique_ptr<IncomingDataStream> stream_;
-  uint64_t alias_ = 0;
-  SubscribeRemoteTrack* alias_track_ = nullptr;
 };
 
 TEST_F(IncomingDataStreamTest, DestructorBeforeTrackAlias) {
diff --git a/quiche/quic/moqt/test_tools/mock_moqt_session.h b/quiche/quic/moqt/test_tools/mock_moqt_session.h
index c09c8c6..8583e26 100644
--- a/quiche/quic/moqt/test_tools/mock_moqt_session.h
+++ b/quiche/quic/moqt/test_tools/mock_moqt_session.h
@@ -8,22 +8,59 @@
 #include <cstdint>
 #include <memory>
 #include <optional>
+#include <utility>
 
+#include "absl/base/nullability.h"
+#include "absl/status/status.h"
 #include "absl/strings/string_view.h"
+#include "absl/types/span.h"
+#include "quiche/quic/core/quic_alarm_factory.h"
+#include "quiche/quic/core/quic_clock.h"
+#include "quiche/quic/core/quic_types.h"
+#include "quiche/quic/moqt/moqt_bidi_stream.h"
 #include "quiche/quic/moqt/moqt_error.h"
 #include "quiche/quic/moqt/moqt_fetch_task.h"
+#include "quiche/quic/moqt/moqt_framer.h"
 #include "quiche/quic/moqt/moqt_key_value_pair.h"
+#include "quiche/quic/moqt/moqt_messages.h"
 #include "quiche/quic/moqt/moqt_names.h"
+#include "quiche/quic/moqt/moqt_parser.h"
+#include "quiche/quic/moqt/moqt_priority.h"
+#include "quiche/quic/moqt/moqt_publisher.h"
 #include "quiche/quic/moqt/moqt_session_callbacks.h"
 #include "quiche/quic/moqt/moqt_session_interface.h"
+#include "quiche/quic/moqt/moqt_subscription.h"
+#include "quiche/quic/moqt/moqt_trace_recorder.h"
 #include "quiche/quic/moqt/moqt_types.h"
 #include "quiche/common/platform/api/quiche_test.h"
-#include "quiche/common/quiche_callbacks.h"
+#include "quiche/common/quiche_mem_slice.h"
 #include "quiche/common/quiche_weak_ptr.h"
+#include "quiche/web_transport/test_tools/mock_web_transport.h"
+#include "quiche/web_transport/web_transport.h"
 
 namespace moqt {
 namespace test {
 
+class MockSessionToPublisherInterface : public SessionToPublisherInterface {
+ public:
+  MockSessionToPublisherInterface() : weak_ptr_factory_(this) {}
+  ~MockSessionToPublisherInterface() override = default;
+  MOCK_METHOD(bool, alternate_delivery_timeout, (), (const, override));
+  MOCK_METHOD(void, UpdateTrackPriority,
+              (uint64_t, std::optional<MoqtTrackPriority>, MoqtTrackPriority),
+              (override));
+  MOCK_METHOD(quic::QuicAlarmFactory*, alarm_factory, (), (override));
+  MOCK_METHOD(std::shared_ptr<MoqtTrackPublisher>, GetTrackPublisher,
+              (const FullTrackName&), (override));
+  MOCK_METHOD(MoqtPublishingMonitorInterface*, ReleaseMonitoringInterface,
+              (const FullTrackName&), (override));
+  MOCK_METHOD(const quic::QuicClock*, clock, (), (override));
+  MOCK_METHOD(MoqtTraceRecorder&, trace_recorder, (), (override));
+  MOCK_METHOD(webtransport::Session*, session, (), (override));
+
+  quiche::QuicheWeakPtrFactory<SessionToPublisherInterface> weak_ptr_factory_;
+};
+
 class MockMoqtSession : public MoqtSessionInterface {
  public:
   MOCK_METHOD(MoqtSessionCallbacks&, callbacks, (), (override));
@@ -37,6 +74,10 @@
               (const FullTrackName&, const MessageParameters&,
                MoqtResponseCallback),
               (override));
+  MOCK_METHOD(bool, PublishUpdate,
+              (const FullTrackName& name, const MessageParameters& parameters,
+               MoqtResponseCallback response_callback),
+              (override));
   MOCK_METHOD(void, Unsubscribe, (const FullTrackName& name), (override));
   MOCK_METHOD(bool, Publish,
               (std::shared_ptr<MoqtTrackPublisher> publisher,
@@ -88,6 +129,45 @@
   quiche::QuicheWeakPtrFactory<MoqtSessionInterface> weak_factory_{this};
 };
 
+inline void ExpectFin(webtransport::test::MockStream& stream,
+                      bool has_data = false) {
+  EXPECT_CALL(stream, Writev)
+      .WillOnce([&](absl::Span<quiche::QuicheMemSlice> data,
+                    const webtransport::StreamWriteOptions& options) {
+        EXPECT_EQ(data.empty(), !has_data);
+        EXPECT_TRUE(options.send_fin());
+        return absl::OkStatus();
+      });
+}
+
+class MockBidiStream : public MoqtBidiStreamBase {
+ public:
+  MockBidiStream()
+      : MoqtBidiStreamBase(reinterpret_cast<MoqtFramer*>(0x1),
+                           MoqtControlMessageParser(
+                               "moqt-00", true, quic::Perspective::IS_CLIENT),
+                           nullptr) {}
+  MockBidiStream(MoqtFramer* absl_nonnull framer,
+                 const MoqtControlMessageParser& message_parser,
+                 SessionErrorCallback session_error_callback)
+      : MoqtBidiStreamBase(framer, message_parser,
+                           std::move(session_error_callback)) {}
+
+  MOCK_METHOD(absl::Status, SendRequestUpdate,
+              (uint64_t request_id, uint64_t existing_request_id,
+               const MessageParameters& parameters,
+               MoqtResponseCallback callback),
+              (override));
+  MOCK_METHOD(void, Detach, (), (override));
+  MOCK_METHOD(void, OnStreamBound, (), (override));
+  MOCK_METHOD(absl::Status, OnRawControlMessage,
+              (const MoqtRawControlMessage& message), (override));
+  MOCK_METHOD(absl::Status, OnControlMessage, (const MoqtRequestOk& message),
+              (override));
+  MOCK_METHOD(absl::Status, OnControlMessage, (const MoqtRequestError& message),
+              (override));
+};
+
 }  // namespace test
 }  // namespace moqt
 
diff --git a/quiche/quic/moqt/test_tools/moqt_framer_utils.cc b/quiche/quic/moqt/test_tools/moqt_framer_utils.cc
index 7214426..a0e6d42 100644
--- a/quiche/quic/moqt/test_tools/moqt_framer_utils.cc
+++ b/quiche/quic/moqt/test_tools/moqt_framer_utils.cc
@@ -31,9 +31,6 @@
   quiche::QuicheBuffer operator()(const MoqtSubscribeOk& message) {
     return framer.SerializeSubscribeOk(message);
   }
-  quiche::QuicheBuffer operator()(const MoqtUnsubscribe& message) {
-    return framer.SerializeUnsubscribe(message);
-  }
   quiche::QuicheBuffer operator()(const MoqtPublishDone& message) {
     return framer.SerializePublishDone(message);
   }
diff --git a/quiche/quic/moqt/test_tools/moqt_framer_utils.h b/quiche/quic/moqt/test_tools/moqt_framer_utils.h
index be2f36b..c5a5d4a 100644
--- a/quiche/quic/moqt/test_tools/moqt_framer_utils.h
+++ b/quiche/quic/moqt/test_tools/moqt_framer_utils.h
@@ -19,13 +19,14 @@
 
 namespace moqt::test {
 
-using AnyMoqtControlMessage = std::variant<
-    MoqtSetup, MoqtRequestOk, MoqtRequestError, MoqtSubscribe, MoqtSubscribeOk,
-    MoqtUnsubscribe, MoqtPublishDone, MoqtRequestUpdate, MoqtPublishNamespace,
-    MoqtPublishNamespaceDone, MoqtPublishNamespaceCancel, MoqtTrackStatus,
-    MoqtGoAway, MoqtSubscribeNamespace, MoqtMaxRequestId, MoqtFetch,
-    MoqtFetchCancel, MoqtFetchOk, MoqtRequestsBlocked, MoqtPublish,
-    MoqtNamespace, MoqtNamespaceDone, MoqtObjectAck>;
+using AnyMoqtControlMessage =
+    std::variant<MoqtSetup, MoqtRequestOk, MoqtRequestError, MoqtSubscribe,
+                 MoqtSubscribeOk, MoqtPublishDone, MoqtRequestUpdate,
+                 MoqtPublishNamespace, MoqtPublishNamespaceDone,
+                 MoqtPublishNamespaceCancel, MoqtTrackStatus, MoqtGoAway,
+                 MoqtSubscribeNamespace, MoqtMaxRequestId, MoqtFetch,
+                 MoqtFetchCancel, MoqtFetchOk, MoqtRequestsBlocked, MoqtPublish,
+                 MoqtNamespace, MoqtNamespaceDone, MoqtObjectAck>;
 
 std::string SerializeGenericMessage(const AnyMoqtControlMessage& frame,
                                     bool use_webtrans = false);
diff --git a/quiche/quic/moqt/test_tools/moqt_session_peer.h b/quiche/quic/moqt/test_tools/moqt_session_peer.h
index 3bb847b..c7bdb73 100644
--- a/quiche/quic/moqt/test_tools/moqt_session_peer.h
+++ b/quiche/quic/moqt/test_tools/moqt_session_peer.h
@@ -193,6 +193,10 @@
                ? std::optional<uint64_t>()
                : session->subscriptions_with_queued_streams_.begin()->second;
   }
+
+  static uint64_t GetLastTrackAlias(MoqtSession* session) {
+    return session->next_local_track_alias_ - 1;
+  }
 };
 
 }  // namespace moqt::test
diff --git a/quiche/quic/moqt/test_tools/moqt_test_message.h b/quiche/quic/moqt/test_tools/moqt_test_message.h
index 4a01ac2..4e1d394 100644
--- a/quiche/quic/moqt/test_tools/moqt_test_message.h
+++ b/quiche/quic/moqt/test_tools/moqt_test_message.h
@@ -112,15 +112,13 @@
  public:
   virtual ~TestMessageBase() = default;
 
-  using MessageStructuredData =
-      std::variant<MoqtSetup, MoqtObject, MoqtRequestOk, MoqtRequestError,
-                   MoqtSubscribe, MoqtSubscribeOk, MoqtUnsubscribe,
-                   MoqtPublishDone, MoqtRequestUpdate, MoqtPublishNamespace,
-                   MoqtPublishNamespaceDone, MoqtPublishNamespaceCancel,
-                   MoqtTrackStatus, MoqtGoAway, MoqtSubscribeNamespace,
-                   MoqtMaxRequestId, MoqtFetch, MoqtFetchCancel, MoqtFetchOk,
-                   MoqtRequestsBlocked, MoqtPublish, MoqtNamespace,
-                   MoqtNamespaceDone, MoqtObjectAck>;
+  using MessageStructuredData = std::variant<
+      MoqtSetup, MoqtObject, MoqtRequestOk, MoqtRequestError, MoqtSubscribe,
+      MoqtSubscribeOk, MoqtPublishDone, MoqtRequestUpdate, MoqtPublishNamespace,
+      MoqtPublishNamespaceDone, MoqtPublishNamespaceCancel, MoqtTrackStatus,
+      MoqtGoAway, MoqtSubscribeNamespace, MoqtMaxRequestId, MoqtFetch,
+      MoqtFetchCancel, MoqtFetchOk, MoqtRequestsBlocked, MoqtPublish,
+      MoqtNamespace, MoqtNamespaceDone, MoqtObjectAck>;
 
   // The total actual size of the message.
   size_t total_message_size() const { return wire_image_size_; }
@@ -872,39 +870,6 @@
   };
 };
 
-class QUICHE_NO_EXPORT UnsubscribeMessage : public TestMessageBase {
- public:
-  UnsubscribeMessage() : TestMessageBase() {
-    SetWireImage(raw_packet_, sizeof(raw_packet_));
-  }
-
-  bool EqualFieldValues(const MessageStructuredData& values) const override {
-    auto cast = std::get<MoqtUnsubscribe>(values);
-    if (cast.request_id != unsubscribe_.request_id) {
-      QUIC_LOG(INFO) << "UNSUBSCRIBE request ID mismatch";
-      return false;
-    }
-    return true;
-  }
-
-  void ExpandVarints() override { ExpandVarintsImpl("v"); }
-
-  MessageStructuredData structured_data() const override {
-    return TestMessageBase::MessageStructuredData(unsubscribe_);
-  }
-
- private:
-  uint8_t raw_packet_[4] = {
-      0x0a,
-      0x00,
-      0x01,
-      0x03,  // request_id = 3
-  };
-
-  MoqtUnsubscribe unsubscribe_ = {
-      /*request_id=*/3,
-  };
-};
 
 class QUICHE_NO_EXPORT PublishDoneMessage : public TestMessageBase {
  public:
@@ -1764,10 +1729,6 @@
 
   bool EqualFieldValues(const MessageStructuredData& values) const override {
     auto cast = std::get<MoqtObjectAck>(values);
-    if (cast.subscribe_id != object_ack_.subscribe_id) {
-      QUIC_LOG(INFO) << "OBJECT_ACK subscribe ID mismatch";
-      return false;
-    }
     if (cast.group_id != object_ack_.group_id) {
       QUIC_LOG(INFO) << "OBJECT_ACK group ID mismatch";
       return false;
@@ -1783,21 +1744,20 @@
     return true;
   }
 
-  void ExpandVarints() override { ExpandVarintsImpl("vvvv"); }
+  void ExpandVarints() override { ExpandVarintsImpl("vvv"); }
 
   MessageStructuredData structured_data() const override {
     return TestMessageBase::MessageStructuredData(object_ack_);
   }
 
  private:
-  uint8_t raw_packet_[8] = {
-      0xb1, 0x84, 0x00, 0x04,  // type
-      0x01, 0x10, 0x20,        // subscribe ID, group, object
+  uint8_t raw_packet_[7] = {
+      0xb1, 0x84, 0x00, 0x03,  // type
+      0x10, 0x20,              // group, object
       0x20,                    // 0x10 time delta
   };
 
   MoqtObjectAck object_ack_ = {
-      /*subscribe_id=*/0x01,
       /*group_id=*/0x10,
       /*object_id=*/0x20,
       /*delta_from_deadline=*/quic::QuicTimeDelta::FromMicroseconds(0x10),
@@ -1817,8 +1777,6 @@
       return std::make_unique<SubscribeMessage>();
     case MoqtMessageType::kSubscribeOk:
       return std::make_unique<SubscribeOkMessage>();
-    case MoqtMessageType::kUnsubscribe:
-      return std::make_unique<UnsubscribeMessage>();
     case MoqtMessageType::kPublishDone:
       return std::make_unique<PublishDoneMessage>();
     case MoqtMessageType::kRequestUpdate: