use std::fmt::Debug;
use tokio_tungstenite::tungstenite::stream::Mode;
use super::types::TcpMessageHandler;
use crate::error::{NetworkConfigError, NetworkConfigResult};
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct SocketHeartbeat {
pub interval_secs: u64,
pub payload: Vec<u8>,
}
#[derive(bon::Builder)]
#[builder(finish_fn(name = build_inner, vis = ""))]
pub struct SocketConfig {
pub url: String,
pub mode: Mode,
pub suffix: Vec<u8>,
pub message_handler: Option<TcpMessageHandler>,
pub heartbeat: Option<SocketHeartbeat>,
pub connect_timeout_ms: Option<u64>,
pub reconnect_delay_initial_ms: Option<u64>,
pub reconnect_delay_max_ms: Option<u64>,
pub reconnect_backoff_factor: Option<f64>,
pub reconnect_jitter_ms: Option<u64>,
pub connection_max_retries: Option<u32>,
pub reconnect_max_attempts: Option<u32>,
pub heartbeat_timeout_secs: Option<u64>,
pub certs_dir: Option<String>,
}
impl<S: socket_config_builder::IsComplete> SocketConfigBuilder<S> {
pub fn build(self) -> NetworkConfigResult<SocketConfig> {
let config = self.build_inner();
config.validate()?;
Ok(config)
}
}
impl SocketConfig {
pub fn validate(&self) -> NetworkConfigResult<()> {
let mut errors = Vec::new();
if self.url.trim().is_empty() {
errors.push(NetworkConfigError::invalid("url", "must not be empty"));
}
if let Some(heartbeat) = &self.heartbeat
&& heartbeat.interval_secs == 0
{
errors.push(NetworkConfigError::invalid(
"heartbeat",
"interval must be positive",
));
}
if let (Some(heartbeat), Some(timeout_secs)) =
(&self.heartbeat, self.heartbeat_timeout_secs)
&& timeout_secs <= heartbeat.interval_secs
{
errors.push(NetworkConfigError::invalid(
"heartbeat_timeout_secs",
format!(
"must exceed heartbeat interval ({}s), was {timeout_secs}s",
heartbeat.interval_secs
),
));
}
for (field, value) in [
("connect_timeout_ms", self.connect_timeout_ms),
(
"reconnect_delay_initial_ms",
self.reconnect_delay_initial_ms,
),
("reconnect_delay_max_ms", self.reconnect_delay_max_ms),
("heartbeat_timeout_secs", self.heartbeat_timeout_secs),
] {
if let Some(value) = value
&& value == 0
{
errors.push(NetworkConfigError::invalid(
field,
format!("must be positive, was {value}"),
));
}
}
if let Some(factor) = self.reconnect_backoff_factor
&& !(1.0..=100.0).contains(&factor)
{
errors.push(NetworkConfigError::invalid(
"reconnect_backoff_factor",
format!("must be in range [1.0, 100.0], was {factor}"),
));
}
if let (Some(initial), Some(max)) =
(self.reconnect_delay_initial_ms, self.reconnect_delay_max_ms)
&& initial > max
{
errors.push(NetworkConfigError::invalid(
"reconnect_delay_initial_ms",
format!("must not exceed reconnect_delay_max_ms ({max}), was {initial}"),
));
}
NetworkConfigError::collect(errors)
}
pub(crate) fn resolved_heartbeat_timeout(&self) -> Option<u64> {
crate::heartbeat::resolve_heartbeat_timeout(
self.heartbeat_timeout_secs,
self.heartbeat
.as_ref()
.map(|heartbeat| heartbeat.interval_secs),
)
}
}
impl Debug for SocketConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct(stringify!(SocketConfig))
.field("url", &self.url)
.field("mode", &self.mode)
.field("suffix", &self.suffix)
.field(
"message_handler",
&self.message_handler.as_ref().map(|_| "<function>"),
)
.field("heartbeat", &self.heartbeat)
.field("connect_timeout_ms", &self.connect_timeout_ms)
.field(
"reconnect_delay_initial_ms",
&self.reconnect_delay_initial_ms,
)
.field("reconnect_delay_max_ms", &self.reconnect_delay_max_ms)
.field("reconnect_backoff_factor", &self.reconnect_backoff_factor)
.field("reconnect_jitter_ms", &self.reconnect_jitter_ms)
.field("connection_max_retries", &self.connection_max_retries)
.field("reconnect_max_attempts", &self.reconnect_max_attempts)
.field("heartbeat_timeout_secs", &self.heartbeat_timeout_secs)
.field("certs_dir", &self.certs_dir)
.finish()
}
}
impl Clone for SocketConfig {
fn clone(&self) -> Self {
Self {
url: self.url.clone(),
mode: self.mode,
suffix: self.suffix.clone(),
message_handler: self.message_handler.clone(),
heartbeat: self.heartbeat.clone(),
connect_timeout_ms: self.connect_timeout_ms,
reconnect_delay_initial_ms: self.reconnect_delay_initial_ms,
reconnect_delay_max_ms: self.reconnect_delay_max_ms,
reconnect_backoff_factor: self.reconnect_backoff_factor,
reconnect_jitter_ms: self.reconnect_jitter_ms,
connection_max_retries: self.connection_max_retries,
reconnect_max_attempts: self.reconnect_max_attempts,
heartbeat_timeout_secs: self.heartbeat_timeout_secs,
certs_dir: self.certs_dir.clone(),
}
}
}
#[cfg(test)]
mod tests {
use rstest::rstest;
use tokio_tungstenite::tungstenite::stream::Mode;
use super::{SocketConfig, SocketHeartbeat};
use crate::error::NetworkConfigError;
fn valid_config() -> SocketConfig {
SocketConfig::builder()
.url("tcp://127.0.0.1:8080".to_string())
.mode(Mode::Plain)
.suffix(vec![b'\n'])
.build()
.expect("baseline socket config should be valid")
}
#[rstest]
fn test_builder_accepts_valid_config() {
let result = SocketConfig::builder()
.url("tcp://127.0.0.1:8080".to_string())
.mode(Mode::Plain)
.suffix(vec![b'\n'])
.build();
assert!(result.is_ok());
}
#[rstest]
fn test_validate_accepts_zero_jitter() {
let mut config = valid_config();
config.reconnect_jitter_ms = Some(0);
assert!(config.validate().is_ok());
}
#[rstest]
fn test_validate_accepts_heartbeat_with_payload() {
let mut config = valid_config();
config.heartbeat = Some(SocketHeartbeat {
interval_secs: 5,
payload: b"ping".to_vec(),
});
assert!(config.validate().is_ok());
}
#[rstest]
#[case::derived(None, Some(15))]
#[case::explicit_wins(Some(20), Some(20))]
fn test_resolve_timeout_from_socket_heartbeat(
#[case] timeout_secs: Option<u64>,
#[case] expected: Option<u64>,
) {
let mut config = valid_config();
config.heartbeat = Some(SocketHeartbeat {
interval_secs: 5,
payload: b"ping".to_vec(),
});
config.heartbeat_timeout_secs = timeout_secs;
assert_eq!(config.resolved_heartbeat_timeout(), expected);
}
#[rstest]
#[case::empty_url(|c: &mut SocketConfig| c.url = String::new(), "url")]
#[case::heartbeat_interval(|c: &mut SocketConfig| { c.heartbeat = Some(SocketHeartbeat { interval_secs: 0, payload: vec![] }); }, "heartbeat")]
#[case::heartbeat_timeout_below_interval(|c: &mut SocketConfig| { c.heartbeat = Some(SocketHeartbeat { interval_secs: 5, payload: vec![b'p'] }); c.heartbeat_timeout_secs = Some(5); }, "heartbeat_timeout_secs")]
#[case::connect_timeout(|c: &mut SocketConfig| c.connect_timeout_ms = Some(0), "connect_timeout_ms")]
#[case::reconnect_delay_initial(|c: &mut SocketConfig| c.reconnect_delay_initial_ms = Some(0), "reconnect_delay_initial_ms")]
#[case::reconnect_delay_max(|c: &mut SocketConfig| c.reconnect_delay_max_ms = Some(0), "reconnect_delay_max_ms")]
#[case::heartbeat_timeout_zero(|c: &mut SocketConfig| c.heartbeat_timeout_secs = Some(0), "heartbeat_timeout_secs")]
fn test_validate_rejects_invalid_field(
#[case] mutate: fn(&mut SocketConfig),
#[case] expected_field: &str,
) {
let mut config = valid_config();
mutate(&mut config);
let err = config
.validate()
.expect_err("invalid value should be rejected");
assert!(
matches!(err, NetworkConfigError::Invalid { field, .. } if field == expected_field)
);
}
#[rstest]
#[case::too_small(0.5)]
#[case::too_large(100.1)]
#[case::nan(f64::NAN)]
#[case::infinite(f64::INFINITY)]
fn test_validate_rejects_invalid_backoff_factor(#[case] factor: f64) {
let mut config = valid_config();
config.reconnect_backoff_factor = Some(factor);
let err = config
.validate()
.expect_err("invalid backoff factor should be rejected");
assert!(
matches!(err, NetworkConfigError::Invalid { field, .. } if field == "reconnect_backoff_factor")
);
}
#[rstest]
fn test_validate_rejects_delay_initial_exceeding_max() {
let mut config = valid_config();
config.reconnect_delay_initial_ms = Some(5_000);
config.reconnect_delay_max_ms = Some(1_000);
let err = config
.validate()
.expect_err("initial delay above max should be rejected");
assert!(
matches!(err, NetworkConfigError::Invalid { field, .. } if field == "reconnect_delay_initial_ms")
);
}
#[rstest]
fn test_validate_collects_multiple_errors() {
let mut config = valid_config();
config.url = String::new();
config.connect_timeout_ms = Some(0);
let err = config.validate().expect_err("multiple invalid fields");
match err {
NetworkConfigError::Multiple { errors } => assert_eq!(errors.len(), 2),
other @ NetworkConfigError::Invalid { .. } => {
panic!("expected Multiple, was {other:?}")
}
}
}
}