mod error;
pub use error::RpcServerError;
mod handle;
pub use handle::RpcServerHandle;
use handle::RpcServerRequest;
#[cfg(feature = "metrics")]
mod metrics;
pub mod mock;
mod early_close;
mod router;
use std::{
borrow::Cow,
cmp,
collections::HashMap,
convert::TryFrom,
future::Future,
io,
io::ErrorKind,
pin::Pin,
sync::Arc,
task::Poll,
time::{Duration, Instant},
};
use futures::{SinkExt, StreamExt, future, stream::FuturesUnordered};
use log::*;
use prost::Message;
use router::Router;
use tokio::{
sync::mpsc,
task::{JoinHandle, JoinSet},
time,
};
use tokio_stream::Stream;
use tower::{Service, make::MakeService};
use tracing::{Instrument, Level, debug, error, instrument, span, trace, warn};
use super::{
Handshake,
RPC_MAX_FRAME_SIZE,
body::Body,
context::{RequestContext, RpcCommsProvider},
error::HandshakeRejectReason,
message::{Request, Response, RpcMessageFlags},
not_found::ProtocolServiceNotFound,
status::RpcStatus,
};
use crate::{
Bytes,
Substream,
bounded_executor::BoundedExecutor,
framing,
framing::CanonicalFraming,
message::MessageExt,
peer_manager::NodeId,
proto,
protocol::{
ProtocolEvent,
ProtocolId,
ProtocolNotification,
ProtocolNotificationRx,
rpc,
rpc::{
body::BodyBytes,
message::{RpcMethod, RpcResponse},
server::early_close::EarlyClose,
},
},
stream_id::{Id, StreamId},
};
const LOG_TARGET: &str = "comms::rpc::server";
pub trait NamedProtocolService {
const PROTOCOL_NAME: &'static [u8];
fn as_protocol_name(&self) -> &'static [u8] {
Self::PROTOCOL_NAME
}
}
pub struct RpcServer {
builder: RpcServerBuilder,
request_tx: mpsc::Sender<RpcServerRequest>,
request_rx: mpsc::Receiver<RpcServerRequest>,
}
impl RpcServer {
pub fn new() -> Self {
Self::builder().finish()
}
pub fn builder() -> RpcServerBuilder {
RpcServerBuilder::new()
}
pub fn add_service<S>(self, service: S) -> Router<S, ProtocolServiceNotFound>
where
S: MakeService<
ProtocolId,
Request<Bytes>,
MakeError = RpcServerError,
Response = Response<Body>,
Error = RpcStatus,
> + NamedProtocolService
+ Send
+ 'static,
S::Future: Send + 'static,
{
Router::new(self, service)
}
pub fn get_handle(&self) -> RpcServerHandle {
RpcServerHandle::new(self.request_tx.clone())
}
pub(super) async fn serve<S, TCommsProvider>(
self,
service: S,
notifications: ProtocolNotificationRx<Substream>,
comms_provider: TCommsProvider,
) -> Result<(), RpcServerError>
where
S: MakeService<
ProtocolId,
Request<Bytes>,
MakeError = RpcServerError,
Response = Response<Body>,
Error = RpcStatus,
> + Send
+ 'static,
S::Service: Send + 'static,
S::Future: Send + 'static,
S::Service: Send + 'static,
<S::Service as Service<Request<Bytes>>>::Future: Send + 'static,
TCommsProvider: RpcCommsProvider + Clone + Send + 'static,
{
PeerRpcServer::new(self.builder, service, notifications, comms_provider, self.request_rx)
.serve()
.await
}
}
impl Default for RpcServer {
fn default() -> Self {
Self::new()
}
}
const DEFAULT_IDLE_SESSION_TIMEOUT: Duration = Duration::from_secs(10 * 60);
const DEFAULT_MAXIMUM_CLIENT_DEADLINE: Duration = Duration::from_secs(10 * 60);
const SESSION_CLOSE_TIMEOUT: Duration = Duration::from_secs(5);
#[derive(Clone)]
pub struct RpcServerBuilder {
maximum_simultaneous_sessions: Option<usize>,
maximum_sessions_per_client: Option<usize>,
minimum_client_deadline: Duration,
maximum_client_deadline: Duration,
handshake_timeout: Duration,
maximum_pending_handshakes: usize,
maximum_pending_handshakes_per_client: usize,
idle_session_timeout: Option<Duration>,
cull_oldest_peer_rpc_connection_on_full: bool,
}
impl RpcServerBuilder {
fn new() -> Self {
Default::default()
}
pub fn with_maximum_simultaneous_sessions(mut self, limit: usize) -> Self {
self.maximum_simultaneous_sessions = Some(cmp::min(limit, BoundedExecutor::max_theoretical_tasks()));
self
}
pub fn with_unlimited_simultaneous_sessions(mut self) -> Self {
self.maximum_simultaneous_sessions = None;
self
}
pub fn with_maximum_sessions_per_client(mut self, limit: usize) -> Self {
self.maximum_sessions_per_client = Some(cmp::min(limit, BoundedExecutor::max_theoretical_tasks()));
self
}
pub fn with_cull_oldest_peer_rpc_connection_on_full(mut self, cull: bool) -> Self {
self.cull_oldest_peer_rpc_connection_on_full = cull;
self
}
pub fn with_unlimited_sessions_per_client(mut self) -> Self {
self.maximum_sessions_per_client = None;
self
}
pub fn with_minimum_client_deadline(mut self, deadline: Duration) -> Self {
self.minimum_client_deadline = deadline;
self
}
pub fn with_maximum_client_deadline(mut self, deadline: Duration) -> Self {
self.maximum_client_deadline = deadline;
self
}
fn clamp_client_deadline(&self, requested: Duration) -> Duration {
let capped = cmp::min(requested, self.maximum_client_deadline);
cmp::max(capped, self.minimum_client_deadline)
}
pub fn with_maximum_pending_handshakes(mut self, limit: usize) -> Self {
self.maximum_pending_handshakes = limit;
self
}
pub fn with_maximum_pending_handshakes_per_client(mut self, limit: usize) -> Self {
self.maximum_pending_handshakes_per_client = limit;
self
}
pub fn with_handshake_timeout(mut self, timeout: Duration) -> Self {
self.handshake_timeout = timeout;
self
}
pub fn with_idle_session_timeout(mut self, timeout: Duration) -> Self {
self.idle_session_timeout = Some(timeout);
self
}
pub fn with_no_idle_session_timeout(mut self) -> Self {
self.idle_session_timeout = None;
self
}
pub fn finish(self) -> RpcServer {
let (request_tx, request_rx) = mpsc::channel(10);
RpcServer {
builder: self,
request_tx,
request_rx,
}
}
}
const DEFAULT_MAXIMUM_SIMULTANEOUS_SESSIONS: usize = 100;
const DEFAULT_MAXIMUM_SESSIONS_PER_CLIENT: usize = 10;
const DEFAULT_HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(10);
const DEFAULT_MAXIMUM_PENDING_HANDSHAKES: usize = 32;
const DEFAULT_MAXIMUM_PENDING_HANDSHAKES_PER_CLIENT: usize = 4;
const HANDSHAKE_WARNING_INTERVAL: Duration = Duration::from_secs(60);
impl Default for RpcServerBuilder {
fn default() -> Self {
Self {
maximum_simultaneous_sessions: Some(DEFAULT_MAXIMUM_SIMULTANEOUS_SESSIONS),
maximum_sessions_per_client: Some(DEFAULT_MAXIMUM_SESSIONS_PER_CLIENT),
minimum_client_deadline: Duration::from_secs(1),
maximum_client_deadline: DEFAULT_MAXIMUM_CLIENT_DEADLINE,
handshake_timeout: DEFAULT_HANDSHAKE_TIMEOUT,
maximum_pending_handshakes: DEFAULT_MAXIMUM_PENDING_HANDSHAKES,
maximum_pending_handshakes_per_client: DEFAULT_MAXIMUM_PENDING_HANDSHAKES_PER_CLIENT,
idle_session_timeout: Some(DEFAULT_IDLE_SESSION_TIMEOUT),
cull_oldest_peer_rpc_connection_on_full: false,
}
}
}
pub(super) struct PeerRpcServer<TSvc, TCommsProvider>
where TSvc: MakeService<ProtocolId, Request<Bytes>>
{
executor: BoundedExecutor,
config: RpcServerBuilder,
service: TSvc,
protocol_notifications: Option<ProtocolNotificationRx<Substream>>,
comms_provider: TCommsProvider,
request_rx: mpsc::Receiver<RpcServerRequest>,
sessions: HashMap<NodeId, Vec<SessionInfo>>,
tasks: FuturesUnordered<JoinHandle<(NodeId, Id)>>,
handshakes: JoinSet<HandshakeOutcome<<TSvc as MakeService<ProtocolId, Request<Bytes>>>::Service>>,
handshake_peers: HashMap<tokio::task::Id, (NodeId, bool)>,
pending_handshakes: HashMap<NodeId, usize>,
accepted_pending: HashMap<NodeId, usize>,
accepted_pending_total: usize,
capacity_rejections: JoinSet<()>,
capacity_warning: RateLimitedWarning,
timeout_warning: RateLimitedWarning,
session_limit_warning: RateLimitedWarning,
}
#[derive(Default)]
struct RateLimitedWarning {
last: Option<Instant>,
suppressed: usize,
}
impl RateLimitedWarning {
fn occurred(&mut self) -> Option<usize> {
let now = Instant::now();
if self
.last
.is_some_and(|last| now.saturating_duration_since(last) < HANDSHAKE_WARNING_INTERVAL)
{
self.suppressed = self.suppressed.saturating_add(1);
return None;
}
self.last = Some(now);
let count = self.suppressed.saturating_add(1);
self.suppressed = 0;
Some(count)
}
}
struct HandshakeOutcome<S> {
protocol: ProtocolId,
node_id: NodeId,
result: Result<(S, CanonicalFraming<Substream>, u32), RpcServerError>,
}
enum SessionDecision<S> {
Accept(S),
Reject(HandshakeRejectReason, RpcServerError),
}
struct SessionInfo {
pub(crate) peer_watch: tokio::sync::watch::Sender<()>,
pub(crate) stream_id: Id,
}
impl<TSvc, TCommsProvider> PeerRpcServer<TSvc, TCommsProvider>
where
TSvc: MakeService<
ProtocolId,
Request<Bytes>,
MakeError = RpcServerError,
Response = Response<Body>,
Error = RpcStatus,
> + Send
+ 'static,
TSvc::Service: Send + 'static,
<TSvc::Service as Service<Request<Bytes>>>::Future: Send + 'static,
TSvc::Future: Send + 'static,
TCommsProvider: RpcCommsProvider + Clone + Send + 'static,
{
fn new(
config: RpcServerBuilder,
service: TSvc,
protocol_notifications: ProtocolNotificationRx<Substream>,
comms_provider: TCommsProvider,
request_rx: mpsc::Receiver<RpcServerRequest>,
) -> Self {
Self {
executor: match config.maximum_simultaneous_sessions {
Some(usize::MAX) => BoundedExecutor::allow_maximum(),
Some(num) => BoundedExecutor::new(num),
None => BoundedExecutor::allow_maximum(),
},
config,
service,
protocol_notifications: Some(protocol_notifications),
comms_provider,
request_rx,
sessions: HashMap::new(),
tasks: FuturesUnordered::new(),
handshakes: JoinSet::new(),
handshake_peers: HashMap::new(),
pending_handshakes: HashMap::new(),
accepted_pending: HashMap::new(),
accepted_pending_total: 0,
capacity_rejections: JoinSet::new(),
capacity_warning: RateLimitedWarning::default(),
timeout_warning: RateLimitedWarning::default(),
session_limit_warning: RateLimitedWarning::default(),
}
}
pub async fn serve(mut self) -> Result<(), RpcServerError> {
let mut protocol_notifs = self
.protocol_notifications
.take()
.expect("PeerRpcServer initialized without protocol_notifications");
loop {
tokio::select! {
maybe_notif = protocol_notifs.recv() => {
match maybe_notif {
Some(notif) => self.handle_protocol_notification(notif).await?,
None => break,
}
}
Some(Ok((node_id, stream_id))) = self.tasks.next() => {
self.on_session_complete(&node_id, stream_id);
},
Some(joined) = self.handshakes.join_next_with_id() => {
self.on_handshake_finished(joined);
},
Some(_) = self.capacity_rejections.join_next() => {},
Some(req) = self.request_rx.recv() => {
self.handle_request(req).await;
},
}
}
debug!(
target: LOG_TARGET,
"Peer RPC server is shut down because the protocol notification stream ended"
);
Ok(())
}
async fn handle_request(&mut self, req: RpcServerRequest) {
#[allow(clippy::enum_glob_use)]
use RpcServerRequest::*;
match req {
GetNumActiveSessions(reply) => {
let max_sessions = self
.config
.maximum_simultaneous_sessions
.unwrap_or_else(BoundedExecutor::max_theoretical_tasks);
let num_active = max_sessions.saturating_sub(self.executor.num_available());
let _ = reply.send(num_active);
},
GetNumActiveSessionsForPeer(node_id, reply) => {
let num_active = self.sessions.get(&node_id).map(|v| v.len()).unwrap_or(0);
let _ = reply.send(num_active);
},
CloseAllSessionsForPeer(node_id, reply) => {
let num_closed = self.close_all_sessions(&node_id);
let _ = reply.send(num_closed);
},
}
}
async fn handle_protocol_notification(
&mut self,
notification: ProtocolNotification<Substream>,
) -> Result<(), RpcServerError> {
match notification.event {
ProtocolEvent::NewInboundSubstream(node_id, substream) => {
trace!(
target: LOG_TARGET,
"New client connection for protocol `{}` from peer `{}`",
String::from_utf8_lossy(¬ification.protocol),
node_id
);
let framed = framing::canonical(substream, RPC_MAX_FRAME_SIZE);
let pending_for_peer = self.pending_handshakes.get(&node_id).copied().unwrap_or(0);
if pending_for_peer >= self.config.maximum_pending_handshakes_per_client {
trace!(
target: LOG_TARGET,
"Rejecting RPC substream from peer `{}`: it already has {} handshake(s) in progress",
node_id,
pending_for_peer
);
self.reject_at_capacity(¬ification.protocol, framed);
return Ok(());
}
if self.handshakes.len() >= self.config.maximum_pending_handshakes {
trace!(
target: LOG_TARGET,
"Rejecting RPC substream from peer `{}`: {} handshake(s) already in progress",
node_id,
self.handshakes.len()
);
self.reject_at_capacity(¬ification.protocol, framed);
return Ok(());
}
let decision = self.decide_session(notification.protocol.clone(), &node_id).await;
self.spawn_handshake(notification.protocol, node_id, framed, decision);
},
}
Ok(())
}
fn reject_at_capacity(&mut self, protocol: &ProtocolId, mut framed: CanonicalFraming<Substream>) {
#[cfg(feature = "metrics")]
metrics::handshake_capacity_rejection_counter(protocol).inc();
#[cfg(not(feature = "metrics"))]
let _ = protocol;
if let Some(count) = self.capacity_warning.occurred() {
warn!(
target: LOG_TARGET,
"Refused {count} RPC substream(s) because too many handshakes were in progress (at most {} per peer, \
{} in total)",
self.config.maximum_pending_handshakes_per_client,
self.config.maximum_pending_handshakes
);
}
if self.capacity_rejections.len() >= self.config.maximum_pending_handshakes {
return;
}
let timeout = self.config.handshake_timeout;
self.capacity_rejections.spawn(async move {
let _result = Handshake::new(&mut framed)
.with_timeout(timeout)
.reject_with_reason(HandshakeRejectReason::NoServerSessionsAvailable(
"Too many handshakes in progress",
))
.await;
});
}
async fn decide_session(&mut self, protocol: ProtocolId, node_id: &NodeId) -> SessionDecision<TSvc::Service> {
let service = match self.service.make_service(protocol).await {
Ok(service) => service,
Err(err) => {
trace!(
target: LOG_TARGET,
"Rejecting RPC session request for peer `{}` because {}",
node_id,
HandshakeRejectReason::ProtocolNotSupported
);
return SessionDecision::Reject(HandshakeRejectReason::ProtocolNotSupported, err);
},
};
let accepted_for_peer = self.accepted_pending.get(node_id).copied().unwrap_or(0);
if let Err(err) = self.new_session_possible_for(node_id, accepted_for_peer) {
return SessionDecision::Reject(
HandshakeRejectReason::NoServerSessionsAvailable("Maximum sessions for client"),
err,
);
}
if self.executor.num_available() <= self.accepted_pending_total {
let msg = format!("Used all {} sessions", self.executor.max_available());
trace!(
target: LOG_TARGET,
"Rejecting RPC session request for peer `{}` because {}",
node_id,
HandshakeRejectReason::NoServerSessionsAvailable("Cannot spawn more sessions")
);
return SessionDecision::Reject(
HandshakeRejectReason::NoServerSessionsAvailable("Cannot spawn more sessions"),
RpcServerError::MaximumSessionsReached(msg),
);
}
SessionDecision::Accept(service)
}
fn spawn_handshake(
&mut self,
protocol: ProtocolId,
node_id: NodeId,
mut framed: CanonicalFraming<Substream>,
decision: SessionDecision<TSvc::Service>,
) {
let timeout = self.config.handshake_timeout;
let task_node_id = node_id.clone();
let accepted = matches!(decision, SessionDecision::Accept(_));
let handle = self.handshakes.spawn(async move {
let result = match decision {
SessionDecision::Reject(reason, err) => {
let _result = Handshake::new(&mut framed)
.with_timeout(timeout)
.reject_with_reason(reason)
.await;
Err(err)
},
SessionDecision::Accept(service) => {
match Handshake::new(&mut framed)
.with_timeout(timeout)
.receive_client_handshake()
.await
{
Ok(version) => {
debug!(
target: LOG_TARGET,
"Server negotiated RPC v{version} with client node `{task_node_id}`"
);
Ok((service, framed, version))
},
Err(err) => Err(err.into()),
}
},
};
HandshakeOutcome {
protocol,
node_id: task_node_id,
result,
}
});
self.handshake_peers.insert(handle.id(), (node_id.clone(), accepted));
if accepted {
let reserved = self.accepted_pending.entry(node_id.clone()).or_insert(0);
*reserved = reserved.saturating_add(1);
self.accepted_pending_total = self.accepted_pending_total.saturating_add(1);
}
let pending = self.pending_handshakes.entry(node_id).or_insert(0);
*pending = pending.saturating_add(1);
}
fn on_handshake_finished(
&mut self,
joined: Result<(tokio::task::Id, HandshakeOutcome<TSvc::Service>), tokio::task::JoinError>,
) {
let task_id = match &joined {
Ok((id, _)) => *id,
Err(err) => err.id(),
};
if let Some((node_id, accepted)) = self.handshake_peers.remove(&task_id) {
decrement(&mut self.pending_handshakes, &node_id);
if accepted {
decrement(&mut self.accepted_pending, &node_id);
self.accepted_pending_total = self.accepted_pending_total.saturating_sub(1);
}
}
let outcome = match joined {
Ok((_, outcome)) => outcome,
Err(err) => {
error!(target: LOG_TARGET, "RPC handshake task failed: {err}");
return;
},
};
let HandshakeOutcome {
protocol,
node_id,
result,
} = outcome;
let (service, framed, version) = match result {
Ok(accepted) => accepted,
Err(err) => {
self.log_handshake_failure(&protocol, &err);
return;
},
};
if let Err(err) = self.start_session(protocol, &node_id, service, framed, version) {
warn!(
target: LOG_TARGET,
"BUG: an accepted RPC handshake from peer `{node_id}` could not start its session: {err}"
);
}
}
fn log_handshake_failure(&mut self, protocol: &ProtocolId, err: &RpcServerError) {
#[cfg(not(feature = "metrics"))]
let _ = protocol;
if matches!(err, RpcServerError::HandshakeError(rpc::RpcHandshakeError::TimedOut)) {
#[cfg(feature = "metrics")]
metrics::handshake_timeout_counter(protocol).inc();
if let Some(count) = self.timeout_warning.occurred() {
warn!(
target: LOG_TARGET,
"{count} RPC handshake(s) timed out after {:.0?} without a handshake from the client",
self.config.handshake_timeout
);
}
}
match err {
err @ RpcServerError::HandshakeError(_) => {
trace!(target: LOG_TARGET, "Handshake error: {}", err);
#[cfg(feature = "metrics")]
metrics::handshake_error_counter(protocol).inc();
},
err => {
trace!(target: LOG_TARGET, "Unable to spawn RPC service: {}", err);
},
}
}
fn new_session_possible_for(&mut self, node_id: &NodeId, reserved: usize) -> Result<usize, RpcServerError> {
let max = match self.config.maximum_sessions_per_client {
Some(max) if max > 0 => max,
Some(_) | None => return Ok(self.sessions.get(node_id).map_or(0, Vec::len)),
};
let running = self.sessions.get(node_id).map_or(0, Vec::len);
let in_use = running.saturating_add(reserved);
if in_use < max {
return Ok(running);
}
if self.config.cull_oldest_peer_rpc_connection_on_full && reserved < max {
let num_to_remove = in_use.saturating_sub(max).saturating_add(1);
if let Some(session_info) = self.sessions.get_mut(node_id) {
for _ in 0..num_to_remove {
let info = session_info.remove(0);
info!(target: LOG_TARGET, "Culling oldest RPC session for peer `{node_id}`");
let _ = info.peer_watch.send(());
}
let running = session_info.len();
if session_info.is_empty() {
self.sessions.remove(node_id);
}
return Ok(running);
}
}
trace!(
target: LOG_TARGET,
"Maximum RPC sessions for peer {node_id} met or exceeded. Max: {max}, running: {running}, accepted \
handshakes: {reserved}"
);
if let Some(count) = self.session_limit_warning.occurred() {
warn!(
target: LOG_TARGET,
"Refused {count} RPC session(s) because the peer was at its limit of {max} session(s)"
);
}
Err(RpcServerError::MaxSessionsPerClientReached {
node_id: node_id.clone(),
max_sessions: max,
})
}
fn close_all_sessions(&mut self, node_id: &NodeId) -> usize {
let mut count = 0usize;
if let Some(session_info) = self.sessions.get_mut(node_id) {
for info in session_info.iter_mut() {
count = count.saturating_add(1);
info!(target: LOG_TARGET, "Closing RPC session {} for peer `{}`", info.stream_id, node_id);
let _ = info.peer_watch.send(());
}
self.sessions.remove(node_id);
}
count
}
fn on_session_complete(&mut self, node_id: &NodeId, stream_id: Id) {
if let Some(session_info) = self.sessions.get_mut(node_id) {
if let Some(info) = session_info.iter_mut().find(|info| info.stream_id == stream_id) {
info!(target: LOG_TARGET, "Session complete for {node_id} (stream id {stream_id})");
let _ = info.peer_watch.send(());
};
session_info.retain(|info| info.stream_id != stream_id);
if session_info.is_empty() {
self.sessions.remove(node_id);
}
}
}
fn start_session(
&mut self,
protocol: ProtocolId,
node_id: &NodeId,
service: TSvc::Service,
mut framed: CanonicalFraming<Substream>,
version: u32,
) -> Result<(), RpcServerError> {
let num_sessions = self.new_session_possible_for(node_id, 0)?;
if !self.executor.can_spawn() {
return Err(RpcServerError::MaximumSessionsReached(format!(
"Used all {} sessions",
self.executor.max_available()
)));
}
info!(
target: LOG_TARGET,
"NEW SESSION for {node_id} ({num_sessions} currently active) "
);
let stream_id = framed.stream_id();
let (stop_tx, stop_rx) = tokio::sync::watch::channel(());
let config = self.config.clone();
let comms_provider = self.comms_provider.clone();
let session_node_id = node_id.clone();
let node_id_clone = node_id.clone();
let handle = self
.executor
.try_spawn(async move {
if let Err(err) = Handshake::new(&mut framed)
.with_timeout(config.handshake_timeout)
.accept_version(version)
.await
{
debug!(
target: LOG_TARGET,
"Could not send the RPC handshake acceptance to `{node_id_clone}`; ending the session: {err}"
);
return (node_id_clone, stream_id);
}
let service = ActivePeerRpcService::new(
config,
protocol,
session_node_id,
service,
framed,
comms_provider,
stop_rx,
);
#[cfg(feature = "metrics")]
let num_sessions = metrics::num_sessions(&service.protocol);
#[cfg(feature = "metrics")]
num_sessions.inc();
service.start().await;
info!(target: LOG_TARGET, "END OF SESSION for {node_id_clone} ");
#[cfg(feature = "metrics")]
num_sessions.dec();
(node_id_clone, stream_id)
})
.map_err(|e| RpcServerError::MaximumSessionsReached(format!("{e:?}")))?;
self.tasks.push(handle);
let mut peer_stop = vec![SessionInfo {
peer_watch: stop_tx,
stream_id,
}];
self.sessions
.entry(node_id.clone())
.and_modify(|entry| entry.append(&mut peer_stop))
.or_insert(peer_stop);
if let Some(info) = self.sessions.get(&node_id.clone()) {
info!(
target: LOG_TARGET,
"NEW SESSION created for {} ({} active) ", node_id.clone(), info.len()
);
if info.iter().filter(|session| session.stream_id == stream_id).count() > 1 {
warn!(
target: LOG_TARGET,
"Stream ID {stream_id} already in use for peer {node_id}. This should not happen."
);
}
}
Ok(())
}
}
fn decrement(counts: &mut HashMap<NodeId, usize>, node_id: &NodeId) {
if let Some(count) = counts.get_mut(node_id) {
*count = count.saturating_sub(1);
if *count == 0 {
counts.remove(node_id);
}
}
}
struct ActivePeerRpcService<TSvc, TCommsProvider> {
config: RpcServerBuilder,
protocol: ProtocolId,
node_id: NodeId,
service: TSvc,
framed: EarlyClose<CanonicalFraming<Substream>>,
comms_provider: TCommsProvider,
logging_context_string: Arc<String>,
stop_rx: tokio::sync::watch::Receiver<()>,
#[cfg(feature = "metrics")]
outbound_response_bytes: tari_metrics::Histogram,
}
impl<TSvc, TCommsProvider> ActivePeerRpcService<TSvc, TCommsProvider>
where
TSvc: Service<Request<Bytes>, Response = Response<Body>, Error = RpcStatus>,
TCommsProvider: RpcCommsProvider + Send + Clone + 'static,
{
pub(self) fn new(
config: RpcServerBuilder,
protocol: ProtocolId,
node_id: NodeId,
service: TSvc,
framed: CanonicalFraming<Substream>,
comms_provider: TCommsProvider,
stop_rx: tokio::sync::watch::Receiver<()>,
) -> Self {
Self {
logging_context_string: Arc::new(format!(
"stream_id: {}, peer: {}, protocol: {}, stream_id: {}",
framed.stream_id(),
node_id,
framed.stream_id(),
String::from_utf8_lossy(&protocol)
)),
#[cfg(feature = "metrics")]
outbound_response_bytes: metrics::outbound_response_bytes(&protocol),
config,
protocol,
node_id,
service,
framed: EarlyClose::new(framed),
comms_provider,
stop_rx,
}
}
async fn start(mut self) {
debug!(
target: LOG_TARGET,
"({}) Rpc server started.", self.logging_context_string,
);
if let Err(err) = self.run().await {
#[cfg(feature = "metrics")]
metrics::error_counter(&self.protocol, &err).inc();
let level = match &err {
RpcServerError::Io(e) => err_to_log_level(e),
RpcServerError::EarlyClose(e) => e.io().map(err_to_log_level).unwrap_or(log::Level::Error),
_ => log::Level::Error,
};
log!(
target: LOG_TARGET,
level,
"({}) Rpc server exited with an error: {}",
self.logging_context_string,
err
);
}
}
async fn send_with_deadline(&mut self, msg: Bytes, deadline: Duration) -> Result<(), RpcServerError> {
match time::timeout(deadline, self.framed.send(msg)).await {
Ok(result) => result.map_err(Into::into),
Err(_elapsed) => {
debug!(
target: LOG_TARGET,
"({}) Peer did not accept a response within the deadline ({:.0?}). Aborting the session.",
self.logging_context_string,
deadline
);
#[cfg(feature = "metrics")]
metrics::error_counter(&self.protocol, &RpcServerError::WriteStreamExceededDeadline).inc();
Err(RpcServerError::WriteStreamExceededDeadline)
},
}
}
async fn close_framed(&mut self) -> Result<(), RpcServerError> {
match time::timeout(SESSION_CLOSE_TIMEOUT, self.framed.close()).await {
Ok(result) => result.map_err(Into::into),
Err(_elapsed) => {
debug!(
target: LOG_TARGET,
"({}) Peer did not accept the substream close within {:.0?}. Dropping it.",
self.logging_context_string,
SESSION_CLOSE_TIMEOUT
);
Ok(())
},
}
}
async fn run(&mut self) -> Result<(), RpcServerError> {
let idle_timeout = self.config.idle_session_timeout;
let idle_sleep = time::sleep(idle_timeout.unwrap_or(DEFAULT_IDLE_SESSION_TIMEOUT));
tokio::pin!(idle_sleep);
loop {
tokio::select! {
_ = self.stop_rx.changed() => {
debug!(target: LOG_TARGET, "({}) Stop signal received, closing substream.", self.logging_context_string);
break;
}
() = &mut idle_sleep, if idle_timeout.is_some() => {
debug!(
target: LOG_TARGET,
"({}) Closing session idle for {:.0?}.",
self.logging_context_string,
idle_timeout.unwrap_or_default()
);
break;
}
result = self.framed.next() => {
match result {
Some(Ok(frame)) => {
#[cfg(feature = "metrics")]
metrics::inbound_requests_bytes(&self.protocol).observe(frame.len() as f64);
let start = Instant::now();
if let Err(err) = self.handle_request(frame.freeze()).await {
let close_result = if matches!(err, RpcServerError::WriteStreamExceededDeadline) {
Ok(())
} else {
self.close_framed().await
};
if let Err(err) = close_result {
let level = err.early_close_io().map(err_to_log_level).unwrap_or(log::Level::Error);
log!(
target: LOG_TARGET,
level,
"({}) Failed to close substream after socket error: {}",
self.logging_context_string,
err,
);
}
let level = err.early_close_io().map(err_to_log_level).unwrap_or(log::Level::Error);
log!(
target: LOG_TARGET,
level,
"(peer: {}, protocol: {}) Failed to handle request: {}",
self.node_id,
self.protocol_name(),
err
);
return Err(err);
}
let elapsed = start.elapsed();
trace!(
target: LOG_TARGET,
"({}) RPC request completed in {:.0?}{}",
self.logging_context_string,
elapsed,
if elapsed.as_secs() > 5 { " (LONG REQUEST)" } else { "" }
);
if let Some(timeout) = idle_timeout {
idle_sleep.as_mut().reset(
time::Instant::now().checked_add(timeout).unwrap_or_else(time::Instant::now),
);
}
},
Some(Err(err)) => {
if let Err(err) = self.close_framed().await {
error!(
target: LOG_TARGET,
"({}) Failed to close substream after socket error: {}", self.logging_context_string, err
);
}
return Err(err.into());
},
None => break,
}
}
}
}
self.close_framed().await?;
Ok(())
}
#[allow(clippy::too_many_lines)]
#[instrument(name = "rpc::server::handle_req", level="trace", skip(self, request), err, fields(request_size = request.len ()))]
async fn handle_request(&mut self, mut request: Bytes) -> Result<(), RpcServerError> {
if let Some(rejection) = reject_oversized_request(&request) {
debug!(
target: LOG_TARGET,
"({}) Rejecting a {} byte request, larger than the maximum of {} bytes",
self.logging_context_string,
request.len(),
rpc::max_request_size()
);
#[cfg(feature = "metrics")]
metrics::status_error_counter(&self.protocol, super::RpcStatusCode::BadRequest).inc();
self.send_with_deadline(rejection.to_encoded_bytes().into(), SESSION_CLOSE_TIMEOUT)
.await?;
return Ok(());
}
let decoded_msg = proto::rpc::RpcRequest::decode(&mut request)?;
let request_id = decoded_msg.request_id;
let method = RpcMethod::from(decoded_msg.method);
let requested_deadline = Duration::from_secs(decoded_msg.deadline);
if requested_deadline < self.config.minimum_client_deadline {
debug!(
target: LOG_TARGET,
"({}) Client has an invalid deadline. {}", self.logging_context_string, decoded_msg
);
let status = RpcStatus::bad_request(&format!(
"Invalid deadline ({:.0?}). The deadline MUST be greater than {:.0?}.",
requested_deadline, self.config.minimum_client_deadline,
));
let bad_request = proto::rpc::RpcResponse {
request_id,
status: status.as_code(),
flags: RpcMessageFlags::FIN.bits().into(),
payload: status.to_details_bytes(),
};
#[cfg(feature = "metrics")]
metrics::status_error_counter(&self.protocol, status.as_status_code()).inc();
self.send_with_deadline(bad_request.to_encoded_bytes().into(), SESSION_CLOSE_TIMEOUT)
.await?;
return Ok(());
}
let deadline = self.config.clamp_client_deadline(requested_deadline);
if deadline < requested_deadline {
debug!(
target: LOG_TARGET,
"({}) Client requested a deadline of {:.0?}, capping it at {:.0?}.",
self.logging_context_string,
requested_deadline,
deadline
);
}
let msg_flags = RpcMessageFlags::from_bits(u8::try_from(decoded_msg.flags).map_err(|_| {
RpcServerError::ProtocolError(format!("invalid message flag: must be less than {}", u8::MAX))
})?)
.ok_or(RpcServerError::ProtocolError(format!(
"invalid message flag, does not match any flags ({})",
decoded_msg.flags
)))?;
if msg_flags.contains(RpcMessageFlags::FIN) {
debug!(target: LOG_TARGET, "({}) Client sent FIN.", self.logging_context_string);
return Ok(());
}
if msg_flags.contains(RpcMessageFlags::ACK) {
debug!(
target: LOG_TARGET,
"({}) sending ACK response.", self.logging_context_string
);
let ack = proto::rpc::RpcResponse {
request_id,
status: RpcStatus::ok().as_code(),
flags: RpcMessageFlags::ACK.bits().into(),
..Default::default()
};
self.send_with_deadline(ack.to_encoded_bytes().into(), deadline).await?;
return Ok(());
}
trace!(
target: LOG_TARGET,
"({}) Request: {}, Method: {}",
self.logging_context_string,
decoded_msg,
method.id()
);
let req = Request::with_context(
self.create_request_context(request_id),
method,
decoded_msg.payload.into(),
);
let service_call = log_timing(
&self.logging_context_string,
request_id,
"service call",
self.service.call(req),
);
let service_result = time::timeout(deadline, service_call).await;
let service_result = match service_result {
Ok(v) => v,
Err(_) => {
warn!(
target: LOG_TARGET,
"{} RPC service was not able to complete within the deadline ({:.0?}). Request aborted",
self.logging_context_string,
deadline,
);
#[cfg(feature = "metrics")]
metrics::error_counter(&self.protocol, &RpcServerError::ServiceCallExceededDeadline).inc();
let status = RpcStatus::timed_out("RPC service did not complete within the deadline");
let timed_out = proto::rpc::RpcResponse {
request_id,
status: status.as_code(),
flags: RpcMessageFlags::FIN.bits().into(),
payload: status.to_details_bytes(),
};
#[cfg(feature = "metrics")]
metrics::status_error_counter(&self.protocol, status.as_status_code()).inc();
self.send_with_deadline(timed_out.to_encoded_bytes().into(), deadline)
.await?;
return Ok(());
},
};
match service_result {
Ok(body) => {
self.process_body(request_id, deadline, body).await?;
},
Err(err) => {
debug!(
target: LOG_TARGET,
"{} Service returned an error: {}", self.logging_context_string, err
);
let resp = proto::rpc::RpcResponse {
request_id,
status: err.as_code(),
flags: RpcMessageFlags::FIN.bits().into(),
payload: err.to_details_bytes(),
};
#[cfg(feature = "metrics")]
metrics::status_error_counter(&self.protocol, err.as_status_code()).inc();
self.send_with_deadline(resp.to_encoded_bytes().into(), deadline)
.await?;
},
}
Ok(())
}
fn protocol_name(&self) -> Cow<'_, str> {
String::from_utf8_lossy(&self.protocol)
}
#[allow(clippy::too_many_lines)]
async fn process_body(
&mut self,
request_id: u32,
deadline: Duration,
body: Response<Body>,
) -> Result<(), RpcServerError> {
trace!(target: LOG_TARGET, "Service call succeeded");
#[cfg(feature = "metrics")]
let protocol = self.protocol.clone();
let mut stream = body
.into_message()
.map(|result| into_response(request_id, result))
.map(move |mut message| {
if message.payload.len() > rpc::max_response_payload_size() {
message = message.exceeded_message_size();
}
#[cfg(feature = "metrics")]
if !message.status.is_ok() {
metrics::status_error_counter(&protocol, message.status).inc();
}
message.to_encoded_bytes()
});
let logging_context = self.logging_context_string.clone();
let mut prefetched: Option<Option<Bytes>> = None;
loop {
let next = match prefetched.take() {
Some(Some(msg)) => {
if let Err(err) = self.check_interruptions().await {
match err {
err @ RpcServerError::ClientInterruptedStream => {
debug!(target: LOG_TARGET, "Stream was interrupted by client: {}", err);
break;
},
err => {
error!(target: LOG_TARGET, "Stream was interrupted: {}", err);
return Err(err);
},
}
}
Some(msg)
},
Some(None) => None,
None => {
let next_item = log_timing(&logging_context, request_id, "message read", stream.next());
let timeout = time::sleep(deadline);
tokio::select! {
Err(err) = self.check_interruptions() => {
match err {
err @ RpcServerError::ClientInterruptedStream => {
debug!(target: LOG_TARGET, "Stream was interrupted by client: {}", err);
break;
},
err => {
error!(target: LOG_TARGET, "Stream was interrupted: {}", err);
return Err(err);
},
}
},
msg = next_item => msg,
_ = timeout => {
debug!(
target: LOG_TARGET,
"({}) Failed to return result within client deadline ({:.0?})",
self.logging_context_string,
deadline
);
#[cfg(feature = "metrics")]
metrics::error_counter(
&self.protocol,
&RpcServerError::ReadStreamExceededDeadline,
)
.inc();
break;
}
} },
};
let Some(msg) = next else {
trace!(target: LOG_TARGET, "{} Request complete", self.logging_context_string,);
break;
};
#[cfg(feature = "metrics")]
self.outbound_response_bytes.observe(msg.len() as f64);
trace!(
target: LOG_TARGET,
"({}) Sending body len = {}",
self.logging_context_string,
msg.len()
);
let send = self.send_with_deadline(msg, deadline);
tokio::pin!(send);
loop {
tokio::select! {
biased;
result = &mut send => {
result?;
break;
},
next = stream.next(), if prefetched.is_none() => {
prefetched = Some(next);
},
}
}
} Ok(())
}
async fn check_interruptions(&mut self) -> Result<(), RpcServerError> {
let check = future::poll_fn(|cx| match Pin::new(&mut self.framed).poll_next(cx) {
Poll::Ready(Some(Ok(mut msg))) => {
if msg.len() > rpc::max_request_size() {
debug!(
target: LOG_TARGET,
"Ignoring a {} byte message received during a streaming response (maximum {} bytes)",
msg.len(),
rpc::max_request_size()
);
return Poll::Ready(None);
}
let decoded_msg = match proto::rpc::RpcRequest::decode(&mut msg) {
Ok(msg) => msg,
Err(err) => {
error!(target: LOG_TARGET, "Client send MALFORMED response: {}", err);
return Poll::Ready(Some(RpcServerError::UnexpectedIncomingMessageMalformed));
},
};
let u8_bits = match u8::try_from(decoded_msg.flags) {
Ok(bits) => bits,
Err(err) => {
error!(target: LOG_TARGET, "Client send MALFORMED flags: {}", err);
return Poll::Ready(Some(RpcServerError::ProtocolError(format!(
"invalid message flag: must be less than {}",
u8::MAX
))));
},
};
let msg_flags = match RpcMessageFlags::from_bits(u8_bits) {
Some(flags) => flags,
None => {
error!(target: LOG_TARGET, "Client send MALFORMED flags: {}", u8_bits);
return Poll::Ready(Some(RpcServerError::ProtocolError(format!(
"invalid message flag, does not match any flags ({u8_bits})"
))));
},
};
if msg_flags.is_fin() {
Poll::Ready(Some(RpcServerError::ClientInterruptedStream))
} else {
Poll::Ready(Some(RpcServerError::UnexpectedIncomingMessage(decoded_msg)))
}
},
Poll::Ready(Some(Err(err))) if err.kind() == io::ErrorKind::WouldBlock => Poll::Ready(None),
Poll::Ready(Some(Err(err))) => Poll::Ready(Some(RpcServerError::from(err))),
Poll::Ready(None) => Poll::Ready(Some(RpcServerError::StreamClosedByRemote)),
Poll::Pending => Poll::Ready(None),
})
.await;
match check {
Some(err) => Err(err),
None => Ok(()),
}
}
fn create_request_context(&self, request_id: u32) -> RequestContext {
RequestContext::new(request_id, self.node_id.clone(), Box::new(self.comms_provider.clone()))
}
}
fn is_trace_enabled() -> bool {
log_enabled!(target: LOG_TARGET, log::Level::Trace) || tracing::enabled!(target: LOG_TARGET, Level::TRACE)
}
async fn log_timing<R, F: Future<Output = R>>(context_str: &str, request_id: u32, tag: &str, fut: F) -> R {
if !is_trace_enabled() {
return fut.await;
}
let t = Instant::now();
let span = span!(Level::TRACE, "rpc::internal::timing", request_id, tag);
let ret = fut.instrument(span).await;
let elapsed = t.elapsed();
trace!(
target: LOG_TARGET,
"({}) RPC TIMING(REQ_ID={}): '{}' took {:.2}s{}",
context_str,
request_id,
tag,
elapsed.as_secs_f32(),
if elapsed.as_secs() >= 5 { " (SLOW)" } else { "" }
);
ret
}
#[derive(Clone, PartialEq, prost::Message)]
struct RpcRequestId {
#[prost(uint32, tag = "1")]
request_id: u32,
}
fn reject_oversized_request(request: &Bytes) -> Option<proto::rpc::RpcResponse> {
if request.len() <= rpc::max_request_size() {
return None;
}
let request_id = RpcRequestId::decode(request.clone()).map(|r| r.request_id).unwrap_or(0);
let status = RpcStatus::bad_request(&format!(
"The request size exceeded the maximum allowed request size. Max = {} bytes, Got = {} bytes",
rpc::max_request_size(),
request.len()
));
Some(proto::rpc::RpcResponse {
request_id,
status: status.as_code(),
flags: RpcMessageFlags::FIN.bits().into(),
payload: status.to_details_bytes(),
})
}
fn into_response(request_id: u32, result: Result<BodyBytes, RpcStatus>) -> RpcResponse {
match result {
Ok(msg) => {
let mut flags = RpcMessageFlags::empty();
if msg.is_finished() {
flags |= RpcMessageFlags::FIN;
}
RpcResponse {
request_id,
status: RpcStatus::ok().as_status_code(),
flags,
payload: msg.into_bytes().unwrap_or_else(Bytes::new),
}
},
Err(err) => {
debug!(target: LOG_TARGET, "Body contained an error: {}", err);
RpcResponse {
request_id,
status: err.as_status_code(),
flags: RpcMessageFlags::FIN,
payload: Bytes::from(err.to_details_bytes()),
}
},
}
}
fn err_to_log_level(err: &io::Error) -> log::Level {
match err.kind() {
ErrorKind::ConnectionReset |
ErrorKind::ConnectionAborted |
ErrorKind::BrokenPipe |
ErrorKind::WriteZero |
ErrorKind::UnexpectedEof => log::Level::Debug,
_ => log::Level::Error,
}
}
#[cfg(test)]
mod test {
use super::*;
#[test]
fn client_deadline_is_capped_from_above() {
let builder = RpcServerBuilder::new();
assert_eq!(builder.maximum_client_deadline, DEFAULT_MAXIMUM_CLIENT_DEADLINE);
assert_eq!(
builder.clamp_client_deadline(Duration::from_secs(1)),
Duration::from_secs(1)
);
assert_eq!(
builder.clamp_client_deadline(DEFAULT_MAXIMUM_CLIENT_DEADLINE),
DEFAULT_MAXIMUM_CLIENT_DEADLINE
);
assert_eq!(
builder.clamp_client_deadline(Duration::from_secs(240)),
Duration::from_secs(240)
);
assert_eq!(
builder.clamp_client_deadline(DEFAULT_MAXIMUM_CLIENT_DEADLINE + Duration::from_secs(1)),
DEFAULT_MAXIMUM_CLIENT_DEADLINE
);
assert_eq!(
builder.clamp_client_deadline(Duration::from_secs(u64::MAX)),
DEFAULT_MAXIMUM_CLIENT_DEADLINE
);
}
#[test]
fn maximum_client_deadline_is_configurable() {
let builder = RpcServerBuilder::new().with_maximum_client_deadline(Duration::from_secs(30));
assert_eq!(
builder.clamp_client_deadline(Duration::from_secs(u64::MAX)),
Duration::from_secs(30)
);
assert_eq!(
builder.clamp_client_deadline(Duration::from_secs(5)),
Duration::from_secs(5)
);
}
#[test]
fn a_maximum_below_the_minimum_does_not_starve_requests() {
let builder = RpcServerBuilder::new()
.with_minimum_client_deadline(Duration::from_secs(1))
.with_maximum_client_deadline(Duration::from_secs(0));
assert_eq!(
builder.clamp_client_deadline(Duration::from_secs(0)),
Duration::from_secs(1)
);
assert_eq!(
builder.clamp_client_deadline(Duration::from_secs(120)),
Duration::from_secs(1)
);
assert_eq!(
builder.clamp_client_deadline(Duration::from_secs(u64::MAX)),
Duration::from_secs(1)
);
}
}