use std::time::Duration;
use pb_mapper_core::checksum::Credential;
use pb_mapper_core::config::ResolvedAddrs;
use pb_mapper_protocol::MessageReader;
use pb_mapper_protocol::command::{
AdminRequest, AdminResponse, MessageSerializer, PbConnRequest, PbConnResponse,
};
use pb_mapper_protocol::secure::ClientHeaderSession;
use snafu::ResultExt;
use super::super::Error;
use super::super::error::{ConnectSnafu, Result};
const REPLAYED_SALT_CODE: &str = "connection_salt_replayed";
const MAX_ATTEMPTS: usize = 2;
pub(super) async fn send_admin_request(
server: &str,
credential: Credential,
request: AdminRequest,
io_timeout: Duration,
) -> Result<AdminResponse> {
let addrs = pb_mapper_core::config::resolve_addrs_async(server)
.await
.map_err(|source| Error::Address {
addr: server.to_string(),
source,
})?;
let encoded = PbConnRequest::Admin(request).encode().map_err(protocol)?;
for attempt in 0..MAX_ATTEMPTS {
let last_attempt = attempt + 1 == MAX_ATTEMPTS;
let mut exchange = Exchange::new();
let outcome = tokio::time::timeout(io_timeout, exchange.run(&addrs, &credential, &encoded))
.await
.unwrap_or(Err(Error::TimedOut {
timeout: io_timeout,
}));
let response = match outcome {
Ok(response) => response,
Err(error) => {
if exchange.sent || last_attempt {
return Err(error);
}
continue;
}
};
match response {
PbConnResponse::Admin(response) => return Ok(response),
PbConnResponse::Error(error)
if error.code == REPLAYED_SALT_CODE && error.retryable && !last_attempt =>
{
continue;
}
PbConnResponse::Error(error) => return Err(Error::from_remote(error)),
other => {
return Err(Error::protocol(format!(
"unexpected administrator response: {other:?}"
)));
}
}
}
Err(Error::protocol(
"connection salt replay retry was exhausted",
))
}
struct Exchange {
sent: bool,
}
impl Exchange {
fn new() -> Self {
Self { sent: false }
}
async fn run(
&mut self,
addrs: &ResolvedAddrs,
credential: &Credential,
encoded: &[u8],
) -> Result<PbConnResponse> {
let mut stream = crate::addr::connect_tcp(addrs)
.await
.context(ConnectSnafu {
addr: addrs.to_string(),
})?;
let session = ClientHeaderSession::new_v2(credential).map_err(protocol)?;
self.write_request(&mut stream, &session, encoded).await?;
let mut reader = session.response_reader(&mut stream).map_err(protocol)?;
let message = reader.read_msg().await.map_err(protocol)?;
PbConnResponse::decode(message).map_err(protocol)
}
async fn write_request<T: tokio::io::AsyncWriteExt + Unpin>(
&mut self,
stream: &mut T,
session: &ClientHeaderSession,
encoded: &[u8],
) -> Result<()> {
self.sent = true;
session
.write_initial(stream, encoded)
.await
.map_err(protocol)
}
}
fn protocol(error: impl std::fmt::Display) -> Error {
Error::protocol(error.to_string())
}
#[cfg(test)]
#[allow(clippy::unwrap_used)]
mod tests {
use super::*;
use tokio::io::AsyncReadExt;
#[tokio::test(start_paused = true)]
async fn cancelled_partial_admin_write_must_not_be_retried() {
let credential = Credential::Admin(*b"0123456789abcdefghijklmnopqrstuv");
let session = ClientHeaderSession::new_v2(&credential).unwrap();
let (mut writer, mut reader) = tokio::io::duplex(1);
let mut exchange = Exchange::new();
assert!(
tokio::time::timeout(
Duration::from_millis(10),
exchange.write_request(&mut writer, &session, b"mutation")
)
.await
.is_err()
);
assert!(
exchange.sent,
"a partial mutation write is not safe to retry"
);
assert!(reader.read_u8().await.is_ok());
}
}