1use std::fmt::Debug;
32
33use nautilus_core::string::secret::REDACTED;
34use tokio_tungstenite::tungstenite::stream::Mode;
35
36use super::types::TcpMessageHandler;
37use crate::error::{NetworkConfigError, NetworkConfigResult};
38
39#[derive(Clone, Debug, PartialEq, Eq)]
44pub struct SocketHeartbeat {
45 pub interval_secs: u64,
49 pub payload: Vec<u8>,
51}
52
53#[derive(Clone, bon::Builder)]
55#[builder(finish_fn(name = build_inner, vis = ""))]
56pub struct SocketConfig {
57 pub url: String,
59 pub mode: Mode,
61 pub suffix: Vec<u8>,
63 pub message_handler: Option<TcpMessageHandler>,
65 pub heartbeat: Option<SocketHeartbeat>,
77 pub connect_timeout_ms: Option<u64>,
83 pub reconnect_delay_initial_ms: Option<u64>,
85 pub reconnect_delay_max_ms: Option<u64>,
87 pub reconnect_backoff_factor: Option<f64>,
89 pub reconnect_jitter_ms: Option<u64>,
91 pub connection_max_retries: Option<u32>,
93 pub reconnect_max_attempts: Option<u32>,
99 pub heartbeat_timeout_secs: Option<u64>,
114 pub certs_dir: Option<String>,
116}
117
118impl<S: socket_config_builder::IsComplete> SocketConfigBuilder<S> {
119 pub fn build(self) -> NetworkConfigResult<SocketConfig> {
126 let config = self.build_inner();
127 config.validate()?;
128 Ok(config)
129 }
130}
131
132impl SocketConfig {
133 pub fn validate(&self) -> NetworkConfigResult<()> {
141 let mut errors = Vec::new();
142
143 if self.url.trim().is_empty() {
144 errors.push(NetworkConfigError::invalid("url", "must not be empty"));
145 }
146
147 if let Some(heartbeat) = &self.heartbeat
148 && heartbeat.interval_secs == 0
149 {
150 errors.push(NetworkConfigError::invalid(
151 "heartbeat",
152 "interval must be positive",
153 ));
154 }
155
156 if let (Some(heartbeat), Some(timeout_secs)) =
159 (&self.heartbeat, self.heartbeat_timeout_secs)
160 && timeout_secs <= heartbeat.interval_secs
161 {
162 errors.push(NetworkConfigError::invalid(
163 "heartbeat_timeout_secs",
164 format!(
165 "must exceed heartbeat interval ({}s), was {timeout_secs}s",
166 heartbeat.interval_secs
167 ),
168 ));
169 }
170
171 for (field, value) in [
174 ("connect_timeout_ms", self.connect_timeout_ms),
175 (
176 "reconnect_delay_initial_ms",
177 self.reconnect_delay_initial_ms,
178 ),
179 ("reconnect_delay_max_ms", self.reconnect_delay_max_ms),
180 ("heartbeat_timeout_secs", self.heartbeat_timeout_secs),
181 ] {
182 if let Some(value) = value
183 && value == 0
184 {
185 errors.push(NetworkConfigError::invalid(
186 field,
187 format!("must be positive, was {value}"),
188 ));
189 }
190 }
191
192 if let Some(factor) = self.reconnect_backoff_factor
193 && !(1.0..=100.0).contains(&factor)
194 {
195 errors.push(NetworkConfigError::invalid(
196 "reconnect_backoff_factor",
197 format!("must be in range [1.0, 100.0], was {factor}"),
198 ));
199 }
200
201 if let (Some(initial), Some(max)) =
202 (self.reconnect_delay_initial_ms, self.reconnect_delay_max_ms)
203 && initial > max
204 {
205 errors.push(NetworkConfigError::invalid(
206 "reconnect_delay_initial_ms",
207 format!("must not exceed reconnect_delay_max_ms ({max}), was {initial}"),
208 ));
209 }
210
211 NetworkConfigError::collect(errors)
212 }
213
214 pub(crate) fn resolved_heartbeat_timeout(&self) -> Option<u64> {
215 crate::heartbeat::resolve_heartbeat_timeout(
216 self.heartbeat_timeout_secs,
217 self.heartbeat
218 .as_ref()
219 .map(|heartbeat| heartbeat.interval_secs),
220 )
221 }
222}
223
224impl Debug for SocketConfig {
225 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
226 f.debug_struct(stringify!(SocketConfig))
227 .field("url", &REDACTED)
228 .field("mode", &self.mode)
229 .field("suffix", &self.suffix)
230 .field(
231 "message_handler",
232 &self.message_handler.as_ref().map(|_| "<function>"),
233 )
234 .field("heartbeat", &self.heartbeat)
235 .field("connect_timeout_ms", &self.connect_timeout_ms)
236 .field(
237 "reconnect_delay_initial_ms",
238 &self.reconnect_delay_initial_ms,
239 )
240 .field("reconnect_delay_max_ms", &self.reconnect_delay_max_ms)
241 .field("reconnect_backoff_factor", &self.reconnect_backoff_factor)
242 .field("reconnect_jitter_ms", &self.reconnect_jitter_ms)
243 .field("connection_max_retries", &self.connection_max_retries)
244 .field("reconnect_max_attempts", &self.reconnect_max_attempts)
245 .field("heartbeat_timeout_secs", &self.heartbeat_timeout_secs)
246 .field("certs_dir", &self.certs_dir)
247 .finish()
248 }
249}
250
251#[cfg(test)]
252mod tests {
253 use rstest::rstest;
254 use tokio_tungstenite::tungstenite::stream::Mode;
255
256 use super::{SocketConfig, SocketHeartbeat};
257 use crate::error::NetworkConfigError;
258
259 fn valid_config() -> SocketConfig {
260 SocketConfig::builder()
261 .url("tcp://127.0.0.1:8080".to_string())
262 .mode(Mode::Plain)
263 .suffix(vec![b'\n'])
264 .build()
265 .expect("baseline socket config should be valid")
266 }
267
268 #[rstest]
269 fn test_builder_accepts_valid_config() {
270 let result = SocketConfig::builder()
271 .url("tcp://127.0.0.1:8080".to_string())
272 .mode(Mode::Plain)
273 .suffix(vec![b'\n'])
274 .build();
275
276 assert!(result.is_ok());
277 }
278
279 #[rstest]
280 fn test_validate_accepts_zero_jitter() {
281 let mut config = valid_config();
282 config.reconnect_jitter_ms = Some(0);
283
284 assert!(config.validate().is_ok());
285 }
286
287 #[rstest]
288 fn test_validate_accepts_heartbeat_with_payload() {
289 let mut config = valid_config();
290 config.heartbeat = Some(SocketHeartbeat {
291 interval_secs: 5,
292 payload: b"ping".to_vec(),
293 });
294
295 assert!(config.validate().is_ok());
296 }
297
298 #[rstest]
299 #[case::derived(None, Some(15))]
300 #[case::explicit_wins(Some(20), Some(20))]
301 fn test_resolve_timeout_from_socket_heartbeat(
302 #[case] timeout_secs: Option<u64>,
303 #[case] expected: Option<u64>,
304 ) {
305 let mut config = valid_config();
306 config.heartbeat = Some(SocketHeartbeat {
307 interval_secs: 5,
308 payload: b"ping".to_vec(),
309 });
310 config.heartbeat_timeout_secs = timeout_secs;
311
312 assert_eq!(config.resolved_heartbeat_timeout(), expected);
313 }
314
315 #[rstest]
316 #[case::empty_url(|c: &mut SocketConfig| c.url = String::new(), "url")]
317 #[case::heartbeat_interval(|c: &mut SocketConfig| { c.heartbeat = Some(SocketHeartbeat { interval_secs: 0, payload: vec![] }); }, "heartbeat")]
318 #[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")]
319 #[case::connect_timeout(|c: &mut SocketConfig| c.connect_timeout_ms = Some(0), "connect_timeout_ms")]
320 #[case::reconnect_delay_initial(|c: &mut SocketConfig| c.reconnect_delay_initial_ms = Some(0), "reconnect_delay_initial_ms")]
321 #[case::reconnect_delay_max(|c: &mut SocketConfig| c.reconnect_delay_max_ms = Some(0), "reconnect_delay_max_ms")]
322 #[case::heartbeat_timeout_zero(|c: &mut SocketConfig| c.heartbeat_timeout_secs = Some(0), "heartbeat_timeout_secs")]
323 fn test_validate_rejects_invalid_field(
324 #[case] mutate: fn(&mut SocketConfig),
325 #[case] expected_field: &str,
326 ) {
327 let mut config = valid_config();
328 mutate(&mut config);
329
330 let err = config
331 .validate()
332 .expect_err("invalid value should be rejected");
333
334 assert!(
335 matches!(err, NetworkConfigError::Invalid { field, .. } if field == expected_field)
336 );
337 }
338
339 #[rstest]
340 #[case::too_small(0.5)]
341 #[case::too_large(100.1)]
342 #[case::nan(f64::NAN)]
343 #[case::infinite(f64::INFINITY)]
344 fn test_validate_rejects_invalid_backoff_factor(#[case] factor: f64) {
345 let mut config = valid_config();
346 config.reconnect_backoff_factor = Some(factor);
347
348 let err = config
349 .validate()
350 .expect_err("invalid backoff factor should be rejected");
351
352 assert!(
353 matches!(err, NetworkConfigError::Invalid { field, .. } if field == "reconnect_backoff_factor")
354 );
355 }
356
357 #[rstest]
358 fn test_validate_rejects_delay_initial_exceeding_max() {
359 let mut config = valid_config();
360 config.reconnect_delay_initial_ms = Some(5_000);
361 config.reconnect_delay_max_ms = Some(1_000);
362
363 let err = config
364 .validate()
365 .expect_err("initial delay above max should be rejected");
366
367 assert!(
368 matches!(err, NetworkConfigError::Invalid { field, .. } if field == "reconnect_delay_initial_ms")
369 );
370 }
371
372 #[rstest]
373 fn test_validate_collects_multiple_errors() {
374 let mut config = valid_config();
375 config.url = String::new();
376 config.connect_timeout_ms = Some(0);
377
378 let err = config.validate().expect_err("multiple invalid fields");
379
380 match err {
381 NetworkConfigError::Multiple { errors } => assert_eq!(errors.len(), 2),
382 other @ NetworkConfigError::Invalid { .. } => {
383 panic!("expected Multiple, was {other:?}")
384 }
385 }
386 }
387
388 #[rstest]
389 fn test_debug_redacts_endpoint_credentials() {
390 const ENDPOINT_PATH_SECRET: &str = "unique-endpoint-path-secret";
391 const ENDPOINT_QUERY_SECRET: &str = "unique-endpoint-query-secret";
392 let mut config = valid_config();
393 config.url =
394 format!("wss://rpc.example.com/{ENDPOINT_PATH_SECRET}?api_key={ENDPOINT_QUERY_SECRET}");
395
396 let debug = format!("{config:?}");
397
398 assert!(debug.contains("url: \"<redacted>\""));
399 assert!(!debug.contains(ENDPOINT_PATH_SECRET));
400 assert!(!debug.contains(ENDPOINT_QUERY_SECRET));
401 }
402}