| // 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/handshake.h" |
| |
| #include <algorithm> |
| #include <cstring> |
| #include <utility> |
| |
| #include "pw_assert/check.h" |
| #include "pw_async2/try.h" |
| #include "pw_bytes/endian.h" |
| #include "pw_rpc2/internal/packet.h" |
| |
| namespace pw::rpc2::internal { |
| |
| void HandshakeFutureBase::Cancel() { |
| io_fut_ = std::monostate{}; |
| if (connection_) { |
| connection_.Close(); |
| connection_ = transport::ReliableDatagramSocket{}; |
| } |
| } |
| |
| async2::Poll<Result<EstablishedConnection>> HandshakeFutureBase::Fail( |
| Status status) { |
| Cancel(); |
| return async2::Ready(Result<EstablishedConnection>(status)); |
| } |
| |
| async2::Poll<Status> HandshakeFutureBase::PendWritePacket( |
| async2::Context& cx, HandshakePacket::Type type) { |
| if (!std::holds_alternative<transport::ReserveWriteFuture>(io_fut_)) { |
| io_fut_.emplace<transport::ReserveWriteFuture>( |
| connection_.ReserveWrite(HandshakePacket::kWireSizeBytes)); |
| } |
| auto& write_fut = std::get<transport::ReserveWriteFuture>(io_fut_); |
| PW_TRY_READY_ASSIGN(auto poll_opt, write_fut.Pend(cx)); |
| io_fut_ = std::monostate{}; |
| if (!poll_opt.has_value()) { |
| return async2::Ready(Status::Cancelled()); |
| } |
| auto reservation = std::move(*poll_opt); |
| // Before the version is negotiated, `negotiated_version` holds the local |
| // maximum, which is what the initiator's SYN advertises. The SYN-ACK and ACK |
| // carry the negotiated version. |
| auto status = HandshakePacket(type, handshake_info_.negotiated_version) |
| .Encode(reservation); |
| if (!status.ok()) { |
| reservation.Cancel(); |
| return async2::Ready(status); |
| } |
| if (!reservation.Commit(HandshakePacket::kWireSizeBytes)) { |
| return async2::Ready(Status::Unavailable()); |
| } |
| return async2::Ready(OkStatus()); |
| } |
| |
| async2::Poll<Result<HandshakePacket>> HandshakeFutureBase::PendReadPacket( |
| async2::Context& cx, HandshakePacket::Type expected_type) { |
| if (!std::holds_alternative<transport::ReadFuture>(io_fut_)) { |
| io_fut_.emplace<transport::ReadFuture>(connection_.Read()); |
| } |
| auto& read_fut = std::get<transport::ReadFuture>(io_fut_); |
| PW_TRY_READY_ASSIGN(ConstBuf read_res, read_fut.Pend(cx)); |
| io_fut_ = std::monostate{}; |
| if (read_res == nullptr) { |
| return async2::Ready(Result<HandshakePacket>(Status::Cancelled())); |
| } |
| auto dec_res = HandshakePacket::Decode(ConstByteSpan(read_res)); |
| if (!dec_res.ok()) { |
| return async2::Ready(dec_res.status()); |
| } |
| if (dec_res->type() != expected_type) { |
| return async2::Ready(Result<HandshakePacket>(Status::DataLoss())); |
| } |
| return async2::Ready(Result<HandshakePacket>(*dec_res)); |
| } |
| |
| std::optional<async2::Poll<Result<EstablishedConnection>>> |
| HandshakeFutureImpl::StepWrite(async2::Context& cx, |
| HandshakePacket::Type type, |
| Stage next_stage) { |
| auto poll = PendWritePacket(cx, type); |
| if (poll.IsPending()) { |
| return async2::Pending(); |
| } |
| if (!poll->ok()) { |
| stage_ = Stage::kCompleted; |
| return Fail(*poll); |
| } |
| stage_ = next_stage; |
| return std::nullopt; |
| } |
| |
| async2::Poll<Result<EstablishedConnection>> InitiatorHandshakeFuture::Pend( |
| async2::Context& cx) { |
| PW_CHECK(is_pendable()); |
| while (stage_ != Stage::kCompleted) { |
| switch (stage_) { |
| case Stage::kEmpty: |
| PW_CRASH("HandshakeStage::kEmpty should be unreachable in Pend"); |
| case Stage::kSyn: { |
| if (auto poll = |
| StepWrite(cx, HandshakePacket::Type::kSyn, Stage::kSynAck)) { |
| return *poll; |
| } |
| break; |
| } |
| case Stage::kSynAck: { |
| if (auto poll = StepRead( |
| cx, |
| HandshakePacket::Type::kSynAck, |
| Stage::kAck, |
| [&](const HandshakePacket& pkt) { |
| // The responder must negotiate down to at most the version |
| // advertised in the SYN. `Decode` has already rejected 0. |
| if (pkt.version() > handshake_info_.negotiated_version) { |
| return Status::DataLoss(); |
| } |
| handshake_info_.negotiated_version = pkt.version(); |
| return OkStatus(); |
| })) { |
| return *poll; |
| } |
| break; |
| } |
| case Stage::kAck: { |
| if (auto poll = |
| StepWrite(cx, HandshakePacket::Type::kAck, Stage::kCompleted)) { |
| return *poll; |
| } |
| break; |
| } |
| case Stage::kCompleted: |
| return async2::Ready( |
| Result<EstablishedConnection>(Status::FailedPrecondition())); |
| } |
| } |
| return Complete(); |
| } |
| |
| async2::Poll<Result<EstablishedConnection>> ResponderHandshakeFuture::Pend( |
| async2::Context& cx) { |
| PW_CHECK(is_pendable()); |
| while (stage_ != Stage::kCompleted) { |
| switch (stage_) { |
| case Stage::kEmpty: |
| PW_CRASH("HandshakeStage::kEmpty should be unreachable in Pend"); |
| case Stage::kSyn: { |
| if (auto poll = StepRead( |
| cx, |
| HandshakePacket::Type::kSyn, |
| Stage::kSynAck, |
| [&](const HandshakePacket& pkt) { |
| handshake_info_.negotiated_version = std::min( |
| handshake_info_.negotiated_version, pkt.version()); |
| return OkStatus(); |
| })) { |
| return *poll; |
| } |
| break; |
| } |
| case Stage::kSynAck: { |
| if (auto poll = |
| StepWrite(cx, HandshakePacket::Type::kSynAck, Stage::kAck)) { |
| return *poll; |
| } |
| break; |
| } |
| case Stage::kAck: { |
| if (auto poll = StepRead(cx, |
| HandshakePacket::Type::kAck, |
| Stage::kCompleted, |
| [&](const HandshakePacket& pkt) { |
| if (pkt.version() != |
| handshake_info_.negotiated_version) { |
| return Status::DataLoss(); |
| } |
| return OkStatus(); |
| })) { |
| return *poll; |
| } |
| break; |
| } |
| case Stage::kCompleted: |
| return async2::Ready( |
| Result<EstablishedConnection>(Status::FailedPrecondition())); |
| } |
| } |
| return Complete(); |
| } |
| |
| } // namespace pw::rpc2::internal |