blob: 3d0fa9cd86f225ba84364ff7cc9366e5627ecc3e [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/server.h"
#include <array>
#include <cstddef>
#include <cstring>
#include <optional>
#include <utility>
#include <variant>
#include "pw_allocator/fault_injecting_allocator.h"
#include "pw_allocator/testing.h"
#include "pw_assert/check.h"
#include "pw_async2/coro.h"
#include "pw_async2/dispatcher_for_test.h"
#include "pw_async2/task.h"
#include "pw_async2/try.h"
#include "pw_async2/value_future.h"
#include "pw_rpc2/internal/handshake.h"
#include "pw_rpc2/internal/method_invoker.h"
#include "pw_rpc2/internal/packet.h"
#include "pw_rpc2/internal/packet_testing.h"
#include "pw_rpc2/internal/server_call.h"
#include "pw_rpc2/internal/test_utils.h"
#include "pw_rpc2/method_type.h"
#include "pw_rpc2/reader.h"
#include "pw_rpc2/service.h"
#include "pw_rpc2/writer.h"
#include "pw_span/span.h"
#include "pw_thread/test_thread_context.h"
#include "pw_thread/thread.h"
#include "pw_transport/transport.h"
#include "pw_unit_test/framework.h"
namespace pw::rpc2 {
namespace {
namespace flags = ::pw::rpc2::internal::flags;
// Awaits a `ControlFuture` on the dispatcher. A future only makes progress
// while a task is pending it, so tests cannot simply run the dispatcher and
// inspect `is_complete()`.
class ControlTask : public async2::Task {
public:
explicit ControlTask(ControlFuture&& future)
: async2::Task(PW_ASYNC_TASK_NAME("ControlTask")),
future_(std::move(future)) {}
[[nodiscard]] bool done() const { return done_; }
// Only valid once `done()`.
[[nodiscard]] Status status() const { return status_; }
private:
async2::Poll<> DoPend(async2::Context& cx) override {
PW_TRY_READY_ASSIGN(status_, future_.Pend(cx));
done_ = true;
return async2::Ready();
}
ControlFuture future_;
bool done_ = false;
Status status_ = Status::Unknown();
};
class MockServerListener : public transport::ReliableDatagramListener {
public:
MockServerListener() = default;
~MockServerListener() override {
accept_provider_.Resolve(Status::Cancelled());
}
transport::ReliableDatagramListener::AcceptFuture Accept() override {
return accept_provider_.Get();
}
void ResolveAccept(Result<transport::ReliableDatagramSocket> res) {
accept_provider_.Resolve(std::move(res));
}
private:
async2::ValueProvider<Result<transport::ReliableDatagramSocket>>
accept_provider_;
};
struct ReadyFuture {
using value_type = void;
bool is_pendable() const { return !completed; }
bool is_complete() const { return completed; }
async2::Poll<> Pend(async2::Context&) {
completed = true;
return async2::Ready();
}
bool completed = false;
};
class TestEchoService : public Service {
public:
explicit TestEchoService(uint32_t service_id)
: Service(service_id, methods_),
methods_({
internal::Method(1u,
MethodType::kUnary,
sizeof(internal::MethodFutureImpl<ReadyFuture>),
&InvokeMethod1),
internal::Method(2u,
MethodType::kUnary,
sizeof(internal::MethodFutureImpl<ReadyFuture>),
&InvokeMethod2),
}) {}
uint32_t last_dispatched_method() const { return last_dispatched_method_; }
size_t last_payload_size() const { return last_payload_size_; }
private:
static internal::ServerError InvokeMethod1(Service& service,
internal::ServerCall& call,
ConstBuf&& request_payload) {
auto& self = static_cast<TestEchoService&>(service);
self.last_dispatched_method_ = 1u;
self.last_payload_size_ = request_payload.size();
auto responder =
internal::CallAccess::Create<RawUnaryWriter>(call.shared_call());
auto res_fut = responder.ReserveFinish(request_payload.size());
return call.EmplaceFutureFromFactory<ReadyFuture>(
[] { return ReadyFuture{}; });
}
static internal::ServerError InvokeMethod2(Service& service,
internal::ServerCall&,
ConstBuf&& request_payload) {
auto& self = static_cast<TestEchoService&>(service);
self.last_dispatched_method_ = 2u;
self.last_payload_size_ = request_payload.size();
return internal::ServerError::kInvalidRequestPayload;
}
std::array<internal::Method, 2> methods_;
uint32_t last_dispatched_method_ = 0;
size_t last_payload_size_ = 0;
};
// A future that never completes, so the `ServerCall` holding it stays open
// until something tears it down.
class StallFuture {
public:
using value_type = void;
// Default constructed futures are empty and never pended.
StallFuture() = default;
explicit StallFuture(int* destroyed) : destroyed_(destroyed) {}
~StallFuture() {
if (destroyed_ != nullptr) {
++(*destroyed_);
}
}
StallFuture(const StallFuture&) = delete;
StallFuture& operator=(const StallFuture&) = delete;
StallFuture(StallFuture&& other) noexcept
: destroyed_(std::exchange(other.destroyed_, nullptr)) {}
StallFuture& operator=(StallFuture&& other) noexcept {
if (this != &other) {
destroyed_ = std::exchange(other.destroyed_, nullptr);
}
return *this;
}
bool is_pendable() const { return true; }
bool is_complete() const { return false; }
async2::Poll<> Pend(async2::Context& cx) {
PW_ASYNC_STORE_WAKER(cx, waker_, "StallFuture");
return async2::Pending();
}
private:
int* destroyed_ = nullptr;
async2::Waker waker_;
};
// Service whose only method parks forever.
class StallingService : public Service {
public:
using Impl = internal::MethodFutureImpl<StallFuture>;
explicit StallingService(uint32_t service_id)
: Service(service_id, methods_),
methods_({
internal::Method(
1u, MethodType::kUnary, sizeof(Impl), &InvokeStall),
}) {}
// Incremented when the parked future is destroyed, which happens only when
// its `ServerCall` is torn down.
int future_destructions() const { return future_destructions_; }
private:
static internal::ServerError InvokeStall(Service& service,
internal::ServerCall& call,
ConstBuf&&) {
auto& self = static_cast<StallingService&>(service);
return call.EmplaceFutureFromFactory<StallFuture>(
[&] { return StallFuture(&self.future_destructions_); });
}
std::array<internal::Method, 1> methods_;
int future_destructions_ = 0;
};
void EstablishServerConnection(Allocator& allocator,
async2::DispatcherForTest& dispatcher,
[[maybe_unused]] Server& server,
MockServerListener& transport,
test::MockConnection* raw_conn,
transport::ReliableDatagramSocket conn) {
dispatcher.RunUntilStalled();
transport.ResolveAccept(conn);
dispatcher.RunUntilStalled();
// Send client SYN
Buf req_buf =
Buf::Allocate(allocator, internal::HandshakePacket::kWireSizeBytes);
auto hs_pkt =
internal::HandshakePacket(internal::HandshakePacket::Type::kSyn);
auto enc_res = hs_pkt.Encode(std::move(req_buf));
PW_CHECK(enc_res.ok());
raw_conn->SetNextRead(std::move(enc_res.value()));
dispatcher.RunUntilStalled();
PW_ASSERT(raw_conn->commit_count() == 1u); // 1 handshake response (SYN-ACK)
// Send client ACK
Buf ack_buf =
Buf::Allocate(allocator, internal::HandshakePacket::kWireSizeBytes);
auto ack_pkt =
internal::HandshakePacket(internal::HandshakePacket::Type::kAck);
auto ack_enc = ack_pkt.Encode(std::move(ack_buf));
PW_CHECK(ack_enc.ok());
raw_conn->SetNextRead(std::move(ack_enc.value()));
dispatcher.RunUntilStalled();
}
TEST(ServerTest, RegisterTransportAndAcceptConnectionWithHandshake) {
allocator::test::AllocatorForTest<16384> allocator;
async2::DispatcherForTest dispatcher;
Server server(allocator, dispatcher);
MockServerListener transport;
EXPECT_EQ(server.RegisterListenerBlocking(transport), OkStatus());
server.Start();
dispatcher.RunUntilStalled();
auto [conn, raw_conn] = test::MakeMockConnection(allocator);
EstablishServerConnection(
allocator, dispatcher, server, transport, raw_conn, conn);
ControlTask close_task(server.Close());
dispatcher.Post(close_task);
dispatcher.RunUntilStalled();
EXPECT_TRUE(close_task.done());
}
TEST(ServerTest, ServiceMethodDispatchRouting) {
allocator::test::AllocatorForTest<16384> allocator;
async2::DispatcherForTest dispatcher;
Server server(allocator, dispatcher);
TestEchoService echo_service(100u);
EXPECT_EQ(server.RegisterServiceBlocking(echo_service), OkStatus());
MockServerListener transport;
EXPECT_EQ(server.RegisterListenerBlocking(transport), OkStatus());
server.Start();
auto [conn, raw_conn] = test::MakeMockConnection(allocator);
EstablishServerConnection(
allocator, dispatcher, server, transport, raw_conn, conn);
// Send a request packet targeting service 100, method 1
std::byte payload[4] = {
std::byte{10}, std::byte{20}, std::byte{30}, std::byte{40}};
auto pkt = internal::PacketFramer::FrameStartUnaryPacket(
allocator, /*call_id=*/5, /*service_id=*/100u, /*method_id=*/1u, payload);
ASSERT_TRUE(pkt.ok());
raw_conn->SetNextRead(std::move(*pkt));
dispatcher.RunUntilStalled();
EXPECT_EQ(echo_service.last_dispatched_method(), 1u);
EXPECT_EQ(echo_service.last_payload_size(), 4u);
ControlTask close_task(server.Close());
dispatcher.Post(close_task);
dispatcher.RunUntilStalled();
EXPECT_TRUE(close_task.done());
}
TEST(ServerTest, ServiceNotFoundRepliesError) {
allocator::test::AllocatorForTest<16384> allocator;
async2::DispatcherForTest dispatcher;
Server server(allocator, dispatcher);
TestEchoService echo_service(100u);
EXPECT_EQ(server.RegisterServiceBlocking(echo_service), OkStatus());
MockServerListener transport;
EXPECT_EQ(server.RegisterListenerBlocking(transport), OkStatus());
server.Start();
auto [conn, raw_conn] = test::MakeMockConnection(allocator);
EstablishServerConnection(
allocator, dispatcher, server, transport, raw_conn, conn);
size_t commits_before = raw_conn->commit_count();
// Send a request packet targeting unregistered service 999
std::byte payload[2] = {std::byte{1}, std::byte{2}};
auto pkt = internal::PacketFramer::FrameStartUnaryPacket(
allocator, /*call_id=*/7, /*service_id=*/999u, /*method_id=*/1u, payload);
ASSERT_TRUE(pkt.ok());
raw_conn->SetNextRead(std::move(*pkt));
dispatcher.RunUntilStalled();
// Server should reply with a server error packet (kUnknownService).
EXPECT_EQ(raw_conn->commit_count(), commits_before + 1u);
auto decode_err =
internal::InboundPacket::Decode(raw_conn->last_written_buf());
ASSERT_TRUE(decode_err.ok());
EXPECT_EQ(
decode_err->type(),
(internal::PacketType::Make<flags::kServer, flags::kErrorTerminal>()));
EXPECT_EQ(decode_err->call_id(), 7u);
EXPECT_EQ(decode_err->server_error(), internal::ServerError::kUnknownService);
ControlTask close_task(server.Close());
dispatcher.Post(close_task);
dispatcher.RunUntilStalled();
EXPECT_TRUE(close_task.done());
}
TEST(ServerTest, MethodNotFoundRepliesError) {
allocator::test::AllocatorForTest<16384> allocator;
async2::DispatcherForTest dispatcher;
Server server(allocator, dispatcher);
TestEchoService echo_service(100u);
EXPECT_EQ(server.RegisterServiceBlocking(echo_service), OkStatus());
MockServerListener transport;
EXPECT_EQ(server.RegisterListenerBlocking(transport), OkStatus());
server.Start();
auto [conn, raw_conn] = test::MakeMockConnection(allocator);
EstablishServerConnection(
allocator, dispatcher, server, transport, raw_conn, conn);
size_t commits_before = raw_conn->commit_count();
// Send a request packet targeting unknown method 999
std::byte payload[2] = {std::byte{1}, std::byte{2}};
auto pkt = internal::PacketFramer::FrameStartUnaryPacket(allocator,
/*call_id=*/8,
/*service_id=*/100u,
/*method_id=*/999u,
payload);
ASSERT_TRUE(pkt.ok());
raw_conn->SetNextRead(std::move(*pkt));
dispatcher.RunUntilStalled();
// Server should reply with ServerError::kUnknownMethod error frame
EXPECT_EQ(raw_conn->commit_count(), commits_before + 1u);
auto decode_err =
internal::InboundPacket::Decode(raw_conn->last_written_buf());
ASSERT_TRUE(decode_err.ok());
EXPECT_EQ(
decode_err->type(),
(internal::PacketType::Make<flags::kServer, flags::kErrorTerminal>()));
EXPECT_EQ(decode_err->call_id(), 8u);
EXPECT_EQ(decode_err->server_error(), internal::ServerError::kUnknownMethod);
ControlTask close_task(server.Close());
dispatcher.Post(close_task);
dispatcher.RunUntilStalled();
EXPECT_TRUE(close_task.done());
}
TEST(ServerTest, UnregisterServiceLifecycle) {
allocator::test::AllocatorForTest<16384> allocator;
async2::DispatcherForTest dispatcher;
Server server(allocator, dispatcher);
TestEchoService echo_service(100u);
EXPECT_EQ(server.RegisterServiceBlocking(echo_service), OkStatus());
EXPECT_EQ(server.UnregisterServiceBlocking(echo_service), OkStatus());
MockServerListener transport;
EXPECT_EQ(server.RegisterListenerBlocking(transport), OkStatus());
server.Start();
auto [conn, raw_conn] = test::MakeMockConnection(allocator);
EstablishServerConnection(
allocator, dispatcher, server, transport, raw_conn, conn);
size_t commits_before = raw_conn->commit_count();
// Request to unregistered service should return NotFound
std::byte payload[1] = {std::byte{1}};
auto pkt = internal::PacketFramer::FrameStartUnaryPacket(
allocator, /*call_id=*/9, /*service_id=*/100u, /*method_id=*/1u, payload);
ASSERT_TRUE(pkt.ok());
raw_conn->SetNextRead(std::move(*pkt));
dispatcher.RunUntilStalled();
EXPECT_EQ(raw_conn->commit_count(), commits_before + 1u);
auto decode_err =
internal::InboundPacket::Decode(raw_conn->last_written_buf());
ASSERT_TRUE(decode_err.ok());
EXPECT_EQ(decode_err->server_error(), internal::ServerError::kUnknownService);
ControlTask close_task(server.Close());
dispatcher.Post(close_task);
dispatcher.RunUntilStalled();
EXPECT_TRUE(close_task.done());
}
TEST(ServerTest, RegisterDuplicateServiceIdCrashes) {
allocator::test::AllocatorForTest<16384> allocator;
async2::DispatcherForTest dispatcher;
Server server(allocator, dispatcher);
TestEchoService echo_service(100u);
TestEchoService duplicate_service(100u);
EXPECT_EQ(server.RegisterServiceBlocking(echo_service), OkStatus());
EXPECT_DEATH_IF_SUPPORTED(
static_cast<void>(server.RegisterService(duplicate_service)), "");
EXPECT_EQ(server.UnregisterServiceBlocking(echo_service), OkStatus());
}
TEST(ServerTest, StartTwiceCrashes) {
allocator::test::AllocatorForTest<16384> allocator;
async2::DispatcherForTest dispatcher;
Server server(allocator, dispatcher);
server.Start();
EXPECT_DEATH_IF_SUPPORTED(server.Start(), "");
// The server in this process really was started, so it has to be closed.
ControlTask close_task(server.Close());
dispatcher.Post(close_task);
dispatcher.RunUntilStalled();
EXPECT_TRUE(close_task.done());
}
TEST(ServerTest, DestroyingUnclosedServerCrashes) {
allocator::test::AllocatorForTest<16384> allocator;
async2::DispatcherForTest dispatcher;
MockServerListener transport;
EXPECT_DEATH_IF_SUPPORTED(
{
Server server(allocator, dispatcher);
PW_CHECK_OK(server.RegisterListenerBlocking(transport));
server.Start();
},
"");
}
TEST(ServerTest, CloseTearsDownListenersAndConnections) {
allocator::test::AllocatorForTest<16384> allocator;
async2::DispatcherForTest dispatcher;
MockServerListener transport;
auto [conn, raw_conn] = test::MakeMockConnection(allocator);
{
Server server(allocator, dispatcher);
EXPECT_EQ(server.RegisterListenerBlocking(transport), OkStatus());
server.Start();
EstablishServerConnection(
allocator, dispatcher, server, transport, raw_conn, conn);
ASSERT_FALSE(raw_conn->is_closed());
ControlTask close_task(server.Close());
dispatcher.Post(close_task);
dispatcher.RunUntilStalled();
EXPECT_TRUE(close_task.done());
EXPECT_TRUE(raw_conn->is_closed());
}
// Nothing the server owned may remain posted to the dispatcher.
dispatcher.RunUntilStalled();
}
TEST(ServerTest, CloseBlockingTearsDownListenersAndConnections) {
allocator::test::AllocatorForTest<16384> allocator;
async2::DispatcherForTest dispatcher;
MockServerListener transport;
auto [conn, raw_conn] = test::MakeMockConnection(allocator);
Server server(allocator, dispatcher);
EXPECT_EQ(server.RegisterListenerBlocking(transport), OkStatus());
server.Start();
EstablishServerConnection(
allocator, dispatcher, server, transport, raw_conn, conn);
ASSERT_FALSE(raw_conn->is_closed());
pw::thread::test::TestThreadContext context;
pw::Thread thread(context.options(), [&server] { server.CloseBlocking(); });
dispatcher.AllowBlocking();
dispatcher.RunToCompletion();
thread.join();
EXPECT_TRUE(raw_conn->is_closed());
}
// The listener teardown path runs on the dispatcher when `Close()` is polled,
// closing every connection task and retiring any parked `ServerCall`s so that
// nothing retains a dangling pointer to the connection.
TEST(ServerTest, CloseDrainsInFlightServerCalls) {
allocator::test::AllocatorForTest<16384> allocator;
async2::DispatcherForTest dispatcher;
Server server(allocator, dispatcher);
StallingService stalling_service(200u);
EXPECT_EQ(server.RegisterServiceBlocking(stalling_service), OkStatus());
MockServerListener transport;
EXPECT_EQ(server.RegisterListenerBlocking(transport), OkStatus());
server.Start();
auto [conn, raw_conn] = test::MakeMockConnection(allocator);
EstablishServerConnection(
allocator, dispatcher, server, transport, raw_conn, conn);
std::byte payload[1] = {std::byte{1}};
auto pkt = internal::PacketFramer::FrameStartUnaryPacket(allocator,
/*call_id=*/11,
/*service_id=*/200u,
/*method_id=*/1u,
payload);
ASSERT_TRUE(pkt.ok());
raw_conn->SetNextRead(std::move(*pkt));
dispatcher.RunUntilStalled();
// The call is dispatched and parked, so the connection still owns it.
ASSERT_EQ(stalling_service.future_destructions(), 0);
ControlTask close_task(server.Close());
dispatcher.Post(close_task);
dispatcher.RunUntilStalled();
EXPECT_TRUE(close_task.done());
// Teardown on the dispatcher drained the parked call task.
EXPECT_EQ(stalling_service.future_destructions(), 1);
EXPECT_TRUE(raw_conn->is_closed());
}
// `UnregisterService()` immediately cancels every in-flight call to the
// service, and its future does not resolve until those calls have been torn
// down, so a resolved future means the service may be destroyed.
TEST(ServerTest, UnregisterServiceAbortsInFlightCallsAndAwaitsQuiescence) {
allocator::test::AllocatorForTest<16384> allocator;
async2::DispatcherForTest dispatcher;
Server server(allocator, dispatcher);
StallingService stalling_service(200u);
EXPECT_EQ(server.RegisterServiceBlocking(stalling_service), OkStatus());
MockServerListener transport;
EXPECT_EQ(server.RegisterListenerBlocking(transport), OkStatus());
server.Start();
auto [conn, raw_conn] = test::MakeMockConnection(allocator);
EstablishServerConnection(
allocator, dispatcher, server, transport, raw_conn, conn);
std::byte payload[1] = {std::byte{1}};
auto pkt = internal::PacketFramer::FrameStartUnaryPacket(allocator,
/*call_id=*/11,
/*service_id=*/200u,
/*method_id=*/1u,
payload);
ASSERT_TRUE(pkt.ok());
raw_conn->SetNextRead(std::move(*pkt));
dispatcher.RunUntilStalled();
// The call is dispatched and parked on a future that never completes.
ASSERT_EQ(stalling_service.future_destructions(), 0);
ControlTask unregister_task(server.UnregisterService(stalling_service));
dispatcher.Post(unregister_task);
dispatcher.RunUntilStalled();
// The parked call was cancelled and retired before the future resolved.
ASSERT_TRUE(unregister_task.done());
EXPECT_EQ(unregister_task.status(), OkStatus());
EXPECT_EQ(stalling_service.future_destructions(), 1);
// Aborting the call does not close the connection; it stays up for other
// services.
EXPECT_FALSE(raw_conn->is_closed());
ControlTask close_task(server.Close());
dispatcher.Post(close_task);
dispatcher.RunUntilStalled();
EXPECT_TRUE(close_task.done());
}
// With nothing in flight there is nothing to wait for, so the unregister
// completes as soon as the dispatcher applies it.
TEST(ServerTest, UnregisterServiceWithNoCallsCompletesImmediately) {
allocator::test::AllocatorForTest<16384> allocator;
async2::DispatcherForTest dispatcher;
Server server(allocator, dispatcher);
StallingService stalling_service(200u);
EXPECT_EQ(server.RegisterServiceBlocking(stalling_service), OkStatus());
MockServerListener transport;
EXPECT_EQ(server.RegisterListenerBlocking(transport), OkStatus());
server.Start();
auto [conn, raw_conn] = test::MakeMockConnection(allocator);
EstablishServerConnection(
allocator, dispatcher, server, transport, raw_conn, conn);
ControlTask unregister_task(server.UnregisterService(stalling_service));
dispatcher.Post(unregister_task);
dispatcher.RunUntilStalled();
ASSERT_TRUE(unregister_task.done());
EXPECT_EQ(unregister_task.status(), OkStatus());
EXPECT_EQ(stalling_service.future_destructions(), 0);
// The service may be re-registered once it has been fully unregistered.
ControlTask register_task(server.RegisterService(stalling_service));
dispatcher.Post(register_task);
dispatcher.RunUntilStalled();
ASSERT_TRUE(register_task.done());
EXPECT_EQ(register_task.status(), OkStatus());
ControlTask close_task(server.Close());
dispatcher.Post(close_task);
dispatcher.RunUntilStalled();
EXPECT_TRUE(close_task.done());
}
// Every thread-safe control operation may be driven from a thread other than
// the dispatcher's, which is what the `*Blocking()` forms are for.
TEST(ServerTest, BlockingControlOperationsRunFromAnotherThread) {
allocator::test::AllocatorForTest<16384> allocator;
async2::DispatcherForTest dispatcher;
Server server(allocator, dispatcher);
TestEchoService echo_service(100u);
MockServerListener transport;
server.Start();
Status register_service = Status::Unknown();
Status register_listener = Status::Unknown();
Status unregister_service = Status::Unknown();
auto run_control_ops = [&] {
register_service = server.RegisterServiceBlocking(echo_service);
register_listener = server.RegisterListenerBlocking(transport);
unregister_service = server.UnregisterServiceBlocking(echo_service);
server.CloseBlocking();
};
pw::thread::test::TestThreadContext context;
pw::Thread thread(context.options(),
[&run_control_ops] { run_control_ops(); });
dispatcher.AllowBlocking();
dispatcher.RunToCompletion();
thread.join();
EXPECT_EQ(register_service, OkStatus());
EXPECT_EQ(register_listener, OkStatus());
EXPECT_EQ(unregister_service, OkStatus());
}
TEST(ServerTest, RegisterListenerFailsWhenRecordAllocationFails) {
allocator::test::AllocatorForTest<16384> backing_allocator;
allocator::test::FaultInjectingAllocator failing_allocator(backing_allocator);
async2::DispatcherForTest dispatcher;
Server server(failing_allocator, dispatcher);
MockServerListener transport;
failing_allocator.DisableAllocate();
EXPECT_EQ(server.RegisterListenerBlocking(transport),
Status::ResourceExhausted());
}
TEST(ServerTest, ConcurrentUnregisterServiceWhileQuiescingDoesNotCrash) {
allocator::test::AllocatorForTest<16384> allocator;
async2::DispatcherForTest dispatcher;
Server server(allocator, dispatcher);
StallingService service(100u);
MockServerListener transport;
auto [conn, raw_conn] = test::MakeMockConnection(allocator);
PW_TEST_EXPECT_OK(server.RegisterServiceBlocking(service));
PW_TEST_EXPECT_OK(server.RegisterListenerBlocking(transport));
server.Start();
EstablishServerConnection(
allocator, dispatcher, server, transport, raw_conn, conn);
// Start a call that stalls forever.
auto req_pkt = internal::PacketFramer::FrameStartUnaryPacket(
allocator, /*call_id=*/1, /*service_id=*/100, /*method_id=*/1, {});
ASSERT_TRUE(req_pkt.ok());
raw_conn->SetNextRead(std::move(*req_pkt));
dispatcher.RunUntilStalled();
EXPECT_EQ(service.future_destructions(), 0);
// Begin unregistering the service. On the first poll, the request is popped
// from ControlQueue (`queued_ = false`, `in_flight_ = true`) and waits in
// `quiescing_services_` for the connection task to retire the cancelled call.
ControlTask unreg1(server.UnregisterService(service));
dispatcher.Post(unreg1);
// Submit a second UnregisterService while the first is in-flight. It must not
// re-queue the already-popped ControlRequest node.
ControlTask unreg2(server.UnregisterService(service));
dispatcher.Post(unreg2);
dispatcher.RunUntilStalled();
EXPECT_EQ(service.future_destructions(), 1);
EXPECT_TRUE(unreg1.done());
EXPECT_EQ(unreg1.status(), OkStatus());
EXPECT_TRUE(unreg2.done());
EXPECT_EQ(unreg2.status(), OkStatus());
unreg1.Deregister();
unreg2.Deregister();
ControlTask close_task(server.Close());
dispatcher.Post(close_task);
dispatcher.RunUntilStalled();
EXPECT_TRUE(close_task.done());
}
TEST(ServerTest, UnregisterServiceNotifiesPeerWithCancelledError) {
allocator::test::AllocatorForTest<16384> allocator;
async2::DispatcherForTest dispatcher;
Server server(allocator, dispatcher);
StallingService stalling_service(200u);
EXPECT_EQ(server.RegisterServiceBlocking(stalling_service), OkStatus());
MockServerListener transport;
EXPECT_EQ(server.RegisterListenerBlocking(transport), OkStatus());
server.Start();
auto [conn, raw_conn] = test::MakeMockConnection(allocator);
EstablishServerConnection(
allocator, dispatcher, server, transport, raw_conn, conn);
raw_conn->clear_written();
std::byte payload[1] = {std::byte{1}};
auto pkt = internal::PacketFramer::FrameStartUnaryPacket(allocator,
/*call_id=*/42u,
/*service_id=*/200u,
/*method_id=*/1u,
payload);
ASSERT_TRUE(pkt.ok());
raw_conn->SetNextRead(std::move(*pkt));
dispatcher.RunUntilStalled();
EXPECT_EQ(raw_conn->written_packet_count(), 0u);
ControlTask unregister_task(server.UnregisterService(stalling_service));
dispatcher.Post(unregister_task);
dispatcher.RunUntilStalled();
ASSERT_TRUE(unregister_task.done());
EXPECT_EQ(unregister_task.status(), OkStatus());
// The client must receive a kServiceUnregistered error packet for the
// aborted call.
ASSERT_GE(raw_conn->written_packet_count(), 1u);
auto decode_res =
internal::InboundPacket::Decode(raw_conn->last_written_buf());
ASSERT_TRUE(decode_res.ok());
EXPECT_EQ(
decode_res->type(),
(internal::PacketType::Make<flags::kServer, flags::kErrorTerminal>()));
EXPECT_EQ(decode_res->call_id(), 42u);
EXPECT_EQ(decode_res->server_error(),
internal::ServerError::kServiceUnregistered);
ControlTask close_task(server.Close());
dispatcher.Post(close_task);
dispatcher.RunUntilStalled();
EXPECT_TRUE(close_task.done());
}
TEST(ServerTest,
ConnectionCloseWhileServiceQuiescingRetiresCallsBeforeUnregisterResolves) {
allocator::test::AllocatorForTest<16384> allocator;
async2::DispatcherForTest dispatcher;
Server server(allocator, dispatcher);
StallingService stalling_service(200u);
EXPECT_EQ(server.RegisterServiceBlocking(stalling_service), OkStatus());
MockServerListener transport;
EXPECT_EQ(server.RegisterListenerBlocking(transport), OkStatus());
server.Start();
auto [conn, raw_conn] = test::MakeMockConnection(allocator);
EstablishServerConnection(
allocator, dispatcher, server, transport, raw_conn, conn);
std::byte payload[1] = {std::byte{1}};
auto pkt = internal::PacketFramer::FrameStartUnaryPacket(allocator,
/*call_id=*/11u,
/*service_id=*/200u,
/*method_id=*/1u,
payload);
ASSERT_TRUE(pkt.ok());
raw_conn->SetNextRead(std::move(*pkt));
dispatcher.RunUntilStalled();
ASSERT_EQ(stalling_service.future_destructions(), 0);
// Submit UnregisterService AND simulate the peer closing the connection at
// the same time so that ServerConnectionTask::DoPend sees is_closed() while
// the service is quiescing.
ControlTask unregister_task(server.UnregisterService(stalling_service));
dispatcher.Post(unregister_task);
conn.Close();
dispatcher.RunUntilStalled();
ASSERT_TRUE(unregister_task.done());
EXPECT_EQ(unregister_task.status(), OkStatus());
// The call's future MUST have been destroyed before UnregisterService
// resolved.
EXPECT_EQ(stalling_service.future_destructions(), 1);
ControlTask close_task(server.Close());
dispatcher.Post(close_task);
dispatcher.RunUntilStalled();
EXPECT_TRUE(close_task.done());
}
TEST(ServerTest, RegisterMultipleListenersWhileRunningDoesNotCrash) {
allocator::test::AllocatorForTest<16384> allocator;
async2::DispatcherForTest dispatcher;
Server server(allocator, dispatcher);
MockServerListener listener1;
MockServerListener listener2;
MockServerListener listener3;
server.Start();
ControlTask reg1(server.RegisterListener(listener1));
ControlTask reg2(server.RegisterListener(listener2));
ControlTask reg3(server.RegisterListener(listener3));
dispatcher.Post(reg1);
dispatcher.Post(reg2);
dispatcher.Post(reg3);
dispatcher.RunUntilStalled();
EXPECT_TRUE(reg1.done());
EXPECT_EQ(reg1.status(), OkStatus());
EXPECT_TRUE(reg2.done());
EXPECT_EQ(reg2.status(), OkStatus());
EXPECT_TRUE(reg3.done());
EXPECT_EQ(reg3.status(), OkStatus());
ControlTask close_task(server.Close());
dispatcher.Post(close_task);
dispatcher.RunUntilStalled();
EXPECT_TRUE(close_task.done());
}
TEST(ServerTest, ImmediateRegistrationAndUnregistrationBeforeStart) {
allocator::test::AllocatorForTest<16384> allocator;
async2::DispatcherForTest dispatcher;
Server server(allocator, dispatcher);
TestEchoService service1(100u);
TestEchoService service2(200u);
MockServerListener listener;
// Blocking registrations before Start() execute immediately with 0
// allocations.
EXPECT_EQ(server.RegisterServiceBlocking(service1), OkStatus());
EXPECT_EQ(server.RegisterListenerBlocking(listener), OkStatus());
// Unregistering before Start() removes immediately.
EXPECT_EQ(server.UnregisterServiceBlocking(service1), OkStatus());
// Async registration before Start() resolves without needing the server
// running.
ControlTask reg_task(server.RegisterService(service2));
dispatcher.Post(reg_task);
dispatcher.RunUntilStalled();
EXPECT_TRUE(reg_task.done());
EXPECT_EQ(reg_task.status(), OkStatus());
server.Start();
ControlTask close_task(server.Close());
dispatcher.Post(close_task);
dispatcher.RunUntilStalled();
EXPECT_TRUE(close_task.done());
}
// Closing unregisters every service, so unregistering while the server is
// closing succeeds, while registering fails.
TEST(ServerTest, UnregisterServiceWhileClosingSucceeds) {
allocator::test::AllocatorForTest<16384> allocator;
async2::DispatcherForTest dispatcher;
Server server(allocator, dispatcher);
StallingService stalling_service(200u);
TestEchoService echo_service(100u);
EXPECT_EQ(server.RegisterServiceBlocking(stalling_service), OkStatus());
MockServerListener transport;
EXPECT_EQ(server.RegisterListenerBlocking(transport), OkStatus());
server.Start();
auto [conn, raw_conn] = test::MakeMockConnection(allocator);
EstablishServerConnection(
allocator, dispatcher, server, transport, raw_conn, conn);
std::byte payload[1] = {std::byte{1}};
auto pkt = internal::PacketFramer::FrameStartUnaryPacket(allocator,
/*call_id=*/11u,
/*service_id=*/200u,
/*method_id=*/1u,
payload);
ASSERT_TRUE(pkt.ok());
raw_conn->SetNextRead(std::move(*pkt));
dispatcher.RunUntilStalled();
ASSERT_EQ(stalling_service.future_destructions(), 0);
ControlTask close_task(server.Close());
ControlTask unregister_task(server.UnregisterService(stalling_service));
ControlTask register_task(server.RegisterService(echo_service));
dispatcher.Post(close_task);
dispatcher.Post(unregister_task);
dispatcher.Post(register_task);
dispatcher.RunUntilStalled();
ASSERT_TRUE(close_task.done());
ASSERT_TRUE(unregister_task.done());
EXPECT_EQ(unregister_task.status(), OkStatus());
ASSERT_TRUE(register_task.done());
EXPECT_EQ(register_task.status(), Status::FailedPrecondition());
EXPECT_EQ(stalling_service.future_destructions(), 1);
}
TEST(ServerTest, UnregisterServiceAfterCloseSucceeds) {
allocator::test::AllocatorForTest<16384> allocator;
async2::DispatcherForTest dispatcher;
Server server(allocator, dispatcher);
TestEchoService service1(100u);
TestEchoService service2(200u);
EXPECT_EQ(server.RegisterServiceBlocking(service1), OkStatus());
EXPECT_EQ(server.RegisterServiceBlocking(service2), OkStatus());
server.Start();
ControlTask close_task(server.Close());
dispatcher.Post(close_task);
dispatcher.RunUntilStalled();
ASSERT_TRUE(close_task.done());
ControlTask unregister_task(server.UnregisterService(service1));
dispatcher.Post(unregister_task);
dispatcher.RunUntilStalled();
ASSERT_TRUE(unregister_task.done());
EXPECT_EQ(unregister_task.status(), OkStatus());
EXPECT_EQ(server.UnregisterServiceBlocking(service2), OkStatus());
EXPECT_EQ(server.RegisterServiceBlocking(service1),
Status::FailedPrecondition());
}
// The async control operations may be submitted from a thread other than the
// dispatcher's; the resulting futures resolve once the dispatcher applies them.
TEST(ServerTest, AsyncControlOperationsSubmittedFromAnotherThread) {
allocator::test::AllocatorForTest<16384> allocator;
async2::DispatcherForTest dispatcher;
Server server(allocator, dispatcher);
TestEchoService echo_service(100u);
MockServerListener transport;
server.Start();
std::optional<ControlFuture> register_service;
std::optional<ControlFuture> register_listener;
std::optional<ControlFuture> unregister_service;
auto submit_control_ops = [&] {
register_service = server.RegisterService(echo_service);
register_listener = server.RegisterListener(transport);
unregister_service = server.UnregisterService(echo_service);
};
pw::thread::test::TestThreadContext context;
pw::Thread thread(context.options(),
[&submit_control_ops] { submit_control_ops(); });
thread.join();
ControlTask register_service_task(std::move(*register_service));
ControlTask register_listener_task(std::move(*register_listener));
ControlTask unregister_service_task(std::move(*unregister_service));
dispatcher.Post(register_service_task);
dispatcher.Post(register_listener_task);
dispatcher.Post(unregister_service_task);
dispatcher.RunUntilStalled();
ASSERT_TRUE(register_service_task.done());
EXPECT_EQ(register_service_task.status(), OkStatus());
ASSERT_TRUE(register_listener_task.done());
EXPECT_EQ(register_listener_task.status(), OkStatus());
ASSERT_TRUE(unregister_service_task.done());
EXPECT_EQ(unregister_service_task.status(), OkStatus());
ControlTask close_task(server.Close());
dispatcher.Post(close_task);
dispatcher.RunUntilStalled();
EXPECT_TRUE(close_task.done());
}
TEST(ServerTest, AsyncControlOperationsFailWhenCommandAllocationFails) {
allocator::test::AllocatorForTest<16384> backing_allocator;
allocator::test::FaultInjectingAllocator failing_allocator(backing_allocator);
async2::DispatcherForTest dispatcher;
Server server(failing_allocator, dispatcher);
TestEchoService echo_service(100u);
MockServerListener transport;
server.Start();
failing_allocator.DisableAllocate();
ControlTask register_service_task(server.RegisterService(echo_service));
ControlTask register_listener_task(server.RegisterListener(transport));
ControlTask unregister_service_task(server.UnregisterService(echo_service));
dispatcher.Post(register_service_task);
dispatcher.Post(register_listener_task);
dispatcher.Post(unregister_service_task);
dispatcher.RunUntilStalled();
ASSERT_TRUE(register_service_task.done());
EXPECT_EQ(register_service_task.status(), Status::ResourceExhausted());
ASSERT_TRUE(register_listener_task.done());
EXPECT_EQ(register_listener_task.status(), Status::ResourceExhausted());
ASSERT_TRUE(unregister_service_task.done());
EXPECT_EQ(unregister_service_task.status(), Status::ResourceExhausted());
failing_allocator.EnableAllocate();
ControlTask close_task(server.Close());
dispatcher.Post(close_task);
dispatcher.RunUntilStalled();
EXPECT_TRUE(close_task.done());
}
TEST(ServerTest, CallAllocationFailureRepliesError) {
allocator::test::AllocatorForTest<16384> backing_allocator;
allocator::test::FaultInjectingAllocator failing_allocator(backing_allocator);
async2::DispatcherForTest dispatcher;
Server server(failing_allocator, dispatcher);
TestEchoService echo_service(100u);
EXPECT_EQ(server.RegisterServiceBlocking(echo_service), OkStatus());
MockServerListener transport;
EXPECT_EQ(server.RegisterListenerBlocking(transport), OkStatus());
server.Start();
auto [conn, raw_conn] = test::MakeMockConnection(backing_allocator);
EstablishServerConnection(
backing_allocator, dispatcher, server, transport, raw_conn, conn);
const size_t commits_before = raw_conn->commit_count();
std::byte payload[1] = {std::byte{1}};
auto pkt = internal::PacketFramer::FrameStartUnaryPacket(backing_allocator,
/*call_id=*/12u,
/*service_id=*/100u,
/*method_id=*/1u,
payload);
ASSERT_TRUE(pkt.ok());
raw_conn->SetNextRead(std::move(*pkt));
failing_allocator.DisableAllocate();
dispatcher.RunUntilStalled();
failing_allocator.EnableAllocate();
EXPECT_EQ(echo_service.last_dispatched_method(), 0u);
EXPECT_EQ(raw_conn->commit_count(), commits_before + 1u);
auto decode_err =
internal::InboundPacket::Decode(raw_conn->last_written_buf());
ASSERT_TRUE(decode_err.ok());
EXPECT_EQ(
decode_err->type(),
(internal::PacketType::Make<flags::kServer, flags::kErrorTerminal>()));
EXPECT_EQ(decode_err->call_id(), 12u);
EXPECT_EQ(decode_err->server_error(),
internal::ServerError::kFailedToAllocateCall);
ControlTask close_task(server.Close());
dispatcher.Post(close_task);
dispatcher.RunUntilStalled();
EXPECT_TRUE(close_task.done());
}
TEST(ServerTest, ListenerKeepsAcceptingAfterConnectionAllocationFailure) {
allocator::test::AllocatorForTest<16384> backing_allocator;
allocator::test::FaultInjectingAllocator failing_allocator(backing_allocator);
async2::DispatcherForTest dispatcher;
Server server(failing_allocator, dispatcher);
MockServerListener transport;
EXPECT_EQ(server.RegisterListenerBlocking(transport), OkStatus());
server.Start();
dispatcher.RunUntilStalled();
// The first connection cannot be served and is dropped.
transport::ReliableDatagramSocket dropped_conn =
test::MakeMockConnection(backing_allocator).first;
failing_allocator.DisableAllocate();
transport.ResolveAccept(std::move(dropped_conn));
dispatcher.RunUntilStalled();
failing_allocator.EnableAllocate();
// With no control operation or other wake in between, the listener must
// already have re-armed `Accept()`, so the next connection is served.
auto [conn, raw_conn] = test::MakeMockConnection(backing_allocator);
transport.ResolveAccept(conn);
dispatcher.RunUntilStalled();
Buf syn_buf = Buf::Allocate(backing_allocator,
internal::HandshakePacket::kWireSizeBytes);
auto syn = internal::HandshakePacket(internal::HandshakePacket::Type::kSyn)
.Encode(std::move(syn_buf));
ASSERT_TRUE(syn.ok());
raw_conn->SetNextRead(std::move(*syn));
dispatcher.RunUntilStalled();
EXPECT_EQ(raw_conn->commit_count(), 1u); // SYN-ACK
ControlTask close_task(server.Close());
dispatcher.Post(close_task);
dispatcher.RunUntilStalled();
EXPECT_TRUE(close_task.done());
}
TEST(ServerTest, InvalidRequestPayloadRepliesError) {
allocator::test::AllocatorForTest<16384> allocator;
async2::DispatcherForTest dispatcher;
Server server(allocator, dispatcher);
TestEchoService echo_service(100u);
EXPECT_EQ(server.RegisterServiceBlocking(echo_service), OkStatus());
MockServerListener transport;
EXPECT_EQ(server.RegisterListenerBlocking(transport), OkStatus());
server.Start();
auto [conn, raw_conn] = test::MakeMockConnection(allocator);
EstablishServerConnection(
allocator, dispatcher, server, transport, raw_conn, conn);
const size_t commits_before = raw_conn->commit_count();
// Method 2 rejects every request payload.
std::byte payload[1] = {std::byte{1}};
auto pkt = internal::PacketFramer::FrameStartUnaryPacket(allocator,
/*call_id=*/13u,
/*service_id=*/100u,
/*method_id=*/2u,
payload);
ASSERT_TRUE(pkt.ok());
raw_conn->SetNextRead(std::move(*pkt));
dispatcher.RunUntilStalled();
EXPECT_EQ(echo_service.last_dispatched_method(), 2u);
EXPECT_EQ(raw_conn->commit_count(), commits_before + 1u);
auto decode_err =
internal::InboundPacket::Decode(raw_conn->last_written_buf());
ASSERT_TRUE(decode_err.ok());
EXPECT_EQ(
decode_err->type(),
(internal::PacketType::Make<flags::kServer, flags::kErrorTerminal>()));
EXPECT_EQ(decode_err->call_id(), 13u);
EXPECT_EQ(decode_err->server_error(),
internal::ServerError::kInvalidRequestPayload);
ControlTask close_task(server.Close());
dispatcher.Post(close_task);
dispatcher.RunUntilStalled();
EXPECT_TRUE(close_task.done());
}
// Service with a bidirectional-streaming method that records the size of every
// message it reads, then drops its writer, which ends the RPC.
class StreamRecordingService : public Service {
public:
explicit StreamRecordingService(uint32_t service_id)
: Service(service_id, methods_),
methods_({
internal::RawMethodInvoker<&StreamRecordingService::Record,
MethodType::kBidirectionalStreaming>::
CreateMethod<StreamRecordingService>(1u),
}) {}
span<const size_t> message_sizes() const {
return span(message_sizes_).first(message_count_);
}
bool finished() const { return finished_; }
Status end_status() const { return end_status_; }
// True once the method's coroutine frame has been destroyed, whether it ran
// to completion or its call was aborted and torn down.
bool frame_destroyed() const { return frame_destroyed_; }
private:
// Sets a flag when the coroutine frame that holds it is destroyed.
class FrameGuard {
public:
explicit FrameGuard(bool& destroyed) : destroyed_(destroyed) {}
~FrameGuard() { destroyed_ = true; }
private:
bool& destroyed_;
};
async2::Coro<void> Record(async2::CoroContext,
RawReader reader,
RawWriter writer) {
FrameGuard guard(frame_destroyed_);
static_cast<void>(writer);
while (true) {
Result<ConstBuf> msg = co_await reader.Read();
if (!msg.ok()) {
end_status_ = msg.status();
break;
}
PW_CHECK_UINT_LT(message_count_, message_sizes_.size());
message_sizes_[message_count_++] = msg->size();
}
finished_ = true;
co_return;
}
std::array<internal::Method, 1> methods_;
std::array<size_t, 4> message_sizes_{};
size_t message_count_ = 0;
bool finished_ = false;
bool frame_destroyed_ = false;
Status end_status_;
};
class ServerStartPacketTest : public ::testing::Test {
protected:
static constexpr uint32_t kEchoServiceId = 100u;
static constexpr uint32_t kStreamServiceId = 200u;
ServerStartPacketTest() {
PW_CHECK_OK(server_.RegisterServiceBlocking(echo_service_));
PW_CHECK_OK(server_.RegisterServiceBlocking(stream_service_));
PW_CHECK_OK(server_.RegisterListenerBlocking(transport_));
server_.Start();
auto [conn, raw_conn] = test::MakeMockConnection(allocator_);
raw_conn_ = raw_conn;
EstablishServerConnection(
allocator_, dispatcher_, server_, transport_, raw_conn_, conn);
}
~ServerStartPacketTest() override {
ControlTask close_task(server_.Close());
dispatcher_.Post(close_task);
dispatcher_.RunUntilStalled();
EXPECT_TRUE(close_task.done());
}
void Receive(Result<Buf>&& packet) {
PW_CHECK_OK(packet.status());
raw_conn_->SetNextRead(std::move(*packet));
dispatcher_.RunUntilStalled();
}
void ReceiveStart(internal::PacketType type,
uint32_t call_id,
uint32_t service_id,
ConstByteSpan payload = {}) {
Receive(internal::PacketFramer::FrameStartPacket(
allocator_, type, call_id, service_id, /*method_id=*/1u, payload));
}
// Decodes the most recent packet the server sent.
internal::InboundPacket LastSent() {
auto decoded =
internal::InboundPacket::Decode(raw_conn_->last_written_buf());
PW_CHECK_OK(decoded.status());
return std::move(*decoded);
}
allocator::test::AllocatorForTest<8192> allocator_;
async2::DispatcherForTest dispatcher_;
Server server_{allocator_, dispatcher_};
TestEchoService echo_service_{kEchoServiceId};
StreamRecordingService stream_service_{kStreamServiceId};
MockServerListener transport_;
test::MockConnection* raw_conn_ = nullptr;
static constexpr std::byte kPayload[4] = {
std::byte{1}, std::byte{2}, std::byte{3}, std::byte{4}};
};
TEST_F(ServerStartPacketTest,
StartStreamToUnaryMethodRepliesMethodTypeMismatch) {
size_t commits_before = raw_conn_->commit_count();
ReceiveStart(internal::PacketType::Make<flags::kStart>(),
/*call_id=*/1,
kEchoServiceId);
EXPECT_EQ(echo_service_.last_dispatched_method(), 0u);
ASSERT_EQ(raw_conn_->commit_count(), commits_before + 1u);
internal::InboundPacket pkt = LastSent();
EXPECT_EQ(
pkt.type(),
(internal::PacketType::Make<flags::kServer, flags::kErrorTerminal>()));
EXPECT_EQ(pkt.call_id(), 1u);
EXPECT_EQ(pkt.server_error(), internal::ServerError::kMethodTypeMismatch);
}
TEST_F(ServerStartPacketTest,
StartStreamWithPayloadToUnaryMethodRepliesMethodTypeMismatch) {
size_t commits_before = raw_conn_->commit_count();
ReceiveStart(internal::PacketType::Make<flags::kStart, flags::kHasPayload>(),
/*call_id=*/2,
kEchoServiceId,
kPayload);
EXPECT_EQ(echo_service_.last_dispatched_method(), 0u);
ASSERT_EQ(raw_conn_->commit_count(), commits_before + 1u);
internal::InboundPacket pkt = LastSent();
EXPECT_EQ(
pkt.type(),
(internal::PacketType::Make<flags::kServer, flags::kErrorTerminal>()));
EXPECT_EQ(pkt.call_id(), 2u);
EXPECT_EQ(pkt.server_error(), internal::ServerError::kMethodTypeMismatch);
}
TEST_F(ServerStartPacketTest, StreamEndWithoutPayloadToUnaryMethodIsMismatch) {
size_t commits_before = raw_conn_->commit_count();
ReceiveStart(internal::PacketType::Make<flags::kStart, flags::kStreamEnd>(),
/*call_id=*/3,
kEchoServiceId);
EXPECT_EQ(echo_service_.last_dispatched_method(), 0u);
ASSERT_EQ(raw_conn_->commit_count(), commits_before + 1u);
EXPECT_EQ(LastSent().server_error(),
internal::ServerError::kMethodTypeMismatch);
}
TEST_F(ServerStartPacketTest, StartStreamThenMessagesToStreamingMethod) {
ReceiveStart(internal::PacketType::Make<flags::kStart>(),
/*call_id=*/4,
kStreamServiceId);
EXPECT_TRUE(stream_service_.message_sizes().empty());
EXPECT_FALSE(stream_service_.finished());
Receive(internal::PacketFramer::FrameClientMessagePacket(
allocator_, /*call_id=*/4, span(kPayload).first(2)));
Receive(internal::PacketFramer::FrameClientStreamEndPacket(allocator_,
/*call_id=*/4));
ASSERT_EQ(stream_service_.message_sizes().size(), 1u);
EXPECT_EQ(stream_service_.message_sizes()[0], 2u);
EXPECT_TRUE(stream_service_.finished());
EXPECT_EQ(stream_service_.end_status(), Status::OutOfRange());
// Dropping the writer terminates the RPC.
internal::InboundPacket pkt = LastSent();
EXPECT_EQ(pkt.type(),
(internal::PacketType::Make<flags::kServer, flags::kOkTerminal>()));
EXPECT_EQ(pkt.call_id(), 4u);
}
TEST_F(ServerStartPacketTest, StartStreamWithPayloadDeliversInitialMessage) {
ReceiveStart(internal::PacketType::Make<flags::kStart, flags::kHasPayload>(),
/*call_id=*/5,
kStreamServiceId,
kPayload);
ASSERT_EQ(stream_service_.message_sizes().size(), 1u);
EXPECT_EQ(stream_service_.message_sizes()[0], sizeof(kPayload));
EXPECT_FALSE(stream_service_.finished());
Receive(internal::PacketFramer::FrameClientMessagePacket(
allocator_, /*call_id=*/5, span(kPayload).first(2)));
Receive(internal::PacketFramer::FrameClientStreamEndPacket(allocator_,
/*call_id=*/5));
ASSERT_EQ(stream_service_.message_sizes().size(), 2u);
EXPECT_EQ(stream_service_.message_sizes()[1], 2u);
EXPECT_TRUE(stream_service_.finished());
EXPECT_EQ(stream_service_.end_status(), Status::OutOfRange());
}
TEST_F(ServerStartPacketTest, RequestToStreamingMethodDeliversMessageAndEnd) {
Receive(internal::PacketFramer::FrameStartUnaryPacket(allocator_,
/*call_id=*/6,
kStreamServiceId,
/*method_id=*/1u,
kPayload));
ASSERT_EQ(stream_service_.message_sizes().size(), 1u);
EXPECT_EQ(stream_service_.message_sizes()[0], sizeof(kPayload));
EXPECT_TRUE(stream_service_.finished());
EXPECT_EQ(stream_service_.end_status(), Status::OutOfRange());
internal::InboundPacket pkt = LastSent();
EXPECT_EQ(pkt.type(),
(internal::PacketType::Make<flags::kServer, flags::kOkTerminal>()));
EXPECT_EQ(pkt.call_id(), 6u);
}
TEST_F(ServerStartPacketTest, StartWithStreamEndClosesStreamImmediately) {
ReceiveStart(internal::PacketType::Make<flags::kStart, flags::kStreamEnd>(),
/*call_id=*/7,
kStreamServiceId);
EXPECT_TRUE(stream_service_.message_sizes().empty());
EXPECT_TRUE(stream_service_.finished());
EXPECT_EQ(stream_service_.end_status(), Status::OutOfRange());
internal::InboundPacket pkt = LastSent();
EXPECT_EQ(pkt.type(),
(internal::PacketType::Make<flags::kServer, flags::kOkTerminal>()));
EXPECT_EQ(pkt.call_id(), 7u);
}
// A client that detects that the server's responses don't match the method's
// type reports it, which aborts the call on the server too. Like any aborted
// call, it is torn down without resuming the method, and nothing is sent back.
TEST_F(ServerStartPacketTest, ClientMethodTypeMismatchAbortsServerCall) {
ReceiveStart(internal::PacketType::Make<flags::kStart>(),
/*call_id=*/8,
kStreamServiceId);
ASSERT_FALSE(stream_service_.frame_destroyed());
const size_t commits_before = raw_conn_->commit_count();
Receive(internal::PacketFramer::FrameClientErrorPacket(
allocator_, /*call_id=*/8, internal::ClientError::kMethodTypeMismatch));
EXPECT_TRUE(stream_service_.frame_destroyed());
EXPECT_FALSE(stream_service_.finished());
EXPECT_EQ(raw_conn_->commit_count(), commits_before);
}
} // namespace
} // namespace pw::rpc2