diff --git a/src/core/BUILD b/src/core/BUILD index e965464e8fc..b4174ba9988 100644 --- a/src/core/BUILD +++ b/src/core/BUILD @@ -6562,6 +6562,7 @@ grpc_cc_library( ], deps = [ "activity", + "cancel_callback", "event_engine_common", "if", "map", diff --git a/src/core/ext/transport/chaotic_good/client/chaotic_good_connector.cc b/src/core/ext/transport/chaotic_good/client/chaotic_good_connector.cc index b6f80e6de46..6be26f47c43 100644 --- a/src/core/ext/transport/chaotic_good/client/chaotic_good_connector.cc +++ b/src/core/ext/transport/chaotic_good/client/chaotic_good_connector.cc @@ -105,16 +105,14 @@ auto ChaoticGoodConnector::DataEndpointReadSettingsFrame( auto ChaoticGoodConnector::DataEndpointWriteSettingsFrame( RefCountedPtr self) { - return [self]() { - // Serialize setting frame. - SettingsFrame frame; - // frame.header set connectiion_type: control - frame.headers = SettingsMetadata{SettingsMetadata::ConnectionType::kData, - self->connection_id_, kDataAlignmentBytes} - .ToMetadataBatch(GetContext()); - auto write_buffer = frame.Serialize(&self->hpack_compressor_); - return self->data_endpoint_.Write(std::move(write_buffer.control)); - }; + // Serialize setting frame. + SettingsFrame frame; + // frame.header set connectiion_type: control + frame.headers = SettingsMetadata{SettingsMetadata::ConnectionType::kData, + self->connection_id_, kDataAlignmentBytes} + .ToMetadataBatch(GetContext()); + auto write_buffer = frame.Serialize(&self->hpack_compressor_); + return self->data_endpoint_.Write(std::move(write_buffer.control)); } auto ChaoticGoodConnector::WaitForDataEndpointSetup( @@ -200,16 +198,14 @@ auto ChaoticGoodConnector::ControlEndpointReadSettingsFrame( auto ChaoticGoodConnector::ControlEndpointWriteSettingsFrame( RefCountedPtr self) { - return [self]() { - // Serialize setting frame. - SettingsFrame frame; - // frame.header set connectiion_type: control - frame.headers = SettingsMetadata{SettingsMetadata::ConnectionType::kControl, - absl::nullopt, absl::nullopt} - .ToMetadataBatch(GetContext()); - auto write_buffer = frame.Serialize(&self->hpack_compressor_); - return self->control_endpoint_.Write(std::move(write_buffer.control)); - }; + // Serialize setting frame. + SettingsFrame frame; + // frame.header set connectiion_type: control + frame.headers = SettingsMetadata{SettingsMetadata::ConnectionType::kControl, + absl::nullopt, absl::nullopt} + .ToMetadataBatch(GetContext()); + auto write_buffer = frame.Serialize(&self->hpack_compressor_); + return self->control_endpoint_.Write(std::move(write_buffer.control)); } void ChaoticGoodConnector::Connect(const Args& args, Result* result, diff --git a/src/core/ext/transport/chaotic_good/server/chaotic_good_server.cc b/src/core/ext/transport/chaotic_good/server/chaotic_good_server.cc index 472e166bf5f..4359d5648ad 100644 --- a/src/core/ext/transport/chaotic_good/server/chaotic_good_server.cc +++ b/src/core/ext/transport/chaotic_good/server/chaotic_good_server.cc @@ -331,37 +331,29 @@ auto ChaoticGoodServerListener::ActiveConnection::HandshakingState:: auto ChaoticGoodServerListener::ActiveConnection::HandshakingState:: ControlEndpointWriteSettingsFrame(RefCountedPtr self) { + self->connection_->NewConnectionID(); + SettingsFrame frame; + frame.headers = + SettingsMetadata{absl::nullopt, self->connection_->connection_id_, + absl::nullopt} + .ToMetadataBatch(GetContext()); + auto write_buffer = frame.Serialize(&self->connection_->hpack_compressor_); return TrySeq( - [self]() { - self->connection_->NewConnectionID(); - SettingsFrame frame; - frame.headers = - SettingsMetadata{absl::nullopt, self->connection_->connection_id_, - absl::nullopt} - .ToMetadataBatch(GetContext()); - auto write_buffer = - frame.Serialize(&self->connection_->hpack_compressor_); - return self->connection_->endpoint_.Write( - std::move(write_buffer.control)); - }, + self->connection_->endpoint_.Write(std::move(write_buffer.control)), WaitForDataEndpointSetup(self)); } auto ChaoticGoodServerListener::ActiveConnection::HandshakingState:: DataEndpointWriteSettingsFrame(RefCountedPtr self) { + // Send data endpoint setting frame + SettingsFrame frame; + frame.headers = + SettingsMetadata{absl::nullopt, self->connection_->connection_id_, + self->connection_->data_alignment_} + .ToMetadataBatch(GetContext()); + auto write_buffer = frame.Serialize(&self->connection_->hpack_compressor_); return TrySeq( - [self]() { - // Send data endpoint setting frame - SettingsFrame frame; - frame.headers = - SettingsMetadata{absl::nullopt, self->connection_->connection_id_, - self->connection_->data_alignment_} - .ToMetadataBatch(GetContext()); - auto write_buffer = - frame.Serialize(&self->connection_->hpack_compressor_); - return self->connection_->endpoint_.Write( - std::move(write_buffer.control)); - }, + self->connection_->endpoint_.Write(std::move(write_buffer.control)), [self]() mutable { MutexLock lock(&self->connection_->listener_->mu_); // Set endpoint to latch @@ -380,8 +372,10 @@ auto ChaoticGoodServerListener::ActiveConnection::HandshakingState:: auto ChaoticGoodServerListener::ActiveConnection::HandshakingState:: EndpointWriteSettingsFrame(RefCountedPtr self, bool is_control_endpoint) { - return If(is_control_endpoint, ControlEndpointWriteSettingsFrame(self), - DataEndpointWriteSettingsFrame(self)); + return If( + is_control_endpoint, + [&self] { return ControlEndpointWriteSettingsFrame(self); }, + [&self] { return DataEndpointWriteSettingsFrame(self); }); } void ChaoticGoodServerListener::ActiveConnection::HandshakingState:: diff --git a/src/core/lib/transport/promise_endpoint.h b/src/core/lib/transport/promise_endpoint.h index 12cd0969284..bcbf13faf8e 100644 --- a/src/core/lib/transport/promise_endpoint.h +++ b/src/core/lib/transport/promise_endpoint.h @@ -39,6 +39,7 @@ #include "src/core/lib/gprpp/sync.h" #include "src/core/lib/iomgr/exec_ctx.h" #include "src/core/lib/promise/activity.h" +#include "src/core/lib/promise/cancel_callback.h" #include "src/core/lib/promise/if.h" #include "src/core/lib/promise/map.h" #include "src/core/lib/promise/poll.h" @@ -69,8 +70,10 @@ class PromiseEndpoint { // `Write()` before the previous write finishes. Doing that results in // undefined behavior. auto Write(SliceBuffer data) { - // Assert previous write finishes. - GPR_ASSERT(!write_state_->complete.load(std::memory_order_relaxed)); + // Start write and assert previous write finishes. + auto prev = write_state_->state.exchange(WriteState::kWriting, + std::memory_order_relaxed); + GPR_ASSERT(prev == WriteState::kIdle); bool completed; if (data.Length() == 0) { completed = true; @@ -92,16 +95,31 @@ class PromiseEndpoint { if (completed) write_state_->waker = Waker(); } return If( - completed, []() { return []() { return absl::OkStatus(); }; }, + completed, + [this]() { + return [write_state = write_state_]() { + auto prev = write_state->state.exchange(WriteState::kIdle, + std::memory_order_relaxed); + GPR_ASSERT(prev == WriteState::kWriting); + return absl::OkStatus(); + }; + }, [this]() { return [write_state = write_state_]() -> Poll { - // If current write isn't finished return `Pending()`, else return - // write result. - if (!write_state->complete.load(std::memory_order_acquire)) { - return Pending(); + // If current write isn't finished return `Pending()`, else + // return write result. + WriteState::State expected = WriteState::kWritten; + if (write_state->state.compare_exchange_strong( + expected, WriteState::kIdle, std::memory_order_acquire, + std::memory_order_relaxed)) { + // State was Written, and we changed it to Idle. We can return + // the result. + return std::move(write_state->result); } - write_state->complete.store(false, std::memory_order_relaxed); - return std::move(write_state->result); + // State was not Written; since we're polling it must be + // Writing. Assert that and return Pending. + GPR_ASSERT(expected == WriteState::kWriting); + return Pending(); }; }); } @@ -228,7 +246,13 @@ class PromiseEndpoint { }; struct WriteState : public RefCounted { - std::atomic complete{false}; + enum State : uint8_t { + kIdle, // Not writing. + kWriting, // Write started, but not completed. + kWritten, // Write completed. + }; + + std::atomic state{kIdle}; // Write buffer used for `EventEngine::Endpoint::Write()` to ensure the // memory behind the buffer is not lost. grpc_event_engine::experimental::SliceBuffer buffer; @@ -239,7 +263,10 @@ class PromiseEndpoint { void Complete(absl::Status status) { result = std::move(status); auto w = std::move(waker); - complete.store(true, std::memory_order_release); + auto prev = state.exchange(kWritten, std::memory_order_release); + // Previous state should be Writing. If we got anything else we've entered + // the callback path twice. + GPR_ASSERT(prev == kWriting); w.Wakeup(); } };