| // Copyright 2026 The Pigweed Authors |
| // |
| // Licensed under the Apache License, Version 2.0 (the "License"); you may not |
| // use this file except in compliance with the License. You may obtain a copy of |
| // the License at |
| // |
| // https://www.apache.org/licenses/LICENSE-2.0 |
| // |
| // Unless required by applicable law or agreed to in writing, software |
| // distributed under the License is distributed on an "AS IS" BASIS, WITHOUT |
| // WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the |
| // License for the specific language governing permissions and limitations under |
| // the License. |
| |
| #include "pw_rpc2/internal/client_connection_task.h" |
| |
| #include <utility> |
| |
| #include "pw_log/log.h" |
| |
| namespace pw::rpc2::internal { |
| |
| ClientConnectionTask::ClientConnectionTask( |
| EstablishedConnection established_connection, Allocator& allocator) |
| : ConnectionTask( |
| std::move(established_connection), allocator, EndpointRole::kClient) { |
| } |
| |
| ClientConnectionTask::~ClientConnectionTask() { |
| // Tear down the connection and unregister from the dispatcher BEFORE calling |
| // `Finish()` so that any caller blocked in `CloseBlocking()` is only released |
| // after this task is no longer running on the dispatcher. |
| Teardown(); |
| Finish(); |
| } |
| |
| uint32_t ClientConnectionTask::NewCallId() { |
| if (next_call_id_ == kMaxCallId) { |
| PW_LOG_ERROR("Client call IDs exhausted; closing the connection"); |
| CloseConnection(Status::ResourceExhausted()); |
| return kMaxCallId; |
| } |
| return next_call_id_++; |
| } |
| |
| bool ClientConnectionTask::RequestClose() { |
| CloseState expected = CloseState::kOpen; |
| if (close_state_.compare_exchange_strong(expected, CloseState::kClosing)) { |
| // The caller holds a reference to this task, so it cannot be destroyed |
| // before the wake completes. |
| Wake(); |
| return true; |
| } |
| return expected != CloseState::kClosed; |
| } |
| |
| ControlFuture ClientConnectionTask::Close() { |
| if (!RequestClose()) { |
| return ControlFuture::Resolved(OkStatus()); |
| } |
| ControlFuture future = close_completion_.Get(); |
| // `Finish()` stores `kClosed` before resolving, so a future obtained after |
| // the resolve (which would never complete) is detected here. |
| if (close_state_.load() == CloseState::kClosed) { |
| return ControlFuture::Resolved(OkStatus()); |
| } |
| return future; |
| } |
| |
| void ClientConnectionTask::CloseBlocking() { |
| static_cast<void>(RequestClose()); |
| closed_.acquire(); |
| closed_.release(); // Pass the token on to the next waiter. |
| } |
| |
| void ClientConnectionTask::ReleaseUserHandle() { |
| if (user_handles_.fetch_sub(1, std::memory_order_acq_rel) != 1) { |
| return; |
| } |
| // The last handle is gone, so nothing can reach the connection to close it. |
| static_cast<void>(RequestClose()); |
| } |
| |
| bool ClientConnectionTask::is_closing_or_closed() const { |
| return close_state_.load() != CloseState::kOpen; |
| } |
| |
| void ClientConnectionTask::Finish() { |
| if (close_state_.exchange(CloseState::kClosed) == CloseState::kClosed) { |
| return; |
| } |
| close_completion_.Resolve(OkStatus()); |
| closed_.release(); |
| } |
| |
| async2::Poll<> ClientConnectionTask::DoPend(async2::Context& cx) { |
| // Checked every poll rather than latched, because the close may be requested |
| // by any thread at any point. |
| if (close_state_.load() == CloseState::kClosing) { |
| CloseConnection(Status::Cancelled()); |
| } |
| |
| const bool progressed = !is_closed() && PollConnection(cx); |
| if (is_closed()) { |
| Finish(); |
| return async2::Ready(); |
| } |
| if (progressed) { |
| cx.ReEnqueue(); |
| } |
| return async2::Pending(); |
| } |
| |
| bool ClientConnectionTask::PollConnection(async2::Context& cx) { |
| StoreWaker(cx); |
| |
| // As in `ServerConnectionTask::PollConnection()`, `progressed` is true after |
| // the loop only when all `kMaxPacketsPerPoll` iterations succeeded and the |
| // task must yield via `cx.ReEnqueue()`; an early break means `PendPacket()` |
| // has already registered wakers on the pending transport futures. |
| bool progressed = false; |
| for (int i = 0; i < kMaxPacketsPerPoll; ++i) { |
| progressed = PendPacket(cx); |
| if (!progressed || is_closed()) { |
| break; |
| } |
| } |
| return progressed; |
| } |
| |
| } // namespace pw::rpc2::internal |