1use futures_util::{
2 SinkExt, StreamExt,
3 stream::{SplitSink, SplitStream},
4};
5use iso_currency::Currency;
6use serde::{Deserialize, Serialize};
7use serde_json::{Value, json};
8use std::{
9 fmt::Debug,
10 sync::{
11 Arc,
12 atomic::{AtomicBool, AtomicU64, Ordering},
13 },
14 time::{Duration, Instant, SystemTime, UNIX_EPOCH},
15};
16use tokio::{
17 net::TcpStream,
18 sync::{Mutex, MutexGuard, RwLock, mpsc},
19 task::JoinHandle,
20 time::timeout,
21};
22use tokio_tungstenite::{
23 MaybeTlsStream, WebSocketStream, connect_async_with_config,
24 tungstenite::{
25 client::IntoClientRequest,
26 protocol::{Message, WebSocketConfig},
27 },
28};
29use tokio_util::sync::CancellationToken;
30use tracing::{debug, error, info, trace, warn};
31use url::Url;
32use ustr::{Ustr, ustr};
33
34use crate::{
35 Error, Interval, MarketAdjustment, Result, SessionType, SocketServerInfo, Timezone,
36 chart::{ChartOptions, options::Range},
37 live::{
38 handler::Handler,
39 models::{
40 DataServer, Socket, SocketMessage, SocketMessageDe, SocketMessageSer,
41 TradingViewDataEvent, WEBSOCKET_HEADERS,
42 },
43 },
44 payload,
45 quote::ALL_QUOTE_FIELDS,
46 study::StudyConfiguration,
47 utils::{parse_packet, symbol_init},
48};
49
50#[derive(Debug, Clone, Copy)]
52pub(crate) struct ErrorRecoveryConfig {
53 max_consecutive_errors: u64,
54 error_reset_interval: Duration,
55 max_recovery_attempts: u32,
56 backoff_base_delay: Duration,
57 backoff_max_delay: Duration,
58 connection_timeout: Duration,
59 ping_interval: Duration,
60 health_check_interval: Duration,
61}
62
63impl Default for ErrorRecoveryConfig {
64 fn default() -> Self {
65 Self {
66 max_consecutive_errors: 5,
67 error_reset_interval: Duration::from_secs(60),
68 max_recovery_attempts: 3,
69 backoff_base_delay: Duration::from_millis(100),
70 backoff_max_delay: Duration::from_secs(30),
71 connection_timeout: Duration::from_secs(30),
72 ping_interval: Duration::from_secs(30),
73 health_check_interval: Duration::from_secs(60),
74 }
75 }
76}
77
78#[derive(Debug, Clone)]
79struct ErrorStats {
80 consecutive_errors: Arc<AtomicU64>,
81 last_error_time: Arc<RwLock<Option<Instant>>>,
82 total_errors: Arc<AtomicU64>,
83 recovery_attempts: Arc<AtomicU64>,
84 connection_drops: Arc<AtomicU64>,
85 last_successful_message: Arc<RwLock<Option<Instant>>>,
86 recent_critical_times: Arc<RwLock<[u64; 4]>>,
87 recent_critical_count: Arc<AtomicU64>,
88}
89
90impl Default for ErrorStats {
91 fn default() -> Self {
92 Self {
93 consecutive_errors: Arc::new(AtomicU64::new(0)),
94 last_error_time: Arc::new(RwLock::new(None)),
95 total_errors: Arc::new(AtomicU64::new(0)),
96 recovery_attempts: Arc::new(AtomicU64::new(0)),
97 connection_drops: Arc::new(AtomicU64::new(0)),
98 last_successful_message: Arc::new(RwLock::new(Some(Instant::now()))),
99 recent_critical_times: Arc::new(RwLock::new([0u64; 4])),
100 recent_critical_count: Arc::new(AtomicU64::new(0)),
101 }
102 }
103}
104
105impl ErrorStats {
106 fn increment_error(&self) -> u64 {
107 let count = self.consecutive_errors.fetch_add(1, Ordering::SeqCst) + 1;
108 self.total_errors.fetch_add(1, Ordering::SeqCst);
109 count
110 }
111
112 fn reset_consecutive(&self) {
113 self.consecutive_errors.store(0, Ordering::SeqCst);
114 }
115
116 async fn update_last_error_time(&self) {
117 let mut last_error = self.last_error_time.write().await;
118 *last_error = Some(Instant::now());
119 }
120
121 async fn update_last_successful_message(&self) {
122 let mut last_success = self.last_successful_message.write().await;
123 *last_success = Some(Instant::now());
124 }
125
126 async fn record_critical_error(&self, now_secs: u64) {
127 let mut times = self.recent_critical_times.write().await;
128 times[0] = times[1];
129 times[1] = times[2];
130 times[2] = times[3];
131 times[3] = now_secs;
132 self.recent_critical_count.fetch_add(1, Ordering::Relaxed);
133 }
134
135 async fn should_reset_consecutive_errors(&self, reset_interval: Duration) -> bool {
136 let last_error = self.last_error_time.read().await;
137 if let Some(last_time) = *last_error {
138 last_time.elapsed() > reset_interval
139 } else {
140 false
141 }
142 }
143
144 async fn is_connection_stale(&self, stale_threshold: Duration) -> bool {
145 let last_success = self.last_successful_message.read().await;
146 if let Some(last_time) = *last_success {
147 last_time.elapsed() > stale_threshold
148 } else {
149 true }
151 }
152
153 fn get_consecutive_errors(&self) -> u64 {
154 self.consecutive_errors.load(Ordering::SeqCst)
155 }
156
157 async fn get_recent_critical_count(&self, window_secs: u64) -> usize {
158 if self.recent_critical_count.load(Ordering::Relaxed) == 0 {
159 return 0;
160 }
161 let now = SystemTime::now()
162 .duration_since(UNIX_EPOCH)
163 .unwrap_or_default()
164 .as_secs();
165 if let Ok(times) = self.recent_critical_times.try_read() {
166 times
167 .iter()
168 .filter(|&&t| t > 0 && now.saturating_sub(t) <= window_secs)
169 .count()
170 } else {
171 0
172 }
173 }
174}
175
176#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
177pub enum ErrorSeverity {
182 Trace,
184 Minor,
186 Moderate,
188 Critical,
190 Fatal,
192}
193
194#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)]
196pub enum ConnectionHealth {
198 Healthy,
199 Degraded,
200 Unstable,
201 Failed,
202}
203
204#[derive(Debug, Clone, Copy)]
205struct HealthMetrics {
206 health: ConnectionHealth,
207 last_ping_time: Option<Instant>,
208 avg_response_time: Duration,
209}
210
211impl Default for HealthMetrics {
212 fn default() -> Self {
213 Self {
214 health: ConnectionHealth::Healthy,
215 last_ping_time: None,
216 avg_response_time: Duration::from_millis(0),
217 }
218 }
219}
220
221#[derive(Debug, Clone, Default, Deserialize, Serialize, Copy)]
227pub struct SeriesInfo {
228 pub chart_session: Ustr,
229 pub options: ChartOptions,
230}
231
232pub struct WebSocketClient<T: Handler> {
253 pub server: DataServer,
254 pub auth_token: Arc<RwLock<Ustr>>,
255
256 handler: T,
257
258 cancellation: CancellationToken,
260 is_closed: Arc<AtomicBool>,
261
262 read: Arc<Mutex<SplitStream<WebSocketStream<MaybeTlsStream<TcpStream>>>>>,
263 write_tx: Arc<RwLock<mpsc::Sender<Message>>>,
270 writer_handle: Arc<Mutex<Option<JoinHandle<()>>>>,
272 buffer_size: usize,
273
274 error_stats: ErrorStats,
276 error_config: ErrorRecoveryConfig,
277 health_metrics: Arc<RwLock<HealthMetrics>>,
278
279 circuit_breaker_open: Arc<AtomicBool>,
281 circuit_breaker_opened_at: Arc<RwLock<Option<Instant>>>,
282}
283
284#[bon::bon]
285#[allow(private_interfaces)]
286impl<T: Handler> WebSocketClient<T> {
287 #[builder]
288 pub async fn new(
289 auth_token: Option<&str>,
290 #[builder(default = DataServer::ProData)] server: DataServer,
291 handler: T,
292 #[builder(default = 1024*1024)] buffer_size: usize,
293 #[builder(default)] error_config: ErrorRecoveryConfig,
294 ) -> Result<Arc<Self>> {
295 let auth_token = Ustr::from(auth_token.unwrap_or("unauthorized_user_token"));
296 let (write, read) = Self::connect(server, Some(buffer_size)).await?;
297
298 let is_closed = Arc::new(AtomicBool::new(false));
299 let auth_token = Arc::new(RwLock::new(auth_token));
300 let read = Arc::new(Mutex::new(read));
301
302 let (write_tx, write_rx) = mpsc::channel(1024);
304 let write_tx = Arc::new(RwLock::new(write_tx));
305 let writer_handle = Arc::new(Mutex::new(None::<JoinHandle<()>>));
306
307 let client = Arc::new(Self {
308 handler,
309 server,
310 read,
311 write_tx: write_tx.clone(),
312 writer_handle: writer_handle.clone(),
313 auth_token,
314 is_closed: is_closed.clone(),
315 buffer_size,
316 cancellation: CancellationToken::new(),
317 error_stats: ErrorStats::default(),
318 error_config,
319 health_metrics: Arc::new(RwLock::new(HealthMetrics::default())),
320 circuit_breaker_open: Arc::new(AtomicBool::new(false)),
321 circuit_breaker_opened_at: Arc::new(RwLock::new(None)),
322 });
323
324 Self::spawn_writer(write, write_rx, writer_handle, is_closed.clone());
326
327 client.spawn_health_monitor();
329
330 Ok(client)
331 }
332
333 pub fn spawn_reader_task(self: Arc<Self>) {
334 tokio::spawn(async move {
335 if let Err(e) = self.subscribe().await {
336 error!("Reader task failed: {}", e);
337 }
338 });
339 }
340
341 fn spawn_writer(
344 mut sink: SplitSink<WebSocketStream<MaybeTlsStream<TcpStream>>, Message>,
345 mut rx: mpsc::Receiver<Message>,
346 handle_storage: Arc<Mutex<Option<JoinHandle<()>>>>,
347 is_closed: Arc<AtomicBool>,
348 ) {
349 let handle = tokio::spawn(async move {
350 while let Some(msg) = rx.recv().await {
351 if is_closed.load(Ordering::Relaxed) {
352 break;
353 }
354 if sink.send(msg).await.is_err() {
355 is_closed.store(true, Ordering::Relaxed);
356 break;
357 }
358 }
359 let _ = sink.close().await;
361 is_closed.store(true, Ordering::Relaxed);
362 });
363
364 tokio::spawn(async move {
366 let mut guard = handle_storage.lock().await;
367 *guard = Some(handle);
368 });
369 }
370
371 async fn connect(
372 server: DataServer,
373 buffer_size: Option<usize>,
374 ) -> Result<(
375 SplitSink<WebSocketStream<MaybeTlsStream<TcpStream>>, Message>,
376 SplitStream<WebSocketStream<MaybeTlsStream<TcpStream>>>,
377 )> {
378 let url = Url::parse(&format!(
379 "wss://{server}.tradingview.com/socket.io/websocket"
380 ))?;
381
382 let buffer_size = buffer_size.unwrap_or(1024 * 1024);
383
384 let mut request = url.into_client_request()?;
385 request.headers_mut().extend((*WEBSOCKET_HEADERS).clone());
386
387 let conf = WebSocketConfig::default()
389 .read_buffer_size(buffer_size)
390 .write_buffer_size(buffer_size);
391
392 let (socket, response) = connect_async_with_config(request, Some(conf), false).await?;
393
394 info!("WebSocket connected with status: {}", response.status());
395
396 let (write, read) = socket.split();
397
398 Ok((write, read))
399 }
400
401 async fn classify_error_severity(&self, error: &Error, context: &str) -> ErrorSeverity {
403 let consecutive_errors = self.error_stats.get_consecutive_errors();
405 let recent_critical = self.error_stats.get_recent_critical_count(300).await;
406
407 let base_severity = match error {
409 Error::WebSocket(msg) => {
410 if msg.contains("ConnectionClosed") || msg.contains("ConnectionReset") {
411 ErrorSeverity::Critical
412 } else if msg.contains("timeout") || msg.contains("WouldBlock") {
413 ErrorSeverity::Moderate
414 } else if msg.contains("Protocol") {
415 ErrorSeverity::Critical
416 } else {
417 ErrorSeverity::Moderate
418 }
419 }
420
421 Error::TradingView { source } => {
422 use crate::error::TradingViewError;
423 match source {
424 TradingViewError::CriticalError => ErrorSeverity::Fatal,
425 TradingViewError::ProtocolError => ErrorSeverity::Critical,
426 TradingViewError::SymbolError | TradingViewError::SeriesError => {
427 ErrorSeverity::Minor
428 }
429 _ => ErrorSeverity::Trace,
430 }
431 }
432
433 Error::JsonParse(_) => {
434 if consecutive_errors > 3 {
435 ErrorSeverity::Moderate
436 } else {
437 ErrorSeverity::Minor
438 }
439 }
440
441 Error::Internal(msg) => {
442 if msg.contains("connection") || msg.contains("timeout") {
443 ErrorSeverity::Critical
444 } else if context.contains("critical") {
445 ErrorSeverity::Fatal
446 } else {
447 ErrorSeverity::Moderate
448 }
449 }
450
451 _ => ErrorSeverity::Moderate,
452 };
453
454 if recent_critical >= 3 {
456 ErrorSeverity::Fatal
457 } else if consecutive_errors >= self.error_config.max_consecutive_errors {
458 std::cmp::max(base_severity, ErrorSeverity::Critical)
459 } else {
460 base_severity
461 }
462 }
463
464 async fn attempt_error_recovery(&self, severity: ErrorSeverity, error: &Error) -> Result<bool> {
465 if self.is_circuit_breaker_open().await {
467 warn!("Circuit breaker is open, skipping recovery attempt");
468 return Ok(false);
469 }
470
471 match severity {
472 ErrorSeverity::Trace => {
473 trace!("Trace level error, no action needed: {}", error);
474 Ok(true)
475 }
476
477 ErrorSeverity::Minor => {
478 debug!("Minor error occurred, continuing: {}", error);
479 Ok(true)
480 }
481
482 ErrorSeverity::Moderate => {
483 warn!(
484 "Moderate error occurred, attempting soft recovery: {}",
485 error
486 );
487
488 if self
490 .error_stats
491 .should_reset_consecutive_errors(self.error_config.error_reset_interval)
492 .await
493 {
494 self.error_stats.reset_consecutive();
495 info!("Reset consecutive error count after timeout period");
496 }
497
498 if let Err(health_err) = self.perform_health_check().await {
500 warn!("Health check failed during recovery: {}", health_err);
501 return Ok(false);
502 }
503
504 Ok(true)
505 }
506
507 ErrorSeverity::Critical => {
508 error!(
509 "Critical error occurred, attempting reconnection: {}",
510 error
511 );
512 self.error_stats
513 .recovery_attempts
514 .fetch_add(1, Ordering::SeqCst);
515
516 for attempt in 1..=self.error_config.max_recovery_attempts {
518 let delay = self.calculate_backoff_delay(attempt);
519 warn!("Recovery attempt {} after {:?} delay", attempt, delay);
520
521 tokio::time::sleep(delay).await;
522
523 match timeout(self.error_config.connection_timeout, self.reconnect()).await {
524 Ok(Ok(_)) => {
525 info!(
526 "Successfully recovered from critical error (attempt {})",
527 attempt
528 );
529 self.error_stats.reset_consecutive();
530 return Ok(true);
531 }
532 Ok(Err(reconnect_err)) => {
533 error!("Reconnection attempt {} failed: {}", attempt, reconnect_err);
534 }
535 Err(_) => {
536 error!("Reconnection attempt {} timed out", attempt);
537 }
538 }
539 }
540
541 self.open_circuit_breaker().await;
543 Ok(false)
544 }
545
546 ErrorSeverity::Fatal => {
547 error!("Fatal error occurred, terminating connection: {}", error);
548 self.is_closed.store(true, Ordering::Relaxed);
549 self.cancellation.cancel();
550 self.open_circuit_breaker().await;
551 Ok(false)
552 }
553 }
554 }
555
556 fn calculate_backoff_delay(&self, attempt: u32) -> Duration {
558 let delay = self
559 .error_config
560 .backoff_base_delay
561 .mul_f64((2_f64).powi(attempt as i32 - 1));
562
563 std::cmp::min(delay, self.error_config.backoff_max_delay)
564 }
565
566 async fn is_circuit_breaker_open(&self) -> bool {
568 if !self.circuit_breaker_open.load(Ordering::Relaxed) {
569 return false;
570 }
571
572 let opened_at = self.circuit_breaker_opened_at.read().await;
574 if let Some(time) = *opened_at
575 && time.elapsed() > Duration::from_secs(300)
576 {
577 self.circuit_breaker_open.store(false, Ordering::Relaxed);
579 info!("Circuit breaker reset after timeout");
580 return false;
581 }
582
583 true
584 }
585
586 async fn open_circuit_breaker(&self) {
587 self.circuit_breaker_open.store(true, Ordering::Relaxed);
588 let mut opened_at = self.circuit_breaker_opened_at.write().await;
589 *opened_at = Some(Instant::now());
590 error!("Circuit breaker opened due to repeated failures");
591 }
592
593 async fn perform_health_check(&self) -> Result<()> {
595 if self.is_closed() {
596 return Err(Error::Internal(ustr("Connection is closed")));
597 }
598
599 if self
601 .error_stats
602 .is_connection_stale(Duration::from_secs(120))
603 .await
604 {
605 warn!("Connection appears stale, performing ping test");
606 }
607
608 let start = Instant::now();
610 self.try_ping().await?;
611 let ping_duration = start.elapsed();
612
613 let mut metrics = self.health_metrics.write().await;
615 metrics.last_ping_time = Some(start);
616
617 if metrics.avg_response_time.is_zero() {
619 metrics.avg_response_time = ping_duration;
620 } else {
621 metrics.avg_response_time = Duration::from_nanos(
622 (metrics.avg_response_time.as_nanos() as f64 * 0.8
623 + ping_duration.as_nanos() as f64 * 0.2) as u64,
624 );
625 }
626
627 metrics.health = if ping_duration > Duration::from_secs(5) {
629 ConnectionHealth::Degraded
630 } else if self.error_stats.get_consecutive_errors() > 2 {
631 ConnectionHealth::Unstable
632 } else {
633 ConnectionHealth::Healthy
634 };
635
636 debug!(
637 "Health check completed: {:?}, ping: {:?}",
638 metrics.health, ping_duration
639 );
640 Ok(())
641 }
642
643 fn spawn_health_monitor(self: &Arc<Self>) {
645 let client = Arc::clone(self);
646 tokio::spawn(async move {
647 let mut interval = tokio::time::interval(client.error_config.health_check_interval);
648
649 loop {
650 tokio::select! {
651 _ = interval.tick() => {
652 if client.is_closed() {
653 break;
654 }
655
656 if let Err(e) = client.perform_health_check().await {
657 warn!("Scheduled health check failed: {}", e);
658 }
659 }
660 _ = client.cancellation.cancelled() => {
661 debug!("Health monitor task cancelled");
662 break;
663 }
664 }
665 }
666 });
667 }
668
669 async fn notify_error_handlers(&self, error: &Error, context: &str, severity: ErrorSeverity) {
671 if severity >= ErrorSeverity::Critical {
673 let now = SystemTime::now()
674 .duration_since(UNIX_EPOCH)
675 .unwrap_or_default()
676 .as_secs();
677 self.error_stats.record_critical_error(now).await;
678 }
679
680 let health_metrics = self.health_metrics.read().await;
682 let error_context = vec![json!({
683 "error_type": format!("{:?}", error),
684 "context": context,
685 "severity": format!("{:?}", severity),
686 "consecutive_errors": self.error_stats.get_consecutive_errors(),
687 "total_errors": self.error_stats.total_errors.load(Ordering::SeqCst),
688 "recovery_attempts": self.error_stats.recovery_attempts.load(Ordering::SeqCst),
689 "connection_health": format!("{:?}", health_metrics.health),
690 "avg_response_time_ms": health_metrics.avg_response_time.as_millis(),
691 "circuit_breaker_open": self.circuit_breaker_open.load(Ordering::Relaxed),
692 "timestamp": chrono::Utc::now().to_rfc3339(),
693 })];
694
695 self.handler.notify_error(*error, &error_context);
697 }
698
699 fn log_error(&self, error: &Error, context: &str, severity: &ErrorSeverity) {
701 let consecutive = self.error_stats.get_consecutive_errors();
702 let total = self.error_stats.total_errors.load(Ordering::SeqCst);
703 let recovery_attempts = self.error_stats.recovery_attempts.load(Ordering::SeqCst);
704
705 let error_info = format!(
706 "{} (consecutive: {}, total: {}, recovery_attempts: {})",
707 error, consecutive, total, recovery_attempts
708 );
709
710 match severity {
711 ErrorSeverity::Trace => {
712 trace!("Trace error in {}: {}", context, error_info);
713 }
714 ErrorSeverity::Minor => {
715 debug!("Minor error in {}: {}", context, error_info);
716 }
717 ErrorSeverity::Moderate => {
718 warn!("Moderate error in {}: {}", context, error_info);
719 }
720 ErrorSeverity::Critical => {
721 error!("Critical error in {}: {}", context, error_info);
722 }
723 ErrorSeverity::Fatal => {
724 error!("FATAL error in {}: {}", context, error_info);
725 }
726 }
727 }
728
729 pub fn is_closed(&self) -> bool {
730 self.is_closed.load(Ordering::Relaxed)
731 }
732
733 pub async fn reconnect(&self) -> Result<()> {
734 let auth_token = self.auth_token.read().await;
735
736 let mut wh = self.writer_handle.lock().await;
738 if let Some(handle) = wh.take() {
739 handle.abort();
740 }
741 drop(wh);
742
743 let (write, read) = Self::connect(self.server, Some(self.buffer_size)).await?;
744
745 let (new_tx, new_rx) = mpsc::channel(1024);
747 Self::spawn_writer(
748 write,
749 new_rx,
750 self.writer_handle.clone(),
751 self.is_closed.clone(),
752 );
753
754 {
756 let mut tx_guard = self.write_tx.write().await;
757 *tx_guard = new_tx;
758 }
759
760 let mut read_guard = self.read.lock().await;
761 *read_guard = read;
762 self.is_closed.store(false, Ordering::Relaxed);
763 self.set_auth_token(&auth_token).await?;
764 Ok(())
765 }
766
767 #[tracing::instrument(skip(self), level = "debug")]
768 pub async fn send_raw_message(&self, message: &str) -> Result<()> {
769 if self.is_closed.load(Ordering::Relaxed) {
770 return Err(Error::Internal("WebSocket is closed".into()));
771 }
772
773 if self.is_circuit_breaker_open().await {
775 return Err(Error::Internal("Circuit breaker is open".into()));
776 }
777
778 let tx = self.write_tx.read().await;
779 match timeout(
780 Duration::from_secs(10),
781 tx.send(Message::Text(message.into())),
782 )
783 .await
784 {
785 Ok(Ok(_)) => {
786 self.error_stats.update_last_successful_message().await;
787 Ok(())
788 }
789 Ok(Err(e)) => Err(Error::WebSocket(e.to_string().into())),
790 Err(_) => Err(Error::Internal("Send timeout".into())),
791 }
792 }
793
794 #[tracing::instrument(skip(self), level = "debug")]
795 pub async fn send(&self, m: &str, p: &[Value]) -> Result<()> {
796 if self.is_closed.load(Ordering::Relaxed) {
797 return Err(Error::Internal("WebSocket is closed".into()));
798 }
799 let tx = self.write_tx.read().await;
800 tx.send(SocketMessageSer::new(m, p).to_message()?)
801 .await
802 .map_err(|e| Error::WebSocket(e.to_string().into()))?;
803 Ok(())
804 }
805
806 #[tracing::instrument(skip(self), level = "debug")]
807 pub async fn ping(&self, ping: &Message) -> Result<()> {
808 let tx = self.write_tx.read().await;
809 tx.send(ping.clone())
810 .await
811 .map_err(|e| Error::WebSocket(e.to_string().into()))?;
812 if ping.is_close() {
813 self.is_closed.store(true, Ordering::Relaxed);
814 tracing::warn!("ping message is close, closing session");
815 }
816 Ok(())
817 }
818
819 pub async fn close(&self) -> Result<()> {
820 self.is_closed.store(true, Ordering::Relaxed);
821 let (dummy_tx, _dummy_rx) = mpsc::channel::<Message>(1);
825 let mut tx_guard = self.write_tx.write().await;
826 let old_tx = std::mem::replace(&mut *tx_guard, dummy_tx);
827 drop(old_tx); drop(tx_guard);
829 Ok(())
830 }
831
832 #[tracing::instrument(skip(self), level = "debug")]
833 pub async fn fast_symbols(&self, quote_session: &str, symbols: &[&str]) -> Result<()> {
834 let mut payloads = payload![quote_session];
835 payloads.extend(symbols.iter().map(|s| Value::from(*s)));
836 self.send("quote_fast_symbols", &payloads).await?;
837 Ok(())
838 }
839
840 #[tracing::instrument(skip(self), level = "debug")]
841 pub async fn create_quote_session(&self, quote_session: &str) -> Result<()> {
842 self.send("quote_create_session", &payload!(quote_session))
843 .await?;
844 Ok(())
845 }
846
847 #[tracing::instrument(skip(self), level = "debug")]
848 pub async fn delete_quote_session(&self, quote_session: &str) -> Result<()> {
849 self.send("quote_delete_session", &payload!(quote_session))
850 .await?;
851 Ok(())
852 }
853
854 #[tracing::instrument(skip(self), level = "debug")]
855 pub async fn set_fields(&self, quote_session: &str) -> Result<()> {
856 let mut quote_fields = payload![quote_session];
857 quote_fields.extend(ALL_QUOTE_FIELDS.iter().copied().map(Value::from));
858 self.send("quote_set_fields", "e_fields).await?;
859 Ok(())
860 }
861
862 #[tracing::instrument(skip(self), level = "debug")]
863 pub async fn add_symbols(&self, quote_session: &str, symbols: &[&str]) -> Result<()> {
864 let mut payloads = payload![quote_session];
865 payloads.extend(symbols.iter().map(|s| Value::from(*s)));
866 self.send("quote_add_symbols", &payloads).await?;
867 info!("Added {} symbols to quote session", symbols.len());
868 Ok(())
869 }
870
871 #[tracing::instrument(skip(self), level = "debug")]
872 pub async fn remove_symbols(&self, quote_session: &str, symbols: &[&str]) -> Result<()> {
873 let mut payloads = payload![quote_session];
874 payloads.extend(symbols.iter().map(|s| Value::from(*s)));
875 self.send("quote_remove_symbols", &payloads).await?;
876 Ok(())
877 }
878
879 #[tracing::instrument(skip(self), level = "debug")]
880 pub async fn set_auth_token(&self, auth_token: &str) -> Result<()> {
881 let mut auth_token_ = self.auth_token.write().await;
882 *auth_token_ = ustr(auth_token);
883 self.send("set_auth_token", &payload!(auth_token)).await?;
884 Ok(())
885 }
886
887 #[tracing::instrument(skip(self), level = "debug")]
889 pub async fn set_locale(&self, language_code: &str, country: &str) -> Result<()> {
890 self.send("set_locale", &payload!(language_code, country))
891 .await?;
892 Ok(())
893 }
894
895 #[tracing::instrument(skip(self), level = "debug")]
896 pub async fn set_data_quality(&self, data_quality: &str) -> Result<()> {
897 self.send("set_data_quality", &payload!(data_quality))
898 .await?;
899
900 Ok(())
901 }
902
903 #[tracing::instrument(skip(self), level = "debug")]
904 pub async fn set_timezone(&self, chart_session: &str, timezone: Timezone) -> Result<()> {
905 self.send(
906 "switch_timezone",
907 &payload!(chart_session, timezone.to_string()),
908 )
909 .await?;
910
911 Ok(())
912 }
913
914 #[tracing::instrument(skip(self), level = "debug")]
915 pub async fn create_chart_session(&self, session: &str) -> Result<()> {
916 self.send("chart_create_session", &payload!(session, ""))
919 .await?;
920 Ok(())
921 }
922
923 #[tracing::instrument(skip(self), level = "debug")]
940 #[builder]
941 pub async fn create_series(
942 &self,
943 chart_session: &str,
944 series_identifier: &str, series_id: &str, symbol_series_id: &str, interval: Interval,
948 bar_count: u64,
949 range: Option<Range>,
950 ) -> Result<()> {
951 if let Some(r) = range {
954 let args = payload!(
955 chart_session,
956 series_identifier,
957 series_id,
958 symbol_series_id,
959 interval.to_string(),
960 0u64, r.to_string() );
963 self.send("create_series", &args).await?;
964 } else {
965 let args: Vec<Value> = vec![
966 Value::from(chart_session),
967 Value::from(series_identifier),
968 Value::from(series_id),
969 Value::from(symbol_series_id),
970 Value::from(interval.to_string()),
971 Value::from(bar_count),
972 ];
973 self.send("create_series", &args).await?;
974 }
975
976 Ok(())
977 }
978
979 #[tracing::instrument(skip(self), level = "debug")]
983 #[builder]
984 pub async fn modify_series(
985 &self,
986 chart_session: &str,
987 series_identifier: &str, series_id: &str, symbol_series_id: &str, interval: Interval,
991 bar_count: u64,
992 range: Option<Range>,
993 ) -> Result<()> {
994 if let Some(r) = range {
995 let args = payload!(
996 chart_session,
997 series_identifier,
998 series_id,
999 symbol_series_id,
1000 interval.to_string(),
1001 0u64,
1002 r.to_string()
1003 );
1004 self.send("modify_series", &args).await?;
1005 } else {
1006 let args: Vec<Value> = vec![
1007 Value::from(chart_session),
1008 Value::from(series_identifier),
1009 Value::from(series_id),
1010 Value::from(symbol_series_id),
1011 Value::from(interval.to_string()),
1012 Value::from(bar_count),
1013 ];
1014 self.send("modify_series", &args).await?;
1015 }
1016
1017 Ok(())
1018 }
1019
1020 #[tracing::instrument(skip(self), level = "debug")]
1021 pub async fn remove_series(&self, chart_session: &str, series_identifier: &str) -> Result<()> {
1022 self.send("remove_series", &payload!(chart_session, series_identifier))
1023 .await?;
1024 Ok(())
1025 }
1026
1027 #[tracing::instrument(skip(self), level = "debug")]
1028 #[builder]
1029 pub async fn resolve_symbol(
1030 &self,
1031 session: &str,
1032 symbol_series_id: &str,
1033 instrument: &str,
1034 adjustment: Option<MarketAdjustment>,
1035 currency: Option<Currency>,
1036 session_type: Option<SessionType>,
1037 replay_session: Option<&str>,
1038 ) -> Result<()> {
1039 self.send(
1040 "resolve_symbol",
1041 &payload!(
1042 session,
1043 symbol_series_id,
1044 symbol_init()
1045 .instrument(instrument)
1046 .maybe_adjustment(adjustment)
1047 .maybe_currency(currency)
1048 .maybe_session_type(session_type)
1049 .maybe_replay(replay_session)
1050 .call()?
1051 ),
1052 )
1053 .await?;
1054 Ok(())
1055 }
1056
1057 #[tracing::instrument(skip(self), level = "debug")]
1058 pub async fn delete_chart_session(&self, session: &str) -> Result<()> {
1059 self.send("chart_delete_session", &payload!(session))
1060 .await?;
1061 Ok(())
1062 }
1063
1064 #[tracing::instrument(skip(self), level = "debug")]
1065 pub async fn request_more_data(
1066 &self,
1067 chart_session: &str,
1068 series_id: &str,
1069 num: u64,
1070 ) -> Result<()> {
1071 self.send(
1072 "request_more_data",
1073 &payload!(chart_session, series_id, num),
1074 )
1075 .await?;
1076 Ok(())
1077 }
1078
1079 #[tracing::instrument(skip(self), level = "debug")]
1080 pub async fn request_more_tickmarks(
1081 &self,
1082 chart_session: &str,
1083 series_id: &str,
1084 num: u64,
1085 ) -> Result<()> {
1086 self.send(
1087 "request_more_tickmarks",
1088 &payload!(chart_session, series_id, num),
1089 )
1090 .await?;
1091 Ok(())
1092 }
1093
1094 #[tracing::instrument(skip(self), level = "debug")]
1095 pub async fn create_replay_session(&self, replay_session: &str) -> Result<()> {
1096 self.send("replay_create_session", &payload!(replay_session))
1097 .await?;
1098 Ok(())
1099 }
1100
1101 #[tracing::instrument(skip(self), level = "debug")]
1102 #[builder]
1103 pub async fn add_replay_series(
1104 &self,
1105 replay_session: &str,
1106 request_id: &str,
1107 instrument: &str, adjustment: Option<MarketAdjustment>,
1109 session_type: Option<SessionType>,
1110 currency: Option<Currency>,
1111 interval: Interval,
1112 ) -> Result<()> {
1113 let sym_init = symbol_init()
1114 .instrument(instrument)
1115 .maybe_adjustment(adjustment)
1116 .maybe_currency(currency)
1117 .maybe_session_type(session_type)
1118 .call()?;
1119 let payloads =
1120 build_add_replay_series_payload(replay_session, request_id, sym_init, interval);
1121 self.send("replay_add_series", &payloads).await?;
1122 Ok(())
1123 }
1124
1125 #[tracing::instrument(skip(self), level = "debug")]
1126 pub async fn delete_replay_session(&self, replay_session: &str) -> Result<()> {
1127 self.send("replay_delete_session", &payload!(replay_session))
1128 .await?;
1129 Ok(())
1130 }
1131
1132 #[tracing::instrument(skip(self), level = "debug")]
1133 pub async fn replay_step(
1134 &self,
1135 replay_session: &str,
1136 request_id: &str,
1137 step: u64,
1138 ) -> Result<()> {
1139 let payloads = build_replay_step_payload(replay_session, request_id, step);
1140 self.send("replay_step", &payloads).await?;
1141 Ok(())
1142 }
1143
1144 #[tracing::instrument(skip(self), level = "debug")]
1145 pub async fn replay_start(
1146 &self,
1147 replay_session: &str,
1148 request_id: &str,
1149 interval: u64,
1150 ) -> Result<()> {
1151 let payloads = build_replay_start_payload(replay_session, request_id, interval);
1152 self.send("replay_start", &payloads).await?;
1153 Ok(())
1154 }
1155
1156 #[tracing::instrument(skip(self), level = "debug")]
1157 pub async fn replay_stop(&self, replay_session: &str, request_id: &str) -> Result<()> {
1158 let payloads = build_replay_stop_payload(replay_session, request_id);
1159 self.send("replay_stop", &payloads).await?;
1160 Ok(())
1161 }
1162
1163 #[tracing::instrument(skip(self), level = "debug")]
1164 pub async fn replay_reset(
1165 &self,
1166 replay_session: &str,
1167 request_id: &str,
1168 timestamp: i64,
1169 ) -> Result<()> {
1170 let payloads = build_replay_reset_payload(replay_session, request_id, timestamp);
1171 self.send("replay_reset", &payloads).await?;
1172 Ok(())
1173 }
1174
1175 #[tracing::instrument(skip(self), level = "debug")]
1176 #[builder]
1177 pub async fn create_study(
1178 &self,
1179 chart_session: &str,
1180 study_ids: &[&str; 2],
1181 chart_series_id: &str,
1182 study: StudyConfiguration,
1183 ) -> Result<()> {
1184 let mut payloads: Vec<Value> = vec![
1185 Value::from(chart_session),
1186 Value::from(study_ids[0]),
1187 Value::from(study_ids[1]),
1188 Value::from(chart_series_id),
1189 ];
1190
1191 match study {
1192 StudyConfiguration::Pine(pine_indicator) => {
1193 payloads.push(Value::from(pine_indicator.script_type.to_string()));
1194 payloads.push(pine_indicator.to_study_inputs()?);
1195 }
1196 StudyConfiguration::Builtin(study_name, study_config) => {
1197 payloads.push(Value::from(study_name));
1198 payloads.push(json!(study_config));
1199 }
1200 }
1201
1202 self.send("create_study", &payloads).await?;
1203 Ok(())
1204 }
1205
1206 #[tracing::instrument(skip(self), level = "debug")]
1207 #[builder]
1208 pub async fn modify_study(
1209 &self,
1210 chart_session: &str,
1211 study_ids: &[&str; 2],
1212 study: StudyConfiguration,
1213 ) -> Result<()> {
1214 let inputs = match study {
1215 StudyConfiguration::Pine(pine_indicator) => pine_indicator.to_study_inputs()?,
1216 StudyConfiguration::Builtin(_study_name, study_config) => json!(study_config),
1217 };
1218 let payloads =
1219 build_modify_study_payload(chart_session, study_ids[0], study_ids[1], inputs);
1220 self.send("modify_study", &payloads).await?;
1221 Ok(())
1222 }
1223
1224 #[tracing::instrument(skip(self), level = "debug")]
1225 pub async fn remove_study(&self, chart_session: &str, study_id: &str) -> Result<()> {
1226 self.send("remove_study", &payload!(chart_session, study_id))
1227 .await?;
1228 Ok(())
1229 }
1230
1231 pub async fn delete(&self) -> Result<()> {
1232 if let Err(e) = self.close().await {
1234 error!("Failed to close socket: {:?}", e);
1235 return Err(e);
1236 }
1237
1238 debug!("WebSocket client deleted successfully");
1239 Ok(())
1240 }
1241
1242 pub async fn subscribe(&self) -> Result<()> {
1243 let read = self.read.lock().await;
1244 if let Err(e) = self.event_loop(read).await {
1245 error!("Event loop failed: {}", e);
1246 self.is_closed.store(true, Ordering::Relaxed);
1247 self.cancellation.cancel();
1248 return Err(e);
1249 }
1250 Ok(())
1251 }
1252
1253 pub async fn closed_notifier(&self) {
1254 self.cancellation.cancelled().await;
1255 }
1256
1257 pub async fn try_ping(&self) -> Result<()> {
1259 if self.is_closed() {
1260 return Ok(());
1261 }
1262 self.ping(&Message::Ping(Vec::new().into()))
1263 .await
1264 .map_err(|e| Error::WebSocket(ustr(&format!("{e}"))))?;
1265 Ok(())
1266 }
1267
1268 pub async fn get_connection_stats(&self) -> Value {
1269 let health_metrics = self.health_metrics.read().await;
1270 let is_closed = self.is_closed();
1271
1272 json!({
1273 "consecutive_errors": self.error_stats.get_consecutive_errors(),
1274 "total_errors": self.error_stats.total_errors.load(Ordering::SeqCst),
1275 "recovery_attempts": self.error_stats.recovery_attempts.load(Ordering::SeqCst),
1276 "connection_drops": self.error_stats.connection_drops.load(Ordering::SeqCst),
1277 "health": format!("{:?}", health_metrics.health),
1278 "avg_response_time_ms": health_metrics.avg_response_time.as_millis(),
1279 "circuit_breaker_open": self.circuit_breaker_open.load(Ordering::Relaxed),
1280 "is_closed": is_closed,
1281 })
1282 }
1283}
1284
1285impl<T: Handler> Socket for WebSocketClient<T> {
1286 async fn event_loop(
1287 &self,
1288 mut read: MutexGuard<'_, SplitStream<WebSocketStream<MaybeTlsStream<TcpStream>>>>,
1289 ) -> Result<()> {
1290 trace!("WebSocket event loop started");
1291
1292 let mut ping_interval = tokio::time::interval(self.error_config.ping_interval);
1293 ping_interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
1294
1295 loop {
1296 if self.is_closed.load(Ordering::Relaxed) {
1297 trace!("WebSocket is closed, ending event loop");
1298 break;
1299 }
1300
1301 tokio::select! {
1302 message_result = timeout(Duration::from_secs(30), read.next()) => {
1304 match message_result {
1305 Ok(Some(Ok(message))) => {
1306 trace!("Received message: {:?}", message);
1307 if let Err(e) = self.handle_raw_messages(message).await {
1308 self.handle_error(e, ustr("handle_raw_messages")).await?;
1309 } else {
1310 if self.error_stats.get_consecutive_errors() > 0 {
1312 self.error_stats.reset_consecutive();
1313 debug!("Reset consecutive errors after successful message processing");
1314 }
1315 self.error_stats.update_last_successful_message().await;
1316 }
1317 }
1318 Ok(Some(Err(e))) => {
1319 error!("Error reading message: {:#?}", e);
1320 self.error_stats.connection_drops.fetch_add(1, Ordering::SeqCst);
1321 self.handle_error(
1322 Error::WebSocket(e.to_string().into()),
1323 ustr("event_loop_read"),
1324 ).await?;
1325
1326 if e.to_string().contains("ConnectionClosed") ||
1328 e.to_string().contains("ConnectionReset") {
1329 break;
1330 }
1331 }
1332 Ok(None) => {
1333 info!("WebSocket stream ended");
1334 self.is_closed.store(true, Ordering::Relaxed);
1335 break;
1336 }
1337 Err(_) => {
1338 warn!("WebSocket read timeout, checking connection health");
1339 if let Err(e) = self.perform_health_check().await {
1340 warn!("Health check failed during timeout: {}", e);
1341 if matches!(self.classify_error_severity(&e, "health_check").await, ErrorSeverity::Fatal) {
1343 break;
1344 }
1345 }
1346 }
1347 }
1348 }
1349
1350 _ = ping_interval.tick() => {
1352 if !self.is_closed()
1353 && let Err(e) = self.try_ping().await
1354 {
1355 warn!("Periodic ping failed: {}", e);
1356 self.handle_error(e, ustr("periodic_ping")).await?;
1357 }
1358 }
1359
1360 _ = self.cancellation.cancelled() => {
1362 info!("Event loop cancelled");
1363 break;
1364 }
1365 }
1366 }
1367
1368 trace!("WebSocket event loop ended");
1369 Ok(())
1370 }
1371
1372 async fn handle_raw_messages(&self, raw: Message) -> Result<()> {
1373 match &raw {
1374 Message::Text(text) => {
1375 trace!("Received text message: {}", text);
1376 self.handle_parsed_messages(parse_packet(text), &raw)
1377 .await?;
1378 }
1379 Message::Close(msg) => {
1380 warn!("Connection closed with code: {:?}", msg);
1381 self.is_closed.store(true, Ordering::Relaxed);
1382 self.cancellation.cancel();
1383 }
1384 Message::Binary(msg) => {
1385 debug!("Received binary message: {:?}", msg);
1386 }
1388 Message::Ping(msg) => {
1389 trace!("Received ping message: {:?}", msg);
1390 }
1391 Message::Pong(msg) => {
1392 trace!("Received pong message: {:?}", msg);
1393 }
1394 Message::Frame(f) => {
1395 debug!("Received frame message: {:?}", f);
1396 }
1397 }
1398 Ok(())
1399 }
1400
1401 async fn handle_parsed_messages(
1402 &self,
1403 messages: Vec<SocketMessage<SocketMessageDe>>,
1404 _raw: &Message,
1405 ) -> Result<()> {
1406 for message in messages {
1407 match message {
1408 SocketMessage::SocketServerInfo(info) => {
1409 trace!("received server info: {:?}", info);
1410 }
1411 SocketMessage::SocketMessage(msg) => {
1412 trace!(
1413 "Processing socket message: method={}, params={:?}",
1414 msg.m, msg.p
1415 );
1416 if let Err(e) = self.handle_message_data(msg).await {
1417 self.handle_error(e, ustr("handle_message_data")).await?;
1418 }
1419 }
1420 SocketMessage::Heartbeat(counter) => {
1421 debug!("handling heartbeat message: {counter}");
1422 let echo = message.heartbeat_echo().unwrap_or_else(|| {
1423 let h = format!("~h~{counter}");
1424 format!("~m~{}~m~{h}", h.len())
1425 });
1426 if let Err(e) = self.send_raw_message(&echo).await {
1427 self.handle_error(e, ustr("heartbeat_echo")).await?;
1428 }
1429 }
1430 SocketMessage::Other(value) => {
1431 trace!("Received other message: {:?}", value);
1432 if let Ok(server_info) = SocketServerInfo::deserialize(&value) {
1433 info!("{}", server_info);
1434 } else {
1435 warn!("Received unrecognized message: {:?}", value);
1436 }
1437 }
1438 SocketMessage::Unknown(s) => {
1439 warn!("unknown message: {:?}", s);
1440 }
1441 }
1442 }
1443 Ok(())
1444 }
1445
1446 #[tracing::instrument(skip(self), level = "trace")]
1447 async fn handle_message_data(&self, message: SocketMessageDe) -> Result<()> {
1448 dispatch_message_data(&self.handler, message);
1449 Ok(())
1450 }
1451
1452 async fn handle_error(&self, error: Error, context: Ustr) -> Result<()> {
1453 let context_str = context.as_str();
1454
1455 let _consecutive_errors = self.error_stats.increment_error();
1457 self.error_stats.update_last_error_time().await;
1458
1459 let severity = self.classify_error_severity(&error, context_str).await;
1461
1462 self.log_error(&error, context_str, &severity);
1464
1465 self.notify_error_handlers(&error, context_str, severity)
1467 .await;
1468
1469 match severity {
1471 ErrorSeverity::Trace | ErrorSeverity::Minor => {
1472 Ok(())
1474 }
1475
1476 ErrorSeverity::Moderate => {
1477 match self.attempt_error_recovery(severity, &error).await {
1479 Ok(true) => Ok(()),
1480 Ok(false) => {
1481 warn!("Moderate error recovery failed, continuing anyway");
1482 Ok(())
1483 }
1484 Err(recovery_err) => {
1485 warn!("Error during moderate recovery: {}", recovery_err);
1486 Ok(()) }
1488 }
1489 }
1490
1491 ErrorSeverity::Critical | ErrorSeverity::Fatal => {
1492 match self.attempt_error_recovery(severity, &error).await {
1494 Ok(recovered) => {
1495 if !recovered {
1496 error!(
1497 "Failed to recover from {} error",
1498 if matches!(severity, ErrorSeverity::Critical) {
1499 "critical"
1500 } else {
1501 "fatal"
1502 }
1503 );
1504 return Err(Error::Internal(ustr("Error recovery failed")));
1505 }
1506 Ok(())
1507 }
1508 Err(recovery_err) => {
1509 error!("Error during recovery attempt: {}", recovery_err);
1510 Err(recovery_err)
1511 }
1512 }
1513 }
1514 }
1515 }
1516}
1517
1518pub fn dispatch_message_data<H: Handler>(handler: &H, message: SocketMessageDe) {
1520 let event = TradingViewDataEvent::from(message.m);
1521 handler.handle_events(event, &message.p);
1522 if event == TradingViewDataEvent::OnQuoteData {
1523 handler.handle_quote_data(&message.p);
1524 }
1525}
1526
1527pub fn build_modify_study_payload(
1531 chart_session: &str,
1532 study_id: &str,
1533 study_sub_id: &str,
1534 inputs: Value,
1535) -> Vec<Value> {
1536 vec![
1537 Value::from(chart_session),
1538 Value::from(study_id),
1539 Value::from(study_sub_id),
1540 inputs,
1541 ]
1542}
1543
1544pub fn build_replay_start_payload(
1548 replay_session: &str,
1549 request_id: &str,
1550 interval_ms: u64,
1551) -> Vec<Value> {
1552 vec![
1553 Value::from(replay_session),
1554 Value::from(request_id),
1555 Value::from(interval_ms),
1556 ]
1557}
1558
1559pub fn build_replay_step_payload(replay_session: &str, request_id: &str, step: u64) -> Vec<Value> {
1563 vec![
1564 Value::from(replay_session),
1565 Value::from(request_id),
1566 Value::from(step),
1567 ]
1568}
1569
1570pub fn build_replay_stop_payload(replay_session: &str, request_id: &str) -> Vec<Value> {
1574 vec![Value::from(replay_session), Value::from(request_id)]
1575}
1576
1577pub fn build_replay_reset_payload(
1581 replay_session: &str,
1582 request_id: &str,
1583 timestamp: i64,
1584) -> Vec<Value> {
1585 vec![
1586 Value::from(replay_session),
1587 Value::from(request_id),
1588 Value::from(timestamp),
1589 ]
1590}
1591
1592pub fn build_add_replay_series_payload(
1596 replay_session: &str,
1597 request_id: &str,
1598 symbol_init: String,
1599 interval: Interval,
1600) -> Vec<Value> {
1601 vec![
1602 Value::from(replay_session),
1603 Value::from(request_id),
1604 Value::from(symbol_init),
1605 Value::from(interval.to_string()),
1606 ]
1607}
1608
1609#[cfg(test)]
1610mod tests {
1611 use super::*;
1612 use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
1613
1614 struct MockHandler {
1615 event_called: AtomicBool,
1616 quote_called: AtomicBool,
1617 quote_count: AtomicUsize,
1618 }
1619
1620 impl MockHandler {
1621 fn new() -> Self {
1622 Self {
1623 event_called: AtomicBool::new(false),
1624 quote_called: AtomicBool::new(false),
1625 quote_count: AtomicUsize::new(0),
1626 }
1627 }
1628 }
1629
1630 impl Handler for MockHandler {
1631 fn handle_events(&self, _event: TradingViewDataEvent, _message: &[Value]) {
1632 self.event_called.store(true, Ordering::SeqCst);
1633 }
1634 fn handle_quote_data(&self, _message: &[Value]) {
1635 self.quote_called.store(true, Ordering::SeqCst);
1636 self.quote_count.fetch_add(1, Ordering::SeqCst);
1637 }
1638 fn handle_series_data(&self, _event: TradingViewDataEvent, _messages: &[Value]) {}
1639 fn notify_error(&self, _error: Error, _message: &[Value]) {}
1640 }
1641
1642 #[test]
1643 fn test_modify_study_payload_has_4_args() {
1644 let inputs = serde_json::json!({ "text": "//@version=5\nindicator('Test')" });
1645 let payload = build_modify_study_payload("cs_123", "st6", "st1", inputs.clone());
1646
1647 assert_eq!(
1648 payload.len(),
1649 4,
1650 "modify_study payload length must be exactly 4"
1651 );
1652 assert_eq!(payload[0], Value::from("cs_123"));
1653 assert_eq!(payload[1], Value::from("st6"));
1654 assert_eq!(payload[2], Value::from("st1"));
1655 assert_eq!(payload[3], inputs);
1656 }
1657
1658 #[test]
1659 fn test_replay_start_payload_interval_is_numeric_millis() {
1660 let payload = build_replay_start_payload("rs_100", "req_replay_1", 250);
1661
1662 assert_eq!(payload.len(), 3, "replay_start payload length must be 3");
1663 assert_eq!(payload[0], Value::from("rs_100"));
1664 assert_eq!(payload[1], Value::from("req_replay_1"));
1665 assert!(
1666 payload[2].is_number(),
1667 "third argument must be numeric milliseconds"
1668 );
1669 assert_eq!(payload[2].as_u64(), Some(250));
1670 }
1671
1672 #[test]
1673 fn test_replay_other_payloads() {
1674 let step_payload = build_replay_step_payload("rs_100", "req_step_1", 5);
1675 assert_eq!(step_payload.len(), 3);
1676 assert_eq!(step_payload[2], Value::from(5u64));
1677
1678 let stop_payload = build_replay_stop_payload("rs_100", "req_stop_1");
1679 assert_eq!(stop_payload.len(), 2);
1680 assert_eq!(stop_payload[0], Value::from("rs_100"));
1681 assert_eq!(stop_payload[1], Value::from("req_stop_1"));
1682
1683 let reset_payload = build_replay_reset_payload("rs_100", "req_reset_1", 1_700_000_000);
1684 assert_eq!(reset_payload.len(), 3);
1685 assert_eq!(reset_payload[2], Value::from(1_700_000_000i64));
1686
1687 let add_payload = build_add_replay_series_payload(
1688 "rs_100",
1689 "req_add_1",
1690 "={\"symbol\":\"NASDAQ:AAPL\"}".to_string(),
1691 Interval::OneDay,
1692 );
1693 assert_eq!(add_payload.len(), 4);
1694 assert_eq!(add_payload[3], Value::from("1D"));
1695 }
1696
1697 #[test]
1698 fn test_heartbeat_echo_framed_format() {
1699 let msg = SocketMessage::<SocketMessageDe>::Heartbeat(42);
1700 assert_eq!(msg.heartbeat_echo(), Some("~m~5~m~~h~42".to_string()));
1701 }
1702
1703 #[test]
1704 fn test_quote_data_handler_routing() {
1705 let handler = MockHandler::new();
1706 let quote_msg = SocketMessageDe {
1707 m: ustr::ustr("qsd"),
1708 p: vec![
1709 serde_json::json!("quote_session_1"),
1710 serde_json::json!({"s": "ok"}),
1711 ],
1712 t: 0,
1713 t_ms: 0,
1714 };
1715
1716 dispatch_message_data(&handler, quote_msg);
1717 assert!(
1718 handler.event_called.load(Ordering::SeqCst),
1719 "handle_events must be called"
1720 );
1721 assert!(
1722 handler.quote_called.load(Ordering::SeqCst),
1723 "handle_quote_data must be called for qsd"
1724 );
1725 assert_eq!(handler.quote_count.load(Ordering::SeqCst), 1);
1726
1727 let non_quote_msg = SocketMessageDe {
1728 m: ustr::ustr("timescale_update"),
1729 p: vec![serde_json::json!("cs_1")],
1730 t: 0,
1731 t_ms: 0,
1732 };
1733 dispatch_message_data(&handler, non_quote_msg);
1734 assert_eq!(
1735 handler.quote_count.load(Ordering::SeqCst),
1736 1,
1737 "quote_data should not be called for non-qsd"
1738 );
1739 }
1740}