| // 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/connection_task.h" |
| |
| #include <cstddef> |
| #include <optional> |
| #include <utility> |
| |
| #include "pw_assert/check.h" |
| #include "pw_async2/await.h" |
| #include "pw_bytes/span.h" |
| #include "pw_log/log.h" |
| #include "pw_rpc2/internal/call.h" |
| #include "pw_rpc2/internal/packet.h" |
| |
| namespace pw::rpc2::internal { |
| namespace { |
| |
| /// Encodes an outbound packet's header into a transport reservation and commits |
| /// it. |
| void EncodeAndCommitHeader(transport::WriteReservation& reservation, |
| const OutboundPacket& packet) { |
| auto encode_res = |
| packet.EncodeHeader(ByteSpan(reservation.data(), reservation.size())); |
| PW_DCHECK(encode_res.ok()); |
| static_cast<void>(reservation.Commit(*encode_res)); |
| } |
| |
| } // namespace |
| |
| async2::Poll<Status> SendControlPacketFuture::Pend(async2::Context& cx) { |
| PW_ASSERT(is_pendable()); |
| |
| PW_AWAIT(auto write_res, reserve_fut_, cx); |
| mark_complete(); |
| |
| if (!write_res.has_value()) { |
| return async2::Ready(Status::Unavailable()); |
| } |
| |
| EncodeAndCommitHeader(*write_res, packet_); |
| return async2::Ready(OkStatus()); |
| } |
| |
| ConnectionTask::ConnectionTask(transport::ReliableDatagramSocket connection, |
| Allocator& allocator, |
| EndpointRole role) |
| : connection_(std::move(connection)), |
| allocator_(allocator), |
| state_(State::kHandshaking), |
| role_(role), |
| pending_control_packets_(allocator) { |
| PW_DCHECK(connection_); |
| } |
| |
| ConnectionTask::ConnectionTask(EstablishedConnection established_connection, |
| Allocator& allocator, |
| EndpointRole role) |
| : connection_(std::move(established_connection.connection)), |
| allocator_(allocator), |
| state_(State::kActive), |
| role_(role), |
| handshake_info_(established_connection.info), |
| pending_control_packets_(allocator) { |
| PW_DCHECK(connection_); |
| } |
| |
| ConnectionTask::~ConnectionTask() { |
| Teardown(); |
| |
| while (!calls_.empty()) { |
| calls_.front().DetachFromConnection(); |
| } |
| } |
| |
| void ConnectionTask::QueueControlPacket(const OutboundPacket& packet) { |
| if (state_ == State::kClosed) { |
| return; |
| } |
| PW_DASSERT(connection_); |
| |
| // 1. Fast path: an immediate synchronous reservation, but only when nothing |
| // is already queued. Writing directly while the queue is non-empty would |
| // let this packet overtake one the dispatcher has not sent yet, and |
| // control packets are terminal for their calls, so their order is |
| // observable by the peer. |
| if (pending_control_packets_.empty() && !send_control_future_.is_pendable()) { |
| const size_t size = packet.payload_offset(); |
| std::optional<transport::WriteReservation> write_res = |
| connection_.TryReserveWrite(size); |
| if (write_res.has_value()) { |
| EncodeAndCommitHeader(*write_res, packet); |
| return; |
| } |
| } |
| |
| // 2. Slow path: queue the packet until the dispatcher can write it. |
| if (pending_control_packets_.try_push_back(packet)) { |
| Wake(); |
| return; |
| } |
| |
| // 3. The allocator is exhausted. This packet is terminal for its call, so |
| // silently dropping it would leave the peer waiting forever. Close the |
| // connection immediately, which completes every call on it and forces the |
| // peer to notice. |
| PW_LOG_ERROR( |
| "Out of memory queueing control packet (type=0x%02x) for call_id %u, " |
| "closing connection", |
| static_cast<unsigned>(packet.type().bits()), |
| static_cast<unsigned>(packet.call_id())); |
| CloseConnection(Status::ResourceExhausted()); |
| } |
| |
| void ConnectionTask::CloseConnection(Status status) { |
| if (state_ == State::kClosed) { |
| return; |
| } |
| PW_DASSERT(connection_); |
| set_state(State::kClosed); |
| close_status_ = status; |
| pending_control_packets_.reset(); |
| connection_.Close(); |
| |
| // Advance `it` before `Complete()` because completing a client call detaches |
| // it from `calls_`. |
| for (auto it = calls_.begin(); it != calls_.end();) { |
| Call& call = *it++; |
| if (!call.is_closed()) { |
| call.Complete(status); |
| } |
| } |
| |
| // Wake this task so that it observes the closed state and retires. This has |
| // to happen here rather than in the loop above: a connection with no open |
| // calls would otherwise never be woken, and one with many would be woken |
| // once per call. |
| Wake(); |
| } |
| |
| Call* ConnectionTask::FindCallById(uint32_t call_id) { |
| for (auto& call : calls_) { |
| if (call.call_id() == call_id) { |
| return &call; |
| } |
| } |
| return nullptr; |
| } |
| |
| void ConnectionTask::StoreWaker(async2::Context& cx) { |
| PW_DASSERT(!is_closed()); |
| PW_ASYNC_STORE_WAKER(cx, waker_, "waiting for connection activity"); |
| } |
| |
| bool ConnectionTask::PendPacket(async2::Context& cx, |
| InboundPacket* incoming_request) { |
| PW_DASSERT(state() == State::kActive); |
| |
| // Stage 1: Outgoing control-packet egress. |
| bool progressed = ProcessOutgoingControlPackets(cx); |
| if (is_closed()) { |
| return true; |
| } |
| |
| // Stage 2: Dispatch stashed ingress packet. |
| progressed |= DispatchPendingIngressPacket(cx); |
| |
| // Stage 3: Read and dispatch incoming packet from transport. |
| progressed |= ReadPacketFromConnection(cx, incoming_request); |
| return progressed; |
| } |
| |
| void ConnectionTask::FinishHandshake(EstablishedConnection&& established) { |
| PW_DASSERT(state_ == State::kHandshaking); |
| PW_DASSERT(established.connection); |
| // The only reassignment of `connection_` after construction. |
| connection_ = std::move(established.connection); |
| handshake_info_ = established.info; |
| set_state(State::kActive); |
| } |
| |
| bool ConnectionTask::ProcessOutgoingControlPackets(async2::Context& cx) { |
| PW_DASSERT(state() == State::kActive); |
| PW_DASSERT(connection_); |
| |
| if (!send_control_future_.is_pendable()) { |
| if (pending_control_packets_.empty()) { |
| return false; |
| } |
| OutboundPacket packet = pending_control_packets_.front(); |
| pending_control_packets_.pop_front(); |
| if (pending_control_packets_.empty() && |
| pending_control_packets_.capacity() > kMaxIdleControlPacketCapacity) { |
| pending_control_packets_.reset(); |
| } |
| send_control_future_ = SendControlPacketFuture(connection_, packet); |
| } |
| |
| auto poll = send_control_future_.Pend(cx); |
| if (!poll.IsReady()) { |
| return false; |
| } |
| |
| const Status send_status = *poll; |
| send_control_future_ = SendControlPacketFuture(); |
| |
| if (!send_status.ok()) { |
| // The transport refuses a reservation only once it can no longer be |
| // written to. Every queued control packet is terminal for its call, so |
| // dropping them silently would strand the peer; close instead, which |
| // completes every call on this connection. |
| PW_LOG_WARN("Failed to send control packet (%s), closing connection", |
| send_status.str()); |
| CloseConnection(send_status); |
| } |
| return true; |
| } |
| |
| bool ConnectionTask::DispatchPendingIngressPacket(async2::Context& cx) { |
| if (pending_dispatch_packet_ == nullptr) { |
| return false; |
| } |
| if (Call* call = FindCallById(pending_dispatch_packet_.call_id()); |
| call != nullptr) { |
| if (!DeliverToCall(*call, pending_dispatch_packet_, cx)) { |
| return false; |
| } |
| } |
| pending_dispatch_packet_ = nullptr; |
| return true; |
| } |
| |
| bool ConnectionTask::ReadPacketFromConnection(async2::Context& cx, |
| InboundPacket* incoming_request) { |
| PW_DASSERT(state() == State::kActive); |
| PW_DASSERT(connection_); |
| |
| if (pending_dispatch_packet_ != nullptr) { |
| return false; |
| } |
| |
| if (!read_future_.is_pendable()) { |
| read_future_ = connection_.Read(); |
| } |
| |
| auto poll = read_future_.Pend(cx); |
| if (!poll.IsReady()) { |
| return false; |
| } |
| |
| auto& result = *poll; |
| if (result == nullptr) { |
| CloseConnection(Status::Cancelled()); |
| return true; |
| } |
| |
| ConstBuf pkt_buf = std::move(result); |
| read_future_ = transport::ReadFuture(); |
| |
| auto decode_result = InboundPacket::Decode(std::move(pkt_buf)); |
| if (!decode_result.ok()) { |
| if (decode_result.status().IsInvalidArgument()) { |
| // Unrecognized packet type. A newer peer may send types this build does |
| // not know about, so drop the packet and keep the connection up. |
| PW_LOG_WARN("Received unrecognized RPC packet type, dropping"); |
| return true; |
| } |
| // Malformed packet: the transport delivered a complete frame that is too |
| // short for the type it claims to be. The peer is not speaking the |
| // protocol, and we cannot attribute the packet to a call, so tearing the |
| // connection down is the only way to avoid stranding calls silently. |
| PW_LOG_ERROR("Received malformed RPC packet (%s), closing connection", |
| decode_result.status().str()); |
| CloseConnection(Status::DataLoss()); |
| return true; |
| } |
| |
| InboundPacket packet = std::move(decode_result.value()); |
| // Replies to packets that cannot be delivered tell the peer to stop sending |
| // for that call. A terminal packet never gets a reply: the peer has already |
| // forgotten the call, and replying to an error with an error could bounce |
| // between the endpoints forever. |
| if (!packet.type().is_for(role_)) { |
| PW_LOG_WARN( |
| "Received packet of type 0x%02x for wrong endpoint role (call_id=%u), " |
| "dropping", |
| static_cast<unsigned>(packet.type().bits()), |
| static_cast<unsigned>(packet.call_id())); |
| if (!packet.type().is_terminal()) { |
| if (role_ == EndpointRole::kServer) { |
| QueueError(packet.call_id(), ServerError::kReceivedPacketForClient); |
| } else { |
| QueueError(packet.call_id(), ClientError::kReceivedPacketForServer); |
| } |
| } |
| return true; |
| } |
| |
| if (Call* call = FindCallById(packet.call_id()); call != nullptr) { |
| if (packet.type().is_start()) { |
| // Only a server receives start packets. The client reused the ID of a |
| // call that is still running, so it has lost track of that call; end |
| // it, and tell the client that neither call will proceed. |
| PW_LOG_WARN( |
| "Received start packet for already active call_id %u, cancelling " |
| "the active call", |
| static_cast<unsigned>(packet.call_id())); |
| QueueError(packet.call_id(), ServerError::kCancelled); |
| call->Complete(Status::Cancelled()); |
| return true; |
| } |
| if (!DeliverToCall(*call, packet, cx)) { |
| pending_dispatch_packet_ = std::move(packet); |
| } |
| return true; |
| } |
| |
| // Nothing is registered for this call ID, so the packet either opens a new |
| // call or is stray. |
| if (!packet.type().is_start()) { |
| if (packet.type().is_terminal()) { |
| PW_LOG_DEBUG( |
| "Dropping terminal packet 0x%02x for closed call %u (likely closed " |
| "by both ends at once)", |
| static_cast<unsigned>(packet.type().bits()), |
| static_cast<unsigned>(packet.call_id())); |
| return true; |
| } |
| PW_LOG_WARN( |
| "Received stray packet of type 0x%02x (call_id=%u) for unknown or " |
| "closed call, cancelling it", |
| static_cast<unsigned>(packet.type().bits()), |
| static_cast<unsigned>(packet.call_id())); |
| QueueCancel(packet.call_id()); |
| return true; |
| } |
| |
| PW_DASSERT(incoming_request != nullptr); |
| *incoming_request = std::move(packet); |
| return true; |
| } |
| |
| bool ConnectionTask::DeliverToCall(Call& call, |
| InboundPacket& packet, |
| async2::Context& cx) { |
| const PacketType type = packet.type(); |
| PW_DASSERT(!type.is_start()); |
| const CloseMode close_mode = type.close_mode(); |
| |
| // A unary or client-streaming call is answered by exactly one packet, which |
| // carries the response and ends the RPC. Any other packet (except an error) |
| // means the server treats the method as server or bidirectional streaming, |
| // so nothing in it can be trusted to be the response. |
| const bool is_single_response = |
| type.has_payload() && close_mode == CloseMode::kOkTerminal; |
| if (call.expects_single_response() && !is_single_response && |
| close_mode != CloseMode::kErrorTerminal) { |
| PW_LOG_WARN( |
| "Received packet of type 0x%02x, which does not answer a " |
| "single-response call (call_id=%u); the client and server disagree " |
| "about the method's type", |
| static_cast<unsigned>(type.bits()), |
| static_cast<unsigned>(packet.call_id())); |
| // Tell the server why the call failed, unless it already ended the RPC. |
| // A server that has ended the RPC has forgotten the call, so the error |
| // would only arrive as a stray packet. |
| if (close_mode != CloseMode::kOkTerminal) { |
| QueueError(packet.call_id(), ClientError::kMethodTypeMismatch); |
| } |
| call.OnError(Status::FailedPrecondition()); |
| return true; |
| } |
| |
| // The message is delivered before any close in the same packet is applied, |
| // so the reader drains it before observing the end of the stream. If the |
| // slot is not available yet, nothing has been applied and the whole packet |
| // is retried later. |
| if (type.has_payload()) { |
| if (call.peer_ended_stream()) { |
| PW_LOG_WARN( |
| "Call %u received a message after the peer ended its stream, " |
| "dropping", |
| static_cast<unsigned>(packet.call_id())); |
| } else { |
| if (!call.ReserveMessageSlot(cx)) { |
| return false; |
| } |
| call.OnMessage(std::move(packet).TakePayload()); |
| } |
| } |
| |
| switch (close_mode) { |
| case CloseMode::kOpen: |
| break; |
| case CloseMode::kStreamEnd: |
| call.OnPeerStreamEnd(); |
| break; |
| case CloseMode::kOkTerminal: |
| call.Complete(OkStatus()); |
| break; |
| case CloseMode::kErrorTerminal: |
| call.OnError(type.is_server() ? ToStatus(packet.server_error()) |
| : ToStatus(packet.client_error())); |
| break; |
| } |
| return true; |
| } |
| |
| } // namespace pw::rpc2::internal |