Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
125 changes: 7 additions & 118 deletions MODULE.bazel.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

34 changes: 23 additions & 11 deletions asio_rpc/server/rpc_server.cc
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
#include "proto/subspace.pb.h"
#include <inttypes.h>
#include <stdio.h>
#include <unistd.h>

namespace subspace::asio_rpc {

Expand Down Expand Up @@ -763,27 +764,36 @@ void RpcServer::SessionStreamingMethodCoroutine(
method_instance->method->name.c_str());

AnyStreamWriter writer(server, session, method_instance, request);
auto stop_pipe_or = toolbelt::Pipe::Create();
if (!stop_pipe_or.ok()) {
server->logger_.Log(toolbelt::LogLevel::kError,
"Failed to create stream cancellation stop pipe: %s",
stop_pipe_or.status().ToString().c_str());
continue;
}
auto stop_pipe =
std::make_shared<toolbelt::Pipe>(std::move(*stop_pipe_or));
auto writer_state = writer.state;
const int32_t session_id = session->session_id;
const int32_t request_id = request.request_id();

// Spawn a coroutine to read the cancellation channel.
boost::asio::spawn(
*server->io_context_,
[server, session, method_instance, &writer,
&request](boost::asio::yield_context yield) {
int dup_interrupt =
::dup(server->interrupt_pipe_.ReadFd().Fd());
toolbelt::FileDescriptor interrupt(dup_interrupt);
while (!writer.IsCancelled()) {
[server, method_instance, stop_pipe, writer_state, session_id,
request_id](boost::asio::yield_context yield) {
while (!writer_state->is_cancelled.load(std::memory_order_acquire)) {
int cancel_fd =
method_instance->cancel_subscriber->GetPollFd().fd;
auto s = async_wait_either(*server->io_context_, cancel_fd,
interrupt.Fd(), yield);
stop_pipe->ReadFd().Fd(), yield);
if (!s.ok()) {
server->logger_.Log(toolbelt::LogLevel::kError,
"Error waiting for cancel: %s",
s.status().ToString().c_str());
return;
}
if (*s == interrupt.Fd()) {
if (*s == stop_pipe->ReadFd().Fd()) {
break;
}
bool cancel_ok = false;
Expand All @@ -805,14 +815,14 @@ void RpcServer::SessionStreamingMethodCoroutine(
msg.status().ToString().c_str());
continue;
}
if (cancel.session_id() == session->session_id &&
cancel.request_id() == request.request_id()) {
if (cancel.session_id() == session_id &&
cancel.request_id() == request_id) {
cancel_ok = true;
break;
}
}
if (cancel_ok) {
writer.Cancel();
writer_state->is_cancelled.store(true, std::memory_order_release);
}
}
},
Expand All @@ -832,6 +842,8 @@ void RpcServer::SessionStreamingMethodCoroutine(
method_status.ToString()),
yield);
}
char stop = 1;
(void)::write(stop_pipe->WriteFd().Fd(), &stop, 1);
}
}

Expand Down
18 changes: 13 additions & 5 deletions asio_rpc/server/rpc_server.h
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
#include "toolbelt/pipe.h"

#include <boost/asio.hpp>
#include <atomic>
#include <boost/asio/spawn.hpp>

namespace subspace::asio_rpc {
Expand All @@ -25,27 +26,34 @@ namespace internal {
struct Session;
struct MethodInstance;

struct StreamWriterState {
std::atomic_bool is_cancelled{false};
};

struct AnyStreamWriter {
AnyStreamWriter(std::shared_ptr<RpcServer> server,
std::shared_ptr<Session> session,
std::shared_ptr<MethodInstance> method_instance,
const RpcRequest &request)
: server(std::move(server)), session(std::move(session)),
method_instance(std::move(method_instance)), request(request) {}
method_instance(std::move(method_instance)), request(request),
state(std::make_shared<StreamWriterState>()) {}

bool Write(std::unique_ptr<google::protobuf::Any> res,
boost::asio::yield_context yield);
void Finish(boost::asio::yield_context yield);

void Cancel() { is_cancelled = true; }
void Cancel() { state->is_cancelled.store(true, std::memory_order_release); }

bool IsCancelled() const { return is_cancelled; }
bool IsCancelled() const {
return state->is_cancelled.load(std::memory_order_acquire);
}

std::shared_ptr<RpcServer> server;
std::shared_ptr<Session> session;
std::shared_ptr<MethodInstance> method_instance;
const RpcRequest &request;
bool is_cancelled = false;
RpcRequest request;
std::shared_ptr<StreamWriterState> state;
};

// An item pushed by either the reply function or the error function of an
Expand Down
12 changes: 12 additions & 0 deletions asio_rpc/server/server_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -133,6 +133,18 @@ static std::shared_ptr<subspace::asio_rpc::RpcServer> BuildServer() {
return server;
}

TEST_F(AsioServerTest, StreamWriterOwnsRequestMetadata) {
std::unique_ptr<subspace::asio_rpc::internal::AnyStreamWriter> writer;
{
subspace::RpcRequest request;
request.set_request_id(1234);
writer = std::make_unique<subspace::asio_rpc::internal::AnyStreamWriter>(
nullptr, nullptr, nullptr, request);
}

EXPECT_EQ(1234, writer->request.request_id());
}

struct ServerContext {
std::shared_ptr<subspace::Client> client;
std::shared_ptr<subspace::Publisher> pub;
Expand Down
28 changes: 20 additions & 8 deletions client/client.cc
Original file line number Diff line number Diff line change
Expand Up @@ -885,9 +885,11 @@ ClientImpl::WaitForReliablePublisher(PublisherImpl *publisher,
uint64_t timeout_ns = timeout.count();
#if SUBSPACE_CORO_BACKEND == SUBSPACE_CORO_BACKEND_ASIO
if (IsCooperative()) {
// The Asio WaitEither has no timeout; it waits until one fd is ready.
absl::StatusOr<int> r = async::WaitEither(
SocketContext(), publisher->GetPollFd().Fd(), fd.Fd());
SocketContext(), publisher->GetPollFd().Fd(), fd.Fd(), timeout);
if (absl::IsDeadlineExceeded(r.status())) {
return absl::InternalError("Timeout waiting for reliable publisher");
}
if (!r.ok()) {
return r.status();
}
Expand Down Expand Up @@ -984,9 +986,11 @@ absl::StatusOr<int> ClientImpl::WaitForSubscriber(
uint64_t timeout_ns = timeout.count();
#if SUBSPACE_CORO_BACKEND == SUBSPACE_CORO_BACKEND_ASIO
if (IsCooperative()) {
// The Asio WaitEither has no timeout; it waits until one fd is ready.
absl::StatusOr<int> r = async::WaitEither(
SocketContext(), subscriber->GetPollFd().Fd(), fd.Fd());
SocketContext(), subscriber->GetPollFd().Fd(), fd.Fd(), timeout);
if (absl::IsDeadlineExceeded(r.status())) {
return absl::InternalError("Timeout waiting for subscriber");
}
if (!r.ok()) {
return r.status();
}
Expand Down Expand Up @@ -1073,8 +1077,12 @@ ClientImpl::WaitForReliablePublisher(PublisherImpl *publisher,
!status.ok()) {
return status;
}
// The Asio WaitEither has no timeout; it waits until one fd is ready.
return async::WaitEither(ctx, publisher->GetPollFd().Fd(), fd.Fd());
absl::StatusOr<int> r =
async::WaitEither(ctx, publisher->GetPollFd().Fd(), fd.Fd(), timeout);
if (absl::IsDeadlineExceeded(r.status())) {
return absl::InternalError("Timeout waiting for reliable publisher");
}
return r;
}

absl::Status ClientImpl::WaitForSubscriber(SubscriberImpl *subscriber,
Expand All @@ -1097,8 +1105,12 @@ absl::StatusOr<int> ClientImpl::WaitForSubscriber(
if (absl::Status status = CheckConnected(); !status.ok()) {
return status;
}
// The Asio WaitEither has no timeout; it waits until one fd is ready.
return async::WaitEither(ctx, subscriber->GetPollFd().Fd(), fd.Fd());
absl::StatusOr<int> r =
async::WaitEither(ctx, subscriber->GetPollFd().Fd(), fd.Fd(), timeout);
if (absl::IsDeadlineExceeded(r.status())) {
return absl::InternalError("Timeout waiting for subscriber");
}
return r;
}
#endif

Expand Down
Loading
Loading