blob: bcc5f39679a9438f89a8ec2f1f7af72e8872afbc [file]
// 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