| // 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 |