Refactor callbacks and request lifetime.

Migrating SUBSCRIBE to a bidi stream has revealed several problems with object lifetimes and a complex web of callbacks. This revises the architecture in advance of the bidi migration.

* MoqtBidiStreamBase has a pure-virtual Detach() function that frees the session of any state related to the request, called when the stream is Reset and other FINed events (like REQUEST_ERROR). Waiting for Reset streams to be cleaned up in the QUIC stack was creating problems for reconnect test cases, and this is more thoughtful about when resources should be freed.

* BidiStreamDeletedCallback was not useful because it has no arguments. Now, each Bidi stream type has its own bespoke callbacks for deletion, and (for response streams) to register the key properties with the session in the first place.

* The specifics of what properties to use as keys are moved back into MoqtSession. The callbacks just present a pointer to the value/context. For example, MoqtNamespaceStream no longer has a notion of SubscribeNamespaceTree. It just asks the Session if a SUBSCRIBE_NAMESPACE is allowed.

* Add a static function for MoqtSession to check if it's still alive, and use it in callbacks generated by the session.

* Condense AI test bloat in moqt_publish_stream_test.

PiperOrigin-RevId: 944543487
diff --git a/quiche/quic/moqt/moqt_bidi_stream.h b/quiche/quic/moqt/moqt_bidi_stream.h
index 23b1d4e..0ebb74c 100644
--- a/quiche/quic/moqt/moqt_bidi_stream.h
+++ b/quiche/quic/moqt/moqt_bidi_stream.h
@@ -35,9 +35,6 @@
 
 using SessionErrorCallback =
     quiche::SingleUseCallback<void(MoqtError, absl::string_view)>;
-// The provider of this callback owns nothing in MoqtBidiStreamBase. This merely
-// deletes the record.
-using BidiStreamDeletedCallback = quiche::SingleUseCallback<void()>;
 
 // MoqtBidiStreamBase is the base class for bidirectional streams in MoQT.  It
 // contains basic methods for handling and dispatching messages.  An instance of
@@ -47,13 +44,11 @@
  public:
   MoqtBidiStreamBase(MoqtFramer* absl_nonnull framer,
                      const MoqtControlMessageParser& message_parser,
-                     BidiStreamDeletedCallback stream_deleted_callback,
                      SessionErrorCallback session_error_callback)
       : framer_(framer),
         message_parser_(message_parser),
-        stream_deleted_callback_(std::move(stream_deleted_callback)),
         session_error_callback_(std::move(session_error_callback)) {}
-  ~MoqtBidiStreamBase() override { std::move(stream_deleted_callback_)(); }
+  ~MoqtBidiStreamBase() = default;
 
   // Binds a WebTransport stream associated with `parser` to this object.
   void BindStream(
@@ -72,8 +67,12 @@
   }
 
   // webtransport::StreamVisitor implementation.
-  void OnResetStreamReceived(webtransport::StreamErrorCode error) override {}
-  void OnStopSendingReceived(webtransport::StreamErrorCode error) override {}
+  void OnResetStreamReceived(webtransport::StreamErrorCode error) override {
+    Reset(error);
+  }
+  void OnStopSendingReceived(webtransport::StreamErrorCode error) override {
+    Reset(error);
+  }
   void OnWriteSideInDataRecvdState() override {}
   void OnCanRead() override;
   void OnCanWrite() override;
@@ -82,7 +81,12 @@
 
   absl::Status SendOrBufferMessage(quiche::QuicheBuffer message,
                                    bool fin = false) {
-    return outgoing_message_queue_.SendOrBufferMessage(std::move(message), fin);
+    absl::Status status =
+        outgoing_message_queue_.SendOrBufferMessage(std::move(message), fin);
+    if (fin) {
+      Detach();
+    }
+    return status;
   }
   void SendOrBufferMessageOrFatal(quiche::QuicheBuffer message,
                                   bool fin = false) {
@@ -99,12 +103,16 @@
   absl::Status SendRequestError(uint64_t request_id, MoqtRequestErrorInfo info,
                                 bool fin = false);
 
-  void Fin() { CheckStatus(outgoing_message_queue_.Fin()); }
+  void Fin() {
+    CheckStatus(outgoing_message_queue_.Fin());
+    Detach();
+  }
   void Reset(webtransport::StreamErrorCode error) {
     webtransport::Stream* stream = stream_parser_->stream();
     if (stream != nullptr) {
       stream->ResetWithUserCode(error);
     }
+    Detach();
   }
 
   // If `status` is not OK, terminates the connection with a fatal error.
@@ -114,6 +122,10 @@
     }
   }
 
+  // Removes any state in MoqtSession related to the stream. Overrides of this
+  // method must be robust to multiple invocations.
+  virtual void Detach() = 0;
+
   // 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_; }
@@ -150,7 +162,6 @@
   std::unique_ptr<MoqtControlStreamParser> absl_nullable stream_parser_;
   MoqtControlMessageParser message_parser_;
   MoqtControlMessageQueue outgoing_message_queue_;
-  BidiStreamDeletedCallback stream_deleted_callback_;
   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 9b8dc62..8b5f011 100644
--- a/quiche/quic/moqt/moqt_bidi_stream_test.cc
+++ b/quiche/quic/moqt/moqt_bidi_stream_test.cc
@@ -5,6 +5,7 @@
 #include "quiche/quic/moqt/moqt_bidi_stream.h"
 
 #include <memory>
+#include <optional>
 
 #include "absl/status/status.h"
 #include "absl/strings/string_view.h"
@@ -14,7 +15,9 @@
 #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/moqt_framer_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/in_memory_stream.h"
 #include "quiche/web_transport/test_tools/mock_web_transport.h"
 
@@ -35,10 +38,10 @@
     return ControlMessageDispatcher::DispatchControlMessage(
         *this, message_parser(), message, "test");
   }
-  int ok_received() const { return ok_received_; }
+  void Detach() override { detached_ = true; }
 
- private:
   int ok_received_ = 0;
+  bool detached_ = false;
 };
 
 class MoqtBidiStreamTest : public quiche::test::QuicheTest {
@@ -50,11 +53,9 @@
             MoqtControlMessageParser(kDefaultMoqtVersion,
                                      /*webtransport=*/true,
                                      quic::Perspective::IS_CLIENT),
-            deleted_callback_.AsStdFunction(),
             error_callback_.AsStdFunction())) {}
 
   MoqtFramer framer_;
-  testing::MockFunction<void()> deleted_callback_;
   testing::StrictMock<testing::MockFunction<void(MoqtError, absl::string_view)>>
       error_callback_;
   std::unique_ptr<TestMoqtBidiStream> stream_;
@@ -65,11 +66,50 @@
   stream_->BindStream(&mock_stream_);
   EXPECT_CALL(mock_stream_, ResetWithUserCode(1234));
   stream_->Reset(1234);
+  EXPECT_TRUE(stream_->detached_);
 }
 
-TEST_F(MoqtBidiStreamTest, DeletedCallback) {
-  EXPECT_CALL(deleted_callback_, Call());
-  stream_.reset();
+TEST_F(MoqtBidiStreamTest, IncomingReset) {
+  stream_->BindStream(&mock_stream_);
+  EXPECT_CALL(mock_stream_, ResetWithUserCode(1234));
+  stream_->OnResetStreamReceived(1234);
+  EXPECT_TRUE(stream_->detached_);
+}
+
+TEST_F(MoqtBidiStreamTest, FinDetaches) {
+  stream_->BindStream(&mock_stream_);
+  stream_->Fin();
+  EXPECT_TRUE(stream_->detached_);
+}
+
+TEST_F(MoqtBidiStreamTest, IncomingStopSending) {
+  stream_->BindStream(&mock_stream_);
+  EXPECT_CALL(mock_stream_, ResetWithUserCode(1234));
+  stream_->OnStopSendingReceived(1234);
+  EXPECT_TRUE(stream_->detached_);
+}
+
+TEST_F(MoqtBidiStreamTest, SendRequestError) {
+  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,
+      MoqtRequestErrorInfo{RequestErrorCode::kUnauthorized,
+                           /*retry_interval=*/std::nullopt, ""},
+      false));
+  EXPECT_FALSE(stream_->detached_);
+  EXPECT_CALL(
+      mock_stream_,
+      Writev(ControlMessageOfType(MoqtMessageType::kRequestError), testing::_));
+  QUICHE_EXPECT_OK(stream_->SendRequestError(
+      1,
+      MoqtRequestErrorInfo{RequestErrorCode::kUnauthorized,
+                           /*retry_interval=*/std::nullopt, ""},
+      true));
+  EXPECT_TRUE(stream_->detached_);
 }
 
 TEST_F(MoqtBidiStreamTest, DispatchControlMessage) {
@@ -78,7 +118,7 @@
   MoqtFramer framer(/*using_webtrans=*/true, quic::Perspective::IS_SERVER);
   stream.Receive(framer.SerializeRequestOk(MoqtRequestOk()).AsStringView());
   stream_->OnCanRead();
-  EXPECT_EQ(stream_->ok_received(), 1u);
+  EXPECT_EQ(stream_->ok_received_, 1u);
 
   stream.Receive(framer.SerializeGoAway(MoqtGoAway()).AsStringView());
   EXPECT_CALL(error_callback_, Call)
diff --git a/quiche/quic/moqt/moqt_control_message_queue.cc b/quiche/quic/moqt/moqt_control_message_queue.cc
index b2b984f..1b7e377 100644
--- a/quiche/quic/moqt/moqt_control_message_queue.cc
+++ b/quiche/quic/moqt/moqt_control_message_queue.cc
@@ -45,10 +45,16 @@
     fin_queued_ = fin;
     return AddToQueue(std::move(message));
   }
+  if (fin) {
+    fin_queued_ = true;
+  }
   return SendMessage(*stream_, std::move(message), fin);
 }
 
 absl::Status MoqtControlMessageQueue::Fin() {
+  if (fin_queued_) {
+    return absl::OkStatus();
+  }
   fin_queued_ = true;
   if (stream_ != nullptr) {
     return OnCanWrite();
@@ -65,6 +71,7 @@
   return absl::OkStatus();
 }
 
+// static
 absl::Status MoqtControlMessageQueue::SendMessage(webtransport::Stream& stream,
                                                   quiche::QuicheBuffer message,
                                                   bool fin) {
diff --git a/quiche/quic/moqt/moqt_namespace_stream.cc b/quiche/quic/moqt/moqt_namespace_stream.cc
index ec3668d..b7a9672 100644
--- a/quiche/quic/moqt/moqt_namespace_stream.cc
+++ b/quiche/quic/moqt/moqt_namespace_stream.cc
@@ -23,7 +23,6 @@
 #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/session_namespace_tree.h"
 #include "quiche/common/platform/api/quiche_logging.h"
 #include "quiche/web_transport/stream_helpers.h"
 
@@ -34,6 +33,7 @@
   if (task != nullptr) {
     task->DeclareEof();
   }
+  Detach();
 }
 absl::Status MoqtNamespaceSubscriberStream::OnRawControlMessage(
     const MoqtRawControlMessage& message) {
@@ -247,26 +247,16 @@
 
 MoqtNamespacePublisherStream::MoqtNamespacePublisherStream(
     MoqtFramer* framer, const MoqtControlMessageParser& message_parser,
+    AddPrefixCallback add_callback, RemovePrefixCallback remove_callback,
     SessionErrorCallback session_error_callback,
-    SessionNamespaceTree* absl_nonnull tree,
     MoqtIncomingSubscribeNamespaceCallback& application)
     // No stream_deleted_callback because there's no state yet.
-    : MoqtBidiStreamBase(
-          framer, message_parser, []() {}, std::move(session_error_callback)),
-      tree_(tree->GetWeakPtr()),
+    : MoqtBidiStreamBase(framer, message_parser,
+                         std::move(session_error_callback)),
+      add_callback_(std::move(add_callback)),
+      remove_callback_(std::move(remove_callback)),
       application_(application) {}
 
-MoqtNamespacePublisherStream::~MoqtNamespacePublisherStream() {
-  if (task_ == nullptr) {
-    return;
-  }
-  SessionNamespaceTree* tree = tree_.GetIfAvailable();
-  if (tree != nullptr) {
-    // Could be null if the stream died early.
-    tree->UnsubscribeNamespace(task_->prefix());
-  }
-}
-
 absl::Status MoqtNamespacePublisherStream::OnRawControlMessage(
     const MoqtRawControlMessage& message) {
   return ControlMessageDispatcher::DispatchControlMessage(
@@ -276,15 +266,15 @@
 absl::Status MoqtNamespacePublisherStream::OnControlMessage(
     const MoqtSubscribeNamespace& message) {
   request_id_ = message.request_id;
-  SessionNamespaceTree* tree = tree_.GetIfAvailable();
-  if (tree == nullptr) {
-    return SendRequestError(request_id_, RequestErrorCode::kInternalError,
-                            std::nullopt, "Session is gone", /*fin=*/true);
+  if (add_callback_ == nullptr) {
+    return absl::InvalidArgumentError("Two SUBSCRIBE_NAMESPACE on one stream");
   }
-  if (!tree->SubscribeNamespace(message.track_namespace_prefix)) {
+  if (!std::move(add_callback_)(message.track_namespace_prefix)) {
+    add_callback_ = nullptr;
     return SendRequestError(request_id_, RequestErrorCode::kPrefixOverlap,
                             std::nullopt, "", /*fin=*/true);
   }
+  add_callback_ = nullptr;
   QUICHE_DCHECK(task_ == nullptr);
   task_ =
       application_(message.track_namespace_prefix, message.subscribe_options,
@@ -305,6 +295,13 @@
   return absl::OkStatus();
 }
 
+void MoqtNamespacePublisherStream::Detach() {
+  if (remove_callback_ != nullptr) {
+    std::move(remove_callback_)(prefix_);
+    remove_callback_ = nullptr;
+  }
+}
+
 void MoqtNamespacePublisherStream::ProcessNamespaces() {
   if (task_ == nullptr) {
     return;
diff --git a/quiche/quic/moqt/moqt_namespace_stream.h b/quiche/quic/moqt/moqt_namespace_stream.h
index 7374260..a80c4a1 100644
--- a/quiche/quic/moqt/moqt_namespace_stream.h
+++ b/quiche/quic/moqt/moqt_namespace_stream.h
@@ -23,26 +23,32 @@
 #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/session_namespace_tree.h"
+#include "quiche/common/quiche_callbacks.h"
 #include "quiche/common/quiche_circular_deque.h"
 #include "quiche/common/quiche_weak_ptr.h"
 #include "quiche/web_transport/web_transport.h"
 
 namespace moqt {
 
+using AddPrefixCallback =
+    quiche::SingleUseCallback<bool(const TrackNamespace&)>;
+using RemovePrefixCallback =
+    quiche::SingleUseCallback<void(const TrackNamespace&)>;
+
 // This class will be owned by the webtransport stream.
 class MoqtNamespaceSubscriberStream : public MoqtBidiStreamBase {
  public:
   // Assumes the caller will send or queue the SUBSCRIBE_NAMESPACE.
-  MoqtNamespaceSubscriberStream(
-      MoqtFramer* framer, const MoqtControlMessageParser& message_parser,
-      uint64_t request_id, BidiStreamDeletedCallback stream_deleted_callback,
-      SessionErrorCallback session_error_callback,
-      MoqtResponseCallback response_callback)
+  MoqtNamespaceSubscriberStream(MoqtFramer* framer,
+                                const MoqtControlMessageParser& message_parser,
+                                uint64_t request_id,
+                                RemovePrefixCallback remove_callback,
+                                SessionErrorCallback session_error_callback,
+                                MoqtResponseCallback response_callback)
       : MoqtBidiStreamBase(framer, message_parser,
-                           std::move(stream_deleted_callback),
                            std::move(session_error_callback)),
         request_id_(request_id),
+        remove_callback_(std::move(remove_callback)),
         response_callback_(std::move(response_callback)) {}
   ~MoqtNamespaceSubscriberStream() override;
 
@@ -58,6 +64,22 @@
   // Send the prefix now so it is only stored in one place (the task).
   std::unique_ptr<MoqtNamespaceTask> CreateTask(const TrackNamespace& prefix);
 
+  void Detach() override {
+    if (remove_callback_ == nullptr) {
+      return;
+    }
+    NamespaceTask* task = task_.GetIfAvailable();
+    // CreateTask() should be called before Detach() can be. If the task is
+    // then destroyed, the destructor should indirectly call this. Either way,
+    // the task should not be null.
+    QUICHE_DCHECK(task != nullptr);
+    if (task != nullptr) {
+      RemovePrefixCallback callback = std::move(remove_callback_);
+      remove_callback_ = nullptr;
+      std::move(callback)(task->prefix());
+    }
+  }
+
  private:
   // The class that will be passed to the application to consume namespace
   // information. Owned by the application.
@@ -118,6 +140,7 @@
   };
 
   const uint64_t request_id_;
+  RemovePrefixCallback remove_callback_;
   MoqtResponseCallback response_callback_;
   absl::flat_hash_set<TrackNamespace> published_suffixes_;
   quiche::QuicheWeakPtr<NamespaceTask> task_;
@@ -128,10 +151,10 @@
   // Constructor for the publisher side.
   MoqtNamespacePublisherStream(
       MoqtFramer* framer, const MoqtControlMessageParser& message_parser,
+      AddPrefixCallback add_callback, RemovePrefixCallback remove_callback,
       SessionErrorCallback session_error_callback,
-      SessionNamespaceTree* absl_nonnull tree,
       MoqtIncomingSubscribeNamespaceCallback& application);
-  ~MoqtNamespacePublisherStream() override;
+  ~MoqtNamespacePublisherStream() override { Detach(); }
 
   void OnStreamBound() override {
     // TODO(martinduke): Set the priority for this stream.
@@ -141,12 +164,16 @@
   absl::Status OnControlMessage(const MoqtSubscribeNamespace& message);
   absl::Status OnControlMessage(const MoqtRequestUpdate& message);
 
+  void Detach() override;
+
  private:
   void ProcessNamespaces();
   MoqtResponseCallback ResponseCallback(uint64_t request_id);
 
   uint64_t request_id_;
-  quiche::QuicheWeakPtr<SessionNamespaceTree> tree_;
+  TrackNamespace prefix_;
+  AddPrefixCallback add_callback_;
+  RemovePrefixCallback remove_callback_;
   MoqtIncomingSubscribeNamespaceCallback& application_;
   std::unique_ptr<MoqtNamespaceTask> task_;
   absl::flat_hash_set<TrackNamespace> published_suffixes_;
diff --git a/quiche/quic/moqt/moqt_namespace_stream_test.cc b/quiche/quic/moqt/moqt_namespace_stream_test.cc
index 63b51dc..e94a45c 100644
--- a/quiche/quic/moqt/moqt_namespace_stream_test.cc
+++ b/quiche/quic/moqt/moqt_namespace_stream_test.cc
@@ -14,6 +14,7 @@
 #include "absl/strings/string_view.h"
 #include "absl/types/span.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"
@@ -24,7 +25,6 @@
 #include "quiche/quic/moqt/moqt_session_callbacks.h"
 #include "quiche/quic/moqt/moqt_session_interface.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/common/platform/api/quiche_test.h"
@@ -72,7 +72,7 @@
   }
 
   MoqtFramer framer_;
-  testing::MockFunction<void()> deleted_callback_;
+  testing::MockFunction<void(const TrackNamespace&)> deleted_callback_;
   testing::MockFunction<void(MoqtError, absl::string_view)> error_callback_;
   testing::MockFunction<void(
       std::variant<MessageParameters, MoqtRequestErrorInfo>)>
@@ -297,11 +297,10 @@
  public:
   MoqtNamespacePublisherStreamTest()
       : framer_(false, quic::Perspective::IS_CLIENT),
-        tree_(),
         application_callback_(mock_application_.AsStdFunction()),
-        stream_(&framer_, ControlMessageParser(),
-                error_callback_.AsStdFunction(), &tree_,
-                application_callback_) {
+        stream_(&framer_, ControlMessageParser(), add_callback_.AsStdFunction(),
+                remove_callback_.AsStdFunction(),
+                error_callback_.AsStdFunction(), application_callback_) {
     stream_.BindStream(&mock_stream_);
     EXPECT_CALL(mock_stream_, CanWrite()).WillRepeatedly(Return(true));
   }
@@ -314,7 +313,8 @@
   MoqtFramer framer_;
   testing::MockFunction<void(MoqtError, absl::string_view)> error_callback_;
   webtransport::test::MockStream mock_stream_;
-  SessionNamespaceTree tree_;
+  testing::MockFunction<bool(const TrackNamespace&)> add_callback_;
+  testing::MockFunction<void(const TrackNamespace&)> remove_callback_;
   testing::MockFunction<std::unique_ptr<MoqtNamespaceTask>(
       const TrackNamespace&, SubscribeNamespaceOption, const MessageParameters&,
       MoqtResponseCallback)>
@@ -334,6 +334,7 @@
   MockNamespaceTask* task_ptr = nullptr;
   MoqtRequestOk ok(kRequestId);
   ok.parameters.expires = quic::QuicTimeDelta::FromSeconds(60);
+  EXPECT_CALL(add_callback_, Call).WillOnce(Return(true));
   EXPECT_CALL(mock_application_, Call)
       .WillOnce([&](const TrackNamespace&, SubscribeNamespaceOption,
                     const MessageParameters&,
@@ -393,6 +394,37 @@
   task_ptr->InvokeCallback();
 }
 
+TEST_F(MoqtNamespacePublisherStreamTest, SubscribeUnsubscribe) {
+  MoqtSubscribeNamespace message = {
+      kRequestId,
+      TrackNamespace({"foo"}),
+      SubscribeNamespaceOption::kNamespace,
+      MessageParameters(),
+  };
+  ObjectsAvailableCallback callback;
+  MockNamespaceTask* task_ptr = nullptr;
+  MoqtRequestOk ok(kRequestId);
+  ok.parameters.expires = quic::QuicTimeDelta::FromSeconds(60);
+  EXPECT_CALL(add_callback_, Call).WillOnce(Return(true));
+  EXPECT_CALL(mock_application_, Call)
+      .WillOnce([&](const TrackNamespace&, SubscribeNamespaceOption,
+                    const MessageParameters&,
+                    MoqtResponseCallback response_callback) {
+        std::move(response_callback)(ok.parameters);
+        auto task =
+            std::make_unique<MockNamespaceTask>(message.track_namespace_prefix);
+        task_ptr = task.get();
+        return task;
+      });
+  EXPECT_CALL(mock_stream_, Writev(SerializedControlMessage(ok), _));
+  ReceiveControlMessage(message);
+  ASSERT_TRUE(task_ptr != nullptr);
+  EXPECT_EQ(task_ptr->prefix(), message.track_namespace_prefix);
+  // Unsubscribe.
+  EXPECT_CALL(remove_callback_, Call);
+  stream_.OnResetStreamReceived(kResetCodeCancelled);
+}
+
 TEST_F(MoqtNamespacePublisherStreamTest, RequestError) {
   MoqtSubscribeNamespace message = {
       kRequestId,
@@ -400,6 +432,7 @@
       SubscribeNamespaceOption::kNamespace,
       MessageParameters(),
   };
+  EXPECT_CALL(add_callback_, Call).WillOnce(Return(true));
   EXPECT_CALL(mock_application_, Call)
       .WillOnce([&](const TrackNamespace&, SubscribeNamespaceOption,
                     const MessageParameters&,
@@ -422,6 +455,7 @@
       MessageParameters(),
   };
   MockNamespaceTask* task_ptr = nullptr;
+  EXPECT_CALL(add_callback_, Call).WillOnce(Return(true));
   EXPECT_CALL(mock_application_, Call)
       .WillOnce([&](const TrackNamespace&, SubscribeNamespaceOption,
                     const MessageParameters&,
@@ -463,6 +497,7 @@
       MessageParameters(),
   };
   MockNamespaceTask* task_ptr = nullptr;
+  EXPECT_CALL(add_callback_, Call).WillOnce(Return(true));
   EXPECT_CALL(mock_application_, Call)
       .WillOnce([&](const TrackNamespace&, SubscribeNamespaceOption,
                     const MessageParameters&,
@@ -505,16 +540,94 @@
       MessageParameters(),
   };
   // The namespace tree already has a subscriber for a prefix of "foo".
-  tree_.SubscribeNamespace(TrackNamespace({"foo", "bar"}));
+  EXPECT_CALL(add_callback_, Call).WillOnce(Return(false));
   EXPECT_CALL(mock_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _));
   ReceiveControlMessage(message);
-  // Try to subscribe to the parent. Also not allowed.
-  message.track_namespace_prefix.PopElement();
-  message.track_namespace_prefix.PopElement();
-  EXPECT_CALL(mock_stream_,
-              Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _));
+}
+
+TEST_F(MoqtNamespacePublisherStreamTest,
+       DuplicateSubscribeNamespaceOnSameStream) {
+  MoqtSubscribeNamespace message = {
+      kRequestId,
+      TrackNamespace({"foo"}),
+      SubscribeNamespaceOption::kNamespace,
+      MessageParameters(),
+  };
+  EXPECT_CALL(add_callback_, Call).WillOnce(Return(true));
+  MoqtRequestOk ok(kRequestId);
+  EXPECT_CALL(mock_stream_, Writev(SerializedControlMessage(ok), _));
+  EXPECT_CALL(mock_application_, Call)
+      .WillOnce([&](const TrackNamespace&, SubscribeNamespaceOption,
+                    const MessageParameters&,
+                    MoqtResponseCallback response_callback) {
+        std::move(response_callback)(MessageParameters());
+        return std::make_unique<MockNamespaceTask>(
+            message.track_namespace_prefix);
+      });
   ReceiveControlMessage(message);
+
+  EXPECT_CALL(error_callback_, Call(MoqtError::kProtocolViolation,
+                                    "Two SUBSCRIBE_NAMESPACE on one stream"));
+  MoqtSubscribeNamespace message2 = {
+      kRequestId + 2,
+      TrackNamespace({"bar"}),
+      SubscribeNamespaceOption::kNamespace,
+      MessageParameters(),
+  };
+  ReceiveControlMessage(message2);
+}
+
+TEST_F(MoqtNamespacePublisherStreamTest,
+       DuplicateSubscribeNamespaceOnDifferentStreams) {
+  MoqtSubscribeNamespace message1 = {
+      kRequestId,
+      TrackNamespace({"foo"}),
+      SubscribeNamespaceOption::kNamespace,
+      MessageParameters(),
+  };
+  EXPECT_CALL(add_callback_, Call).WillOnce(Return(true));
+  MoqtRequestOk ok1(kRequestId);
+  EXPECT_CALL(mock_stream_, Writev(SerializedControlMessage(ok1), _));
+  EXPECT_CALL(mock_application_, Call)
+      .WillOnce([&](const TrackNamespace&, SubscribeNamespaceOption,
+                    const MessageParameters&,
+                    MoqtResponseCallback response_callback) {
+        std::move(response_callback)(MessageParameters());
+        return std::make_unique<MockNamespaceTask>(
+            message1.track_namespace_prefix);
+      });
+  ReceiveControlMessage(message1);
+
+  testing::MockFunction<void(MoqtError, absl::string_view)> error_callback2;
+  webtransport::test::MockStream mock_stream2;
+  MoqtNamespacePublisherStream stream2(
+      &framer_, ControlMessageParser(), add_callback_.AsStdFunction(),
+      remove_callback_.AsStdFunction(), error_callback2.AsStdFunction(),
+      application_callback_);
+  stream2.BindStream(&mock_stream2);
+  EXPECT_CALL(mock_stream2, CanWrite()).WillRepeatedly(Return(true));
+
+  MoqtSubscribeNamespace message2 = {
+      kRequestId + 2,
+      TrackNamespace({"foo"}),
+      SubscribeNamespaceOption::kNamespace,
+      MessageParameters(),
+  };
+  EXPECT_CALL(add_callback_, Call).WillOnce(Return(false));
+  EXPECT_CALL(mock_stream2,
+              Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _));
+  stream2.CheckStatus(stream2.OnControlMessage(message2));
+
+  EXPECT_CALL(error_callback2, Call(MoqtError::kProtocolViolation,
+                                    "Two SUBSCRIBE_NAMESPACE on one stream"));
+  MoqtSubscribeNamespace message3 = {
+      kRequestId + 4,
+      TrackNamespace({"foo"}),
+      SubscribeNamespaceOption::kNamespace,
+      MessageParameters(),
+  };
+  stream2.CheckStatus(stream2.OnControlMessage(message3));
 }
 
 }  // namespace
diff --git a/quiche/quic/moqt/moqt_publish_stream.cc b/quiche/quic/moqt/moqt_publish_stream.cc
index 1e10548..dd7c36e 100644
--- a/quiche/quic/moqt/moqt_publish_stream.cc
+++ b/quiche/quic/moqt/moqt_publish_stream.cc
@@ -30,15 +30,15 @@
 MoqtPublishPublisherStream::MoqtPublishPublisherStream(
     MoqtFramer* absl_nonnull framer,
     const MoqtControlMessageParser& message_parser,
-    BidiStreamDeletedCallback stream_deleted_callback,
+    SubscriptionPublisher::RemoveCallback stream_deleted_callback,
     SessionErrorCallback session_error_callback,
     MoqtResponseCallback response_callback)
     : MoqtBidiStreamBase(framer, message_parser,
-                         std::move(stream_deleted_callback),
                          std::move(session_error_callback)),
-      response_callback_(std::move(response_callback)) {}
+      response_callback_(std::move(response_callback)),
+      stream_deleted_callback_(std::move(stream_deleted_callback)) {}
 
-MoqtPublishPublisherStream::~MoqtPublishPublisherStream() {}
+MoqtPublishPublisherStream::~MoqtPublishPublisherStream() { Detach(); }
 
 void MoqtPublishPublisherStream::OnStreamBound() {
   stream_parser()->set_allow_fin(true);
@@ -110,21 +110,17 @@
     quic::QuicAlarmFactory* absl_nonnull alarm_factory,
     SessionErrorCallback session_error_callback,
     const MoqtIncomingPublishCallback* absl_nonnull incoming_publish_callback,
-    SubscribeRemoteTrack::SubscribeCallbacks callbacks)
-    : MoqtBidiStreamBase(
-          framer, message_parser,
-          /*stream_deleted_callback=*/+[]() {},
-          std::move(session_error_callback)),
+    SubscribeRemoteTrack::AddCallback add_callback,
+    SubscribeRemoteTrack::RemoveCallback remove_callback)
+    : MoqtBidiStreamBase(framer, message_parser,
+                         std::move(session_error_callback)),
       clock_(clock),
       alarm_factory_(alarm_factory),
       incoming_publish_callback_(incoming_publish_callback),
-      callbacks_(std::move(callbacks)),
+      add_callback_(std::move(add_callback)),
+      remove_callback_(std::move(remove_callback)),
       weak_ptr_factory_(this) {}
 
-MoqtPublishSubscriberStream::~MoqtPublishSubscriberStream() {
-  in_destructor_ = true;
-}
-
 absl::Status MoqtPublishSubscriberStream::OnRawControlMessage(
     const MoqtRawControlMessage& message) {
   return ControlMessageDispatcher::DispatchControlMessage(
@@ -133,29 +129,22 @@
 
 absl::Status MoqtPublishSubscriberStream::OnControlMessage(
     const MoqtPublish& message) {
-  if (incoming_publish_callback_ == nullptr) {
+  if (add_callback_ == nullptr) {
     // Two PUBLISH messages for the same stream.
     return absl::InvalidArgumentError("Multiple PUBLISH on the same stream");
   }
-  SubscribeVisitor* visitor = nullptr;
-  SubscribeRemoteTrack* existing_track =
-      std::move(callbacks_.query_name)(message.full_track_name);
-  callbacks_.query_name = nullptr;
-  if (existing_track != nullptr) {
-    // Track already exists.
-    if (!existing_track->ErrorIsAllowed()) {
-      // It's not a pending SUBSCRIBE; refuse this PUBLISH.
-      return SendRequestError(message.request_id,
-                              RequestErrorCode::kDuplicateSubscription,
-                              /*retry_interval=*/std::nullopt, "",
-                              /*fin=*/true);
-    }
-    // It's a pending SUBSCRIBE. Transition it and accept the PUBLISH.
-    visitor = existing_track->ReleaseVisitor();
-    existing_track->Destroy();
-  } else {
-    // No existing SUBSCRIBE, get a new visitor from the application callback.
-    visitor = (*incoming_publish_callback_)(
+  subscriber_ = std::make_unique<SubscribeRemoteTrack>(message, nullptr, this);
+  if (!std::move(add_callback_)(subscriber_.get())) {
+    add_callback_ = nullptr;
+    return SendRequestError(message.request_id,
+                            RequestErrorCode::kDuplicateSubscription,
+                            /*retry_interval=*/std::nullopt, "",
+                            /*fin=*/true);
+  }
+  add_callback_ = nullptr;
+  if (subscriber_->visitor() == nullptr) {
+    // There was no existing SUBSCRIBE, so invoke the callback.
+    subscriber_->set_visitor((*incoming_publish_callback_)(
         message.full_track_name, message.parameters, message.extensions,
         [weakptr = weak_ptr_factory_.Create(), request_id = message.request_id](
             const std::variant<MessageParameters, MoqtRequestErrorInfo>
@@ -168,36 +157,27 @@
               absl::Overload{[&](const MessageParameters& parameters) {
                                stream->subscriber_->Update(parameters);
                                stream->CheckStatus(stream->SendRequestOk(
-                                   request_id, parameters));
+                                   request_id, parameters, /*fin=*/false));
                              },
                              [&](const MoqtRequestErrorInfo& error_info) {
                                stream->CheckStatus(stream->SendRequestError(
-                                   request_id, error_info));
+                                   request_id, error_info, /*fin=*/true));
                              }},
               response);
-        });
+        }));
   }
   incoming_publish_callback_ = nullptr;
-  if (visitor == nullptr) {
+  if (subscriber_->visitor() == nullptr) {
+    // The application doesn't care.
     CheckStatus(SendRequestError(message.request_id,
                                  RequestErrorCode::kUninterested,
                                  /*retry_interval=*/std::nullopt, "",
                                  /*fin=*/true));
     return absl::OkStatus();
   }
-  subscriber_ = std::make_unique<SubscribeRemoteTrack>(
-      message, visitor,
-      [this]() {
-        if (!in_destructor_) {
-          subscriber_.reset();
-          stream()->ResetWithUserCode(kResetCodeCancelled);
-        }
-      },
-      std::move(callbacks_));
-  bool success = subscriber_->set_track_alias(message.track_alias);
-  if (!success) {
-    OnFatalError(absl::AlreadyExistsError(""));
-  }
+  // Notify the visitor.
+  subscriber_->OnObjectOrOk(
+      SubscribeOkData{message.parameters, message.extensions});
   return absl::OkStatus();
 }
 
diff --git a/quiche/quic/moqt/moqt_publish_stream.h b/quiche/quic/moqt/moqt_publish_stream.h
index b10f5fd..a3cef79 100644
--- a/quiche/quic/moqt/moqt_publish_stream.h
+++ b/quiche/quic/moqt/moqt_publish_stream.h
@@ -33,11 +33,12 @@
   // 2. Call SetPublisher()
   // 3. Call Webtransport::Stream::SetVisitor()
   // 4. Call this::BindStream()
-  MoqtPublishPublisherStream(MoqtFramer* absl_nonnull framer,
-                             const MoqtControlMessageParser& message_parser,
-                             BidiStreamDeletedCallback stream_deleted_callback,
-                             SessionErrorCallback session_error_callback,
-                             MoqtResponseCallback response_callback);
+  MoqtPublishPublisherStream(
+      MoqtFramer* absl_nonnull framer,
+      const MoqtControlMessageParser& message_parser,
+      SubscriptionPublisher::RemoveCallback stream_deleted_callback,
+      SessionErrorCallback session_error_callback,
+      MoqtResponseCallback response_callback);
   ~MoqtPublishPublisherStream();
 
   // MoqtBidiStreamBase overrides.
@@ -52,10 +53,21 @@
     publisher_ = std::move(publisher);
   }
 
+  void Detach() override {
+    if (stream_deleted_callback_ == nullptr) {
+      return;
+    }
+    SubscriptionPublisher::RemoveCallback callback =
+        std::move(stream_deleted_callback_);
+    stream_deleted_callback_ = nullptr;
+    std::move(callback)(publisher_.get());
+  }
+
  private:
   MoqtResponseCallback response_callback_;
   std::unique_ptr<SubscriptionPublisher> publisher_;
   absl::flat_hash_map<uint64_t, MoqtResponseCallback> pending_updates_;
+  SubscriptionPublisher::RemoveCallback stream_deleted_callback_;
 };
 
 class MoqtPublishSubscriberStream : public MoqtBidiStreamBase {
@@ -67,8 +79,9 @@
       quic::QuicAlarmFactory* absl_nonnull alarm_factory,
       SessionErrorCallback session_error_callback,
       const MoqtIncomingPublishCallback* absl_nonnull incoming_publish_callback,
-      SubscribeRemoteTrack::SubscribeCallbacks callbacks);
-  ~MoqtPublishSubscriberStream();
+      SubscribeRemoteTrack::AddCallback add_callback,
+      SubscribeRemoteTrack::RemoveCallback remove_callback);
+  ~MoqtPublishSubscriberStream() { Detach(); }
 
   // MoqtBidiStreamBase overrides.
   void OnStreamBound() override {
@@ -83,6 +96,15 @@
   absl::Status OnControlMessage(const MoqtRequestError& message);
   absl::Status OnControlMessage(const MoqtPublishDone& message);
 
+  void Detach() override {
+    if (remove_callback_ != nullptr) {
+      SubscribeRemoteTrack::RemoveCallback callback =
+          std::move(remove_callback_);
+      remove_callback_ = nullptr;
+      std::move(callback)(subscriber_.get());
+    }
+  }
+
  private:
   uint64_t request_id_;
   SubscribeVisitor* absl_nullable subscribe_visitor_ = nullptr;
@@ -92,7 +114,8 @@
   const quic::QuicClock* clock_;
   quic::QuicAlarmFactory* alarm_factory_;
   const MoqtIncomingPublishCallback* incoming_publish_callback_;
-  SubscribeRemoteTrack::SubscribeCallbacks callbacks_;
+  SubscribeRemoteTrack::AddCallback add_callback_;
+  SubscribeRemoteTrack::RemoveCallback remove_callback_;
   quiche::QuicheWeakPtrFactory<MoqtPublishSubscriberStream> weak_ptr_factory_;
 };
 
diff --git a/quiche/quic/moqt/moqt_publish_stream_test.cc b/quiche/quic/moqt/moqt_publish_stream_test.cc
index 6c863a3..e59cb95 100644
--- a/quiche/quic/moqt/moqt_publish_stream_test.cc
+++ b/quiche/quic/moqt/moqt_publish_stream_test.cc
@@ -12,8 +12,8 @@
 #include <variant>
 
 #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"
@@ -31,11 +31,13 @@
 #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/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/web_transport/test_tools/in_memory_stream.h"
+#include "quiche/common/quiche_mem_slice.h"
+#include "quiche/common/test_tools/quiche_test_utils.h"
 #include "quiche/web_transport/test_tools/mock_web_transport.h"
 #include "quiche/web_transport/web_transport.h"
 
@@ -76,21 +78,6 @@
   MOCK_METHOD(webtransport::Session*, session, (), (override));
 };
 
-class BidiStreamWithReset
-    : public webtransport::test::InMemoryStreamWithWriteBuffer {
- public:
-  using InMemoryStreamWithWriteBuffer::InMemoryStreamWithWriteBuffer;
-  void ResetWithUserCode(webtransport::StreamErrorCode error) override {
-    last_reset_code_ = error;
-  }
-  std::optional<webtransport::StreamErrorCode> last_reset_code() const {
-    return last_reset_code_;
-  }
-
- private:
-  std::optional<webtransport::StreamErrorCode> last_reset_code_;
-};
-
 constexpr uint64_t kRequestId = 1;
 constexpr uint64_t kTrackAlias = 10;
 const FullTrackName kTrackName("foo", "bar");
@@ -103,7 +90,7 @@
                         quic::Perspective::IS_CLIENT),
         track_publisher_(std::make_shared<TestTrackPublisher>(kTrackName)) {
     // Construct the stream visitor.
-    stream_visitor_ = std::make_unique<MoqtPublishPublisherStream>(
+    stream_ = std::make_unique<MoqtPublishPublisherStream>(
         &framer_, message_parser_, deleted_callback_.AsStdFunction(),
         error_callback_.AsStdFunction(),
         [this](std::variant<MessageParameters, MoqtRequestErrorInfo> response) {
@@ -117,18 +104,20 @@
 
     EXPECT_CALL(visitor_, session).WillRepeatedly(Return(&webtrans_));
     auto publisher = std::make_unique<SubscriptionPublisher>(
-        framer_, track_publisher_, stream_visitor_.get(), kRequestId,
-        kTrackAlias, parameters_, &visitor_, /*monitoring_interface=*/nullptr,
-        &mock_clock_, trace_recorder_, /*is_publish=*/true);
+        framer_, track_publisher_, stream_.get(), kRequestId, kTrackAlias,
+        parameters_, &visitor_, /*monitoring_interface=*/nullptr, &mock_clock_,
+        trace_recorder_, /*is_publish=*/true);
 
     publisher_ = publisher.get();  // Keep raw pointer for testing
-    stream_visitor_->SetPublisher(std::move(publisher));
+    stream_->SetPublisher(std::move(publisher));
+    EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true));
   }
 
   MoqtFramer framer_;
   MoqtControlMessageParser message_parser_;
+  webtransport::test::MockStream mock_stream_;
   std::shared_ptr<TestTrackPublisher> track_publisher_;
-  testing::MockFunction<void()> deleted_callback_;
+  testing::MockFunction<void(SubscriptionPublisher*)> deleted_callback_;
   testing::StrictMock<testing::MockFunction<void(MoqtError, absl::string_view)>>
       error_callback_;
   MockSessionToPublisherInterface visitor_;
@@ -137,49 +126,30 @@
   MoqtTraceRecorder trace_recorder_;
   MessageParameters parameters_;
 
-  std::unique_ptr<MoqtPublishPublisherStream> stream_visitor_;
+  std::unique_ptr<MoqtPublishPublisherStream> stream_;
   SubscriptionPublisher* publisher_;  // Raw pointer
   std::optional<std::variant<MessageParameters, MoqtRequestErrorInfo>>
       response_;
 };
 
 TEST_F(MoqtPublishPublisherStreamTest, OnStreamBoundSendsPublish) {
-  webtransport::test::InMemoryStreamWithWriteBuffer stream(0);
-  stream_visitor_->BindStream(&stream);  // Calls OnStreamBound
-
-  // Verify PUBLISH message was sent.
-  std::string& written = stream.write_buffer();
-  MoqtControlStreamParser parser(&stream);
-  // Feed the written data back to a parser to verify it.
-  webtransport::test::InMemoryStream read_stream(0);
-  read_stream.Receive(written);
-  MoqtControlStreamParser read_parser(&read_stream);
-  absl::StatusOr<MoqtRawControlMessage> message = read_parser.ReadNextMessage();
-  ASSERT_TRUE(message.ok());
-  EXPECT_EQ(message->type, MoqtMessageType::kPublish);
-
-  MoqtControlMessageParser cmp(kDefaultMoqtVersion, true,
-                               quic::Perspective::IS_CLIENT);
-  absl::StatusOr<MoqtPublish> publish = cmp.ProcessPublish(message->payload);
-  ASSERT_TRUE(publish.ok());
-  EXPECT_EQ(publish->request_id, kRequestId);
-  EXPECT_EQ(publish->full_track_name, kTrackName);
-  EXPECT_EQ(publish->track_alias, kTrackAlias);
-  EXPECT_EQ(publish->parameters.delivery_timeout, parameters_.delivery_timeout);
-  EXPECT_EQ(publish->parameters.group_order, parameters_.group_order);
+  EXPECT_CALL(mock_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kPublish), _))
+      .WillOnce(Return(absl::OkStatus()));
+  stream_->BindStream(&mock_stream_);  // Calls OnStreamBound
 }
 
 TEST_F(MoqtPublishPublisherStreamTest, ReceiveRequestOk) {
-  webtransport::test::InMemoryStreamWithWriteBuffer stream(0);
-  stream_visitor_->BindStream(&stream);
+  EXPECT_CALL(mock_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kPublish), _))
+      .WillOnce(Return(absl::OkStatus()));
+  stream_->BindStream(&mock_stream_);  // Calls OnStreamBound
 
   MoqtRequestOk request_ok;
   request_ok.request_id = kRequestId;
   request_ok.parameters.delivery_timeout = quic::QuicTimeDelta::FromSeconds(2);
   request_ok.parameters.group_order = MoqtDeliveryOrder::kDescending;
-
-  stream.Receive(framer_.SerializeRequestOk(request_ok).AsStringView());
-  stream_visitor_->OnCanRead();
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(request_ok));
 
   // Verify response callback was called.
   ASSERT_TRUE(response_.has_value());
@@ -199,17 +169,17 @@
 }
 
 TEST_F(MoqtPublishPublisherStreamTest, ReceiveRequestError) {
-  webtransport::test::InMemoryStreamWithWriteBuffer stream(0);
-  stream_visitor_->BindStream(&stream);
+  EXPECT_CALL(mock_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kPublish), _))
+      .WillOnce(Return(absl::OkStatus()));
+  stream_->BindStream(&mock_stream_);  // Calls OnStreamBound
 
   MoqtRequestError request_error;
   request_error.request_id = kRequestId;
   request_error.error_code = RequestErrorCode::kUnauthorized;
   request_error.retry_interval = quic::QuicTimeDelta::FromSeconds(5);
   request_error.reason_phrase = "Unauthorized";
-
-  stream.Receive(framer_.SerializeRequestError(request_error).AsStringView());
-  stream_visitor_->OnCanRead();
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(request_error));
 
   // Verify response callback was called with error.
   ASSERT_TRUE(response_.has_value());
@@ -221,10 +191,12 @@
 }
 
 TEST_F(MoqtPublishPublisherStreamTest, ReceiveRequestUpdate) {
-  webtransport::test::InMemoryStreamWithWriteBuffer stream(0);
-  stream_visitor_->BindStream(&stream);
-  stream.write_buffer().clear();  // Clear initial PUBLISH
+  EXPECT_CALL(mock_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kPublish), _))
+      .WillOnce(Return(absl::OkStatus()));
+  stream_->BindStream(&mock_stream_);
 
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(MoqtRequestOk{kRequestId}));
   // Set largest location on publisher
   track_publisher_->AddObject(Location(1, 2), 0, "payload", true);
 
@@ -236,9 +208,10 @@
   request_update.parameters.subscriber_priority = 5;
   request_update.parameters.subscription_filter.emplace(
       MoqtFilterType::kLargestObject);
-
-  stream.Receive(framer_.SerializeRequestUpdate(request_update).AsStringView());
-  stream_visitor_->OnCanRead();
+  EXPECT_CALL(mock_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kRequestOk), _))
+      .WillOnce(Return(absl::OkStatus()));
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(request_update));
 
   // Verify publisher parameters were updated.
   const MessageParameters& pub_params =
@@ -254,51 +227,6 @@
   EXPECT_EQ(pub_params.subscription_filter->type(),
             MoqtFilterType::kAbsoluteStart);
   EXPECT_EQ(pub_params.subscription_filter->start(), Location(1, 3));
-
-  // Verify REQUEST_OK response was sent.
-  std::string& written = stream.write_buffer();
-  webtransport::test::InMemoryStream read_stream(0);
-  read_stream.Receive(written);
-  MoqtControlStreamParser read_parser(&read_stream);
-  absl::StatusOr<MoqtRawControlMessage> message = read_parser.ReadNextMessage();
-  ASSERT_TRUE(message.ok());
-  EXPECT_EQ(message->type, MoqtMessageType::kRequestOk);
-
-  MoqtControlMessageParser cmp(kDefaultMoqtVersion, true,
-                               quic::Perspective::IS_CLIENT);
-  absl::StatusOr<MoqtRequestOk> request_ok =
-      cmp.ProcessRequestOk(message->payload);
-  ASSERT_TRUE(request_ok.ok());
-  EXPECT_EQ(request_ok->request_id, request_update.request_id);
-}
-
-TEST_F(MoqtPublishPublisherStreamTest, ReceiveRequestOkMismatchedId) {
-  BidiStreamWithReset stream(0);
-  stream_visitor_->BindStream(&stream);
-
-  // Receive REQUEST_OK with mismatched ID.
-  MoqtRequestOk request_ok;
-  request_ok.request_id = kRequestId + 1;  // Mismatched
-  EXPECT_CALL(error_callback_,
-              Call(MoqtError::kProtocolViolation,
-                   "REQUEST_OK does not match PUBLISH request ID"));
-  stream.Receive(framer_.SerializeRequestOk(request_ok).AsStringView());
-  stream_visitor_->OnCanRead();
-}
-
-TEST_F(MoqtPublishPublisherStreamTest, ReceiveRequestErrorMismatchedId) {
-  BidiStreamWithReset stream(0);
-  stream_visitor_->BindStream(&stream);
-
-  // Receive REQUEST_ERROR with mismatched ID.
-  MoqtRequestError request_error;
-  request_error.request_id = kRequestId + 1;  // Mismatched
-  request_error.error_code = RequestErrorCode::kUninterested;
-  EXPECT_CALL(error_callback_,
-              Call(MoqtError::kProtocolViolation,
-                   "REQUEST_OK does not match PUBLISH request ID"));
-  stream.Receive(framer_.SerializeRequestError(request_error).AsStringView());
-  stream_visitor_->OnCanRead();
 }
 
 class MoqtPublishSubscriberStreamTest : public quiche::test::QuicheTest {
@@ -309,33 +237,18 @@
                         quic::Perspective::IS_SERVER),
         incoming_publish_callback_(
             incoming_publish_callback_mock_.AsStdFunction()) {
-    SubscribeRemoteTrack::SubscribeCallbacks callbacks;
-    callbacks.query_name = [this](const FullTrackName& name) {
-      return query_name_mock_.Call(name);
-    };
-    callbacks.register_name = [this](const FullTrackName& name,
-                                     SubscribeRemoteTrack* track) {
-      register_name_mock_.Call(name, track);
-    };
-    callbacks.register_alias = [this](uint64_t alias,
-                                      SubscribeRemoteTrack* track) {
-      return register_alias_mock_.Call(alias, track);
-    };
-    callbacks.unregister = [this](const FullTrackName& name,
-                                  std::optional<uint64_t> alias) {
-      unregister_mock_.Call(name, alias);
-    };
-
-    stream_visitor_ = std::make_unique<MoqtPublishSubscriberStream>(
+    stream_ = std::make_unique<MoqtPublishSubscriberStream>(
         &framer_, message_parser_, &mock_clock_, &mock_alarm_factory_,
         error_callback_.AsStdFunction(), &incoming_publish_callback_,
-        std::move(callbacks));
+        mock_add_callback_.AsStdFunction(),
+        mock_remove_callback_.AsStdFunction());
+    stream_->BindStream(&mock_stream_);
+    EXPECT_CALL(mock_stream_, CanWrite).WillRepeatedly(Return(true));
   }
 
-  void ExpectSubscriberDestruction() {
-    EXPECT_CALL(unregister_mock_,
-                Call(kTrackName, std::optional<uint64_t>(kTrackAlias)))
-        .Times(1);
+  MoqtPublish DefaultPublish() {
+    return MoqtPublish{kRequestId, kTrackName, kTrackAlias, MessageParameters(),
+                       TrackExtensions()};
   }
 
   MoqtFramer framer_;
@@ -351,263 +264,113 @@
       incoming_publish_callback_mock_;
   MoqtIncomingPublishCallback incoming_publish_callback_;
 
-  testing::MockFunction<SubscribeRemoteTrack*(const FullTrackName&)>
-      query_name_mock_;
-  testing::MockFunction<void(const FullTrackName&, SubscribeRemoteTrack*)>
-      register_name_mock_;
-  testing::MockFunction<bool(uint64_t, SubscribeRemoteTrack*)>
-      register_alias_mock_;
-  testing::MockFunction<void(const FullTrackName&, std::optional<uint64_t>)>
-      unregister_mock_;
+  testing::MockFunction<bool(SubscribeRemoteTrack*)> mock_add_callback_;
+  testing::MockFunction<void(SubscribeRemoteTrack*)> mock_remove_callback_;
 
   StrictMock<MockSubscribeRemoteTrackVisitor> mock_subscribe_visitor_;
   MoqtResponseCallback captured_response_callback_;
-  std::unique_ptr<MoqtPublishSubscriberStream> stream_visitor_;
+  webtransport::test::MockStream mock_stream_;
+  std::unique_ptr<MoqtPublishSubscriberStream> stream_;
 };
 
 TEST_F(MoqtPublishSubscriberStreamTest, ReceivePublishAndAccept) {
-  BidiStreamWithReset stream(0);
-  stream_visitor_->BindStream(&stream);
-
   EXPECT_CALL(mock_subscribe_visitor_, OnReply(kTrackName, _))
       .WillOnce(
           [](const FullTrackName&,
              const std::variant<SubscribeOkData, MoqtRequestErrorInfo>& reply) {
             EXPECT_TRUE(std::holds_alternative<SubscribeOkData>(reply));
           });
-  EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(kTrackName));
-  ExpectSubscriberDestruction();
-
-  MoqtPublish publish;
-  publish.request_id = kRequestId;
-  publish.full_track_name = kTrackName;
-  publish.track_alias = kTrackAlias;
-  publish.parameters.delivery_timeout = quic::QuicTimeDelta::FromSeconds(1);
-
   EXPECT_CALL(incoming_publish_callback_mock_, Call(kTrackName, _, _, _))
       .WillOnce([this](const FullTrackName&, const MessageParameters&,
                        const TrackExtensions&, MoqtResponseCallback callback) {
         captured_response_callback_ = std::move(callback);
         return &mock_subscribe_visitor_;
       });
-
-  EXPECT_CALL(query_name_mock_, Call(kTrackName)).WillOnce(Return(nullptr));
-  EXPECT_CALL(register_name_mock_, Call(kTrackName, NotNull())).Times(1);
   SubscribeRemoteTrack* captured_subscriber = nullptr;
-  EXPECT_CALL(register_alias_mock_, Call(kTrackAlias, NotNull()))
-      .WillOnce([&](uint64_t, SubscribeRemoteTrack* track) {
-        captured_subscriber = track;
+  EXPECT_CALL(mock_add_callback_, Call(NotNull()))
+      .WillOnce([&](SubscribeRemoteTrack* subscriber) {
+        captured_subscriber = subscriber;
         return true;
       });
-
-  stream.Receive(framer_.SerializePublish(publish).AsStringView());
-  stream_visitor_->OnCanRead();
+  MoqtPublish publish = DefaultPublish();
+  publish.parameters.delivery_timeout = quic::QuicTimeDelta::FromSeconds(1);
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(publish));
 
   // Verify subscriber is created.
   ASSERT_NE(captured_subscriber, nullptr);
   EXPECT_EQ(captured_subscriber->track_alias(), kTrackAlias);
   EXPECT_EQ(captured_subscriber->visitor(), &mock_subscribe_visitor_);
 
-  // Now call the response callback with success.
-  stream.write_buffer().clear();
+  // Verify REQUEST_OK response was sent.
+  EXPECT_CALL(mock_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kRequestOk), _))
+      .WillOnce(Return(absl::OkStatus()));
   MessageParameters response_parameters;
   response_parameters.delivery_timeout = quic::QuicTimeDelta::FromSeconds(2);
   std::move(captured_response_callback_)(response_parameters);
 
-  // Verify REQUEST_OK response was sent.
-  std::string& written = stream.write_buffer();
-  webtransport::test::InMemoryStream read_stream(0);
-  read_stream.Receive(written);
-  MoqtControlStreamParser read_parser(&read_stream);
-  absl::StatusOr<MoqtRawControlMessage> message = read_parser.ReadNextMessage();
-  ASSERT_TRUE(message.ok());
-  EXPECT_EQ(message->type, MoqtMessageType::kRequestOk);
-
-  MoqtControlMessageParser cmp(kDefaultMoqtVersion, true,
-                               quic::Perspective::IS_SERVER);
-  absl::StatusOr<MoqtRequestOk> request_ok =
-      cmp.ProcessRequestOk(message->payload);
-  ASSERT_TRUE(request_ok.ok());
-  EXPECT_EQ(request_ok->request_id, kRequestId);
-  EXPECT_EQ(request_ok->parameters.delivery_timeout,
-            response_parameters.delivery_timeout);
-
   // Verify subscriber parameters were updated.
   const MessageParameters& sub_params =
       SubscribeRemoteTrackPeer::parameters(*captured_subscriber);
   EXPECT_EQ(sub_params.delivery_timeout, response_parameters.delivery_timeout);
+  EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone);
 }
 
 TEST_F(MoqtPublishSubscriberStreamTest, ReceivePublishAndReject) {
-  BidiStreamWithReset stream(0);
-  stream_visitor_->BindStream(&stream);
-
-  MoqtPublish publish;
-  publish.request_id = kRequestId;
-  publish.full_track_name = kTrackName;
-  publish.track_alias = kTrackAlias;
-
+  MoqtPublish publish = DefaultPublish();
   // Callback returns nullptr (rejection).
+  EXPECT_CALL(mock_add_callback_, Call(NotNull())).WillOnce(Return(true));
   EXPECT_CALL(incoming_publish_callback_mock_, Call(kTrackName, _, _, _))
       .WillOnce(Return(nullptr));
-
-  stream.Receive(framer_.SerializePublish(publish).AsStringView());
-  stream_visitor_->OnCanRead();
-
-  // Verify REQUEST_ERROR was sent.
-  std::string& written = stream.write_buffer();
-  webtransport::test::InMemoryStream read_stream(0);
-  read_stream.Receive(written);
-  MoqtControlStreamParser read_parser(&read_stream);
-  absl::StatusOr<MoqtRawControlMessage> message = read_parser.ReadNextMessage();
-  ASSERT_TRUE(message.ok());
-  EXPECT_EQ(message->type, MoqtMessageType::kRequestError);
-
-  MoqtControlMessageParser cmp(kDefaultMoqtVersion, true,
-                               quic::Perspective::IS_SERVER);
-  absl::StatusOr<MoqtRequestError> request_error =
-      cmp.ProcessRequestError(message->payload);
-  ASSERT_TRUE(request_error.ok());
-  EXPECT_EQ(request_error->request_id, kRequestId);
-  EXPECT_EQ(request_error->error_code, RequestErrorCode::kUninterested);
-  EXPECT_TRUE(stream.fin_sent());
+  EXPECT_CALL(mock_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _))
+      .WillOnce(Return(absl::OkStatus()));
+  EXPECT_CALL(mock_remove_callback_, Call);
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(publish));
 }
 
-TEST_F(MoqtPublishSubscriberStreamTest, ReceivePublishDuplicate) {
-  BidiStreamWithReset stream(0);
-  stream_visitor_->BindStream(&stream);
-
-  MoqtPublish publish;
-  publish.request_id = kRequestId;
-  publish.full_track_name = kTrackName;
-  publish.track_alias = kTrackAlias;
-
+TEST_F(MoqtPublishSubscriberStreamTest, ReceiveTwoPublishOnStream) {
+  MoqtPublish publish = DefaultPublish();
+  EXPECT_CALL(mock_add_callback_, Call(NotNull())).WillOnce(Return(true));
   EXPECT_CALL(incoming_publish_callback_mock_, Call(kTrackName, _, _, _))
       .WillOnce(Return(&mock_subscribe_visitor_));
-  EXPECT_CALL(query_name_mock_, Call(kTrackName)).WillOnce(Return(nullptr));
-  EXPECT_CALL(register_name_mock_, Call(kTrackName, NotNull())).Times(1);
-  EXPECT_CALL(register_alias_mock_, Call(kTrackAlias, 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));
+
+  // Receive second PUBLISH on same stream.
+  publish.request_id = kRequestId + 2;
+  publish.full_track_name = FullTrackName("dead", "beef");
+  publish.track_alias = kTrackAlias + 1;
+  EXPECT_FALSE(stream_->OnControlMessage(publish).ok());
   EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(kTrackName));
-  ExpectSubscriberDestruction();
-
-  stream.Receive(framer_.SerializePublish(publish).AsStringView());
-  stream_visitor_->OnCanRead();
-
-  // Send second PUBLISH on same stream.
-  stream.Receive(framer_.SerializePublish(publish).AsStringView());
-
-  // It should return InvalidArgumentError, which calls OnFatalError in
-  // MoqtBidiStreamBase.
-  EXPECT_CALL(error_callback_, Call(MoqtError::kProtocolViolation,
-                                    "Multiple PUBLISH on the same stream"));
-  stream_visitor_->OnCanRead();
 }
 
-TEST_F(MoqtPublishSubscriberStreamTest, ReceivePublishDuplicateName) {
-  BidiStreamWithReset stream(0);
-  stream_visitor_->BindStream(&stream);
-
-  MoqtPublish publish;
+TEST_F(MoqtPublishSubscriberStreamTest, ReceivePublishDuplicate) {
+  MoqtPublish publish = DefaultPublish();
   publish.request_id = kRequestId + 2;
   publish.full_track_name = kTrackName;
   publish.track_alias = kTrackAlias;
 
-  // Simulate an existing established track.
-  MoqtSubscribe sub;
-  sub.full_track_name = kTrackName;
-  sub.request_id = kRequestId;
-  StrictMock<MockSubscribeRemoteTrackVisitor> existing_visitor;
-  SubscribeRemoteTrack existing_track(sub, &existing_visitor, []() {}, {});
-  existing_track.OnObjectOrOk();
-
-  EXPECT_CALL(existing_visitor, OnPublishDone(kTrackName)).Times(1);
-
-  EXPECT_CALL(query_name_mock_, Call(kTrackName))
-      .WillOnce(Return(&existing_track));
-
-  stream.Receive(framer_.SerializePublish(publish).AsStringView());
-  stream_visitor_->OnCanRead();
-
-  // Verify REQUEST_ERROR was sent.
-  std::string& written = stream.write_buffer();
-  webtransport::test::InMemoryStream read_stream(0);
-  read_stream.Receive(written);
-  MoqtControlStreamParser read_parser(&read_stream);
-  absl::StatusOr<MoqtRawControlMessage> message = read_parser.ReadNextMessage();
-  ASSERT_TRUE(message.ok());
-  EXPECT_EQ(message->type, MoqtMessageType::kRequestError);
-
-  MoqtControlMessageParser cmp(kDefaultMoqtVersion, true,
-                               quic::Perspective::IS_SERVER);
-  absl::StatusOr<MoqtRequestError> request_error =
-      cmp.ProcessRequestError(message->payload);
-  ASSERT_TRUE(request_error.ok());
-  EXPECT_EQ(request_error->request_id, publish.request_id);
-  EXPECT_EQ(request_error->error_code,
-            RequestErrorCode::kDuplicateSubscription);
-  EXPECT_TRUE(stream.fin_sent());
-}
-
-TEST_F(MoqtPublishSubscriberStreamTest, ReceivePublishDuplicateAlias) {
-  BidiStreamWithReset stream(0);
-  stream_visitor_->BindStream(&stream);
-
-  MoqtPublish publish;
-  publish.request_id = kRequestId;
-  publish.full_track_name = kTrackName;
-  publish.track_alias = kTrackAlias;
-
-  EXPECT_CALL(incoming_publish_callback_mock_, Call(kTrackName, _, _, _))
-      .WillOnce(Return(&mock_subscribe_visitor_));
-
-  EXPECT_CALL(query_name_mock_, Call(kTrackName)).WillOnce(Return(nullptr));
-  EXPECT_CALL(register_name_mock_, Call(kTrackName, NotNull())).Times(1);
-  // Return duplicate alias error (false) from alias callback.
-  EXPECT_CALL(register_alias_mock_, Call(kTrackAlias, NotNull()))
-      .WillOnce(Return(false));
-  EXPECT_CALL(mock_subscribe_visitor_, OnReply(kTrackName, _))
-      .WillOnce(
-          [](const FullTrackName&,
-             const std::variant<SubscribeOkData, MoqtRequestErrorInfo>& reply) {
-            EXPECT_TRUE(std::holds_alternative<SubscribeOkData>(reply));
-          });
-  EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(kTrackName));
-  ExpectSubscriberDestruction();
-
-  // It should call OnFatalError, which calls error_callback_.
-  // Note: The error message is now empty because we pass
-  // AlreadyExistsError("").
-  EXPECT_CALL(error_callback_, Call(MoqtError::kDuplicateTrackAlias, ""));
-
-  stream.Receive(framer_.SerializePublish(publish).AsStringView());
-  stream_visitor_->OnCanRead();
+  EXPECT_CALL(mock_add_callback_, Call).WillOnce(Return(false));
+  EXPECT_CALL(mock_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _))
+      .WillOnce(Return(absl::OkStatus()));
+  EXPECT_CALL(mock_remove_callback_, Call);
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(publish));
 }
 
 TEST_F(MoqtPublishSubscriberStreamTest, ReceiveRequestUpdate) {
-  // First, establish subscription.
-  BidiStreamWithReset stream(0);
-  stream_visitor_->BindStream(&stream);
-
-  MoqtPublish publish;
-  publish.request_id = kRequestId;
-  publish.full_track_name = kTrackName;
-  publish.track_alias = kTrackAlias;
-
+  MoqtPublish publish = DefaultPublish();
   EXPECT_CALL(incoming_publish_callback_mock_, Call(kTrackName, _, _, _))
       .WillOnce(Return(&mock_subscribe_visitor_));
-
-  EXPECT_CALL(query_name_mock_, Call(kTrackName)).WillOnce(Return(nullptr));
-  EXPECT_CALL(register_name_mock_, Call(kTrackName, NotNull())).Times(1);
   SubscribeRemoteTrack* captured_subscriber = nullptr;
-  EXPECT_CALL(register_alias_mock_, Call(kTrackAlias, NotNull()))
-      .WillOnce([&](uint64_t, SubscribeRemoteTrack* track) {
+  EXPECT_CALL(mock_add_callback_, Call(NotNull()))
+      .WillOnce([&](SubscribeRemoteTrack* track) {
         captured_subscriber = track;
         return true;
       });
@@ -617,12 +380,7 @@
              const std::variant<SubscribeOkData, MoqtRequestErrorInfo>& reply) {
             EXPECT_TRUE(std::holds_alternative<SubscribeOkData>(reply));
           });
-  EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(kTrackName));
-  ExpectSubscriberDestruction();
-
-  stream.Receive(framer_.SerializePublish(publish).AsStringView());
-  stream_visitor_->OnCanRead();
-  stream.write_buffer().clear();
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(publish));
 
   // Now receive REQUEST_UPDATE.
   MoqtRequestUpdate request_update;
@@ -630,9 +388,10 @@
   request_update.existing_request_id = kRequestId;
   request_update.parameters.delivery_timeout =
       quic::QuicTimeDelta::FromSeconds(3);
-
-  stream.Receive(framer_.SerializeRequestUpdate(request_update).AsStringView());
-  stream_visitor_->OnCanRead();
+  EXPECT_CALL(mock_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kRequestOk), _))
+      .WillOnce(Return(absl::OkStatus()));
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(request_update));
 
   // Verify subscriber parameters were updated.
   ASSERT_NE(captured_subscriber, nullptr);
@@ -640,139 +399,166 @@
       SubscribeRemoteTrackPeer::parameters(*captured_subscriber);
   EXPECT_EQ(sub_params.delivery_timeout,
             request_update.parameters.delivery_timeout);
-
-  // Verify REQUEST_OK response was sent.
-  std::string& written = stream.write_buffer();
-  webtransport::test::InMemoryStream read_stream(0);
-  read_stream.Receive(written);
-  MoqtControlStreamParser read_parser(&read_stream);
-  absl::StatusOr<MoqtRawControlMessage> message = read_parser.ReadNextMessage();
-  ASSERT_TRUE(message.ok());
-  EXPECT_EQ(message->type, MoqtMessageType::kRequestOk);
+  EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(kTrackName));
 }
 
 TEST_F(MoqtPublishSubscriberStreamTest, ReceivePublishDone) {
-  // First, establish subscription.
-  BidiStreamWithReset stream(0);
-  stream_visitor_->BindStream(&stream);
-
-  MoqtPublish publish;
-  publish.request_id = kRequestId;
-  publish.full_track_name = kTrackName;
-  publish.track_alias = kTrackAlias;
-
+  MoqtPublish publish = DefaultPublish();
   EXPECT_CALL(incoming_publish_callback_mock_, Call(kTrackName, _, _, _))
       .WillOnce(Return(&mock_subscribe_visitor_));
-  EXPECT_CALL(query_name_mock_, Call(kTrackName)).WillOnce(Return(nullptr));
-  EXPECT_CALL(register_name_mock_, Call(kTrackName, NotNull())).Times(1);
-  EXPECT_CALL(register_alias_mock_, Call(kTrackAlias, NotNull()))
-      .WillOnce(Return(true));
+  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));
           });
-  EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(kTrackName));
-  ExpectSubscriberDestruction();
-
-  stream.Receive(framer_.SerializePublish(publish).AsStringView());
-  stream_visitor_->OnCanRead();
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(publish));
 
   // Now receive PUBLISH_DONE.
   MoqtPublishDone publish_done;
   publish_done.request_id = kRequestId;
   publish_done.status_code = PublishDoneCode::kTrackEnded;
   publish_done.stream_count = 0;  // Trigger immediate Destroy
-
-  stream.Receive(framer_.SerializePublishDone(publish_done).AsStringView());
-  stream_visitor_->OnCanRead();
-
-  EXPECT_EQ(stream.last_reset_code(), kResetCodeCancelled);
+  EXPECT_CALL(mock_stream_, Writev)
+      .WillOnce([](absl::Span<quiche::QuicheMemSlice> data,
+                   const webtransport::StreamWriteOptions& options) {
+        EXPECT_TRUE(data.empty());
+        EXPECT_TRUE(options.send_fin());
+        return absl::OkStatus();
+      });
+  EXPECT_CALL(mock_remove_callback_, Call);
+  EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(kTrackName));
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(publish_done));
 }
 
 TEST_F(MoqtPublishSubscriberStreamTest, ReceivePublishAndRejectCallback) {
-  BidiStreamWithReset stream(0);
-  stream_visitor_->BindStream(&stream);
-
-  MoqtPublish publish;
-  publish.request_id = kRequestId;
-  publish.full_track_name = kTrackName;
-  publish.track_alias = kTrackAlias;
-
+  MoqtPublish publish = DefaultPublish();
   EXPECT_CALL(incoming_publish_callback_mock_, Call(kTrackName, _, _, _))
       .WillOnce([this](const FullTrackName&, const MessageParameters&,
                        const TrackExtensions&, MoqtResponseCallback callback) {
         captured_response_callback_ = std::move(callback);
         return &mock_subscribe_visitor_;
       });
-
-  EXPECT_CALL(query_name_mock_, Call(kTrackName)).WillOnce(Return(nullptr));
-  EXPECT_CALL(register_name_mock_, Call(kTrackName, NotNull())).Times(1);
-  EXPECT_CALL(register_alias_mock_, Call(kTrackAlias, NotNull()))
-      .WillOnce(Return(true));
+  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));
           });
-  EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(kTrackName));
-  ExpectSubscriberDestruction();
-
-  stream.Receive(framer_.SerializePublish(publish).AsStringView());
-  stream_visitor_->OnCanRead();
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(publish));
 
   // Now call the response callback with error (reject).
-  stream.write_buffer().clear();
+  EXPECT_CALL(mock_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _))
+      .WillOnce(Return(absl::OkStatus()));
   MoqtRequestErrorInfo error_info{RequestErrorCode::kUninterested, std::nullopt,
                                   "rejected by app"};
+  EXPECT_CALL(mock_remove_callback_, Call);
+  EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(kTrackName));
   std::move(captured_response_callback_)(error_info);
-
-  // Verify REQUEST_ERROR response was sent.
-  std::string& written = stream.write_buffer();
-  webtransport::test::InMemoryStream read_stream(0);
-  read_stream.Receive(written);
-  MoqtControlStreamParser read_parser(&read_stream);
-  absl::StatusOr<MoqtRawControlMessage> message = read_parser.ReadNextMessage();
-  ASSERT_TRUE(message.ok());
-  EXPECT_EQ(message->type, MoqtMessageType::kRequestError);
-
-  MoqtControlMessageParser cmp(kDefaultMoqtVersion, true,
-                               quic::Perspective::IS_SERVER);
-  absl::StatusOr<MoqtRequestError> request_error =
-      cmp.ProcessRequestError(message->payload);
-  ASSERT_TRUE(request_error.ok());
-  EXPECT_EQ(request_error->request_id, kRequestId);
-  EXPECT_EQ(request_error->error_code, RequestErrorCode::kUninterested);
-  EXPECT_EQ(request_error->reason_phrase, "rejected by app");
 }
 
 TEST_F(MoqtPublishSubscriberStreamTest, ReceivePublishDoneOnRejectedStream) {
-  BidiStreamWithReset stream(0);
-  stream_visitor_->BindStream(&stream);
-
-  MoqtPublish publish;
-  publish.request_id = kRequestId;
-  publish.full_track_name = kTrackName;
-  publish.track_alias = kTrackAlias;
-
+  MoqtPublish publish = DefaultPublish();
   // Callback returns nullptr (rejection).
+  EXPECT_CALL(mock_add_callback_, Call(NotNull())).WillOnce(Return(true));
   EXPECT_CALL(incoming_publish_callback_mock_, Call(kTrackName, _, _, _))
       .WillOnce(Return(nullptr));
-
-  stream.Receive(framer_.SerializePublish(publish).AsStringView());
-  stream_visitor_->OnCanRead();
+  EXPECT_CALL(mock_stream_,
+              Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _))
+      .WillOnce(Return(absl::OkStatus()));
+  EXPECT_CALL(mock_remove_callback_, Call);
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(publish));
 
   // Now receive PUBLISH_DONE.
   MoqtPublishDone publish_done;
   publish_done.request_id = kRequestId;
   publish_done.status_code = PublishDoneCode::kTrackEnded;
   publish_done.stream_count = 0;
+  QUICHE_EXPECT_OK(stream_->OnControlMessage(publish_done));
+}
 
-  EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(kTrackName)).Times(0);
-  stream.Receive(framer_.SerializePublishDone(publish_done).AsStringView());
-  stream_visitor_->OnCanRead();
+TEST_F(MoqtPublishSubscriberStreamTest, DuplicatePublishOnSameStream) {
+  MoqtPublish publish = DefaultPublish();
+  EXPECT_CALL(mock_add_callback_, Call(NotNull())).WillOnce(Return(true));
+  EXPECT_CALL(incoming_publish_callback_mock_, Call(kTrackName, _, _, _))
+      .WillOnce(Return(&mock_subscribe_visitor_));
+  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));
+
+  // Receive second PUBLISH on same stream. Should cause a session error.
+  EXPECT_CALL(error_callback_, Call(MoqtError::kProtocolViolation,
+                                    "Multiple PUBLISH on the same stream"));
+  MoqtPublish publish2 = DefaultPublish();
+  publish2.request_id = kRequestId + 2;
+  publish2.full_track_name = FullTrackName("dead", "beef");
+  publish2.track_alias = kTrackAlias + 1;
+  stream_->CheckStatus(stream_->OnControlMessage(publish2));
+  EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(kTrackName));
+}
+
+TEST_F(MoqtPublishSubscriberStreamTest, DuplicatePublishOnDifferentStreams) {
+  MoqtPublish publish1 = DefaultPublish();
+  EXPECT_CALL(mock_add_callback_, Call(NotNull())).WillOnce(Return(true));
+  EXPECT_CALL(incoming_publish_callback_mock_, Call(kTrackName, _, _, _))
+      .WillOnce(Return(&mock_subscribe_visitor_));
+  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(publish1));
+
+  // Second stream
+  testing::MockFunction<void(MoqtError, absl::string_view)> error_callback2;
+  testing::MockFunction<bool(SubscribeRemoteTrack*)> mock_add_callback2;
+  testing::MockFunction<void(SubscribeRemoteTrack*)> mock_remove_callback2;
+  testing::MockFunction<SubscribeVisitor*(
+      const FullTrackName&, const MessageParameters&, const TrackExtensions&,
+      MoqtResponseCallback)>
+      incoming_publish_callback_mock2;
+  MoqtIncomingPublishCallback incoming_publish_callback2 =
+      incoming_publish_callback_mock2.AsStdFunction();
+
+  webtransport::test::MockStream mock_stream2;
+  MoqtPublishSubscriberStream stream2(
+      &framer_, message_parser_, &mock_clock_, &mock_alarm_factory_,
+      error_callback2.AsStdFunction(), &incoming_publish_callback2,
+      mock_add_callback2.AsStdFunction(),
+      mock_remove_callback2.AsStdFunction());
+  stream2.BindStream(&mock_stream2);
+  EXPECT_CALL(mock_stream2, CanWrite).WillRepeatedly(Return(true));
+
+  // Duplicate PUBLISH (same track name) on stream2 is rejected.
+  MoqtPublish publish2 = DefaultPublish();
+  publish2.request_id = kRequestId + 2;
+  publish2.full_track_name = kTrackName;
+  EXPECT_CALL(mock_add_callback2, Call(NotNull())).WillOnce(Return(false));
+  EXPECT_CALL(mock_stream2,
+              Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _))
+      .WillOnce(Return(absl::OkStatus()));
+  EXPECT_CALL(mock_remove_callback2, Call);
+
+  QUICHE_EXPECT_OK(stream2.OnControlMessage(publish2));
+
+  // Finally, a PUBLISH on the second stream triggers a session error.
+  EXPECT_CALL(error_callback2, Call(MoqtError::kProtocolViolation,
+                                    "Multiple PUBLISH on the same stream"));
+  MoqtPublish publish3 = DefaultPublish();
+  publish3.request_id = kRequestId + 4;
+  publish3.full_track_name = kTrackName;
+  stream2.CheckStatus(stream2.OnControlMessage(publish3));
+
+  // Test teardown expectations
+  EXPECT_CALL(mock_subscribe_visitor_, OnPublishDone(kTrackName));
 }
 
 }  // namespace
diff --git a/quiche/quic/moqt/moqt_session.cc b/quiche/quic/moqt/moqt_session.cc
index e9c62b6..f7f8333 100644
--- a/quiche/quic/moqt/moqt_session.cc
+++ b/quiche/quic/moqt/moqt_session.cc
@@ -296,18 +296,18 @@
   std::unique_ptr<MoqtNamespaceSubscriberStream> state =
       std::make_unique<MoqtNamespaceSubscriberStream>(
           &framer_, ControlMessageParser(), next_request_id_,
-          [session_weak_ptr = GetWeakPtr(), this, pref = prefix]() {
-            if (!session_weak_ptr.IsValid() || is_closing_) {
-              return;
+          [weakptr = GetWeakPtr()](const TrackNamespace& prefix) {
+            MoqtSession* session = MoqtSessionFromWeakPtr(weakptr);
+            if (session != nullptr) {
+              session->outgoing_subscribe_namespace_.UnsubscribeNamespace(
+                  prefix);
             }
-            outgoing_subscribe_namespace_.UnsubscribeNamespace(pref);
           },
-          [session_weak_ptr = GetWeakPtr(), this](MoqtError error,
-                                                  absl::string_view reason) {
-            if (!session_weak_ptr.IsValid() || is_closing_) {
-              return;
+          [weakptr = GetWeakPtr()](MoqtError error, absl::string_view reason) {
+            MoqtSession* session = MoqtSessionFromWeakPtr(weakptr);
+            if (session != nullptr) {
+              session->Error(error, reason);
             }
-            Error(error, reason);
           },
           std::move(response_callback));
   MoqtNamespaceSubscriberStream* state_ptr = state.get();
@@ -489,8 +489,26 @@
                   << message.full_track_name;
   auto track = std::make_unique<SubscribeRemoteTrack>(
       message, visitor,
-      [this, id = message.request_id]() { upstream_by_id_.erase(id); },
-      GetSubscribeCallbacks());
+      [weakptr = GetWeakPtr()](SubscribeRemoteTrack* track) {
+        MoqtSession* session = MoqtSessionFromWeakPtr(weakptr);
+        if (session == nullptr || !track->track_alias().has_value()) {
+          return false;
+        }
+        auto [it, success] = session->subscribe_by_alias_.try_emplace(
+            *track->track_alias(), track);
+        return success;
+      },
+      [weakptr = GetWeakPtr()](SubscribeRemoteTrack* track) {
+        MoqtSession* session = MoqtSessionFromWeakPtr(weakptr);
+        if (session == nullptr) {
+          return;
+        }
+        session->subscribe_by_name_.erase(track->full_track_name());
+        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));
   return true;
@@ -565,14 +583,14 @@
   }
   auto stream_visitor = std::make_unique<MoqtPublishPublisherStream>(
       &framer_, ControlMessageParser(),
-      [weak_session = GetWeakPtr(), track_name = name]() {
-        // Stream deleted callback.
-        MoqtSession* session =
-            absl::down_cast<MoqtSession*>(weak_session.GetIfAvailable());
+      [weak_session = GetWeakPtr()](SubscriptionPublisher* publisher) {
+        MoqtSession* session = MoqtSessionFromWeakPtr(weak_session);
         if (session == nullptr) {
           return;
         }
-        session->subscribed_track_names_.erase(track_name);
+        session->subscribed_track_names_.erase(
+            publisher->publisher().GetTrackName());
+        session->published_subscriptions_.erase(publisher->request_id());
       },
       [weak_session = GetWeakPtr()](MoqtError code, absl::string_view reason) {
         MoqtSessionInterface* session = weak_session.GetIfAvailable();
@@ -801,35 +819,6 @@
   return it->second;
 }
 
-SubscribeRemoteTrack::SubscribeCallbacks MoqtSession::GetSubscribeCallbacks() {
-  return {
-      [this](const FullTrackName& name) {  // query_name_
-        return RemoteTrackByName(name);
-      },
-      [this](const FullTrackName& name, SubscribeRemoteTrack* track) {
-        // register_name_
-        subscribe_by_name_[name] = track;
-      },
-      [this](const uint64_t alias, SubscribeRemoteTrack* track) {
-        // register_alias_
-        return subscribe_by_alias_.try_emplace(alias, track).second;
-      },
-      [weaksession = GetWeakPtr()](const FullTrackName& name,
-                                   std::optional<uint64_t> alias) {
-        // unregister_
-        MoqtSession* session =
-            absl::down_cast<MoqtSession*>(weaksession.GetIfAvailable());
-        if (session == nullptr) {
-          return;
-        }
-        session->subscribe_by_name_.erase(name);
-        if (alias.has_value()) {
-          session->subscribe_by_alias_.erase(*alias);
-        }
-      },
-  };
-}
-
 void MoqtSession::OnCanCreateNewOutgoingUnidirectionalStream() {
   while (!subscriptions_with_queued_streams_.empty() &&
          session_->CanOpenNextOutgoingUnidirectionalStream()) {
@@ -925,10 +914,28 @@
     case MoqtMessageType::kSubscribeNamespace: {
       auto namespace_stream = std::make_unique<MoqtNamespacePublisherStream>(
           &session_->framer_, session_->ControlMessageParser(),
-          [session = session_](MoqtError code, absl::string_view reason) {
-            session->Error(code, reason);
+          [weakptr = session_->GetWeakPtr()](const TrackNamespace& prefix) {
+            MoqtSession* session = MoqtSessionFromWeakPtr(weakptr);
+            if (session != nullptr) {
+              return session->incoming_subscribe_namespace_.SubscribeNamespace(
+                  prefix);
+            }
+            return true;
           },
-          &session_->incoming_subscribe_namespace_,
+          [weakptr = session_->GetWeakPtr()](const TrackNamespace& prefix) {
+            MoqtSession* session = MoqtSessionFromWeakPtr(weakptr);
+            if (session != nullptr) {
+              session->incoming_subscribe_namespace_.UnsubscribeNamespace(
+                  prefix);
+            }
+          },
+          [weakptr = session_->GetWeakPtr()](MoqtError code,
+                                             absl::string_view reason) {
+            MoqtSession* session = MoqtSessionFromWeakPtr(weakptr);
+            if (session != nullptr) {
+              session->Error(code, reason);
+            }
+          },
           session_->callbacks_.incoming_subscribe_namespace_callback);
       namespace_stream->BindStream(std::move(parser_));
       MoqtNamespacePublisherStream* temp_stream = namespace_stream.get();
@@ -942,11 +949,54 @@
       auto publish_stream = std::make_unique<MoqtPublishSubscriberStream>(
           &session_->framer_, session_->ControlMessageParser(),
           session_->callbacks_.clock, session_->alarm_factory(),
-          [session = session_](MoqtError code, absl::string_view reason) {
-            session->Error(code, reason);
+          [weakptr = session_->GetWeakPtr()](MoqtError code,
+                                             absl::string_view reason) {
+            MoqtSession* session = MoqtSessionFromWeakPtr(weakptr);
+            if (session != nullptr) {
+              session->Error(code, reason);
+            }
           },
           &session_->callbacks_.incoming_publish_callback,
-          session_->GetSubscribeCallbacks());
+          [weakptr = session_->GetWeakPtr()](SubscribeRemoteTrack* track) {
+            MoqtSession* session = MoqtSessionFromWeakPtr(weakptr);
+            if (session == nullptr) {
+              return false;
+            }
+            QUICHE_BUG_IF(quiche_bug_publish_no_track_alias,
+                          !track->track_alias().has_value())
+                << "PUBLISH with no track alias";
+            if (!track->track_alias().has_value()) {
+              return false;
+            }
+            auto [alias_it, alias_inserted] =
+                session->subscribe_by_alias_.try_emplace(*track->track_alias(),
+                                                         track);
+            if (!alias_inserted) {
+              // Already a PUBLISH or an established SUBSCRIBE.
+              return false;
+            }
+            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.
+              track->set_visitor(it->second->ReleaseVisitor());
+              session->Unsubscribe(it->second->full_track_name());
+            }
+            auto [name_it, name_inserted] =
+                session->subscribe_by_name_.try_emplace(
+                    track->full_track_name(), track);
+            QUICHE_DCHECK(name_inserted);
+            return true;
+          },
+          [weakptr = session_->GetWeakPtr()](SubscribeRemoteTrack* track) {
+            MoqtSession* session = MoqtSessionFromWeakPtr(weakptr);
+            if (session != nullptr) {
+              session->subscribe_by_name_.erase(track->full_track_name());
+              if (track->track_alias().has_value()) {
+                session->subscribe_by_alias_.erase(*track->track_alias());
+              }
+            }
+          });
       publish_stream->BindStream(std::move(parser_));
       MoqtPublishSubscriberStream* temp_stream = publish_stream.get();
       stream_->SetVisitor(std::move(publish_stream));
diff --git a/quiche/quic/moqt/moqt_session.h b/quiche/quic/moqt/moqt_session.h
index b47935f..d0bcb21 100644
--- a/quiche/quic/moqt/moqt_session.h
+++ b/quiche/quic/moqt/moqt_session.h
@@ -11,6 +11,7 @@
 #include <string>
 #include <utility>
 
+#include "absl/base/casts.h"
 #include "absl/base/nullability.h"
 #include "absl/cleanup/cleanup.h"
 #include "absl/container/btree_map.h"
@@ -245,9 +246,6 @@
     explicit ControlStream(MoqtSession* session)
         : MoqtBidiStreamBase(
               &session->framer_, session->ControlMessageParser(),
-              // Do nothing on deletion. It threw an error on RESET_STREAM or
-              // FIN, and we're here because the session is being destroyed.
-              []() {},
               [session](MoqtError code, absl::string_view reason) {
                 session->control_stream_ =
                     quiche::QuicheWeakPtr<ControlStream>();
@@ -284,6 +282,9 @@
     quiche::QuicheWeakPtr<ControlStream> GetWeakPtr() {
       return weak_ptr_factory_.Create();
     }
+    void Detach() override {
+      session_->Error(MoqtError::kProtocolViolation, "Control stream closed");
+    }
 
    private:
     friend class test::MoqtSessionPeer;
@@ -401,8 +402,6 @@
   RemoteTrack* RemoteTrackById(uint64_t request_id);
   SubscribeRemoteTrack* RemoteTrackByName(const FullTrackName& name);
 
-  SubscribeRemoteTrack::SubscribeCallbacks GetSubscribeCallbacks();
-
   // Checks that a subscribe ID from a SUBSCRIBE or FETCH is valid, and throws
   // a session error if is not.
   bool ValidateRequestId(uint64_t request_id);
@@ -595,6 +594,11 @@
   std::shared_ptr<Empty> liveness_token_;
 };
 
+static MoqtSession* absl_nullable MoqtSessionFromWeakPtr(
+    const quiche::QuicheWeakPtr<MoqtSessionInterface>& weak_ptr) {
+  return absl::down_cast<MoqtSession*>(weak_ptr.GetIfAvailable());
+}
+
 }  // namespace moqt
 
 #endif  // QUICHE_QUIC_MOQT_MOQT_SESSION_H_
diff --git a/quiche/quic/moqt/moqt_session_test.cc b/quiche/quic/moqt/moqt_session_test.cc
index 79f9e73..c327f5e 100644
--- a/quiche/quic/moqt/moqt_session_test.cc
+++ b/quiche/quic/moqt/moqt_session_test.cc
@@ -2355,7 +2355,7 @@
       &framer,
       MoqtControlMessageParser(kDefaultMoqtVersion, true,
                                quic::Perspective::IS_CLIENT),
-      nullptr, &tree, callback);
+      [](const TrackNamespace&) { return true; }, nullptr, nullptr, callback);
   namespace_stream.BindStream(&mock_stream_);
   EXPECT_CALL(mock_stream_,
               Writev(ControlMessageOfType(MoqtMessageType::kRequestError), _));
@@ -2967,6 +2967,9 @@
                       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);
   EXPECT_FALSE(incoming_publish_callback_called);
diff --git a/quiche/quic/moqt/moqt_subscription.cc b/quiche/quic/moqt/moqt_subscription.cc
index b3d5416..9bf4513 100644
--- a/quiche/quic/moqt/moqt_subscription.cc
+++ b/quiche/quic/moqt/moqt_subscription.cc
@@ -142,8 +142,8 @@
 }
 
 void SubscriptionPublisher::OnSubscribeRejected(MoqtRequestErrorInfo info) {
-  bidi_stream_->CheckStatus(bidi_stream_->SendRequestError(request_id_, info,
-                                                           /*fin=*/true));
+  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_);
   }
diff --git a/quiche/quic/moqt/moqt_subscription.h b/quiche/quic/moqt/moqt_subscription.h
index b058beb..d32ed4c 100644
--- a/quiche/quic/moqt/moqt_subscription.h
+++ b/quiche/quic/moqt/moqt_subscription.h
@@ -28,6 +28,7 @@
 #include "quiche/quic/moqt/moqt_types.h"
 #include "quiche/quic/moqt/moqt_uni_stream.h"
 #include "quiche/common/platform/api/quiche_export.h"
+#include "quiche/common/quiche_callbacks.h"
 #include "quiche/common/quiche_weak_ptr.h"
 #include "quiche/web_transport/web_transport.h"
 
@@ -94,6 +95,11 @@
 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.
+  using RemoveCallback =
+      quiche::SingleUseCallback<void(SubscriptionPublisher*)>;
   SubscriptionPublisher(MoqtFramer framer,
                         std::shared_ptr<MoqtTrackPublisher> track_publisher,
                         MoqtBidiStreamBase* absl_nonnull bidi_stream,
diff --git a/quiche/quic/moqt/moqt_subscription_test.cc b/quiche/quic/moqt/moqt_subscription_test.cc
index 791c05f..154725b 100644
--- a/quiche/quic/moqt/moqt_subscription_test.cc
+++ b/quiche/quic/moqt/moqt_subscription_test.cc
@@ -87,10 +87,8 @@
  public:
   TestMoqtBidiStream(MoqtFramer* absl_nonnull framer,
                      const MoqtControlMessageParser& message_parser,
-                     BidiStreamDeletedCallback stream_deleted_callback,
                      SessionErrorCallback session_error_callback)
       : MoqtBidiStreamBase(framer, message_parser,
-                           std::move(stream_deleted_callback),
                            std::move(session_error_callback)) {
     set_control_stream();  // TODO(martinduke): Delete
   }
@@ -100,6 +98,8 @@
       const MoqtRawControlMessage& message) override {
     return absl::OkStatus();
   }
+  void Detach() override { detached_ = true; }
+  bool detached_ = false;
 };
 
 std::optional<PublishedObject> DefaultPublishedObject(
@@ -124,9 +124,8 @@
   SubscriptionPublisherTest()
       : track_publisher_(
             std::make_shared<MockTrackPublisher>(FullTrackName("foo", "bar"))),
-        bidi_stream_(
-            &framer_, message_parser_, [] {},
-            [](MoqtError, absl::string_view) {}),
+        bidi_stream_(&framer_, message_parser_,
+                     [](MoqtError, absl::string_view) {}),
         trace_recorder_(nullptr) {
     bidi_stream_.BindStream(&mock_bidi_stream_);
     parameters_.set_forward(true);
diff --git a/quiche/quic/moqt/moqt_track.cc b/quiche/quic/moqt/moqt_track.cc
index 4530cd9..36e06cf 100644
--- a/quiche/quic/moqt/moqt_track.cc
+++ b/quiche/quic/moqt/moqt_track.cc
@@ -17,7 +17,6 @@
 #include "quiche/quic/core/quic_alarm_factory.h"
 #include "quiche/quic/core/quic_clock.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"
@@ -40,24 +39,18 @@
 
 }  // namespace
 
-void RemoteTrack::Destroy() {
-  if (delete_callback_ == nullptr) {
-    return;
-  }
-  BidiStreamDeletedCallback delete_callback = std::move(delete_callback_);
-  delete_callback_ = nullptr;
-  std::move(delete_callback)();
-}
-
 SubscribeRemoteTrack::~SubscribeRemoteTrack() {
   if (publish_done_alarm_ != nullptr) {
     publish_done_alarm_->PermanentCancel();
   }
-  if (callbacks_.unregister) {
-    std::move(callbacks_.unregister)(full_track_name(), track_alias_);
+  if (remove_callback_ != nullptr) {
+    RemoveCallback callback = std::move(remove_callback_);
+    remove_callback_ = nullptr;
+    std::move(callback)(this);
   }
   if (visitor_ != nullptr) {
     visitor_->OnPublishDone(full_track_name());
+    visitor_ = nullptr;
   }
 }
 
@@ -191,6 +184,10 @@
     // If this has already been called, UpstreamFetchTask will ignore it.
     task->OnStreamAndFetchClosed(kResetCodeCancelled, "");
   }
+  if (remove_callback_ != nullptr) {
+    std::move(remove_callback_)();
+    remove_callback_ = nullptr;
+  }
 }
 
 void UpstreamFetch::OnFetchResult(Location largest_location,
diff --git a/quiche/quic/moqt/moqt_track.h b/quiche/quic/moqt/moqt_track.h
index 1d85bbb..54716d1 100644
--- a/quiche/quic/moqt/moqt_track.h
+++ b/quiche/quic/moqt/moqt_track.h
@@ -11,7 +11,6 @@
 #include <memory>
 #include <optional>
 #include <utility>
-#include <variant>
 
 #include "absl/status/status.h"
 #include "absl/strings/string_view.h"
@@ -46,13 +45,13 @@
  public:
   RemoteTrack(const FullTrackName& full_track_name, uint64_t id,
               const MessageParameters& parameters,
-              BidiStreamDeletedCallback callback)
+              MoqtBidiStreamBase* request_stream)
       : full_track_name_(full_track_name),
         request_id_(id),
+        request_stream_(request_stream),
         parameters_(parameters),
-        delete_callback_(std::move(callback)),
         weak_ptr_factory_(this) {}
-  virtual ~RemoteTrack() { Destroy(); }
+  virtual ~RemoteTrack() {}
 
   const FullTrackName& full_track_name() const { return full_track_name_; }
   // If REQUEST_ERROR arrives after OK or an object, it is a protocol violation.
@@ -70,13 +69,15 @@
 
   virtual bool is_fetch() const = 0;
 
-  void Destroy();
+  virtual void Destroy() = 0;
 
   // A REQUEST_UPDATE changes any field that is present in |parameters|.
   void Update(const MessageParameters& parameters) {
     parameters_.Update(parameters);
   }
 
+  MoqtBidiStreamBase* request_stream() { return request_stream_; }
+
  protected:
   const MessageParameters& const_parameters() const { return parameters_; }
   MessageParameters& parameters() { return parameters_; }
@@ -84,11 +85,11 @@
  private:
   const FullTrackName full_track_name_;
   const uint64_t request_id_;
+  MoqtBidiStreamBase* request_stream_;
   MessageParameters parameters_;
   // If false, an object or OK message has been received, so any ERROR message
   // is a protocol violation.
   bool error_is_allowed_ = true;
-  BidiStreamDeletedCallback delete_callback_;
 
   // Must be last.
   quiche::QuicheWeakPtrFactory<RemoteTrack> weak_ptr_factory_;
@@ -97,63 +98,25 @@
 // A track on the peer to which the session has subscribed.
 class SubscribeRemoteTrack : public RemoteTrack {
  public:
-  struct SubscribeCallbacks {
-    quiche::SingleUseCallback<SubscribeRemoteTrack*(const FullTrackName&)>
-        query_name;
-    quiche::SingleUseCallback<void(const FullTrackName&, SubscribeRemoteTrack*)>
-        register_name;
-    quiche::SingleUseCallback<bool(uint64_t, SubscribeRemoteTrack*)>
-        register_alias;
-    quiche::SingleUseCallback<void(const FullTrackName&,
-                                   std::optional<uint64_t>)>
-        unregister;
-  };
-  // Tells the session about changes to a track's subscription status.
-  // If SubscribeRemoteTrack* is null, the subscription is gone and the callback
-  // will always return true.
-  // If non-null, try to add the track. If the name is new, return true. If
-  // ready present but for a pending SUBSCRIBE, return the visitor for that
-  // SUBSCRIBE. Otherwise, return false.
-  using RegisterNameCallback =
-      quiche::MultiUseCallback<std::variant<bool, SubscribeVisitor*>(
-          const FullTrackName&, SubscribeRemoteTrack*)>;
-  // When SubscribeRemoteTrack* is non-null, this callback informs the session
-  // of the track alias after the receipt of SUBSCRIBE_OK or PUBLISH.
-  //
-  // If the second argument is null, it means the subscription to the track
-  // alias has ended, and always returns absl::OkStatus().
-  //
-  // Returns true if the operation was successful, It can only fail on
-  // registration because there is already a track with that alias.
-  using RegisterTrackAliasCallback =
-      quiche::MultiUseCallback<bool(uint64_t, SubscribeRemoteTrack*)>;
+  // Returns the existing subscription, if present.
+  using AddCallback = quiche::SingleUseCallback<bool(SubscribeRemoteTrack*)>;
+  using RemoveCallback = quiche::SingleUseCallback<void(SubscribeRemoteTrack*)>;
   SubscribeRemoteTrack(const MoqtSubscribe& subscribe,
-                       SubscribeVisitor* visitor,
-                       BidiStreamDeletedCallback callback,
-                       SubscribeCallbacks callbacks)
+                       SubscribeVisitor* visitor, AddCallback add_callback,
+                       RemoveCallback remove_callback)
       : RemoteTrack(subscribe.full_track_name, subscribe.request_id,
-                    subscribe.parameters, std::move(callback)),
+                    subscribe.parameters, /*request_stream=*/nullptr),
         visitor_(visitor),
-        callbacks_(std::move(callbacks)) {
-    if (callbacks_.register_name) {
-      std::move(callbacks_.register_name)(full_track_name(), this);
-    }
-  }
-
+        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.
   SubscribeRemoteTrack(const MoqtPublish& publish, SubscribeVisitor* visitor,
-                       BidiStreamDeletedCallback callback,
-                       SubscribeCallbacks callbacks)
+                       MoqtBidiStreamBase* request_stream)
       : RemoteTrack(publish.full_track_name, publish.request_id,
-                    publish.parameters, std::move(callback)),
-        visitor_(visitor),
-        callbacks_(std::move(callbacks)) {
-    OnObjectOrOk();
-    visitor_->OnReply(publish.full_track_name,
-                      SubscribeOkData(publish.parameters, publish.extensions));
-    if (callbacks_.register_name) {
-      std::move(callbacks_.register_name)(full_track_name(), this);
-      callbacks_.register_name = nullptr;
-    }
+                    publish.parameters, request_stream),
+        visitor_(visitor) {
+    track_alias_.emplace(publish.track_alias);
   }
   ~SubscribeRemoteTrack() override;
 
@@ -164,10 +127,10 @@
   std::optional<uint64_t> track_alias() const { return track_alias_; }
   // Returns false if the callback returns false, meaning the session has been
   // destroyed.
-  [[nodiscard]] bool set_track_alias(uint64_t track_alias) {
+  bool set_track_alias(uint64_t track_alias) {
     track_alias_.emplace(track_alias);
-    if (callbacks_.register_alias) {
-      return std::move(callbacks_.register_alias)(track_alias, this);
+    if (add_callback_ != nullptr) {
+      return std::move(add_callback_)(this);
     }
     return true;
   }
@@ -204,6 +167,21 @@
     visitor_ = nullptr;
     return temp;
   }
+  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);
+    }
+  }
 
  private:
   friend class test::MoqtSessionPeer;
@@ -238,7 +216,9 @@
   int currently_open_streams_ = 0;
   // Every stream that has received FIN or RESET_STREAM.
   uint64_t streams_closed_ = 0;
-  SubscribeCallbacks callbacks_;
+  // 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_;
@@ -257,47 +237,51 @@
 
 // Class for upstream FETCH. It will notify the application using |callback|
 // when a FETCH_OK or REQUEST_ERROR is received.
+using RemoveFetchCallback = quiche::SingleUseCallback<void()>;
 class UpstreamFetch : public RemoteTrack {
  public:
   // Standalone Fetch constructor
   UpstreamFetch(const MoqtFetch& fetch, const StandaloneFetch standalone,
                 FetchResponseCallback callback,
-                BidiStreamDeletedCallback delete_callback)
+                RemoveFetchCallback delete_callback)
       : RemoteTrack(standalone.full_track_name, fetch.request_id,
-                    fetch.parameters, std::move(delete_callback)),
+                    fetch.parameters, /*request_stream=*/nullptr),
         group_order_(fetch.parameters.group_order.value_or(
             MoqtDeliveryOrder::kAscending)),
         start_(standalone.start_location),
         end_(standalone.end_location),
         subscriber_priority_(fetch.parameters.subscriber_priority.value_or(
             kDefaultSubscriberPriority)),
-        ok_callback_(std::move(callback)) {}
+        ok_callback_(std::move(callback)),
+        remove_callback_(std::move(delete_callback)) {}
   // Relative Joining Fetch constructor
   UpstreamFetch(const MoqtFetch& fetch, FullTrackName full_track_name,
                 FetchResponseCallback callback,
-                BidiStreamDeletedCallback delete_callback)
+                RemoveFetchCallback delete_callback)
       : RemoteTrack(full_track_name, fetch.request_id, fetch.parameters,
-                    std::move(delete_callback)),
+                    /*request_stream=*/nullptr),
         group_order_(fetch.parameters.group_order.value_or(
             MoqtDeliveryOrder::kAscending)),
         relative_groups_(
             std::get<JoiningFetchRelative>(fetch.fetch).joining_start),
         subscriber_priority_(fetch.parameters.subscriber_priority.value_or(
             kDefaultSubscriberPriority)),
-        ok_callback_(std::move(callback)) {}
+        ok_callback_(std::move(callback)),
+        remove_callback_(std::move(delete_callback)) {}
   // Absolute Joining Fetch constructor
   UpstreamFetch(const MoqtFetch& fetch, FullTrackName full_track_name,
                 JoiningFetchAbsolute absolute_joining,
                 FetchResponseCallback callback,
-                BidiStreamDeletedCallback delete_callback)
+                RemoveFetchCallback delete_callback)
       : RemoteTrack(full_track_name, fetch.request_id, fetch.parameters,
-                    std::move(delete_callback)),
+                    /*request_stream=*/nullptr),
         group_order_(fetch.parameters.group_order.value_or(
             MoqtDeliveryOrder::kAscending)),
         start_(Location(absolute_joining.joining_start, 0)),
         subscriber_priority_(fetch.parameters.subscriber_priority.value_or(
             kDefaultSubscriberPriority)),
-        ok_callback_(std::move(callback)) {}
+        ok_callback_(std::move(callback)),
+        remove_callback_(std::move(delete_callback)) {}
   UpstreamFetch(const UpstreamFetch&) = delete;
   ~UpstreamFetch();
 
@@ -308,6 +292,14 @@
   // Called when the data stream is destroyed.
   void OnStreamClosed() { Destroy(); }
 
+  void Destroy() override {
+    if (remove_callback_) {
+      RemoveFetchCallback callback = std::move(remove_callback_);
+      remove_callback_ = nullptr;
+      std::move(callback)();
+    }
+  }
+
   class UpstreamFetchTask : public MoqtFetchTask {
    public:
     // If the UpstreamFetch is destroyed, it will call OnStreamAndFetchClosed
@@ -429,6 +421,8 @@
 
   // Initial values from Fetch() call.
   FetchResponseCallback ok_callback_;  // Will be destroyed on FETCH_OK.
+
+  RemoveFetchCallback remove_callback_;
 };
 
 }  // namespace moqt
diff --git a/quiche/quic/moqt/moqt_track_test.cc b/quiche/quic/moqt/moqt_track_test.cc
index 2e3b52f..010cda9 100644
--- a/quiche/quic/moqt/moqt_track_test.cc
+++ b/quiche/quic/moqt/moqt_track_test.cc
@@ -4,7 +4,6 @@
 
 #include "quiche/quic/moqt/moqt_track.h"
 
-#include <cstdint>
 #include <memory>
 #include <optional>
 #include <utility>
@@ -22,7 +21,6 @@
 #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/common/test_tools/quiche_test_utils.h"
 #include "quiche/web_transport/web_transport.h"
 
 namespace moqt {
@@ -56,19 +54,12 @@
  public:
   SubscribeRemoteTrackTest()
       : track_(
-            subscribe_, &visitor_, [this]() { deleted_ = true; },
-            SubscribeRemoteTrack::SubscribeCallbacks{
-                /*query_name_=*/nullptr,
-                /*register_name_=*/nullptr,
-                /*register_alias_=*/
-                [this](uint64_t, SubscribeRemoteTrack* track) {
-                  alias_registered_ = (track != nullptr);
-                  if (alias_registered_) {
-                    EXPECT_EQ(track, &track_);
-                  }
-                  return true;
-                },
-                /*unregister_=*/nullptr}) {}
+            subscribe_, &visitor_,
+            [&](SubscribeRemoteTrack*) {
+              alias_registered_ = true;
+              return true;
+            },
+            [this](SubscribeRemoteTrack*) { deleted_ = true; }) {}
 
   MockSubscribeRemoteTrackVisitor visitor_;
   MoqtSubscribe subscribe_ = {/*request_id=*/1, FullTrackName("foo", "bar"),
@@ -88,6 +79,7 @@
   EXPECT_FALSE(track_.is_fetch());
   EXPECT_TRUE(track_.set_track_alias(1));
   EXPECT_EQ(track_.track_alias(), 1);
+  EXPECT_TRUE(alias_registered_);
 }
 
 TEST_F(SubscribeRemoteTrackTest, AllowError) {
diff --git a/quiche/quic/moqt/moqt_uni_stream_test.cc b/quiche/quic/moqt/moqt_uni_stream_test.cc
index f56c5c7..9efcabd 100644
--- a/quiche/quic/moqt/moqt_uni_stream_test.cc
+++ b/quiche/quic/moqt/moqt_uni_stream_test.cc
@@ -492,17 +492,13 @@
     EXPECT_CALL(session_, deliver_partial_objects())
         .WillRepeatedly(Return(false));
     track_ = std::make_unique<SubscribeRemoteTrack>(
-        subscribe_message_, &visitor_, []() {},
-        SubscribeRemoteTrack::SubscribeCallbacks{
-            /*query_name_=*/nullptr,
-            /*register_name_=*/nullptr,
-            /*register_alias_=*/
-            [this](uint64_t alias, SubscribeRemoteTrack* track) {
-              alias_ = alias;
-              alias_track_ = track;
-              return true;
-            },
-            /*unregister_=*/nullptr});
+        subscribe_message_, &visitor_,
+        [this](SubscribeRemoteTrack* track) {
+          alias_track_ = track;
+          alias_ = track->track_alias().value();
+          return true;
+        },
+        nullptr);
     EXPECT_TRUE(track_->set_track_alias(2));
     CreateStream();
   }