blob: 228fcbc44366f878bde77c98ef0444f9212a1fdc [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/server_connection_task.h"
#include <utility>
#include "pw_assert/check.h"
#include "pw_log/log.h"
#include "pw_rpc2/internal/call.h"
#include "pw_rpc2/internal/packet.h"
#include "pw_rpc2/internal/server_call.h"
#include "pw_rpc2/internal/server_task.h"
#include "pw_rpc2/server.h"
namespace pw::rpc2::internal {
ServerConnectionTask::ServerConnectionTask(
transport::ReliableDatagramSocket connection,
Allocator& allocator,
ServerTask& server_task)
: ConnectionTask(std::move(connection), allocator, EndpointRole::kServer),
server_task_(server_task),
handshake_future_(this->connection()) {}
ServerConnectionTask::ServerConnectionTask(
transport::ReliableDatagramSocket connection,
Allocator& allocator,
Server& server)
: ServerConnectionTask(std::move(connection), allocator, server.task_) {}
ServerConnectionTask::ServerConnectionTask(
EstablishedConnection established_connection,
Allocator& allocator,
Server& server)
: ConnectionTask(
std::move(established_connection), allocator, EndpointRole::kServer),
server_task_(server.task_) {}
ServerConnectionTask::~ServerConnectionTask() {
Teardown();
RetireAllServerCalls();
UnregisterFromServer();
}
void ServerConnectionTask::UnregisterFromServer() {
server_task_.RemoveConnection(*this);
}
void ServerConnectionTask::RetireServerCall(ServerCall& call) {
// Reclaim this connection's reference into `ref` so that `call` stays alive
// while `Retire()` destroys the user's future, and is recycled when `ref`
// goes out of scope unless a handle outlived the method body.
IntrusivePtr<Call> ref = call.Retire();
// An `UnregisterService()` may be waiting for exactly this call to go away.
// Asking the server first keeps the common case --- nothing waiting --- to a
// single bool read instead of a wake per completed RPC.
if (server_task_.is_awaiting_quiescence()) {
server_task_.Wake();
}
}
bool ServerConnectionTask::CancelCallsForService(const Service& service) {
bool any_found = false;
for (Call& call : calls()) {
ServerCall& server_call = static_cast<ServerCall&>(call);
if (&server_call.service() != &service) {
continue;
}
any_found = true;
if (!server_call.is_closed()) {
// Send a kServiceUnregistered error packet before `Complete()`
// closes the write side and disarms the responder/writer destructors.
if (!server_call.is_write_closed()) {
server_call.QueueError(ServerError::kServiceUnregistered);
}
// Completing is enough to abort the call: the next poll sees a closed
// call and retires it without consulting the user's future. `Complete()`
// wakes this connection, so that poll is already scheduled.
server_call.Complete(Status::Cancelled());
}
}
return any_found;
}
bool ServerConnectionTask::HasCallsForService(const Service& service) const {
for (const Call& call : calls()) {
const ServerCall& server_call = static_cast<const ServerCall&>(call);
if (&server_call.service() == &service) {
return true;
}
}
return false;
}
ServerCall* ServerConnectionTask::FindRetirableServerCall() {
for (Call& call : calls()) {
ServerCall& server_call = static_cast<ServerCall&>(call);
// A closed call is retirable even if its task has not run since: the
// method will not be consulted again, and retiring it here is what stops
// its task from ever being polled.
if (server_call.is_finished() || server_call.is_closed()) {
return &server_call;
}
}
return nullptr;
}
void ServerConnectionTask::RetireFinishedServerCalls() {
// Retiring unlists the call and may destroy it, so rescan from the start
// each time rather than holding an iterator across it. Quadratic only in the
// number of calls retired on the same poll, and the scan is skipped entirely
// once none are.
while (ServerCall* finished = FindRetirableServerCall()) {
RetireServerCall(*finished);
}
}
void ServerConnectionTask::RetireAllServerCalls() {
while (!calls().empty()) {
RetireServerCall(static_cast<ServerCall&>(calls().front()));
}
}
async2::Poll<> ServerConnectionTask::DoPend(async2::Context& cx) {
const bool progressed = !is_closed() && PollConnection(cx);
if (is_closed()) {
// Retire all server calls before unregistering from the server so that
// `UnregisterService()` never resolves while any call's user future is
// still alive on this connection.
RetireAllServerCalls();
UnregisterFromServer();
return async2::Ready();
}
if (progressed) {
cx.ReEnqueue();
}
return async2::Pending();
}
bool ServerConnectionTask::PollConnection(async2::Context& cx) {
StoreWaker(cx);
// Stage 1: Complete responder handshake before accepting packets.
if (state() == State::kHandshaking && !PollHandshake(cx)) {
return false;
}
// Stage 2: Drain outgoing control packets and read/dispatch up to
// `kMaxPacketsPerPoll` incoming packets.
//
// Note: `progressed` is intentionally overwritten on each iteration rather
// than OR-accumulated. If `PendPacket()` returns false on iteration `i <
// kMaxPacketsPerPoll`, it has already polled the underlying transport futures
// to `Pending()` and registered wakers with `cx`, so re-enqueuing the task
// would only cause a redundant poll. `progressed` remains true after the loop
// only when all `kMaxPacketsPerPoll` iterations succeeded, meaning more work
// may still be ready in the transport and the task must yield via
// `cx.ReEnqueue()`.
bool progressed = false;
for (int i = 0; i < kMaxPacketsPerPoll; ++i) {
InboundPacket request;
progressed = PendPacket(cx, &request);
if (request != nullptr) {
HandleIncomingRequest(std::move(request));
}
if (!progressed || is_closed()) {
break;
}
}
if (is_closed()) {
return false;
}
// Stage 3: Reap the calls whose methods have finished. Their tasks ran on
// this same dispatcher; all that is left is to destroy their futures and
// drop this connection's reference to them.
RetireFinishedServerCalls();
return progressed;
}
bool ServerConnectionTask::PollHandshake(async2::Context& cx) {
auto poll = handshake_future_.Pend(cx);
if (!poll.IsReady()) {
return false;
}
if (!poll->ok()) {
PW_LOG_WARN("Handshake failed with status %s, closing connection",
poll->status().str());
CloseConnection(poll->status());
return false;
}
FinishHandshake(std::move(poll->value()));
handshake_future_ = ResponderHandshakeFuture();
return true;
}
void ServerConnectionTask::HandleIncomingRequest(InboundPacket&& packet) {
const PacketType type = packet.type();
PW_DCHECK(type.is_start());
// Runs on the dispatcher thread, which also owns the registry, so the
// service resolved here cannot be unregistered underneath the dispatch.
Service* target_service = server_task_.FindService(packet.service_id());
if (target_service == nullptr) {
PW_LOG_WARN("RPC request for unknown service_id 0x%08x",
static_cast<unsigned>(packet.service_id()));
QueueError(packet.call_id(), ServerError::kUnknownService);
return;
}
const uint32_t method_id = packet.method_id();
const Method* method = target_service->FindMethod(method_id);
if (method == nullptr) {
PW_LOG_WARN("RPC request for unknown method_id 0x%08x in service 0x%08x",
static_cast<unsigned>(method_id),
static_cast<unsigned>(packet.service_id()));
QueueError(packet.call_id(), ServerError::kUnknownMethod);
return;
}
// A unary or server-streaming method takes exactly one request, so it must
// be started by the packet that carries that request and closes the client's
// stream. Anything else means the client disagrees about the method's type.
const bool streaming_request = HasClientStream(method->type());
const bool is_single_request =
type.has_payload() && type.close_mode() == CloseMode::kStreamEnd;
if (!streaming_request && !is_single_request) {
PW_LOG_WARN(
"RPC start packet of type 0x%02x does not match the type of method "
"0x%08x in service 0x%08x",
static_cast<unsigned>(type.bits()),
static_cast<unsigned>(method_id),
static_cast<unsigned>(packet.service_id()));
QueueError(packet.call_id(), ServerError::kMethodTypeMismatch);
return;
}
Result<ServerCall*> call_res = ServerCall::Allocate(
*this, packet.call_id(), *target_service, *method, allocator());
if (!call_res.ok()) {
PW_LOG_ERROR("Failed to allocate ServerCall: %s", call_res.status().str());
QueueError(packet.call_id(), ServerError::kFailedToAllocateCall);
return;
}
ServerCall& call = **call_res;
// A streaming request may begin in the start packet. Queue its message for
// the method's reader, and close the stream if the client already did, in
// that order so that the message is read before the end of the stream.
ConstBuf request_payload;
if (!streaming_request) {
request_payload = std::move(packet).TakePayload();
} else {
if (type.has_payload()) {
call.DeliverInitialMessage(std::move(packet).TakePayload());
}
if (type.close_mode() == CloseMode::kStreamEnd) {
call.OnPeerStreamEnd();
}
}
const ServerError invocation_error =
method->Invoke(*target_service, call, std::move(request_payload));
if (invocation_error != ServerError::kOk) {
PW_LOG_WARN("Method invocation failed with protocol error: %s",
pw::EnumToString(invocation_error));
if (!call.is_write_closed()) {
QueueError(call.call_id(), invocation_error);
call.CloseWrite();
}
RetireServerCall(call);
return;
}
// The method is now running. Give it its own task on the server's
// dispatcher --- the one polling this connection --- so that the wakers it
// stores wake this call alone. A failed dispatch is retired above without
// ever being posted.
server_task_.dispatcher().Post(call);
}
} // namespace pw::rpc2::internal