use std::fmt::Debug;
use nautilus_core::string::secret::REDACTED;
use serde::{Deserialize, Serialize};
use crate::error::{NetworkConfigError, NetworkConfigResult};
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[cfg_attr(
feature = "python",
pyo3::pyclass(
module = "nautilus_trader.network",
eq,
from_py_object,
rename_all = "SCREAMING_SNAKE_CASE"
)
)]
#[cfg_attr(
feature = "python",
pyo3_stub_gen::derive::gen_stub_pyclass_enum(module = "nautilus_trader.network")
)]
#[allow(
clippy::unsafe_derive_deserialize,
reason = "network configuration requires strict serde decoding"
)]
pub enum TransportBackend {
#[cfg_attr(not(feature = "transport-sockudo"), default)]
Tungstenite,
#[cfg_attr(feature = "transport-sockudo", default)]
Sockudo,
}
#[allow(
clippy::unsafe_derive_deserialize,
reason = "network configuration requires strict serde decoding"
)]
#[derive(Clone, Serialize, Deserialize, bon::Builder)]
#[builder(finish_fn(name = build_inner, vis = ""))]
#[serde(deny_unknown_fields)]
pub struct WebSocketConfig {
pub url: String,
#[serde(default)]
#[builder(default)]
pub headers: Vec<(String, String)>,
#[serde(default)]
pub heartbeat_interval_secs: Option<u64>,
#[serde(default)]
pub heartbeat_payload: Option<String>,
#[serde(default)]
pub connect_timeout_ms: Option<u64>,
#[serde(default)]
pub reconnect_delay_initial_ms: Option<u64>,
#[serde(default)]
pub reconnect_delay_max_ms: Option<u64>,
#[serde(default)]
pub reconnect_backoff_factor: Option<f64>,
#[serde(default)]
pub reconnect_jitter_ms: Option<u64>,
#[serde(default)]
pub reconnect_max_attempts: Option<u32>,
#[serde(default)]
pub heartbeat_timeout_secs: Option<u64>,
#[serde(default)]
pub idle_timeout_ms: Option<u64>,
#[serde(default)]
#[builder(default)]
pub backend: TransportBackend,
#[serde(default)]
pub proxy_url: Option<String>,
}
impl Debug for WebSocketConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct(stringify!(WebSocketConfig))
.field("url", &REDACTED)
.field(
"headers",
&format_args!("<{} header(s)>", self.headers.len()),
)
.field("heartbeat_interval_secs", &self.heartbeat_interval_secs)
.field("heartbeat_payload", &self.heartbeat_payload)
.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("reconnect_max_attempts", &self.reconnect_max_attempts)
.field("heartbeat_timeout_secs", &self.heartbeat_timeout_secs)
.field("idle_timeout_ms", &self.idle_timeout_ms)
.field("backend", &self.backend)
.field("proxy_url", &self.proxy_url.as_ref().map(|_| REDACTED))
.finish()
}
}
impl<S: web_socket_config_builder::IsComplete> WebSocketConfigBuilder<S> {
pub fn build(self) -> NetworkConfigResult<WebSocketConfig> {
let config = self.build_inner();
config.validate()?;
Ok(config)
}
}
impl WebSocketConfig {
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(interval) = self.heartbeat_interval_secs
&& interval == 0
{
errors.push(NetworkConfigError::invalid(
"heartbeat_interval_secs",
"interval must be positive",
));
}
if let (Some(interval_secs), Some(timeout_secs)) =
(self.heartbeat_interval_secs, self.heartbeat_timeout_secs)
&& timeout_secs <= interval_secs
{
errors.push(NetworkConfigError::invalid(
"heartbeat_timeout_secs",
format!(
"must exceed heartbeat_interval_secs ({interval_secs}s), was {timeout_secs}s"
),
));
}
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),
("idle_timeout_ms", self.idle_timeout_ms),
] {
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_interval_secs,
)
}
}
#[cfg(test)]
mod tests {
use rstest::rstest;
use serde_json::json;
use super::WebSocketConfig;
use crate::error::NetworkConfigError;
#[rstest]
fn test_deserialize_websocket_config_rejects_unknown_field() {
let config = json!({
"url": "wss://example.com/ws",
"unexpected": true,
});
let error = serde_json::from_value::<WebSocketConfig>(config).unwrap_err();
assert!(error.to_string().contains("unknown field `unexpected`"));
}
fn valid_config() -> WebSocketConfig {
WebSocketConfig::builder()
.url("wss://example.com/ws".to_string())
.build()
.expect("baseline websocket config should be valid")
}
#[rstest]
fn test_builder_accepts_valid_config() {
let result = WebSocketConfig::builder()
.url("wss://example.com/ws".to_string())
.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]
#[case::empty_url(|c: &mut WebSocketConfig| c.url = String::new(), "url")]
#[case::heartbeat_interval(|c: &mut WebSocketConfig| c.heartbeat_interval_secs = Some(0), "heartbeat_interval_secs")]
#[case::heartbeat_timeout_below_interval(|c: &mut WebSocketConfig| { c.heartbeat_interval_secs = Some(30); c.heartbeat_timeout_secs = Some(30); }, "heartbeat_timeout_secs")]
#[case::connect_timeout(|c: &mut WebSocketConfig| c.connect_timeout_ms = Some(0), "connect_timeout_ms")]
#[case::reconnect_delay_initial(|c: &mut WebSocketConfig| c.reconnect_delay_initial_ms = Some(0), "reconnect_delay_initial_ms")]
#[case::reconnect_delay_max(|c: &mut WebSocketConfig| c.reconnect_delay_max_ms = Some(0), "reconnect_delay_max_ms")]
#[case::heartbeat_timeout_zero(|c: &mut WebSocketConfig| c.heartbeat_timeout_secs = Some(0), "heartbeat_timeout_secs")]
#[case::idle_timeout(|c: &mut WebSocketConfig| c.idle_timeout_ms = Some(0), "idle_timeout_ms")]
fn test_validate_rejects_invalid_field(
#[case] mutate: fn(&mut WebSocketConfig),
#[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:?}")
}
}
}
#[rstest]
#[case::derived(Some(30), None, Some(90))]
#[case::explicit_wins(Some(30), Some(45), Some(45))]
fn test_resolve_timeout_from_websocket_heartbeat(
#[case] interval_secs: Option<u64>,
#[case] timeout_secs: Option<u64>,
#[case] expected: Option<u64>,
) {
let mut config = valid_config();
config.heartbeat_interval_secs = interval_secs;
config.heartbeat_timeout_secs = timeout_secs;
assert_eq!(config.resolved_heartbeat_timeout(), expected);
}
#[rstest]
fn test_debug_redacts_endpoint_and_proxy_credentials() {
const ENDPOINT_PATH_SECRET: &str = "unique-endpoint-path-secret";
const ENDPOINT_QUERY_SECRET: &str = "unique-endpoint-query-secret";
const PROXY_SECRET: &str = "unique-proxy-secret";
let mut config = valid_config();
config.url =
format!("wss://rpc.example.com/{ENDPOINT_PATH_SECRET}?api_key={ENDPOINT_QUERY_SECRET}");
config.proxy_url = Some(format!(
"http://proxytest:{PROXY_SECRET}@proxy.example.com:8080"
));
let debug = format!("{config:?}");
assert!(debug.contains("url: \"<redacted>\""));
assert!(debug.contains("proxy_url: Some(\"<redacted>\")"));
assert!(!debug.contains(ENDPOINT_PATH_SECRET));
assert!(!debug.contains(ENDPOINT_QUERY_SECRET));
assert!(!debug.contains(PROXY_SECRET));
}
}