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 tokio::net::TcpStream;
use uni_stream::addr::each_addr;
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 = each_addr(addrs.as_slice(), TcpStream::connect)
.await
.context(ConnectSnafu {
addr: addrs.to_string(),
})?;
let session = ClientHeaderSession::new_v2(credential).map_err(protocol)?;
session
.write_initial(&mut stream, encoded)
.await
.map_err(protocol)?;
self.sent = true;
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)
}
}
fn protocol(error: impl std::fmt::Display) -> Error {
Error::protocol(error.to_string())
}