use std::net::{AddrParseError, SocketAddr};
use std::sync::{Arc, Once};
use std::time::Duration;
use clap::ValueEnum;
use snafu::ResultExt;
use tracing_subscriber::layer::SubscriberExt;
use tracing_subscriber::{EnvFilter, Layer, fmt};
use crate::error::{CfgPbServerEnvNotExistSnafu, Result};
#[derive(ValueEnum, Debug, Clone, Copy)]
pub enum StatusOp {
RemoteId,
Keys,
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct ResolvedAddrs {
addrs: Arc<[SocketAddr]>,
}
impl ResolvedAddrs {
fn new(name: &str, addrs: Vec<SocketAddr>, parse_error: AddrParseError) -> Result<Self> {
if addrs.is_empty() {
return Err(crate::error::Error::CfgParseSockAddr {
string: name.to_string(),
source: parse_error,
});
}
Ok(Self {
addrs: Arc::from(addrs),
})
}
#[must_use]
pub fn from_candidates(addrs: Vec<SocketAddr>) -> Option<Self> {
if addrs.is_empty() {
return None;
}
Some(Self {
addrs: Arc::from(addrs),
})
}
#[inline]
#[must_use]
pub fn as_slice(&self) -> &[SocketAddr] {
&self.addrs
}
#[inline]
#[must_use]
pub fn primary(&self) -> SocketAddr {
self.addrs[0]
}
}
impl std::fmt::Display for ResolvedAddrs {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(formatter, "{}", self.primary())?;
match self.addrs.len() {
1 => Ok(()),
more => write!(formatter, " (+{} more)", more - 1),
}
}
}
impl From<SocketAddr> for ResolvedAddrs {
fn from(addr: SocketAddr) -> Self {
Self {
addrs: Arc::from(vec![addr]),
}
}
}
#[inline]
pub fn resolve_addrs(addr: &str) -> Result<ResolvedAddrs> {
let parse_error = match addr.parse::<SocketAddr>() {
Ok(socket_addr) => return Ok(ResolvedAddrs::from(socket_addr)),
Err(error) => error,
};
let system = || match std::net::ToSocketAddrs::to_socket_addrs(addr) {
Ok(addrs) => addrs.collect(),
Err(_) => Vec::new(),
};
let addrs = if addr.starts_with("localhost:") {
system()
} else {
match crate::addr::get_socket_addrs(addr) {
Ok(addrs) if !addrs.is_empty() => addrs,
_ => system(),
}
};
ResolvedAddrs::new(addr, addrs, parse_error)
}
pub async fn resolve_addrs_async(addr: &str) -> Result<ResolvedAddrs> {
let parse_error = match addr.parse::<SocketAddr>() {
Ok(socket_addr) => return Ok(ResolvedAddrs::from(socket_addr)),
Err(error) => error,
};
async fn system(addr: &str) -> Vec<SocketAddr> {
match tokio::net::lookup_host(addr).await {
Ok(addrs) => addrs.collect(),
Err(_) => Vec::new(),
}
}
let addrs = if addr.starts_with("localhost:") {
system(addr).await
} else {
match crate::addr::get_socket_addrs_async(addr).await {
Ok(addrs) if !addrs.is_empty() => addrs,
_ => system(addr).await,
}
};
ResolvedAddrs::new(addr, addrs, parse_error)
}
const PB_MAPPER_SERVER: &str = "PB_MAPPER_SERVER";
pub const PB_MAPPER_KEEP_ALIVE: &str = "PB_MAPPER_KEEP_ALIVE";
pub const PB_MAPPER_CONTROL_IO_TIMEOUT: &str = "PB_MAPPER_CONTROL_IO_TIMEOUT";
pub const PB_MAPPER_STREAM_ACK_TIMEOUT: &str = "PB_MAPPER_STREAM_ACK_TIMEOUT";
pub const PB_MAPPER_STREAM_READY_TIMEOUT: &str = "PB_MAPPER_STREAM_READY_TIMEOUT";
pub const PB_MAPPER_STREAM_RECOVERY_TIMEOUT: &str = "PB_MAPPER_STREAM_RECOVERY_TIMEOUT";
pub const PB_MAPPER_CONTROL_CONN_POOL_SIZE: &str = "PB_MAPPER_CONTROL_CONN_POOL_SIZE";
pub const PB_MAPPER_CONTROL_HEARTBEAT_INTERVAL: &str = "PB_MAPPER_CONTROL_HEARTBEAT_INTERVAL";
pub const PB_MAPPER_CONTROL_HEARTBEAT_TOLERANCE: &str = "PB_MAPPER_CONTROL_HEARTBEAT_TOLERANCE";
pub const PB_MAPPER_CONTROL_SUSPECT_GRACE: &str = "PB_MAPPER_CONTROL_SUSPECT_GRACE";
pub const PB_MAPPER_REGISTRATION_PROBE_TIMEOUT: &str = "PB_MAPPER_REGISTRATION_PROBE_TIMEOUT";
pub const PB_MAPPER_REGISTRATION_REJECT_BACKOFF_MIN: &str =
"PB_MAPPER_REGISTRATION_REJECT_BACKOFF_MIN";
pub const PB_MAPPER_REGISTRATION_REJECT_BACKOFF_MAX: &str =
"PB_MAPPER_REGISTRATION_REJECT_BACKOFF_MAX";
pub const PB_MAPPER_SERVER_LEASE_TIMEOUT: &str = "PB_MAPPER_SERVER_LEASE_TIMEOUT";
pub const PB_MAPPER_SERVER_LEASE_SWEEP_INTERVAL: &str = "PB_MAPPER_SERVER_LEASE_SWEEP_INTERVAL";
pub const PB_MAPPER_CLIENT_HEALTH_CHECK_INTERVAL: &str = "PB_MAPPER_CLIENT_HEALTH_CHECK_INTERVAL";
pub const PB_MAPPER_CLIENT_HEALTH_CHECK_TIMEOUT: &str = "PB_MAPPER_CLIENT_HEALTH_CHECK_TIMEOUT";
pub const PB_MAPPER_CLIENT_HEALTH_FAILURE_THRESHOLD: &str =
"PB_MAPPER_CLIENT_HEALTH_FAILURE_THRESHOLD";
pub const PB_MAPPER_LOG_FORMAT: &str = "PB_MAPPER_LOG_FORMAT";
const DEFAULT_CONTROL_IO_TIMEOUT: Duration = Duration::from_secs(30);
const DEFAULT_STREAM_ACK_TIMEOUT: Duration = Duration::from_millis(300);
const DEFAULT_STREAM_READY_TIMEOUT: Duration = Duration::from_secs(1);
const DEFAULT_STREAM_RECOVERY_TIMEOUT: Duration = Duration::from_secs(2);
const DEFAULT_CONTROL_CONN_POOL_SIZE: usize = 2;
const DEFAULT_CONTROL_HEARTBEAT_INTERVAL: Duration = Duration::from_secs(2);
const DEFAULT_CONTROL_HEARTBEAT_TOLERANCE: Duration = Duration::from_secs(6);
const DEFAULT_CONTROL_SUSPECT_GRACE: Duration = Duration::from_secs(2);
const DEFAULT_REGISTRATION_PROBE_TIMEOUT: Duration = Duration::from_secs(1);
const DEFAULT_REGISTRATION_REJECT_BACKOFF_MIN: Duration = Duration::from_secs(5);
const DEFAULT_REGISTRATION_REJECT_BACKOFF_MAX: Duration = Duration::from_secs(80);
const DEFAULT_SERVER_LEASE_TIMEOUT: Duration = Duration::from_secs(15);
const DEFAULT_SERVER_LEASE_SWEEP_INTERVAL: Duration = Duration::from_secs(5);
const DEFAULT_CLIENT_HEALTH_CHECK_INTERVAL: Duration = Duration::from_secs(15);
const DEFAULT_CLIENT_HEALTH_CHECK_TIMEOUT: Duration = Duration::from_secs(5);
const DEFAULT_CLIENT_HEALTH_FAILURE_THRESHOLD: usize = 3;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum LogFormat {
Pretty,
Compact,
Json,
}
pub fn parse_log_format(value: &str) -> LogFormat {
match value.trim().to_ascii_lowercase().as_str() {
"compact" => LogFormat::Compact,
"json" => LogFormat::Json,
_ => LogFormat::Pretty,
}
}
fn log_format_from_env() -> LogFormat {
std::env::var(PB_MAPPER_LOG_FORMAT)
.ok()
.map(|value| parse_log_format(&value))
.unwrap_or(LogFormat::Pretty)
}
fn default_env_filter() -> EnvFilter {
EnvFilter::builder()
.with_default_directive(tracing::level_filters::LevelFilter::INFO.into())
.from_env_lossy()
}
pub fn parse_duration(value: &str) -> Option<Duration> {
let value = value.trim();
if value.is_empty() {
return None;
}
if let Some(raw) = value.strip_suffix("ms") {
return raw.trim().parse::<u64>().ok().map(Duration::from_millis);
}
if let Some(raw) = value.strip_suffix('s') {
return raw.trim().parse::<u64>().ok().map(Duration::from_secs);
}
if let Some(raw) = value.strip_suffix('m') {
return raw
.trim()
.parse::<u64>()
.ok()
.and_then(|minutes| minutes.checked_mul(60))
.map(Duration::from_secs);
}
if let Some(raw) = value.strip_suffix('h') {
return raw
.trim()
.parse::<u64>()
.ok()
.and_then(|hours| hours.checked_mul(60 * 60))
.map(Duration::from_secs);
}
value.parse::<u64>().ok().map(Duration::from_secs)
}
pub fn duration_from_env(name: &str, default: Duration) -> Duration {
std::env::var(name)
.ok()
.and_then(|value| parse_duration(&value))
.unwrap_or(default)
}
fn positive_duration_from_env(name: &str, default: Duration) -> Duration {
let value = duration_from_env(name, default);
if value.is_zero() {
tracing::warn!(
event = "config_zero_duration_ignored",
variable = name,
default = ?default,
"ignoring a zero duration and using the default instead"
);
return default;
}
value
}
pub fn control_io_timeout() -> Duration {
duration_from_env(PB_MAPPER_CONTROL_IO_TIMEOUT, DEFAULT_CONTROL_IO_TIMEOUT)
}
pub fn stream_ack_timeout() -> Duration {
duration_from_env(PB_MAPPER_STREAM_ACK_TIMEOUT, DEFAULT_STREAM_ACK_TIMEOUT)
}
pub fn stream_ready_timeout() -> Duration {
duration_from_env(PB_MAPPER_STREAM_READY_TIMEOUT, DEFAULT_STREAM_READY_TIMEOUT)
}
pub fn stream_recovery_timeout() -> Duration {
duration_from_env(
PB_MAPPER_STREAM_RECOVERY_TIMEOUT,
DEFAULT_STREAM_RECOVERY_TIMEOUT,
)
}
pub fn control_conn_pool_size() -> usize {
std::env::var(PB_MAPPER_CONTROL_CONN_POOL_SIZE)
.ok()
.and_then(|value| value.trim().parse::<usize>().ok())
.filter(|size| *size > 0)
.map(|size| size.min(16))
.unwrap_or(DEFAULT_CONTROL_CONN_POOL_SIZE)
}
pub fn control_heartbeat_interval() -> Duration {
duration_from_env(
PB_MAPPER_CONTROL_HEARTBEAT_INTERVAL,
DEFAULT_CONTROL_HEARTBEAT_INTERVAL,
)
}
pub fn control_heartbeat_tolerance() -> Duration {
duration_from_env(
PB_MAPPER_CONTROL_HEARTBEAT_TOLERANCE,
DEFAULT_CONTROL_HEARTBEAT_TOLERANCE,
)
}
pub fn control_suspect_grace() -> Duration {
duration_from_env(
PB_MAPPER_CONTROL_SUSPECT_GRACE,
DEFAULT_CONTROL_SUSPECT_GRACE,
)
}
pub fn registration_probe_timeout() -> Duration {
duration_from_env(
PB_MAPPER_REGISTRATION_PROBE_TIMEOUT,
DEFAULT_REGISTRATION_PROBE_TIMEOUT,
)
}
pub fn registration_reject_backoff() -> (Duration, Duration) {
let min = positive_duration_from_env(
PB_MAPPER_REGISTRATION_REJECT_BACKOFF_MIN,
DEFAULT_REGISTRATION_REJECT_BACKOFF_MIN,
);
let max = positive_duration_from_env(
PB_MAPPER_REGISTRATION_REJECT_BACKOFF_MAX,
DEFAULT_REGISTRATION_REJECT_BACKOFF_MAX,
)
.max(min);
(min, max)
}
pub fn server_lease_timeout() -> Duration {
duration_from_env(PB_MAPPER_SERVER_LEASE_TIMEOUT, DEFAULT_SERVER_LEASE_TIMEOUT)
}
pub fn server_lease_sweep_interval() -> Duration {
positive_duration_from_env(
PB_MAPPER_SERVER_LEASE_SWEEP_INTERVAL,
DEFAULT_SERVER_LEASE_SWEEP_INTERVAL,
)
}
pub fn client_health_check_interval() -> Duration {
duration_from_env(
PB_MAPPER_CLIENT_HEALTH_CHECK_INTERVAL,
DEFAULT_CLIENT_HEALTH_CHECK_INTERVAL,
)
}
pub fn client_health_check_timeout() -> Duration {
duration_from_env(
PB_MAPPER_CLIENT_HEALTH_CHECK_TIMEOUT,
DEFAULT_CLIENT_HEALTH_CHECK_TIMEOUT,
)
}
pub fn client_health_failure_threshold() -> usize {
std::env::var(PB_MAPPER_CLIENT_HEALTH_FAILURE_THRESHOLD)
.ok()
.and_then(|value| value.trim().parse::<usize>().ok())
.filter(|threshold| *threshold > 0)
.map(|threshold| threshold.min(100))
.unwrap_or(DEFAULT_CLIENT_HEALTH_FAILURE_THRESHOLD)
}
pub fn keep_alive_from_env() -> bool {
match std::env::var(PB_MAPPER_KEEP_ALIVE) {
Ok(value) => matches!(
value.trim().to_ascii_lowercase().as_str(),
"on" | "1" | "true" | "yes"
),
Err(_) => false,
}
}
pub fn pb_mapper_server_addr(addr: Option<&str>) -> Result<String> {
match addr {
Some(addr) => Ok(addr.to_string()),
None => std::env::var(PB_MAPPER_SERVER).context(CfgPbServerEnvNotExistSnafu),
}
}
#[inline]
pub fn resolve_pb_mapper_server(addr: Option<&str>) -> Result<ResolvedAddrs> {
match addr {
Some(addr) => resolve_addrs(addr),
None => {
let addr = std::env::var(PB_MAPPER_SERVER).context(CfgPbServerEnvNotExistSnafu)?;
resolve_addrs(&addr)
}
}
}
pub async fn resolve_pb_mapper_server_async(addr: Option<&str>) -> Result<ResolvedAddrs> {
match addr {
Some(addr) => resolve_addrs_async(addr).await,
None => {
let addr = std::env::var(PB_MAPPER_SERVER).context(CfgPbServerEnvNotExistSnafu)?;
resolve_addrs_async(&addr).await
}
}
}
pub fn init_tracing() {
static INIT_TRACING: Once = Once::new();
INIT_TRACING.call_once(|| {
let result = match log_format_from_env() {
LogFormat::Pretty => {
let subscriber = tracing_subscriber::registry().with(
fmt::layer()
.pretty()
.with_writer(std::io::stdout)
.with_filter(default_env_filter()),
);
tracing::subscriber::set_global_default(subscriber)
}
LogFormat::Compact => {
let subscriber = tracing_subscriber::registry().with(
fmt::layer()
.compact()
.with_writer(std::io::stdout)
.with_filter(default_env_filter()),
);
tracing::subscriber::set_global_default(subscriber)
}
LogFormat::Json => {
let subscriber = tracing_subscriber::registry().with(
fmt::layer()
.json()
.flatten_event(true)
.with_writer(std::io::stdout)
.with_filter(default_env_filter()),
);
tracing::subscriber::set_global_default(subscriber)
}
};
if let Err(e) = result {
eprintln!("failed to initialize tracing subscriber: {e}");
}
});
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parse_log_format_accepts_supported_values() {
assert_eq!(parse_log_format("pretty"), LogFormat::Pretty);
assert_eq!(parse_log_format("compact"), LogFormat::Compact);
assert_eq!(parse_log_format("json"), LogFormat::Json);
assert_eq!(parse_log_format(" JSON "), LogFormat::Json);
assert_eq!(parse_log_format("unknown"), LogFormat::Pretty);
}
#[test]
fn keep_alive_reads_the_environment_every_time() {
let restore = std::env::var(PB_MAPPER_KEEP_ALIVE).ok();
unsafe {
std::env::remove_var(PB_MAPPER_KEEP_ALIVE);
}
assert!(!keep_alive_from_env(), "absent means off");
unsafe {
std::env::set_var(PB_MAPPER_KEEP_ALIVE, "ON");
}
assert!(keep_alive_from_env(), "the documented spelling");
unsafe {
std::env::set_var(PB_MAPPER_KEEP_ALIVE, "OFF");
}
assert!(
!keep_alive_from_env(),
"OFF must mean off; the old check was `is_ok()`, so any value at \
all — OFF included — turned keep-alive on"
);
for truthy in ["on", "1", "true", "yes", " ON "] {
unsafe {
std::env::set_var(PB_MAPPER_KEEP_ALIVE, truthy);
}
assert!(keep_alive_from_env(), "{truthy:?} should enable");
}
for falsy in ["", "off", "0", "false", "no"] {
unsafe {
std::env::set_var(PB_MAPPER_KEEP_ALIVE, falsy);
}
assert!(!keep_alive_from_env(), "{falsy:?} should not enable");
}
unsafe {
match restore {
Some(value) => std::env::set_var(PB_MAPPER_KEEP_ALIVE, value),
None => std::env::remove_var(PB_MAPPER_KEEP_ALIVE),
}
}
}
#[test]
fn resolved_addrs_keeps_every_candidate() {
let first: SocketAddr = "127.0.0.1:7666".parse().expect("literal");
let second: SocketAddr = "[::1]:7666".parse().expect("literal");
let addrs = ResolvedAddrs::from_candidates(vec![first, second]).expect("non-empty");
assert_eq!(addrs.as_slice(), [first, second]);
assert_eq!(addrs.primary(), first, "order is the resolver's order");
}
#[test]
fn resolved_addrs_rejects_an_empty_candidate_list() {
assert!(ResolvedAddrs::from_candidates(Vec::new()).is_none());
}
#[test]
fn resolved_addrs_renders_the_primary_and_the_rest_as_a_count() {
let single: SocketAddr = "127.0.0.1:7666".parse().expect("literal");
assert_eq!(ResolvedAddrs::from(single).to_string(), "127.0.0.1:7666");
let second: SocketAddr = "[::1]:7666".parse().expect("literal");
let both = ResolvedAddrs::from_candidates(vec![single, second]).expect("non-empty");
assert_eq!(both.to_string(), "127.0.0.1:7666 (+1 more)");
}
#[test]
fn resolve_addrs_passes_a_literal_through() {
let addrs = resolve_addrs("127.0.0.1:7666").expect("a literal always resolves");
assert_eq!(
addrs.as_slice(),
["127.0.0.1:7666".parse().expect("literal")]
);
}
#[test]
fn resolve_addrs_keeps_every_localhost_record() {
let addrs = resolve_addrs("localhost:7666").expect("localhost always resolves");
assert!(
addrs.as_slice().iter().all(|addr| addr.ip().is_loopback()),
"localhost must resolve to loopback only, got {:?}",
addrs.as_slice()
);
}
#[test]
fn resolve_addrs_fails_when_nothing_can_be_resolved() {
assert!(resolve_addrs("127.0.0.1").is_err(), "no port");
assert!(resolve_addrs("localhost").is_err(), "no port");
}
#[test]
fn reject_backoff_and_sweep_interval_survive_a_bad_environment() {
let names = [
PB_MAPPER_REGISTRATION_REJECT_BACKOFF_MIN,
PB_MAPPER_REGISTRATION_REJECT_BACKOFF_MAX,
PB_MAPPER_SERVER_LEASE_SWEEP_INTERVAL,
];
let restore = names.map(|name| (name, std::env::var(name).ok()));
unsafe {
for name in names {
std::env::remove_var(name);
}
}
assert_eq!(
registration_reject_backoff(),
(
DEFAULT_REGISTRATION_REJECT_BACKOFF_MIN,
DEFAULT_REGISTRATION_REJECT_BACKOFF_MAX
),
"absent means the defaults"
);
assert_eq!(
server_lease_sweep_interval(),
DEFAULT_SERVER_LEASE_SWEEP_INTERVAL
);
unsafe {
std::env::set_var(PB_MAPPER_REGISTRATION_REJECT_BACKOFF_MIN, "30s");
std::env::set_var(PB_MAPPER_REGISTRATION_REJECT_BACKOFF_MAX, "2m");
}
assert_eq!(
registration_reject_backoff(),
(Duration::from_secs(30), Duration::from_secs(120)),
"both ends come from the environment"
);
unsafe {
std::env::set_var(PB_MAPPER_REGISTRATION_REJECT_BACKOFF_MIN, "0s");
std::env::set_var(PB_MAPPER_REGISTRATION_REJECT_BACKOFF_MAX, "0s");
}
assert_eq!(
registration_reject_backoff(),
(
DEFAULT_REGISTRATION_REJECT_BACKOFF_MIN,
DEFAULT_REGISTRATION_REJECT_BACKOFF_MAX
),
"a zero at either end selects the default, not a millisecond"
);
unsafe {
std::env::set_var(PB_MAPPER_REGISTRATION_REJECT_BACKOFF_MIN, "0ms");
std::env::set_var(PB_MAPPER_REGISTRATION_REJECT_BACKOFF_MAX, "3m");
}
assert_eq!(
registration_reject_backoff(),
(
DEFAULT_REGISTRATION_REJECT_BACKOFF_MIN,
Duration::from_secs(180)
)
);
unsafe {
std::env::set_var(PB_MAPPER_REGISTRATION_REJECT_BACKOFF_MIN, "1m");
std::env::set_var(PB_MAPPER_REGISTRATION_REJECT_BACKOFF_MAX, "1s");
}
let (min, max) = registration_reject_backoff();
assert_eq!(min, Duration::from_secs(60));
assert_eq!(max, min, "an inverted range collapses to a fixed delay");
unsafe {
std::env::set_var(PB_MAPPER_SERVER_LEASE_SWEEP_INTERVAL, "0s");
}
assert_eq!(
server_lease_sweep_interval(),
DEFAULT_SERVER_LEASE_SWEEP_INTERVAL,
"a zero period would panic `tokio::time::interval`, and a millisecond \
one would queue a full scan every millisecond"
);
unsafe {
for (name, value) in restore {
match value {
Some(value) => std::env::set_var(name, value),
None => std::env::remove_var(name),
}
}
}
}
#[tokio::test]
async fn resolve_addrs_async_matches_the_blocking_path_on_a_literal() {
let expected = resolve_addrs("127.0.0.1:7666").expect("literal");
let actual = resolve_addrs_async("127.0.0.1:7666")
.await
.expect("literal");
assert_eq!(actual, expected);
}
}